refactor: select in account_service (RegisterService class) (#34500)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Renzo 2026-04-03 08:21:26 +02:00 committed by GitHub
parent 4fedd43af5
commit da3b0caf5e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 29 additions and 47 deletions

View File

@ -8,7 +8,7 @@ from hashlib import sha256
from typing import Any, TypedDict, cast from typing import Any, TypedDict, cast
from pydantic import BaseModel, TypeAdapter from pydantic import BaseModel, TypeAdapter
from sqlalchemy import func, select from sqlalchemy import delete, func, select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@ -1392,10 +1392,10 @@ class RegisterService:
db.session.add(dify_setup) db.session.add(dify_setup)
db.session.commit() db.session.commit()
except Exception as e: except Exception as e:
db.session.query(DifySetup).delete() db.session.execute(delete(DifySetup))
db.session.query(TenantAccountJoin).delete() db.session.execute(delete(TenantAccountJoin))
db.session.query(Account).delete() db.session.execute(delete(Account))
db.session.query(Tenant).delete() db.session.execute(delete(Tenant))
db.session.commit() db.session.commit()
logger.exception("Setup account failed, email: %s, name: %s", email, name) logger.exception("Setup account failed, email: %s, name: %s", email, name)
@ -1496,7 +1496,11 @@ class RegisterService:
TenantService.switch_tenant(account, tenant.id) TenantService.switch_tenant(account, tenant.id)
else: else:
TenantService.check_member_permission(tenant, inviter, account, "add") TenantService.check_member_permission(tenant, inviter, account, "add")
ta = db.session.query(TenantAccountJoin).filter_by(tenant_id=tenant.id, account_id=account.id).first() ta = db.session.scalar(
select(TenantAccountJoin)
.where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == account.id)
.limit(1)
)
if not ta: if not ta:
TenantService.create_tenant_member(tenant, account, role) TenantService.create_tenant_member(tenant, account, role)
@ -1553,21 +1557,18 @@ class RegisterService:
if not invitation_data: if not invitation_data:
return None return None
tenant = ( tenant = db.session.scalar(
db.session.query(Tenant) select(Tenant).where(Tenant.id == invitation_data["workspace_id"], Tenant.status == "normal").limit(1)
.where(Tenant.id == invitation_data["workspace_id"], Tenant.status == "normal")
.first()
) )
if not tenant: if not tenant:
return None return None
tenant_account = ( tenant_account = db.session.execute(
db.session.query(Account, TenantAccountJoin.role) select(Account, TenantAccountJoin.role)
.join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id)
.where(Account.email == invitation_data["email"], TenantAccountJoin.tenant_id == tenant.id) .where(Account.email == invitation_data["email"], TenantAccountJoin.tenant_id == tenant.id)
.first() ).first()
)
if not tenant_account: if not tenant_account:
return None return None

View File

@ -1034,7 +1034,7 @@ class TestRegisterService:
) )
# Verify rollback operations were called # Verify rollback operations were called
mock_db_dependencies["db"].session.query.assert_called() mock_db_dependencies["db"].session.execute.assert_called()
# ==================== Registration Tests ==================== # ==================== Registration Tests ====================
@ -1599,10 +1599,8 @@ class TestRegisterService:
mock_session_class.return_value.__exit__.return_value = None mock_session_class.return_value.__exit__.return_value = None
mock_lookup.return_value = mock_existing_account mock_lookup.return_value = mock_existing_account
# Mock the db.session.query for TenantAccountJoin # Mock scalar for TenantAccountJoin lookup - no existing member
mock_db_query = MagicMock() mock_db_dependencies["db"].session.scalar.return_value = None
mock_db_query.filter_by.return_value.first.return_value = None # No existing member
mock_db_dependencies["db"].session.query.return_value = mock_db_query
# Mock TenantService methods # Mock TenantService methods
with ( with (
@ -1777,14 +1775,9 @@ class TestRegisterService:
} }
mock_get_invitation_by_token.return_value = invitation_data mock_get_invitation_by_token.return_value = invitation_data
# Mock database queries - complex query mocking # Mock scalar for tenant lookup, execute for account+role lookup
mock_query1 = MagicMock() mock_db_dependencies["db"].session.scalar.return_value = mock_tenant
mock_query1.where.return_value.first.return_value = mock_tenant mock_db_dependencies["db"].session.execute.return_value.first.return_value = (mock_account, "normal")
mock_query2 = MagicMock()
mock_query2.join.return_value.where.return_value.first.return_value = (mock_account, "normal")
mock_db_dependencies["db"].session.query.side_effect = [mock_query1, mock_query2]
# Execute test # Execute test
result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123") result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123")
@ -1816,10 +1809,8 @@ class TestRegisterService:
} }
mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode() mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode()
# Mock database queries - no tenant found # Mock scalar for tenant lookup - not found
mock_query = MagicMock() mock_db_dependencies["db"].session.scalar.return_value = None
mock_query.filter.return_value.first.return_value = None
mock_db_dependencies["db"].session.query.return_value = mock_query
# Execute test # Execute test
result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123") result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123")
@ -1842,14 +1833,9 @@ class TestRegisterService:
} }
mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode() mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode()
# Mock database queries # Mock scalar for tenant, execute for account+role
mock_query1 = MagicMock() mock_db_dependencies["db"].session.scalar.return_value = mock_tenant
mock_query1.filter.return_value.first.return_value = mock_tenant mock_db_dependencies["db"].session.execute.return_value.first.return_value = None # No account found
mock_query2 = MagicMock()
mock_query2.join.return_value.where.return_value.first.return_value = None # No account found
mock_db_dependencies["db"].session.query.side_effect = [mock_query1, mock_query2]
# Execute test # Execute test
result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123") result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123")
@ -1875,14 +1861,9 @@ class TestRegisterService:
} }
mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode() mock_redis_dependencies.get.return_value = json.dumps(invitation_data).encode()
# Mock database queries # Mock scalar for tenant, execute for account+role
mock_query1 = MagicMock() mock_db_dependencies["db"].session.scalar.return_value = mock_tenant
mock_query1.filter.return_value.first.return_value = mock_tenant mock_db_dependencies["db"].session.execute.return_value.first.return_value = (mock_account, "normal")
mock_query2 = MagicMock()
mock_query2.join.return_value.where.return_value.first.return_value = (mock_account, "normal")
mock_db_dependencies["db"].session.query.side_effect = [mock_query1, mock_query2]
# Execute test # Execute test
result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123") result = RegisterService.get_invitation_if_token_valid("tenant-456", "test@example.com", "token-123")