From 908dc703ec3024040974fedceb8777a91b5a0aa2 Mon Sep 17 00:00:00 2001 From: Parman Mohammadalizadeh Date: Thu, 30 Jul 2026 20:29:50 +0330 Subject: [PATCH] chore(api): consolidate chained .where() calls into single calls (#39819) --- api/controllers/mcp/mcp.py | 6 +++--- api/controllers/service_api/wraps.py | 11 ++++++----- api/core/app/task_pipeline/message_cycle_manager.py | 6 +----- api/core/rag/retrieval/dataset_retrieval.py | 7 +++++-- api/extensions/ext_login.py | 9 +++++---- api/models/account.py | 7 ++----- api/services/account_service.py | 3 +-- api/services/annotation_service.py | 4 ++-- api/services/conversation_service.py | 12 +++++------- ...est_sqlalchemy_workflow_trigger_log_repository.py | 3 +-- 10 files changed, 31 insertions(+), 37 deletions(-) diff --git a/api/controllers/mcp/mcp.py b/api/controllers/mcp/mcp.py index 45a5a4e6899..95aa2be8218 100644 --- a/api/controllers/mcp/mcp.py +++ b/api/controllers/mcp/mcp.py @@ -240,9 +240,9 @@ class MCPAppApi(Resource): with sessionmaker(db.engine, expire_on_commit=False).begin() as session: return session.scalar( select(EndUser) - .where(EndUser.tenant_id == tenant_id) - .where(EndUser.session_id == mcp_server_id) - .where(EndUser.type == EndUserType.MCP) + .where( + EndUser.tenant_id == tenant_id, EndUser.session_id == mcp_server_id, EndUser.type == EndUserType.MCP + ) .limit(1) ) diff --git a/api/controllers/service_api/wraps.py b/api/controllers/service_api/wraps.py index c3c8e02e438..1f84b533bd4 100644 --- a/api/controllers/service_api/wraps.py +++ b/api/controllers/service_api/wraps.py @@ -322,11 +322,12 @@ def validate_dataset_token[R](view: Callable[..., R]) -> Callable[..., R]: raise Forbidden("Dataset api access is not enabled.") tenant_account_join = db.session.execute( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == api_token.tenant_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.role.in_(["owner"])) - .where(Tenant.status == TenantStatus.NORMAL) + select(Tenant, TenantAccountJoin).where( + Tenant.id == api_token.tenant_id, + TenantAccountJoin.tenant_id == Tenant.id, + TenantAccountJoin.role.in_(["owner"]), + Tenant.status == TenantStatus.NORMAL, + ) ).one_or_none() # TODO: only owner information is required, so only one is returned. if tenant_account_join: tenant, ta = tenant_account_join diff --git a/api/core/app/task_pipeline/message_cycle_manager.py b/api/core/app/task_pipeline/message_cycle_manager.py index 3bffe83ee28..91b67a03e47 100644 --- a/api/core/app/task_pipeline/message_cycle_manager.py +++ b/api/core/app/task_pipeline/message_cycle_manager.py @@ -66,11 +66,7 @@ class MessageCycleManager: # Use SQLAlchemy 2.x style session.scalar(select(...)) with session_factory.create_session() as session: message_file = session.scalar( - select(MessageFile) - .where( - MessageFile.message_id == message_id, - ) - .where(MessageFile.belongs_to == "assistant") + select(MessageFile).where(MessageFile.message_id == message_id, MessageFile.belongs_to == "assistant") ) if message_file: diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index 317e8d98e38..5c1f2b6e6d9 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -2003,8 +2003,11 @@ class DatasetRetrieval: results = session.scalars( select(Dataset) .outerjoin(subquery, Dataset.id == subquery.c.dataset_id) - .where(Dataset.tenant_id == tenant_id, Dataset.id.in_(dataset_ids)) - .where((subquery.c.available_document_count > 0) | (Dataset.provider == "external")) + .where( + Dataset.tenant_id == tenant_id, + Dataset.id.in_(dataset_ids), + (subquery.c.available_document_count > 0) | (Dataset.provider == "external"), + ) ).all() available_datasets = [] diff --git a/api/extensions/ext_login.py b/api/extensions/ext_login.py index fddefb14f52..caee32d3174 100644 --- a/api/extensions/ext_login.py +++ b/api/extensions/ext_login.py @@ -72,10 +72,11 @@ def _load_user_from_request(request_from_flask_login: Request, session: Session) workspace_id = request.headers.get("X-WORKSPACE-ID") if workspace_id: tenant_account_join = session.execute( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == workspace_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.role == "owner") + select(Tenant, TenantAccountJoin).where( + Tenant.id == workspace_id, + TenantAccountJoin.tenant_id == Tenant.id, + TenantAccountJoin.role == "owner", + ) ).one_or_none() if tenant_account_join: tenant, ta = tenant_account_join diff --git a/api/models/account.py b/api/models/account.py index 919ee7da820..5bf4e3e641e 100644 --- a/api/models/account.py +++ b/api/models/account.py @@ -165,11 +165,8 @@ class Account(UserMixin, TypeBase): def set_tenant_id_with_session(self, tenant_id: str, *, session: Session) -> None: """Set the current tenant by id using the caller-owned session.""" - query = ( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == tenant_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.account_id == self.id) + query = select(Tenant, TenantAccountJoin).where( + Tenant.id == tenant_id, TenantAccountJoin.tenant_id == Tenant.id, TenantAccountJoin.account_id == self.id ) tenant_account_join = session.execute(query).first() if not tenant_account_join: diff --git a/api/services/account_service.py b/api/services/account_service.py index cd89ddba2ab..d58f18702c3 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -1624,8 +1624,7 @@ class TenantService: select(Account, TenantAccountJoin.role) .select_from(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) - .where(TenantAccountJoin.tenant_id == tenant.id) - .where(TenantAccountJoin.role == "dataset_operator") + .where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.role == "dataset_operator") ) # Initialize an empty list to store the updated accounts diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 947c35fc0a0..087bbd9be2b 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -230,12 +230,12 @@ class AppAnnotationService: escaped_keyword = escape_like_pattern(keyword) stmt = ( select(MessageAnnotation) - .where(MessageAnnotation.app_id == app_id) .where( + MessageAnnotation.app_id == app_id, or_( MessageAnnotation.question.ilike(f"%{escaped_keyword}%", escape="\\"), MessageAnnotation.content.ilike(f"%{escaped_keyword}%", escape="\\"), - ) + ), ) .order_by(MessageAnnotation.created_at.desc(), MessageAnnotation.id.desc()) ) diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index 3cd53f8249f..406af068194 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -256,8 +256,7 @@ class ConversationService: stmt = ( select(ConversationVariable) - .where(ConversationVariable.app_id == app_model.id) - .where(ConversationVariable.conversation_id == conversation.id) + .where(ConversationVariable.app_id == app_model.id, ConversationVariable.conversation_id == conversation.id) .order_by(ConversationVariable.created_at) ) @@ -342,11 +341,10 @@ class ConversationService: conversation = cls.get_conversation(app_model, conversation_id, user, session=session) # Get the existing conversation variable - stmt = ( - select(ConversationVariable) - .where(ConversationVariable.app_id == app_model.id) - .where(ConversationVariable.conversation_id == conversation.id) - .where(ConversationVariable.id == variable_id) + stmt = select(ConversationVariable).where( + ConversationVariable.app_id == app_model.id, + ConversationVariable.conversation_id == conversation.id, + ConversationVariable.id == variable_id, ) existing_variable = session.scalar(stmt) diff --git a/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py b/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py index 0c4d75359e4..e7ac10fc986 100644 --- a/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py +++ b/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py @@ -125,8 +125,7 @@ def test_delete_by_run_ids_empty_short_circuits(db_session_with_containers: Sess remaining_count = db_session_with_containers.scalar( select(func.count()) .select_from(WorkflowTriggerLog) - .where(WorkflowTriggerLog.tenant_id == tenant_id) - .where(WorkflowTriggerLog.workflow_run_id == run_id) + .where(WorkflowTriggerLog.tenant_id == tenant_id, WorkflowTriggerLog.workflow_run_id == run_id) ) assert remaining_count == 1 finally: