From 68d8328b9c39e8aea5e257368e9d2437b0c02a08 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Wed, 8 Jul 2026 12:07:27 +0900 Subject: [PATCH] chore: clean Db session from service (#38227) Co-authored-by: chariri Co-authored-by: WH-2099 --- api/commands/account.py | 8 +- api/commands/data_migration.py | 15 +- api/commands/plugin.py | 1 + api/commands/rbac.py | 35 +- api/controllers/common/app_access.py | 3 +- api/controllers/console/agent/composer.py | 20 +- api/controllers/console/agent/roster.py | 16 +- api/controllers/console/app/agent.py | 12 +- .../console/app/agent_app_feature.py | 2 +- .../console/app/agent_app_sandbox.py | 4 + .../console/app/agent_config_inspector.py | 6 +- .../console/app/agent_drive_inspector.py | 35 +- api/controllers/console/app/annotation.py | 25 +- api/controllers/console/app/app.py | 30 +- api/controllers/console/app/audio.py | 2 +- api/controllers/console/app/conversation.py | 4 +- api/controllers/console/app/message.py | 7 +- api/controllers/console/app/ops_trace.py | 17 +- .../console/app/permission_keys.py | 5 +- api/controllers/console/app/workflow.py | 63 +- .../console/app/workflow_comment.py | 2 +- .../console/app/workflow_draft_variable.py | 10 +- .../app/workflow_node_output_inspector.py | 10 +- api/controllers/console/auth/activate.py | 8 +- .../console/auth/data_source_bearer_auth.py | 6 +- .../console/auth/email_register.py | 8 +- .../console/auth/forgot_password.py | 10 +- api/controllers/console/auth/login.py | 26 +- api/controllers/console/auth/oauth.py | 16 +- api/controllers/console/auth/oauth_server.py | 5 +- api/controllers/console/billing/billing.py | 4 +- .../console/datasets/data_source.py | 10 +- api/controllers/console/datasets/datasets.py | 56 +- .../console/datasets/datasets_document.py | 101 +-- .../console/datasets/datasets_segments.py | 105 +-- api/controllers/console/datasets/external.py | 13 +- .../console/datasets/hit_testing_base.py | 4 +- api/controllers/console/datasets/metadata.py | 36 +- .../datasets/rag_pipeline/datasource_auth.py | 11 +- .../datasource_content_preview.py | 3 +- .../datasets/rag_pipeline/rag_pipeline.py | 20 +- .../rag_pipeline/rag_pipeline_datasets.py | 6 +- .../rag_pipeline_draft_variable.py | 6 +- .../rag_pipeline/rag_pipeline_workflow.py | 54 +- api/controllers/console/datasets/wraps.py | 8 +- api/controllers/console/explore/audio.py | 2 +- .../console/explore/conversation.py | 8 +- .../console/explore/installed_app.py | 2 +- api/controllers/console/explore/message.py | 9 +- api/controllers/console/explore/parameter.py | 3 +- .../console/explore/recommended_app.py | 6 +- .../console/explore/saved_message.py | 10 +- api/controllers/console/explore/trial.py | 18 +- api/controllers/console/extension.py | 16 +- api/controllers/console/init_validate.py | 2 +- api/controllers/console/setup.py | 4 +- api/controllers/console/socketio/workflow.py | 4 +- api/controllers/console/tag/tags.py | 14 +- api/controllers/console/workspace/account.py | 20 +- .../workspace/load_balancing_config.py | 3 + api/controllers/console/workspace/members.py | 22 +- .../console/workspace/model_providers.py | 2 +- api/controllers/console/workspace/models.py | 3 + api/controllers/console/workspace/plugin.py | 14 +- api/controllers/console/workspace/rbac.py | 5 +- api/controllers/console/workspace/snippets.py | 2 +- .../console/workspace/tool_providers.py | 2 + .../console/workspace/workspace.py | 19 +- api/controllers/files/agent_drive_archive.py | 2 + api/controllers/inner_api/app/dsl.py | 1 + .../inner_api/plugin/agent_drive.py | 5 +- .../inner_api/workspace/workspace.py | 6 +- api/controllers/openapi/account.py | 14 +- api/controllers/openapi/app_dsl.py | 1 + api/controllers/openapi/apps.py | 12 +- .../openapi/apps_permitted_external.py | 4 +- api/controllers/openapi/auth/prepare.py | 10 +- api/controllers/openapi/auth/verify.py | 4 +- api/controllers/openapi/oauth_device.py | 4 +- api/controllers/openapi/oauth_device_sso.py | 6 +- api/controllers/openapi/workspaces.py | 26 +- api/controllers/service_api/app/annotation.py | 10 +- api/controllers/service_api/app/app.py | 3 +- api/controllers/service_api/app/audio.py | 2 +- .../service_api/app/conversation.py | 14 +- api/controllers/service_api/app/message.py | 14 +- .../service_api/dataset/dataset.py | 46 +- .../service_api/dataset/document.py | 28 +- .../service_api/dataset/metadata.py | 38 +- .../rag_pipeline/rag_pipeline_workflow.py | 10 +- .../service_api/dataset/segment.py | 48 +- api/controllers/web/app.py | 11 +- api/controllers/web/audio.py | 2 +- api/controllers/web/completion.py | 6 +- api/controllers/web/conversation.py | 8 +- api/controllers/web/forgot_password.py | 4 +- api/controllers/web/login.py | 11 +- api/controllers/web/message.py | 10 +- api/controllers/web/passport.py | 2 +- api/controllers/web/saved_message.py | 6 +- api/controllers/web/wraps.py | 10 +- .../easy_ui_based_app/dataset/manager.py | 2 +- .../app/apps/advanced_chat/app_generator.py | 2 +- api/core/app/apps/agent_app/app_generator.py | 4 +- api/core/app/apps/agent_chat/app_generator.py | 2 +- api/core/app/apps/chat/app_generator.py | 2 +- .../annotation_reply/annotation_reply.py | 5 +- api/core/app/llm/quota.py | 2 + .../task_pipeline/message_cycle_manager.py | 2 +- .../index_tool_callback_handler.py | 6 +- api/core/llm_generator/llm_generator.py | 7 +- api/core/mcp/server/streamable_http.py | 10 +- api/core/provider_manager.py | 2 + api/core/rag/datasource/retrieval_service.py | 12 +- .../processor/paragraph_index_processor.py | 18 +- .../processor/parent_child_index_processor.py | 6 +- .../processor/qa_index_processor.py | 4 +- api/core/rag/summary_index/summary_index.py | 7 +- .../dataset_multi_retriever_tool.py | 4 +- .../dataset_retriever_tool.py | 4 +- .../nodes/agent_v2/dify_tools_builder.py | 3 +- .../update_provider_when_message_created.py | 1 + api/extensions/ext_login.py | 4 +- api/services/account_service.py | 160 ++-- api/services/agent/composer_service.py | 428 +++++++---- api/services/agent/roster_service.py | 1 + .../agent/skill_standardize_service.py | 4 + .../agent/skill_tool_inference_service.py | 11 +- api/services/agent_app_feature_service.py | 4 +- api/services/agent_app_sandbox_service.py | 12 +- api/services/agent_drive_service.py | 337 ++++----- api/services/agent_service.py | 12 +- api/services/agent_tool_inner_service.py | 2 +- api/services/annotation_service.py | 114 +-- api/services/api_based_extension_service.py | 8 +- api/services/app_dsl_service.py | 28 +- api/services/app_generate_service.py | 59 +- api/services/app_service.py | 131 ++-- api/services/async_workflow_service.py | 13 +- api/services/audio_service.py | 6 +- api/services/auth/api_key_auth_service.py | 8 +- api/services/billing_service.py | 4 +- api/services/conversation_service.py | 131 ++-- api/services/credential_permission_service.py | 6 +- api/services/credit_pool_service.py | 75 +- api/services/data_migration/export_service.py | 30 +- api/services/data_migration/import_service.py | 114 ++- api/services/dataset_service.py | 265 +++---- api/services/datasource_provider_service.py | 61 +- .../enterprise/account_deletion_sync.py | 9 +- api/services/enterprise/rbac_service.py | 96 +-- api/services/external_knowledge_service.py | 37 +- api/services/file_service.py | 2 +- api/services/hit_testing_service.py | 18 +- api/services/message_service.py | 57 +- api/services/metadata_service.py | 17 +- api/services/model_load_balancing_service.py | 63 +- api/services/oauth_device_flow.py | 26 +- api/services/oauth_server.py | 6 +- api/services/ops_service.py | 38 +- .../plugin/plugin_auto_upgrade_service.py | 221 +++--- .../plugin/plugin_permission_service.py | 39 +- .../rag_pipeline/pipeline_generate_service.py | 23 +- .../built_in/built_in_retrieval.py | 7 +- .../customized/customized_retrieval.py | 12 +- .../database/database_retrieval.py | 12 +- .../pipeline_template_base.py | 4 +- .../remote/remote_retrieval.py | 10 +- api/services/rag_pipeline/rag_pipeline.py | 689 ++++++++---------- .../rag_pipeline/rag_pipeline_dsl_service.py | 4 +- .../rag_pipeline_transform_service.py | 7 +- .../buildin/buildin_retrieval.py | 11 +- .../database/database_retrieval.py | 38 +- .../recommend_app/recommend_app_base.py | 8 +- .../recommend_app/remote/remote_retrieval.py | 11 +- api/services/recommended_app_service.py | 18 +- api/services/saved_message_service.py | 15 +- api/services/snippet_service.py | 10 +- api/services/summary_index_service.py | 602 +++++++-------- api/services/tag_service.py | 25 +- .../tools/builtin_tools_manage_service.py | 13 +- api/services/trigger/schedule_service.py | 23 +- .../trigger/trigger_provider_service.py | 2 +- .../trigger_subscription_operator_service.py | 7 +- api/services/trigger/webhook_service.py | 6 +- api/services/vector_service.py | 49 +- api/services/web_conversation_service.py | 19 +- api/services/webapp_auth_service.py | 32 +- .../workflow/node_output_inspector_service.py | 58 +- api/services/workflow/workflow_converter.py | 40 +- .../workflow_collaboration_service.py | 23 +- api/services/workflow_service.py | 131 ++-- api/services/workspace_service.py | 12 +- .../batch_create_segment_to_index_task.py | 2 +- api/tasks/regenerate_summary_index_task.py | 3 +- api/tasks/retry_document_indexing_task.py | 5 +- api/tasks/workflow_schedule_tasks.py | 2 +- api/tests/integration_tests/conftest.py | 2 +- .../services/plugin/test_plugin_lifecycle.py | 43 +- .../test_node_output_inspector_service.py | 51 +- .../controllers/console/app/test_app_apis.py | 6 +- .../auth/test_data_source_bearer_auth.py | 2 +- .../console/auth/test_email_register.py | 2 +- .../console/auth/test_forgot_password.py | 2 +- .../controllers/console/auth/test_oauth.py | 2 +- .../rag_pipeline/test_rag_pipeline.py | 36 +- .../console/test_api_based_extension.py | 2 +- .../openapi/test_account_sessions.py | 2 +- .../controllers/openapi/test_app_dsl.py | 2 +- .../controllers/openapi/test_app_run.py | 2 +- .../controllers/openapi/test_apps.py | 2 +- .../controllers/openapi/test_files.py | 2 +- .../service_api/dataset/test_dataset.py | 5 +- .../web/test_web_forgot_password.py | 4 +- .../controllers/web/test_wraps.py | 2 + .../auth/test_api_key_auth_service.py | 34 +- .../services/auth/test_auth_integration.py | 20 +- .../enterprise/test_account_deletion_sync.py | 20 +- .../plugin/test_plugin_permission_service.py | 10 +- .../test_rag_pipeline_service_db.py | 24 +- .../recommend_app/test_database_retrieval.py | 54 +- .../services/test_account_service.py | 6 +- .../services/test_agent_service.py | 30 +- .../services/test_annotation_service.py | 169 +++-- .../test_api_based_extension_service.py | 72 +- .../services/test_app_dsl_service.py | 40 +- .../services/test_app_generate_service.py | 82 ++- .../services/test_app_service.py | 107 +-- .../services/test_billing_service.py | 4 +- .../services/test_conversation_service.py | 80 +- .../test_conversation_service_variables.py | 18 +- .../services/test_credit_pool_service.py | 132 ++-- .../services/test_dataset_service.py | 7 +- .../test_dataset_service_permissions.py | 22 +- .../test_dataset_service_update_dataset.py | 26 +- .../test_file_service_zip_and_lookup.py | 8 +- .../services/test_hit_testing_service.py | 16 +- .../test_human_input_delivery_test.py | 3 + .../services/test_message_service.py | 136 +++- ...message_service_execution_extra_content.py | 2 + .../services/test_metadata_partial_update.py | 14 +- .../services/test_metadata_service.py | 82 ++- .../test_model_load_balancing_service.py | 18 +- .../services/test_oauth_server_service.py | 10 +- .../services/test_ops_service.py | 63 +- .../services/test_recommended_app_service.py | 48 +- .../services/test_saved_message_service.py | 44 +- .../services/test_tag_service.py | 18 +- .../services/test_web_conversation_service.py | 22 +- .../services/test_webapp_auth_service.py | 52 +- .../test_webhook_service_relationships.py | 7 +- .../services/test_workflow_app_service.py | 4 +- .../services/test_workflow_run_service.py | 8 +- .../services/test_workflow_service.py | 39 +- .../services/test_workspace_service.py | 40 +- .../test_workflow_tools_manage_service.py | 2 +- .../workflow/test_workflow_converter.py | 15 +- .../trigger/test_trigger_e2e.py | 4 +- .../commands/test_data_migration_commands.py | 8 +- .../controllers/common/test_app_access.py | 2 +- .../console/agent/test_agent_controllers.py | 68 +- .../console/app/test_agent_app_sandbox.py | 3 + .../console/app/test_annotation_api.py | 6 +- .../console/app/test_annotation_security.py | 14 +- .../console/app/test_app_response_models.py | 25 +- .../controllers/console/app/test_workflow.py | 4 +- .../test_workflow_human_input_debug_api.py | 5 +- .../test_workflow_node_output_inspector.py | 9 +- .../console/auth/test_account_activation.py | 2 +- .../auth/test_data_source_bearer_auth.py | 6 +- .../rag_pipeline/test_datasource_auth.py | 3 +- .../rag_pipeline/test_rag_pipeline.py | 84 ++- .../test_rag_pipeline_workflow.py | 18 +- .../console/datasets/test_datasets.py | 18 +- .../console/datasets/test_external.py | 7 +- .../console/explore/test_recommended_app.py | 12 +- .../console/explore/test_saved_message.py | 6 +- .../console/snippets/test_snippet_workflow.py | 5 +- .../controllers/console/tag/test_tags.py | 12 +- .../controllers/console/test_extension.py | 12 +- .../console/test_workspace_account.py | 2 +- .../workspace/test_load_balancing_config.py | 4 +- .../console/workspace/test_tool_providers.py | 4 +- .../controllers/inner_api/app/test_dsl.py | 12 +- .../inner_api/plugin/test_agent_drive.py | 10 +- .../inner_api/workspace/test_workspace.py | 2 +- .../openapi/test_workspaces_members.py | 46 +- .../service_api/app/test_annotation.py | 6 +- .../controllers/service_api/app/test_app.py | 4 +- .../service_api/app/test_completion.py | 28 +- .../service_api/app/test_conversation.py | 1 + .../service_api/app/test_message.py | 31 +- .../service_api/app/test_workflow.py | 34 +- .../test_rag_pipeline_workflow.py | 37 +- .../dataset/test_dataset_segment.py | 2 +- .../service_api/dataset/test_document.py | 2 +- .../service_api/dataset/test_metadata.py | 4 +- .../unit_tests/controllers/web/test_app.py | 4 +- .../controllers/web/test_message_list.py | 4 +- .../controllers/web/test_web_login.py | 6 +- .../apps/advanced_chat/test_app_generator.py | 2 +- .../apps/test_advanced_chat_app_generator.py | 2 +- .../unit_tests/core/app/test_llm_quota.py | 21 +- .../test_llm_generator_missing.py | 6 +- .../datasource/test_datasource_retrieval.py | 23 +- .../test_paragraph_index_processor.py | 4 +- .../test_parent_child_index_processor.py | 4 +- .../processor/test_qa_index_processor.py | 4 +- .../events/test_app_event_signals.py | 33 +- ...st_update_provider_when_message_created.py | 18 +- .../services/agent/test_agent_services.py | 275 +++++-- .../agent/test_skill_standardize_service.py | 1 + .../test_skill_tool_inference_service.py | 21 +- .../data_migration/test_export_service.py | 14 +- .../data_migration/test_import_service.py | 142 ++-- .../services/enterprise/test_rbac_service.py | 85 +-- api/tests/unit_tests/services/hit_service.py | 37 +- .../test_plugin_auto_upgrade_service.py | 103 +-- .../test_built_in_retrieval.py | 12 +- .../test_customized_retrieval.py | 6 +- .../test_database_retrieval.py | 6 +- .../test_pipeline_template_base.py | 10 +- .../test_remote_retrieval.py | 8 +- .../test_pipeline_generate_service.py | 74 +- .../rag_pipeline/test_rag_pipeline_service.py | 84 ++- .../test_rag_pipeline_transform_service.py | 71 +- .../recommend_app/test_buildin_retrieval.py | 9 +- .../recommend_app/test_remote_retrieval.py | 16 +- .../services/test_account_service.py | 74 +- .../test_agent_app_sandbox_service.py | 5 + .../services/test_agent_drive_service.py | 130 +++- .../services/test_agent_tool_inner_service.py | 12 +- .../services/test_annotation_service.py | 104 +-- .../services/test_app_generate_service.py | 91 ++- .../unit_tests/services/test_app_service.py | 26 +- .../services/test_async_workflow_service.py | 22 +- .../services/test_billing_service.py | 8 +- .../services/test_conversation_service.py | 7 +- .../test_credential_permission_service.py | 4 +- .../services/test_credit_pool_service.py | 107 ++- .../services/test_dataset_service_dataset.py | 73 +- .../services/test_dataset_service_document.py | 2 +- .../services/test_dataset_service_segment.py | 42 +- .../test_datasource_provider_service.py | 28 +- .../services/test_external_dataset_service.py | 480 +++++++----- .../unit_tests/services/test_file_service.py | 6 +- .../services/test_message_service.py | 69 +- .../services/test_metadata_bug_complete.py | 6 +- .../services/test_metadata_nullable_bug.py | 4 +- .../test_model_load_balancing_service.py | 41 +- .../services/test_oauth_device_flow.py | 10 +- .../services/test_summary_index_service.py | 182 ++--- .../services/test_trigger_provider_service.py | 6 +- .../services/test_vector_service.py | 128 ++-- .../test_workflow_collaboration_service.py | 21 +- .../services/test_workflow_service.py | 101 ++- .../test_builtin_tools_manage_service.py | 6 +- .../test_node_output_inspector_service.py | 210 +++--- .../test_workflow_converter_additional.py | 18 +- .../test_workflow_human_input_delivery.py | 4 + 360 files changed, 6607 insertions(+), 4975 deletions(-) diff --git a/api/commands/account.py b/api/commands/account.py index dfd57d43142..9ea52dfd248 100644 --- a/api/commands/account.py +++ b/api/commands/account.py @@ -25,7 +25,7 @@ def reset_password(email, new_password, password_confirm): return normalized_email = email.strip().lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email.strip()) + account = AccountService.get_account_by_email_with_case_fallback(email.strip(), session=db.session()) if not account: click.echo(click.style(f"Account not found for email: {email}", fg="red")) @@ -67,7 +67,7 @@ def reset_email(email, new_email, email_confirm): return normalized_new_email = new_email.strip().lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email.strip()) + account = AccountService.get_account_by_email_with_case_fallback(email.strip(), session=db.session()) if not account: click.echo(click.style(f"Account not found for email: {email}", fg="red")) @@ -133,9 +133,9 @@ def create_tenant(email: str, language: str | None = None, name: str | None = No password=new_password, language=language, create_workspace_required=False, - session=db.session, + session=db.session(), ) - TenantService.create_owner_tenant_if_not_exist(account, name, session=db.session) + TenantService.create_owner_tenant_if_not_exist(account, name, session=db.session()) click.echo( click.style( diff --git a/api/commands/data_migration.py b/api/commands/data_migration.py index bd56c41ea44..8c2627601a6 100644 --- a/api/commands/data_migration.py +++ b/api/commands/data_migration.py @@ -9,6 +9,7 @@ from uuid import UUID import click import sqlalchemy as sa import yaml +from sqlalchemy.orm import Session from core.db.session_factory import session_factory from extensions.ext_database import db @@ -108,7 +109,7 @@ def export_migration_data(input_file: str | None, output_file: str | None, overw raw_config = _load_json_object(input_file, "Export config") selection = ExportConfigParser().parse(raw_config) with session_factory.create_session() as session: - result = MigrationExportService().export(session, selection) + result = MigrationExportService().export(selection, session=session) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _render_report(result.report_items, context=_with_output_path(result.report_context, output_file)) @@ -157,7 +158,6 @@ def import_migration_data( package = MigrationPackageService().load_package(input_file) with session_factory.create_session() as session: result = MigrationImportService().import_package( - session, ImportRequest( package=package, cli_target_tenant=target_tenant, @@ -169,6 +169,7 @@ def import_migration_data( create_app_api_token_on_import=create_app_api_token_on_import, ), ), + session=session, ) _render_report(result.report_items, context=result.report_context) except MigrationDataError as exc: @@ -217,7 +218,9 @@ def migration_data_wizard() -> None: default=True, show_default=False, ) - auto_tools = _discover_auto_tools([app for app in apps if app.id in set(app_ids)], include_referenced_tools) + auto_tools = _discover_auto_tools( + [app for app in apps if app.id in set(app_ids)], include_referenced_tools, session=db.session() + ) auto_tools = _resolve_auto_tool_names(tenant.id, auto_tools) _print_auto_tools(auto_tools) additional_tools = _prompt_additional_tools(tenant.id, auto_tools) @@ -253,7 +256,7 @@ def migration_data_wizard() -> None: output_file=output_file, ) with session_factory.create_session() as session: - result = MigrationExportService().export(session, selection) + result = MigrationExportService().export(selection, session=session) MigrationPackageService().save_package(result.package, output_file, overwrite=overwrite) click.echo(click.style(f"Output written to {output_file}", fg="green")) _print_wizard_step("Report") @@ -394,13 +397,13 @@ def _prompt_import_options() -> tuple[bool, bool, str, str]: return include_secrets, create_tokens, id_strategy, conflict_strategy -def _discover_auto_tools(apps: list[App], include_referenced_tools: bool) -> WizardToolMap: +def _discover_auto_tools(apps: list[App], include_referenced_tools: bool, *, session: Session) -> WizardToolMap: auto_tools: WizardToolMap = {"api_tools": {}, "workflow_tools": {}, "mcp_tools": {}} if not include_referenced_tools: return auto_tools discovery_service = DependencyDiscoveryService() for app in apps: - dsl_content = AppDslService.export_dsl(app_model=app, include_secret=False) + dsl_content = AppDslService.export_dsl(app_model=app, session=session, include_secret=False) raw_dsl = yaml.safe_load(dsl_content) if dsl_content else {} dsl = raw_dsl if isinstance(raw_dsl, dict) else {} for dependency in discovery_service.discover_from_dsl(dsl): diff --git a/api/commands/plugin.py b/api/commands/plugin.py index 718fa60761c..3695c742921 100644 --- a/api/commands/plugin.py +++ b/api/commands/plugin.py @@ -472,6 +472,7 @@ def backfill_plugin_auto_upgrade( try: result = PluginAutoUpgradeService.backfill_strategy_categories( current_tenant_id, + session=db.session(), ) except Exception as e: failed_count += 1 diff --git a/api/commands/rbac.py b/api/commands/rbac.py index 0793d11cbb2..be4993920ad 100644 --- a/api/commands/rbac.py +++ b/api/commands/rbac.py @@ -6,6 +6,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed import click from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.db.session_factory import session_factory @@ -131,16 +132,35 @@ def _replace_member_role( operator_account_id: str, member_account_id: str, role_id: str, + *, + session: Session, ) -> str: RBACService.MemberRoles.replace( tenant_id=tenant_id, account_id=operator_account_id, member_account_id=member_account_id, role_ids=[role_id], + session=session, ) return member_account_id +def _replace_member_role_with_new_session( + tenant_id: str, + operator_account_id: str, + member_account_id: str, + role_id: str, +) -> str: + with session_factory.create_session() as session: + return _replace_member_role( + tenant_id=tenant_id, + operator_account_id=operator_account_id, + member_account_id=member_account_id, + role_id=role_id, + session=session, + ) + + @click.command( "rbac-migrate-member-roles", help="Migrate legacy workspace member roles into RBAC member-role bindings." ) @@ -217,14 +237,21 @@ def migrate_member_roles_to_rbac( if replace_jobs: if workers == 1: - for member_account_id, resolved_role_id in replace_jobs: - _replace_member_role(workspace_id, owner_account_id, member_account_id, resolved_role_id) - migrated_count += 1 + with session_factory.create_session() as session: + for member_account_id, resolved_role_id in replace_jobs: + _replace_member_role( + workspace_id, + owner_account_id, + member_account_id, + resolved_role_id, + session=session, + ) + migrated_count += 1 else: with ThreadPoolExecutor(max_workers=workers) as executor: futures = [ executor.submit( - _replace_member_role, + _replace_member_role_with_new_session, workspace_id, owner_account_id, member_account_id, diff --git a/api/controllers/common/app_access.py b/api/controllers/common/app_access.py index 863b69d2339..214d2de71b4 100644 --- a/api/controllers/common/app_access.py +++ b/api/controllers/common/app_access.py @@ -4,6 +4,7 @@ from collections.abc import Sequence from dataclasses import dataclass from typing import TYPE_CHECKING +from extensions.ext_database import db from services.enterprise import rbac_service as enterprise_rbac_service if TYPE_CHECKING: @@ -76,7 +77,7 @@ def resolve_app_access_filter( inner-API round trip; otherwise it is fetched here. """ if permissions is None: - permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id) + permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id, session=db.session()) whitelist_scope = enterprise_rbac_service.RBACService.AppAccess.whitelist_resources(tenant_id, account_id) can_manage_own_apps = _MANAGE_OWN_APPS_PERMISSION_KEY in permissions.workspace.permission_keys diff --git a/api/controllers/console/agent/composer.py b/api/controllers/console/agent/composer.py index d089772e3ab..f5d71990ade 100644 --- a/api/controllers/console/agent/composer.py +++ b/api/controllers/console/agent/composer.py @@ -16,6 +16,7 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user_id, ) +from extensions.ext_database import db from fields.agent_fields import ( AgentAppComposerResponse, AgentComposerCandidatesResponse, @@ -69,6 +70,7 @@ class WorkflowAgentComposerApi(Resource): node_id=node_id, account_id=account_id, snapshot_id=query.snapshot_id, + session=db.session(), ), ) @@ -94,6 +96,7 @@ class WorkflowAgentComposerApi(Resource): node_id=node_id, account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -126,6 +129,7 @@ class WorkflowAgentComposerCopyFromRosterApi(Resource): source_agent_id=payload.source_agent_id, source_snapshot_id=payload.source_snapshot_id, idempotency_key=payload.idempotency_key, + session=db.session(), ), ) @@ -149,8 +153,9 @@ class WorkflowAgentComposerValidateApi(Resource): tenant_id=tenant_id, payload=payload, agent_id=AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ), + session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -174,6 +179,7 @@ class WorkflowAgentComposerCandidatesApi(Resource): app_id=app_model.id, node_id=node_id, user_id=current_user_id, + session=db.session(), ), ) @@ -196,7 +202,9 @@ class WorkflowAgentComposerImpactApi(Resource): ) return dump_response( AgentComposerImpactResponse, - AgentComposerService.calculate_impact(tenant_id=tenant_id, current_snapshot_id=current_snapshot_id), + AgentComposerService.calculate_impact( + tenant_id=tenant_id, current_snapshot_id=current_snapshot_id, session=db.session() + ), ) @@ -224,6 +232,7 @@ class WorkflowAgentComposerSaveToRosterApi(Resource): node_id=node_id, account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -238,7 +247,7 @@ class AgentComposerApi(Resource): def get(self, tenant_id: str, agent_id: UUID): return dump_response( AgentAppComposerResponse, - AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id)), + AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()), ) @console_ns.expect(console_ns.models[ComposerSavePayload.__name__]) @@ -259,6 +268,7 @@ class AgentComposerApi(Resource): agent_id=str(agent_id), account_id=account_id, payload=payload, + session=db.session(), ), ) @@ -274,7 +284,7 @@ class AgentComposerValidateApi(Resource): @account_initialization_required @with_current_tenant_id def post(self, tenant_id: str, agent_id: UUID): - AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id)) + AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()) payload = ComposerSavePayload.model_validate(console_ns.payload or {}) ComposerConfigValidator.validate_publish_payload(payload) AgentComposerService.validate_knowledge_datasets(tenant_id=tenant_id, agent_soul=payload.agent_soul) @@ -282,6 +292,7 @@ class AgentComposerValidateApi(Resource): tenant_id=tenant_id, payload=payload, agent_id=str(agent_id), + session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -303,5 +314,6 @@ class AgentComposerCandidatesApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), user_id=current_user_id, + session=db.session(), ), ) diff --git a/api/controllers/console/agent/roster.py b/api/controllers/console/agent/roster.py index 349826e54d7..1467cc0c246 100644 --- a/api/controllers/console/agent/roster.py +++ b/api/controllers/console/agent/roster.py @@ -534,7 +534,7 @@ class AgentAppListApi(Resource): status="normal", ) - app_pagination = AppService().get_paginate_apps(current_user.id, current_tenant_id, params, db.session) + app_pagination = AppService().get_paginate_apps(current_user.id, current_tenant_id, params, db.session()) if app_pagination is None: empty = AgentAppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return empty.model_dump(mode="json") @@ -567,7 +567,7 @@ class AgentAppListApi(Resource): icon_background=args.icon_background, ) - app = AppService().create_app(current_tenant_id, params, current_user) + app = AppService().create_app(current_tenant_id, params, current_user, session=db.session()) return _serialize_agent_app_detail(app, current_user=current_user), 201 @@ -607,7 +607,7 @@ class AgentAppApi(Resource): "max_active_requests": args.max_active_requests or 0, "role": args.role, } - updated = AppService().update_app(app_model, args_dict) + updated = AppService().update_app(app_model, args_dict, session=db.session()) return _serialize_agent_app_detail(updated, current_user=current_user) @console_ns.response(204, "Agent app deleted successfully") @@ -619,7 +619,7 @@ class AgentAppApi(Resource): @with_current_tenant_id def delete(self, tenant_id: str, agent_id: UUID): app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=db.session()) return "", 204 @@ -668,6 +668,7 @@ class AgentPublishApi(Resource): agent_id=str(agent_id), account_id=current_user.id, version_note=args.version_note, + session=db.session(), ) @@ -688,6 +689,7 @@ class AgentBuildDraftCheckoutApi(Resource): agent_id=str(agent_id), account_id=current_user.id, force=args.force, + session=db.session(), ) @@ -705,6 +707,7 @@ class AgentBuildDraftApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @console_ns.expect(console_ns.models[ComposerSavePayload.__name__]) @@ -722,6 +725,7 @@ class AgentBuildDraftApi(Resource): agent_id=str(agent_id), account_id=current_user.id, payload=payload, + session=db.session(), ) @console_ns.response(200, "Agent build draft discarded", console_ns.models[AgentSimpleResultResponse.__name__]) @@ -736,6 +740,7 @@ class AgentBuildDraftApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @@ -753,6 +758,7 @@ class AgentBuildDraftApplyApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, + session=db.session(), ) @@ -810,7 +816,7 @@ class AgentApiStatusApi(Resource): def post(self, tenant_id: str, agent_id: UUID): app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) args = AgentApiStatusPayload.model_validate(console_ns.payload) - app_model = AppService().update_app_api_status(app_model, args.enable_api) + app_model = AppService().update_app_api_status(app_model, args.enable_api, session=db.session()) return _serialize_agent_api_access(app_model) diff --git a/api/controllers/console/app/agent.py b/api/controllers/console/app/agent.py index 99164b4755a..81d17ace37a 100644 --- a/api/controllers/console/app/agent.py +++ b/api/controllers/console/app/agent.py @@ -172,7 +172,7 @@ register_response_schema_models( def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: if node_id and app_model.mode != AppMode.AGENT: return AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ) return app_model.bound_agent_id @@ -202,6 +202,7 @@ def _upload_skill_for_app(*, current_user: Account, app_model: App): tenant_id=app_model.tenant_id, user_id=current_user.id, agent_id=agent_id, + session=db.session(), ) except (SkillPackageError, AgentDriveError) as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -240,6 +241,7 @@ def _commit_drive_file_for_app(*, current_user: Account, app_model: App, allow_n value_owned_by_drive=True, ) ], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -273,6 +275,7 @@ def _delete_drive_file_for_app(*, current_user: Account, app_model: App, allow_n user_id=current_user.id, agent_id=agent_id, items=[DriveCommitItem(key=key, file_ref=None)], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -298,6 +301,7 @@ def _delete_skill_for_app(*, current_user: Account, app_model: App, slug: str, a DriveCommitItem(key=f"{slug}/SKILL.md", file_ref=None), DriveCommitItem(key=f"{slug}/.DIFY-SKILL-FULL.zip", file_ref=None), ], + session=db.session(), ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -313,7 +317,9 @@ def _infer_skill_tools_for_app(*, app_model: App, slug: str): if "/" in slug or not slug.strip(): return {"code": "drive_key_invalid", "message": "skill slug must be a single path segment"}, 400 try: - return SkillToolInferenceService().infer(tenant_id=app_model.tenant_id, agent_id=agent_id, slug=slug) + return SkillToolInferenceService().infer( + tenant_id=app_model.tenant_id, agent_id=agent_id, slug=slug, session=db.session() + ) except SkillToolInferenceError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -335,7 +341,7 @@ class AgentLogApi(Resource): """Get agent logs""" args = AgentLogQuery.model_validate(request.args.to_dict(flat=True)) - return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id) + return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id, db.session()) @console_ns.route("/agent//skills/upload") diff --git a/api/controllers/console/app/agent_app_feature.py b/api/controllers/console/app/agent_app_feature.py index 6990886a511..edd2f31f75f 100644 --- a/api/controllers/console/app/agent_app_feature.py +++ b/api/controllers/console/app/agent_app_feature.py @@ -93,7 +93,7 @@ class AgentAppFeatureConfigResource(Resource): app_model=app_model, account=current_user, config=args.model_dump(exclude_none=True), - session=db.session, + session=db.session(), ) app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config) diff --git a/api/controllers/console/app/agent_app_sandbox.py b/api/controllers/console/app/agent_app_sandbox.py index 4324f425a08..6f3811ccdc8 100644 --- a/api/controllers/console/app/agent_app_sandbox.py +++ b/api/controllers/console/app/agent_app_sandbox.py @@ -25,6 +25,7 @@ from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.app.wraps import get_app_model from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models.model import App, AppMode @@ -269,6 +270,7 @@ class WorkflowAgentSandboxListResource(Resource): node_id=node_id, node_execution_id=query.node_execution_id, path=query.path, + session=db.session(), ) except Exception as exc: return _handle(exc) @@ -305,6 +307,7 @@ class WorkflowAgentSandboxReadResource(Resource): node_id=node_id, node_execution_id=query.node_execution_id, path=query.path, + session=db.session(), ) except Exception as exc: return _handle(exc) @@ -334,6 +337,7 @@ class WorkflowAgentSandboxUploadResource(Resource): node_id=node_id, node_execution_id=payload.node_execution_id, path=payload.path, + session=db.session(), ) except Exception as exc: return _handle(exc) diff --git a/api/controllers/console/app/agent_config_inspector.py b/api/controllers/console/app/agent_config_inspector.py index 83824d6434f..0f7aa80ca78 100644 --- a/api/controllers/console/app/agent_config_inspector.py +++ b/api/controllers/console/app/agent_config_inspector.py @@ -253,6 +253,7 @@ def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, + session=db.session(), ) return app_model.bound_agent_id @@ -288,13 +289,16 @@ def _resolve_console_version( tenant_id=tenant_id, agent_id=agent_id, account_id=account_id, + session=db.session(), ) draft = state.get("draft") or {} draft_id = draft.get("id") if isinstance(draft_id, str) and draft_id: return draft_id, AgentConfigVersionKind.BUILD_DRAFT else: - state = AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=agent_id) + state = AgentComposerService.load_agent_composer( + tenant_id=tenant_id, agent_id=agent_id, session=db.session() + ) draft = state.get("draft") or {} draft_id = draft.get("id") if isinstance(draft_id, str) and draft_id: diff --git a/api/controllers/console/app/agent_drive_inspector.py b/api/controllers/console/app/agent_drive_inspector.py index 473e7364b3e..5166393b3d9 100644 --- a/api/controllers/console/app/agent_drive_inspector.py +++ b/api/controllers/console/app/agent_drive_inspector.py @@ -28,6 +28,7 @@ from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.app.wraps import get_app_model from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models.model import App, AppMode @@ -147,7 +148,7 @@ def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: """Agent identity for the drive: app-bound agent, or the workflow node binding.""" if node_id: return AgentComposerService.resolve_workflow_node_agent_id( - tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id + tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, session=db.session() ) return app_model.bound_agent_id @@ -184,7 +185,9 @@ class AgentDriveListByAgentApi(Resource): query = query_params_from_request(AgentDriveListByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - items = AgentDriveService().manifest(tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix) + items = AgentDriveService().manifest( + tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"items": [{k: v for k, v in item.items() if k != "file_id"} for item in items]} @@ -203,7 +206,7 @@ class AgentDriveSkillListByAgentApi(Resource): def get(self, tenant_id: str, agent_id: UUID): resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id)) + items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -227,6 +230,7 @@ class AgentDriveSkillInspectByAgentApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), skill_path=skill_path, + session=db.session(), ) ) except AgentDriveError as exc: @@ -247,7 +251,9 @@ class AgentDrivePreviewByAgentApi(Resource): query = query_params_from_request(AgentDriveFileByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - return AgentDriveService().preview(tenant_id=tenant_id, agent_id=str(agent_id), key=query.key) + return AgentDriveService().preview( + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) @@ -266,7 +272,9 @@ class AgentDriveDownloadByAgentApi(Resource): query = query_params_from_request(AgentDriveFileByAgentQuery) resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) try: - url = AgentDriveService().download_url(tenant_id=tenant_id, agent_id=str(agent_id), key=query.key) + url = AgentDriveService().download_url( + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"url": url} @@ -288,7 +296,9 @@ class AgentDriveListApi(Resource): if not agent_id: return _agent_not_bound() try: - items = AgentDriveService().manifest(tenant_id=app_model.tenant_id, agent_id=agent_id, prefix=query.prefix) + items = AgentDriveService().manifest( + tenant_id=app_model.tenant_id, agent_id=agent_id, prefix=query.prefix, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) # the inner manifest exposes file_id for agent-side pulls; the console @@ -312,7 +322,9 @@ class AgentDriveSkillListApi(Resource): if not agent_id: return _agent_not_bound() try: - items = AgentDriveService().list_skills(tenant_id=app_model.tenant_id, agent_id=agent_id) + items = AgentDriveService().list_skills( + tenant_id=app_model.tenant_id, agent_id=agent_id, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -345,6 +357,7 @@ class AgentDriveSkillInspectApi(Resource): tenant_id=app_model.tenant_id, agent_id=agent_id, skill_path=skill_path, + session=db.session(), ) ) except AgentDriveError as exc: @@ -367,7 +380,9 @@ class AgentDrivePreviewApi(Resource): if not agent_id: return _agent_not_bound() try: - return AgentDriveService().preview(tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key) + return AgentDriveService().preview( + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) @@ -388,7 +403,9 @@ class AgentDriveDownloadApi(Resource): if not agent_id: return _agent_not_bound() try: - url = AgentDriveService().download_url(tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key) + url = AgentDriveService().download_url( + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=db.session() + ) except AgentDriveError as exc: return _handle(exc) return {"url": url} diff --git a/api/controllers/console/app/annotation.py b/api/controllers/console/app/annotation.py index d14c7d2a7dc..961f9e2f1d8 100644 --- a/api/controllers/console/app/annotation.py +++ b/api/controllers/console/app/annotation.py @@ -211,7 +211,7 @@ class AppAnnotationSettingDetailApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) def get(self, app_id: UUID): - result = AppAnnotationService.get_app_annotation_setting_by_app_id(str(app_id)) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(str(app_id), session=db.session()) return dump_response(AnnotationSettingResponse, result), 200 @@ -235,7 +235,7 @@ class AppAnnotationSettingUpdateApi(Resource): setting_args: UpdateAnnotationSettingArgs = {"score_threshold": args.score_threshold} result = AppAnnotationService.update_app_annotation_setting( - str(app_id), annotation_setting_id_str, setting_args + str(app_id), annotation_setting_id_str, setting_args, session=db.session() ) return dump_response(AnnotationSettingResponse, result), 200 @@ -292,7 +292,9 @@ class AnnotationApi(Resource): limit = args.limit keyword = args.keyword - annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id(str(app_id), page, limit, keyword) + annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( + str(app_id), page, limit, keyword, session=db.session() + ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return AnnotationList( data=annotation_models, has_more=len(annotation_list) == limit, limit=limit, total=total, page=page @@ -321,7 +323,9 @@ class AnnotationApi(Resource): upsert_args["message_id"] = args.message_id if args.question is not None: upsert_args["question"] = args.question - annotation = AppAnnotationService.up_insert_app_annotation_from_message(upsert_args, str(app_id)) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + upsert_args, str(app_id), session=db.session() + ) return dump_response(Annotation, annotation), 201 @setup_required @@ -345,11 +349,11 @@ class AnnotationApi(Resource): }, 400 app_ref = _get_app_ref(str(app_id)) - AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids) + AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids, session=db.session()) return "", 204 # If no annotation_ids are provided, handle clearing all annotations else: - AppAnnotationService.clear_all_annotations(str(app_id)) + AppAnnotationService.clear_all_annotations(str(app_id), session=db.session()) return "", 204 @@ -370,7 +374,7 @@ class AnnotationExportApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) def get(self, app_id: UUID): - annotation_list = AppAnnotationService.export_annotation_list_by_app_id(str(app_id)) + annotation_list = AppAnnotationService.export_annotation_list_by_app_id(str(app_id), session=db.session()) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return ( AnnotationExportList(data=annotation_models).model_dump(mode="json"), @@ -406,7 +410,7 @@ class AnnotationUpdateDeleteApi(Resource): update_args["question"] = args.question app_ref = _get_app_ref(str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session()) return Annotation.model_validate(annotation, from_attributes=True).model_dump(mode="json") @setup_required @@ -418,7 +422,7 @@ class AnnotationUpdateDeleteApi(Resource): def delete(self, app_id: UUID, annotation_id: UUID): app_ref = _get_app_ref(str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - AppAnnotationService.delete_app_annotation(annotation_ref, db.session) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session()) return "", 204 @@ -477,7 +481,7 @@ class AnnotationBatchImportApi(Resource): return dump_response( AnnotationBatchImportResponse, - AppAnnotationService.batch_import_app_annotations(str(app_id), file), + AppAnnotationService.batch_import_app_annotations(str(app_id), file, session=db.session()), ) @@ -538,6 +542,7 @@ class AnnotationHitHistoryListApi(Resource): annotation_ref, page, limit, + session=db.session(), ) history_models = TypeAdapter(list[AnnotationHitHistory]).validate_python( annotation_hit_history_list, from_attributes=True diff --git a/api/controllers/console/app/app.py b/api/controllers/console/app/app.py index a1423318c72..4f0022b37b4 100644 --- a/api/controllers/console/app/app.py +++ b/api/controllers/console/app/app.py @@ -584,6 +584,7 @@ class AppListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user_id, + session=db.session(), ) if dify_config.RBAC_ENABLED: access_filter = resolve_app_access_filter( @@ -595,7 +596,7 @@ class AppListApi(Resource): # get app list app_service = AppService() - app_pagination = app_service.get_paginate_apps(current_user_id, current_tenant_id, params, db.session) + app_pagination = app_service.get_paginate_apps(current_user_id, current_tenant_id, params, session) if not app_pagination: response = AppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return response.model_dump(mode="json"), 200 @@ -643,11 +644,12 @@ class AppListApi(Resource): ) app_service = AppService() - app = app_service.create_app(current_tenant_id, params, current_user) + app = app_service.create_app(current_tenant_id, params, current_user, session=db.session()) permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( str(current_tenant_id), current_user.id, [str(app.id)], + session=db.session(), ) app_detail = AppDetailWithSite.model_validate(app, from_attributes=True).model_copy( update={"permission_keys": permission_keys_map.get(str(app.id), [])} @@ -681,7 +683,7 @@ class StarredAppListApi(Resource): is_created_by_me=args.is_created_by_me, ) - app_pagination = AppService().get_paginate_starred_apps(current_user_id, current_tenant_id, params, db.session) + app_pagination = AppService().get_paginate_starred_apps(current_user_id, current_tenant_id, params, session) if not app_pagination: empty = AppPagination(page=args.page, limit=args.limit, total=0, has_more=False, data=[]) return empty.model_dump(mode="json"), 200 @@ -705,7 +707,7 @@ class AppStarApi(Resource): @with_session @get_app_model(mode=None) def post(self, session: Session, current_user_id: str, app_model: App): - AppService.star_app(session, app=app_model, account_id=current_user_id) + AppService.star_app(app=app_model, account_id=current_user_id, session=session) return SimpleResultResponse(result="success").model_dump(mode="json") @console_ns.doc("unstar_app") @@ -721,7 +723,7 @@ class AppStarApi(Resource): @with_session @get_app_model(mode=None) def delete(self, session: Session, current_user_id: str, app_model: App): - AppService.unstar_app(session, app=app_model, account_id=current_user_id) + AppService.unstar_app(app=app_model, account_id=current_user_id, session=session) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -753,6 +755,7 @@ class AppApi(Resource): str(current_tenant_id), current_user.id, app_id=str(app_model.id), + session=db.session(), ) permission_keys_map = permissions.app.permission_keys_by_resource_ids([str(app_model.id)]) @@ -789,7 +792,7 @@ class AppApi(Resource): "use_icon_as_answer_icon": args.use_icon_as_answer_icon or False, "max_active_requests": args.max_active_requests or 0, } - app_model = app_service.update_app(app_model, args_dict) + app_model = app_service.update_app(app_model, args_dict, session=db.session()) return dump_response(AppDetailWithSite, app_model) @console_ns.doc("delete_app") @@ -806,7 +809,7 @@ class AppApi(Resource): def delete(self, app_model: App): """Delete app""" app_service = AppService() - app_service.delete_app(app_model) + app_service.delete_app(app_model, session=db.session()) return "", 204 @@ -835,7 +838,7 @@ class AppCopyApi(Resource): with Session(db.engine, expire_on_commit=False) as session: import_service = AppDslService(session) - yaml_content = import_service.export_dsl(app_model=app_model, include_secret=True) + yaml_content = import_service.export_dsl(app_model=app_model, session=session, include_secret=True) result = import_service.import_app( account=current_user, import_mode=ImportMode.YAML_CONTENT, @@ -877,6 +880,7 @@ class AppCopyApi(Resource): str(current_tenant_id), current_user.id, [str(app.id)], + session=db.session(), ) response_model = AppDetailWithSite.model_validate(app, from_attributes=True).model_copy( update={"permission_keys": permission_keys_map.get(str(app.id), [])} @@ -905,6 +909,7 @@ class AppExportApi(Resource): response = AppExportResponse( data=AppDslService.export_dsl( app_model=app_model, + session=db.session(), include_secret=args.include_secret, workflow_id=args.workflow_id, ) @@ -929,7 +934,7 @@ class AppPublishToCreatorsPlatformApi(Resource): if not dify_config.CREATORS_PLATFORM_FEATURES_ENABLED: return {"error": "Creators Platform features are not enabled"}, 403 - dsl_content = AppDslService.export_dsl(app_model=app_model, include_secret=False) + dsl_content = AppDslService.export_dsl(app_model=app_model, session=db.session(), include_secret=False) dsl_bytes = dsl_content.encode("utf-8") claim_code = upload_dsl(dsl_bytes) @@ -955,7 +960,7 @@ class AppNameApi(Resource): args = AppNamePayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_name(app_model, args.name) + app_model = app_service.update_app_name(app_model, args.name, session=db.session()) return dump_response(AppDetail, app_model) @@ -982,6 +987,7 @@ class AppIconApi(Resource): args.icon or "", args.icon_background or "", args.icon_type, + session=db.session(), ) return dump_response(AppDetail, app_model) @@ -1004,7 +1010,7 @@ class AppSiteStatus(Resource): args = AppSiteStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_site_status(app_model, args.enable_site) + app_model = app_service.update_app_site_status(app_model, args.enable_site, session=db.session()) return dump_response(AppDetail, app_model) @@ -1026,7 +1032,7 @@ class AppApiStatus(Resource): args = AppApiStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_api_status(app_model, args.enable_api) + app_model = app_service.update_app_api_status(app_model, args.enable_api, session=db.session()) return dump_response(AppDetail, app_model) diff --git a/api/controllers/console/app/audio.py b/api/controllers/console/app/audio.py index c6cd71f30f1..0c9ed786a1e 100644 --- a/api/controllers/console/app/audio.py +++ b/api/controllers/console/app/audio.py @@ -161,7 +161,7 @@ class ChatMessageTextApi(Resource): # response-contract:ignore return AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=payload.text, voice=payload.voice, message_ref=message_ref, diff --git a/api/controllers/console/app/conversation.py b/api/controllers/console/app/conversation.py index a80935e5e33..b7d422d30b1 100644 --- a/api/controllers/console/app/conversation.py +++ b/api/controllers/console/app/conversation.py @@ -200,7 +200,7 @@ class CompletionConversationDetailApi(Resource): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user) + ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -354,7 +354,7 @@ class ChatConversationDetailApi(Resource): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user) + ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") diff --git a/api/controllers/console/app/message.py b/api/controllers/console/app/message.py index 6a44ca3db8a..958b356de94 100644 --- a/api/controllers/console/app/message.py +++ b/api/controllers/console/app/message.py @@ -363,6 +363,7 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): app_model=app_model, conversation_id=args.conversation_id, user=current_user, + session=db.session(), ) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -474,7 +475,11 @@ def _get_message_suggested_questions(*, current_user: Account, app_model: App, m try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, message_id=message_id_str, user=current_user, invoke_from=InvokeFrom.DEBUGGER + app_model=app_model, + message_id=message_id_str, + user=current_user, + invoke_from=InvokeFrom.DEBUGGER, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") diff --git a/api/controllers/console/app/ops_trace.py b/api/controllers/console/app/ops_trace.py index 46d5ea56e20..e86f65fc035 100644 --- a/api/controllers/console/app/ops_trace.py +++ b/api/controllers/console/app/ops_trace.py @@ -17,6 +17,7 @@ from controllers.console.wraps import ( rbac_permission_required, setup_required, ) +from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models import App @@ -78,7 +79,7 @@ class TraceAppConfigApi(Resource): try: trace_config = OpsService.get_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider + app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() ) if not trace_config: return {"has_not_configured": True} @@ -109,7 +110,10 @@ class TraceAppConfigApi(Resource): try: result = OpsService.create_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, tracing_config=args.tracing_config + app_id=app_model.id, + tracing_provider=args.tracing_provider, + tracing_config=args.tracing_config, + session=db.session(), ) if not result: raise TracingConfigIsExist() @@ -142,7 +146,10 @@ class TraceAppConfigApi(Resource): try: result = OpsService.update_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, tracing_config=args.tracing_config + app_id=app_model.id, + tracing_provider=args.tracing_provider, + tracing_config=args.tracing_config, + session=db.session(), ) if not result: raise TracingConfigNotExist() @@ -168,7 +175,9 @@ class TraceAppConfigApi(Resource): args = TraceProviderQuery.model_validate(request.args.to_dict(flat=True)) try: - result = OpsService.delete_tracing_app_config(app_id=app_model.id, tracing_provider=args.tracing_provider) + result = OpsService.delete_tracing_app_config( + app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() + ) if not result: raise TracingConfigNotExist() return "", 204 diff --git a/api/controllers/console/app/permission_keys.py b/api/controllers/console/app/permission_keys.py index 810ea04e377..be10f904021 100644 --- a/api/controllers/console/app/permission_keys.py +++ b/api/controllers/console/app/permission_keys.py @@ -1,6 +1,9 @@ +from extensions.ext_database import db from services.enterprise import rbac_service as enterprise_rbac_service def get_app_permission_keys(tenant_id: str, account_id: str | None, app_id: str) -> list[str]: - permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get(tenant_id, account_id, [app_id]) + permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( + tenant_id, account_id, [app_id], session=db.session() + ) return permission_keys_map.get(app_id, []) diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 609dbfb82c5..53c7c6ea788 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -2,7 +2,7 @@ import json import logging from collections.abc import Sequence from datetime import datetime -from typing import Any, NotRequired, TypedDict, cast +from typing import Any, NotRequired, TypedDict from flask import abort, request from flask_restx import Resource, fields @@ -522,7 +522,7 @@ class DraftWorkflowApi(Resource): """ # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=db.session()) if not workflow: raise DraftWorkflowNotExist() @@ -533,7 +533,7 @@ class DraftWorkflowApi(Resource): # front-end can treat draft graph node data as the editing source. response = WorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json") response["graph"] = WorkflowAgentPublishService.project_draft_bindings_to_graph( - session=cast(Session, db.session), + session=db.session(), draft_workflow=workflow, ) return response @@ -602,6 +602,7 @@ class DraftWorkflowApi(Resource): account=current_user, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db.session(), ) except WorkflowHashNotEqualError: raise DraftWorkflowNotSync() @@ -695,7 +696,12 @@ class AdvancedChatDraftRunIterationNodeApi(Resource): try: response = AppGenerateService.generate_single_iteration( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -738,7 +744,12 @@ class WorkflowDraftRunIterationNodeApi(Resource): try: response = AppGenerateService.generate_single_iteration( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -777,7 +788,12 @@ class AdvancedChatDraftRunLoopNodeApi(Resource): try: response = AppGenerateService.generate_single_loop( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -820,7 +836,12 @@ class WorkflowDraftRunLoopNodeApi(Resource): try: response = AppGenerateService.generate_single_loop( - app_model=app_model, user=current_user, node_id=node_id, args=args, streaming=True + app_model=app_model, + user=current_user, + node_id=node_id, + args=args, + session=db.session(), + streaming=True, ) return helper.compact_generate_response(response) @@ -897,6 +918,7 @@ class AdvancedChatDraftHumanInputFormPreviewApi(Resource): account=current_user, node_id=node_id, inputs=inputs, + session=db.session(), ) return jsonable_encoder(preview) @@ -932,6 +954,7 @@ class AdvancedChatDraftHumanInputFormRunApi(Resource): form_inputs=args.form_inputs, inputs=args.inputs, action=args.action, + session=db.session(), ) return jsonable_encoder(result) @@ -963,6 +986,7 @@ class WorkflowDraftHumanInputFormPreviewApi(Resource): account=current_user, node_id=node_id, inputs=inputs, + session=db.session(), ) return jsonable_encoder(preview) @@ -998,6 +1022,7 @@ class WorkflowDraftHumanInputFormRunApi(Resource): form_inputs=args.form_inputs, inputs=args.inputs, action=args.action, + session=db.session(), ) return jsonable_encoder(result) @@ -1028,6 +1053,7 @@ class WorkflowDraftHumanInputDeliveryTestApi(Resource): node_id=node_id, delivery_method_id=args.delivery_method_id, inputs=args.inputs, + session=db.session(), ) return jsonable_encoder({}) @@ -1138,7 +1164,7 @@ class DraftWorkflowNodeRunApi(Resource): workflow_srv = WorkflowService() # fetch draft workflow by app_model - draft_workflow = workflow_srv.get_draft_workflow(app_model=app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model=app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not initialized") files = _parse_file(draft_workflow, args.get("files")) @@ -1181,7 +1207,7 @@ class PublishedWorkflowApi(Resource): """ # fetch published workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=db.session()) # return workflow, if not found, return None if workflow is None: @@ -1323,7 +1349,9 @@ class ConvertToWorkflowApi(Resource): # convert to workflow mode workflow_service = WorkflowService() - new_app_model = workflow_service.convert_to_workflow(app_model=app_model, account=current_user, args=args) + new_app_model = workflow_service.convert_to_workflow( + app_model=app_model, account=current_user, args=args, session=db.session() + ) # return app id return { @@ -1358,7 +1386,9 @@ class WorkflowFeaturesApi(Resource): features = args.features.model_dump(mode="json", exclude_unset=True) workflow_service = WorkflowService() - workflow_service.update_draft_workflow_features(app_model=app_model, features=features, account=current_user) + workflow_service.update_draft_workflow_features( + app_model=app_model, features=features, account=current_user, session=db.session() + ) return {"result": "success"} @@ -1439,6 +1469,7 @@ class DraftWorkflowRestoreApi(Resource): app_model=app_model, workflow_id=workflow_id, account=current_user, + session=db.session(), ) except IsDraftWorkflowError as exc: raise BadRequest(RESTORE_SOURCE_WORKFLOW_MUST_BE_PUBLISHED_MESSAGE) from exc @@ -1553,7 +1584,7 @@ class DraftWorkflowNodeLastRunApi(Resource): @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) def get(self, app_model: App, node_id: str): srv = WorkflowService() - workflow = srv.get_draft_workflow(app_model) + workflow = srv.get_draft_workflow(app_model, session=db.session()) if not workflow: raise NotFound("Workflow not found") node_exec = srv.get_node_last_run( @@ -1606,7 +1637,7 @@ class DraftWorkflowTriggerRunApi(Resource): args = DraftWorkflowTriggerRunPayload.model_validate(console_ns.payload or {}) node_id = args.node_id workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1675,7 +1706,7 @@ class DraftWorkflowTriggerNodeApi(Resource): """ workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1759,7 +1790,7 @@ class DraftWorkflowTriggerRunAllApi(Resource): args = DraftWorkflowTriggerRunAllPayload.model_validate(console_ns.payload or {}) node_ids = args.node_ids workflow_service = WorkflowService() - draft_workflow = workflow_service.get_draft_workflow(app_model) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=db.session()) if not draft_workflow: raise ValueError("Workflow not found") @@ -1828,7 +1859,7 @@ class WorkflowOnlineUsersApi(Resource): return {"data": []} workflow_service = WorkflowService() - accessible_app_ids = workflow_service.get_accessible_app_ids(app_ids, current_tenant_id) + accessible_app_ids = workflow_service.get_accessible_app_ids(app_ids, current_tenant_id, session=db.session()) ordered_accessible_app_ids = [app_id for app_id in app_ids if app_id in accessible_app_ids] users_json_by_app_id: dict[str, Any] = {} diff --git a/api/controllers/console/app/workflow_comment.py b/api/controllers/console/app/workflow_comment.py index 64df78f3748..9de91ac59d5 100644 --- a/api/controllers/console/app/workflow_comment.py +++ b/api/controllers/console/app/workflow_comment.py @@ -490,7 +490,7 @@ class WorkflowCommentMentionUsersApi(Resource): current_tenant = current_user.current_tenant # need the tenant object here if current_tenant is None: raise ValueError("current tenant is required") - members = TenantService.get_tenant_members(current_tenant, session=db.session) + members = TenantService.get_tenant_members(current_tenant, session=db.session()) users = TypeAdapter(list[AccountWithRole]).validate_python(members, from_attributes=True) response = WorkflowCommentMentionUsersPayload(users=users) return response.model_dump(mode="json"), 200 diff --git a/api/controllers/console/app/workflow_draft_variable.py b/api/controllers/console/app/workflow_draft_variable.py index 0ccc67f642d..1ffc01e3cef 100644 --- a/api/controllers/console/app/workflow_draft_variable.py +++ b/api/controllers/console/app/workflow_draft_variable.py @@ -337,7 +337,7 @@ class WorkflowVariableCollectionApi(Resource): # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow_exist = workflow_service.is_workflow_exist(app_model=app_model) + workflow_exist = workflow_service.is_workflow_exist(app_model=app_model, session=db.session()) if not workflow_exist: raise DraftWorkflowNotExist() @@ -553,7 +553,7 @@ class VariableResetApi(Resource): ) workflow_srv = WorkflowService() - draft_workflow = workflow_srv.get_draft_workflow(app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model, session=db.session()) if draft_workflow is None: raise NotFoundError( f"Draft workflow not found, app_id={app_model.id}", @@ -606,7 +606,7 @@ class ConversationVariableCollectionApi(Resource): # NOTE(QuantumGhost): Prefill conversation variables into the draft variables table # so their IDs can be returned to the caller. workflow_srv = WorkflowService() - draft_workflow = workflow_srv.get_draft_workflow(app_model) + draft_workflow = workflow_srv.get_draft_workflow(app_model, session=db.session()) if draft_workflow is None: raise NotFoundError(description=f"draft workflow not found, id={app_model.id}") draft_var_srv = WorkflowDraftVariableService(db.session()) @@ -646,6 +646,7 @@ class ConversationVariableCollectionApi(Resource): app_model=app_model, account=current_user, conversation_variables=conversation_variables, + session=db.session(), ) return {"result": "success"} @@ -683,7 +684,7 @@ class EnvironmentVariableCollectionApi(Resource): """ # fetch draft workflow by app_model workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=db.session()) if workflow is None: raise DraftWorkflowNotExist() @@ -740,6 +741,7 @@ class EnvironmentVariableCollectionApi(Resource): app_model=app_model, account=current_user, environment_variables=environment_variables, + session=db.session(), ) return {"result": "success"} diff --git a/api/controllers/console/app/workflow_node_output_inspector.py b/api/controllers/console/app/workflow_node_output_inspector.py index 6ed59d6c566..ea45a718a02 100644 --- a/api/controllers/console/app/workflow_node_output_inspector.py +++ b/api/controllers/console/app/workflow_node_output_inspector.py @@ -41,6 +41,7 @@ from controllers.console.wraps import ( rbac_permission_required, setup_required, ) +from extensions.ext_database import db from libs.exception import BaseHTTPException from libs.login import login_required from models import App, AppMode @@ -92,7 +93,9 @@ def _serve_snapshot(app_model: App, run_id: UUID) -> dict: Flask request context. """ try: - snapshot = _service().snapshot_workflow_run(app_model=app_model, workflow_run_id=str(run_id)) + snapshot = _service().snapshot_workflow_run( + app_model=app_model, workflow_run_id=str(run_id), session=db.session() + ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error return snapshot.model_dump(mode="json") @@ -105,6 +108,7 @@ def _serve_node_detail(app_model: App, run_id: UUID, node_id: str) -> dict: app_model=app_model, workflow_run_id=str(run_id), node_id=node_id, + session=db.session(), ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -119,6 +123,7 @@ def _serve_output_preview(app_model: App, run_id: UUID, node_id: str, output_nam workflow_run_id=str(run_id), node_id=node_id, output_name=output_name, + session=db.session(), ) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -245,7 +250,7 @@ def _stream_inspector_events(app_model: App, run_id: UUID) -> Iterator[str]: # if the run is gone (raised before yielding any bytes, so Flask turns it # into the normal HTTP 404 path). try: - snapshot = service.snapshot_workflow_run(app_model=app_model, workflow_run_id=run_id_str) + snapshot = service.snapshot_workflow_run(app_model=app_model, workflow_run_id=run_id_str, session=db.session()) except NodeOutputInspectorError as error: raise _InspectorNotFound(error) from error @@ -308,6 +313,7 @@ def _stream_inspector_events(app_model: App, run_id: UUID) -> Iterator[str]: app_model=app_model, workflow_run_id=run_id_str, node_id=message.node_id, + session=db.session(), ) except NodeOutputInspectorError: # Node may not appear in the graph yet (race with persistence); skip. diff --git a/api/controllers/console/auth/activate.py b/api/controllers/console/auth/activate.py index b6045685b55..1f58dbe910f 100644 --- a/api/controllers/console/auth/activate.py +++ b/api/controllers/console/auth/activate.py @@ -90,7 +90,7 @@ class ActivateCheckApi(Resource): token = args.token invitation = RegisterService.get_invitation_with_case_fallback( - workspaceId, args.email, token, session=db.session + workspaceId, args.email, token, session=db.session() ) if invitation: data = invitation.get("data", {}) @@ -140,7 +140,7 @@ class ActivateApi(Resource): normalized_request_email = args.email.lower() if args.email else None invitation = RegisterService.get_invitation_with_case_fallback( - args.workspace_id, args.email, args.token, session=db.session + args.workspace_id, args.email, args.token, session=db.session() ) if invitation is None: raise AlreadyActivateError() @@ -178,7 +178,7 @@ class ActivateApi(Resource): RegisterService.revoke_token(args.workspace_id, normalized_request_email, args.token) if membership_id is None: - TenantService.create_tenant_member(tenant, account, db.session, role=role) + TenantService.create_tenant_member(tenant, account, db.session(), role=role) if setup_fields: account.name = setup_fields[0] @@ -188,6 +188,6 @@ class ActivateApi(Resource): account.status = AccountStatus.ACTIVE account.initialized_at = naive_utc_now() - TenantService.switch_tenant(account, tenant.id, session=db.session) + TenantService.switch_tenant(account, tenant.id, session=db.session()) return {"result": "success"} diff --git a/api/controllers/console/auth/data_source_bearer_auth.py b/api/controllers/console/auth/data_source_bearer_auth.py index 11fab84a831..fac725e8534 100644 --- a/api/controllers/console/auth/data_source_bearer_auth.py +++ b/api/controllers/console/auth/data_source_bearer_auth.py @@ -59,7 +59,7 @@ class ApiKeyAuthDataSource(Resource): @account_initialization_required @with_current_tenant_id def get(self, current_tenant_id: str): - data_source_api_key_bindings = ApiKeyAuthService.get_provider_auth_list(db.session(), current_tenant_id) + data_source_api_key_bindings = ApiKeyAuthService.get_provider_auth_list(current_tenant_id, session=db.session()) if data_source_api_key_bindings: return { "sources": [ @@ -93,7 +93,7 @@ class ApiKeyAuthDataSourceBinding(Resource): data = payload.model_dump() ApiKeyAuthService.validate_api_key_auth_args(data) try: - ApiKeyAuthService.create_provider_auth(db.session(), current_tenant_id, data) + ApiKeyAuthService.create_provider_auth(current_tenant_id, data, session=db.session()) except Exception as e: raise ApiKeyAuthFailedError(str(e)) return {"result": "success"}, 200 @@ -110,6 +110,6 @@ class ApiKeyAuthDataSourceBindingDelete(Resource): @with_current_tenant_id def delete(self, current_tenant_id: str, binding_id: UUID): # The role of the current user in the table must be admin or owner - ApiKeyAuthService.delete_provider_auth(db.session(), current_tenant_id, str(binding_id)) + ApiKeyAuthService.delete_provider_auth(current_tenant_id, str(binding_id), session=db.session()) return "", 204 diff --git a/api/controllers/console/auth/email_register.py b/api/controllers/console/auth/email_register.py index ba4fc1275d9..d89caa9224f 100644 --- a/api/controllers/console/auth/email_register.py +++ b/api/controllers/console/auth/email_register.py @@ -101,7 +101,7 @@ class EmailRegisterSendEmailApi(Resource): if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): raise AccountInFreezeError() - account = AccountService.get_account_by_email_with_case_fallback(db.session, args.email) + account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language) return {"result": "success", "data": token} @@ -176,7 +176,7 @@ class EmailRegisterResetApi(Resource): email = register_data.get("email", "") normalized_email = email.lower() - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: raise EmailAlreadyInUseError() @@ -187,7 +187,7 @@ class EmailRegisterResetApi(Resource): timezone=args.timezone, language=args.language, ) - token_pair = AccountService.login(account=account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account=account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(normalized_email) return {"result": "success", "data": token_pair.model_dump()} @@ -206,7 +206,7 @@ class EmailRegisterResetApi(Resource): password=password, interface_language=get_valid_language(language), timezone=timezone, - session=db.session, + session=db.session(), ) except AccountRegisterError: raise AccountInFreezeError() diff --git a/api/controllers/console/auth/forgot_password.py b/api/controllers/console/auth/forgot_password.py index 8df9600070c..6456bb480f4 100644 --- a/api/controllers/console/auth/forgot_password.py +++ b/api/controllers/console/auth/forgot_password.py @@ -82,7 +82,7 @@ class ForgotPasswordSendEmailApi(Resource): else: language = "en-US" - account = AccountService.get_account_by_email_with_case_fallback(db.session, args.email) + account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) token = AccountService.send_reset_password_email( account=account, @@ -180,7 +180,7 @@ class ForgotPasswordResetApi(Resource): password_hashed = hash_password(args.new_password, salt) email = reset_data.get("email", "") - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: account = db.session.merge(account) @@ -198,10 +198,10 @@ class ForgotPasswordResetApi(Resource): # Create workspace if needed if ( - not TenantService.get_join_tenants(account, session=db.session) + not TenantService.get_join_tenants(account, session=db.session()) and FeatureService.get_system_features().is_allow_create_workspace ): - tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(tenant, account, db.session, role="owner") + tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(tenant, account, db.session(), role="owner") account.current_tenant = tenant tenant_was_created.send(tenant) diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index 5165fc3003a..486f79bcae2 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -126,7 +126,7 @@ class LoginApi(Resource): invitation_data: InvitationDetailDict | None = None if invite_token: invitation_data = RegisterService.get_invitation_with_case_fallback( - None, request_email, invite_token, session=db.session + None, request_email, invite_token, session=db.session() ) if invitation_data is None: invite_token = None @@ -153,7 +153,7 @@ class LoginApi(Resource): _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.INVALID_CREDENTIALS) raise AuthenticationFailedError() from exc # SELF_HOSTED only have one workspace - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if len(tenants) == 0: system_features = FeatureService.get_system_features() @@ -165,7 +165,7 @@ class LoginApi(Resource): data="workspace not found, please contact system admin to invite you to join in a workspace", ).model_dump(mode="json") - token_pair = AccountService.login(account=account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account=account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(normalized_email) # Create response with cookies instead of returning tokens in body @@ -301,7 +301,7 @@ class EmailCodeLoginApi(Resource): _log_console_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_IN_FREEZE) raise AccountInFreezeError() if account: - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if not tenants: workspaces = FeatureService.get_system_features().license.workspaces if not workspaces.is_available(): @@ -309,8 +309,8 @@ class EmailCodeLoginApi(Resource): if not FeatureService.get_system_features().is_allow_create_workspace: raise NotAllowedCreateWorkspace() else: - new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(new_tenant, account, db.session, role="owner") + new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(new_tenant, account, db.session(), role="owner") account.current_tenant = new_tenant tenant_was_created.send(new_tenant) @@ -321,7 +321,7 @@ class EmailCodeLoginApi(Resource): name=user_email, interface_language=get_valid_language(language), timezone=args.timezone, - session=db.session, + session=db.session(), ) except WorkSpaceNotAllowedCreateError: raise NotAllowedCreateWorkspace() @@ -330,7 +330,7 @@ class EmailCodeLoginApi(Resource): raise AccountInFreezeError() except WorkspacesLimitExceededError: raise WorkspacesLimitExceeded() - token_pair = AccountService.login(account, session=db.session, ip_address=extract_remote_ip(request)) + token_pair = AccountService.login(account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(user_email) # Create response with cookies instead of returning tokens in body @@ -358,7 +358,7 @@ class RefreshTokenApi(Resource): ), 401 try: - new_token_pair = AccountService.refresh_token(refresh_token, session=db.session) + new_token_pair = AccountService.refresh_token(refresh_token, session=db.session()) except Unauthorized as exc: return SimpleResultMessageResponse(result="fail", message=exc.description or "Unauthorized.").model_dump( mode="json" @@ -378,22 +378,22 @@ class RefreshTokenApi(Resource): def _get_account_with_case_fallback(email: str): - account = AccountService.get_user_through_email(email, session=db.session) + account = AccountService.get_user_through_email(email, session=db.session()) if account or email == email.lower(): return account - return AccountService.get_user_through_email(email.lower(), session=db.session) + return AccountService.get_user_through_email(email.lower(), session=db.session()) def _authenticate_account_with_case_fallback( original_email: str, normalized_email: str, password: str, invite_token: str | None ): try: - return AccountService.authenticate(original_email, password, invite_token, session=db.session) + return AccountService.authenticate(original_email, password, invite_token, session=db.session()) except services.errors.account.AccountPasswordError: if original_email == normalized_email: raise - return AccountService.authenticate(normalized_email, password, invite_token, session=db.session) + return AccountService.authenticate(normalized_email, password, invite_token, session=db.session()) def _log_console_login_failure(*, email: str, reason: LoginFailureReason) -> None: diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index 65f3a5addde..5afafd43131 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -195,7 +195,7 @@ class OAuthCallback(Resource): db.session.commit() try: - TenantService.create_owner_tenant_if_not_exist(account, session=db.session) + TenantService.create_owner_tenant_if_not_exist(account, session=db.session()) except Unauthorized: return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Workspace not found.") except WorkSpaceNotAllowedCreateError: @@ -206,7 +206,7 @@ class OAuthCallback(Resource): token_pair = AccountService.login( account=account, - session=db.session, + session=db.session(), ip_address=extract_remote_ip(request), ) @@ -225,7 +225,7 @@ def _get_account_by_openid_or_email(provider: str, user_info: OAuthUserInfo) -> account: Account | None = Account.get_by_openid(provider, user_info.id) if not account: - account = AccountService.get_account_by_email_with_case_fallback(db.session, user_info.email) + account = AccountService.get_account_by_email_with_case_fallback(user_info.email, session=db.session()) return account @@ -241,13 +241,13 @@ def _generate_account( oauth_new_user = False if account: - tenants = TenantService.get_join_tenants(account, session=db.session) + tenants = TenantService.get_join_tenants(account, session=db.session()) if not tenants: if not FeatureService.get_system_features().is_allow_create_workspace: raise WorkSpaceNotAllowedCreateError() else: - new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session) - TenantService.create_tenant_member(new_tenant, account, db.session, role="owner") + new_tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=db.session()) + TenantService.create_tenant_member(new_tenant, account, db.session(), role="owner") account.current_tenant = new_tenant tenant_was_created.send(new_tenant) @@ -273,10 +273,10 @@ def _generate_account( provider=provider, language=interface_language, timezone=timezone, - session=db.session, + session=db.session(), ) # Link account - AccountService.link_account_integrate(provider, user_info.id, account, session=db.session) + AccountService.link_account_integrate(provider, user_info.id, account, session=db.session()) return account, oauth_new_user diff --git a/api/controllers/console/auth/oauth_server.py b/api/controllers/console/auth/oauth_server.py index 46e2983c12b..d068fb0785e 100644 --- a/api/controllers/console/auth/oauth_server.py +++ b/api/controllers/console/auth/oauth_server.py @@ -10,6 +10,7 @@ from werkzeug.exceptions import BadRequest, NotFound from controllers.common.schema import register_response_schema_models, register_schema_models from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from extensions.ext_database import db from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.login import login_required from models import Account @@ -131,7 +132,9 @@ def oauth_server_access_token_required[T, **P, R]( response.headers["WWW-Authenticate"] = "Bearer" return response - account = OAuthServerService.validate_oauth_access_token(oauth_provider_app.client_id, access_token) + account = OAuthServerService.validate_oauth_access_token( + oauth_provider_app.client_id, access_token, db.session() + ) if not account: response = jsonify({"error": "access_token or client_id is invalid"}) response.status_code = 401 diff --git a/api/controllers/console/billing/billing.py b/api/controllers/console/billing/billing.py index d6974fe129c..3a983b50176 100644 --- a/api/controllers/console/billing/billing.py +++ b/api/controllers/console/billing/billing.py @@ -56,7 +56,7 @@ class Subscription(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account): args = SubscriptionQuery.model_validate(request.args.to_dict(flat=True)) - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) return BillingService.get_subscription(args.plan, args.interval, current_user.email, current_tenant_id) @@ -70,7 +70,7 @@ class Invoices(Resource): @with_current_user @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account): - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) return BillingService.get_invoices(current_user.email, current_tenant_id) diff --git a/api/controllers/console/datasets/data_source.py b/api/controllers/console/datasets/data_source.py index b2c8bda0581..17f027df9b3 100644 --- a/api/controllers/console/datasets/data_source.py +++ b/api/controllers/console/datasets/data_source.py @@ -245,7 +245,7 @@ class DataSourceNotionListApi(Resource): exist_page_ids = [] # import notion in the exist dataset if query.dataset_id: - dataset = DatasetService.get_dataset(query.dataset_id, db.session) + dataset = DatasetService.get_dataset(query.dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") if dataset.data_source_type != "notion_import": @@ -400,11 +400,11 @@ class DataSourceNotionDatasetSyncApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - documents = DocumentService.get_document_by_dataset_id(dataset_id_str, db.session) + documents = DocumentService.get_document_by_dataset_id(dataset_id_str, db.session()) for document in documents: document_indexing_sync_task.delay(dataset_id_str, document.id) return {"result": "success"}, 200 @@ -420,11 +420,11 @@ class DataSourceNotionDocumentSyncApi(Resource): def get(self, dataset_id: UUID, document_id: UUID) -> tuple[dict[str, str], int]: dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if document is None: raise NotFound("Document not found.") document_indexing_sync_task.delay(dataset_id_str, document_id_str) diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 0bef535d82b..c8ca1d621c4 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -418,6 +418,7 @@ class DatasetListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user.id, + session=db.session(), ) accessible_dataset_ids: list[str] | None = None @@ -461,7 +462,7 @@ class DatasetListApi(Resource): datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session, + db.session(), current_tenant_id, current_user, query.keyword, @@ -573,6 +574,7 @@ class DatasetListApi(Resource): current_tenant_id, current_user.id, [dataset.id], + session=session, ) item = DatasetDetailWithPartialMembersResponse.model_validate(dataset, from_attributes=True).model_dump( @@ -602,17 +604,18 @@ class DatasetApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) permissions = enterprise_rbac_service.RBACService.MyPermissions.get( current_tenant_id, current_user.id, dataset_id=dataset_id_str, + session=db.session(), ) permission_keys_map = permissions.dataset.permission_keys_by_resource_ids([dataset_id_str]) data = dump_response(DatasetDetailResponse, dataset) @@ -622,7 +625,7 @@ class DatasetApi(Resource): provider_id = ModelProviderID(dataset.embedding_model_provider) data["embedding_model_provider"] = str(provider_id) if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) data.update({"partial_member_list": part_users_list}) # check embedding setting @@ -666,7 +669,7 @@ class DatasetApi(Resource): @with_session def patch(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -685,10 +688,10 @@ class DatasetApi(Resource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not dify_config.RBAC_ENABLED: DatasetPermissionService.check_permission( - session, current_user, dataset, payload.permission, payload.partial_member_list + current_user, dataset, payload.permission, payload.partial_member_list, session=session ) - dataset = DatasetService.update_dataset(session, dataset_id_str, payload_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user, session=session) if dataset is None: raise NotFound("Dataset not found.") @@ -697,6 +700,7 @@ class DatasetApi(Resource): current_tenant_id, current_user.id, [dataset_id_str], + session=session, ) result_data = dump_response(DatasetDetailResponse, dataset) result_data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) @@ -704,13 +708,13 @@ class DatasetApi(Resource): if payload.partial_member_list is not None and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: DatasetPermissionService.update_partial_member_list( - tenant_id, dataset_id_str, payload.partial_member_list, db.session + tenant_id, dataset_id_str, payload.partial_member_list, db.session() ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) result_data.update({"partial_member_list": partial_member_list}) return result_data, 200 @@ -729,8 +733,8 @@ class DatasetApi(Resource): raise Forbidden() try: - if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session()): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) return "", 204 else: raise NotFound("Dataset not found.") @@ -755,7 +759,7 @@ class DatasetUseCheckApi(Resource): def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, db.session) + dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, db.session()) return {"is_using": dataset_is_using}, 200 @@ -776,12 +780,12 @@ class DatasetQueryApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -917,16 +921,16 @@ class DatasetRelatedAppListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - app_dataset_joins = DatasetService.get_related_apps(dataset.id, db.session) + app_dataset_joins = DatasetService.get_related_apps(dataset.id, db.session()) related_apps = [] for app_dataset_join in app_dataset_joins: @@ -1101,7 +1105,7 @@ class DatasetEnableApiApi(Resource): def post(self, dataset_id: UUID, status: str): dataset_id_str = str(dataset_id) - DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", db.session) + DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", db.session()) return {"result": "success"}, 200 @@ -1170,10 +1174,10 @@ class DatasetErrorDocs(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, db.session) + results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, db.session()) return dump_response(ErrorDocsResponse, {"data": results, "total": len(results)}), 200 @@ -1197,15 +1201,15 @@ class DatasetPermissionUserListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_members_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) return dump_response(PartialMemberListResponse, {"data": partial_members_list}), 200 @@ -1227,8 +1231,8 @@ class DatasetAutoDisableLogApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, db.session) + auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, db.session()) return dump_response(AutoDisableLogsResponse, auto_disable_logs), 200 diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index a6263c8e2e3..ee441704b20 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -183,16 +183,16 @@ class DocumentResource(Resource): def get_document( self, dataset_id: str, document_id: str, current_user: Account, current_tenant_id: str ) -> Document: - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id, document_id, session=db.session) + document = DocumentService.get_document(dataset_id, document_id, session=db.session()) if not document: raise NotFound("Document not found.") @@ -203,16 +203,16 @@ class DocumentResource(Resource): return document def get_batch_documents(self, dataset_id: str, batch: str, current_user: Account) -> Sequence[Document]: - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - documents = DocumentService.get_batch_documents(dataset_id, batch, db.session) + documents = DocumentService.get_batch_documents(dataset_id, batch, db.session()) if not documents: raise NotFound("Documents not found.") @@ -243,13 +243,13 @@ class GetProcessRuleApi(Resource): # get the latest process rule document = db.get_or_404(Document, document_id) - dataset = DatasetService.get_dataset(document.dataset_id, db.session) + dataset = DatasetService.get_dataset(document.dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -319,12 +319,12 @@ class DatasetDocumentListApi(Resource): ) except (ArgumentTypeError, ValueError, Exception): fetch = False - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -376,6 +376,7 @@ class DatasetDocumentListApi(Resource): documents=documents, dataset=dataset, tenant_id=current_tenant_id, + session=db.session(), ) if fetch: @@ -423,7 +424,7 @@ class DatasetDocumentListApi(Resource): def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") @@ -433,7 +434,7 @@ class DatasetDocumentListApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -447,9 +448,9 @@ class DatasetDocumentListApi(Resource): try: documents, batch = DocumentService.save_document_with_dataset_id( - dataset, knowledge_config, current_user, session=db.session + dataset, knowledge_config, current_user, session=db.session() ) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -468,7 +469,7 @@ class DatasetDocumentListApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def delete(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -477,7 +478,7 @@ class DatasetDocumentListApi(Resource): try: document_ids = request.args.getlist("document_id") dataset_ref = DatasetRefService.create_dataset_ref(dataset) - DocumentService.delete_documents(dataset_ref, document_ids, dataset.doc_form, db.session) + DocumentService.delete_documents(dataset_ref, document_ids, dataset.doc_form, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -536,7 +537,7 @@ class DatasetInitApi(Resource): tenant_id=current_tenant_id, knowledge_config=knowledge_config, account=current_user, - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -873,7 +874,7 @@ class DocumentApi(DocumentResource): if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -907,7 +908,7 @@ class DocumentApi(DocumentResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} response = { "id": document.id, @@ -956,7 +957,7 @@ class DocumentApi(DocumentResource): def delete(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -965,7 +966,7 @@ class DocumentApi(DocumentResource): document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) try: - DocumentService.delete_document(document, db.session) + DocumentService.delete_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -989,7 +990,7 @@ class DocumentDownloadApi(DocumentResource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID) -> dict[str, Any]: # Reuse the shared permission/tenant checks implemented in DocumentResource. document = self.get_document(str(dataset_id), str(document_id), current_user, current_tenant_id) - return {"url": DocumentService.get_document_download_url(document, db.session)} + return {"url": DocumentService.get_document_download_url(document, db.session())} @console_ns.route("/datasets//documents/download-zip") @@ -1019,7 +1020,7 @@ class DocumentBatchDownloadZipApi(DocumentResource): document_ids=document_ids, tenant_id=current_tenant_id, current_user=current_user, - session=db.session, + session=db.session(), ) # Delegate ZIP packing to FileService, but keep Flask response+cleanup in the route. @@ -1168,7 +1169,7 @@ class DocumentStatusApi(DocumentResource): self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable", "archive", "un_archive"] ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -1180,12 +1181,12 @@ class DocumentStatusApi(DocumentResource): DatasetService.check_dataset_model_setting(dataset) # check user's permission - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) document_ids = request.args.getlist("document_id") try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session()) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -1209,11 +1210,11 @@ class DocumentPauseApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1225,7 +1226,7 @@ class DocumentPauseApi(DocumentResource): try: # pause document - DocumentService.pause_document(document, db.session) + DocumentService.pause_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot pause completed document.") @@ -1244,10 +1245,10 @@ class DocumentRecoverApi(DocumentResource): """recover document.""" dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1258,7 +1259,7 @@ class DocumentRecoverApi(DocumentResource): raise ArchivedDocumentImmutableError() try: # pause document - DocumentService.recover_document(document, db.session) + DocumentService.recover_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Document is not in paused status.") @@ -1278,13 +1279,13 @@ class DocumentRetryApi(DocumentResource): """retry document.""" payload = DocumentRetryPayload.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) retry_documents = [] if not dataset: raise NotFound("Dataset not found.") for document_id in payload.document_ids: try: - document = DocumentService.get_document(dataset.id, document_id, session=db.session) + document = DocumentService.get_document(dataset.id, document_id, session=db.session()) # 404 if document not found if document is None: @@ -1302,7 +1303,7 @@ class DocumentRetryApi(DocumentResource): logger.exception("Failed to retry document, document id: %s", document_id) continue # retry document - DocumentService.retry_document(dataset_id_str, retry_documents, db.session) + DocumentService.retry_document(dataset_id_str, retry_documents, db.session()) return "", 204 @@ -1320,14 +1321,14 @@ class DocumentRenameApi(DocumentResource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not current_user.is_dataset_editor: raise Forbidden() - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: raise NotFound("Dataset not found.") - DatasetService.check_dataset_operator_permission(current_user, dataset, session=db.session) + DatasetService.check_dataset_operator_permission(current_user, dataset, session=db.session()) payload = DocumentRenamePayload.model_validate(console_ns.payload or {}) try: - document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, db.session) + document = DocumentService.rename_document(str(dataset_id), str(document_id), payload.name, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -1345,11 +1346,11 @@ class WebsiteDocumentSyncApi(DocumentResource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID): """sync website document.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if document.tenant_id != current_tenant_id: @@ -1360,7 +1361,7 @@ class WebsiteDocumentSyncApi(DocumentResource): if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() # sync document - DocumentService.sync_website_document(dataset_id_str, document, db.session) + DocumentService.sync_website_document(dataset_id_str, document, db.session()) return {"result": "success"}, 200 @@ -1380,10 +1381,10 @@ class DocumentPipelineExecutionLogApi(DocumentResource): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") log = db.session.scalar( @@ -1438,7 +1439,7 @@ class DocumentGenerateSummaryApi(Resource): dataset_id_str = str(dataset_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") @@ -1447,7 +1448,7 @@ class DocumentGenerateSummaryApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1472,7 +1473,7 @@ class DocumentGenerateSummaryApi(Resource): raise ValueError("Summary index is not enabled for this dataset. Please enable it in the dataset settings.") # Verify all documents exist and belong to the dataset - documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, db.session) + documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, db.session()) if len(documents) != len(document_list): found_ids = {doc.id for doc in documents} @@ -1488,7 +1489,7 @@ class DocumentGenerateSummaryApi(Resource): DocumentService.update_documents_need_summary( dataset_id=dataset_id_str, document_ids=document_ids_to_update, - session=db.session, + session=db.session(), need_summary=True, ) @@ -1539,13 +1540,13 @@ class DocumentSummaryStatusApi(DocumentResource): document_id_str = str(document_id) # Get dataset - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # Check permissions try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1555,7 +1556,7 @@ class DocumentSummaryStatusApi(DocumentResource): result = SummaryIndexService.get_document_summary_status_detail( document_id=document_id_str, dataset_id=dataset_id_str, - session=db.session, + session=db.session(), ) return result, 200 diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 5cccd2453dc..e4f2abeb844 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -173,7 +173,7 @@ def _get_segment_for_document( raise NotFound("Document not found.") segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) - segment = SegmentService.get_segment_by_ref(segment_ref) + segment = SegmentService.get_segment_by_ref(segment_ref, db.session()) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -193,16 +193,16 @@ class DatasetDocumentSegmentListApi(Resource): def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -278,7 +278,7 @@ class DatasetDocumentSegmentListApi(Resource): summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: summary.summary_content for chunk_id, summary in summary_records.items()} @@ -303,14 +303,14 @@ class DatasetDocumentSegmentListApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_ids = request.args.getlist("segment_id") @@ -319,10 +319,10 @@ class DatasetDocumentSegmentListApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - SegmentService.delete_segments(segment_ids, document, dataset, db.session) + SegmentService.delete_segments(segment_ids, document, dataset, db.session()) return "", 204 @@ -348,11 +348,11 @@ class DatasetDocumentSegmentApi(Resource): action: Literal["enable", "disable"], ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # check user's model setting @@ -362,7 +362,7 @@ class DatasetDocumentSegmentApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -388,7 +388,7 @@ class DatasetDocumentSegmentApi(Resource): if cache_result is not None: raise InvalidActionError("Document is being indexed, please try again later") try: - SegmentService.update_segments_status(segment_ids, action, dataset, document, db.session) + SegmentService.update_segments_status(segment_ids, action, dataset, document, db.session()) except Exception as e: raise InvalidActionError(str(e)) return dump_response(SimpleResultResponse, {"result": "success"}), 200 @@ -411,12 +411,12 @@ class DatasetDocumentSegmentAddApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: @@ -438,15 +438,20 @@ class DatasetDocumentSegmentAddApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args payload = SegmentCreatePayload.model_validate(console_ns.payload or {}) payload_dict = payload.model_dump(exclude_none=True) SegmentService.segment_create_args_validate(payload_dict, document) - segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset, db.session)) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) + segment = type_cast( + DocumentSegment, + SegmentService.create_segment(payload_dict, document, dataset, db.session()), + ) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -472,21 +477,21 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -518,9 +523,11 @@ class DatasetDocumentSegmentUpdateApi(Resource): segment, document, dataset, - db.session, + db.session(), + ) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() ) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -541,26 +548,26 @@ class DatasetDocumentSegmentUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - SegmentService.delete_segment(segment, document, dataset, db.session) + SegmentService.delete_segment(segment, document, dataset, db.session()) return "", 204 @@ -583,12 +590,12 @@ class DatasetDocumentSegmentBatchImportApi(Resource): def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -658,18 +665,18 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # check embedding model setting @@ -693,7 +700,7 @@ class ChildChunkAddApi(Resource): # validate args try: payload = ChildChunkCreatePayload.model_validate(console_ns.payload or {}) - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 @@ -709,14 +716,14 @@ class ChildChunkAddApi(Resource): def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) @@ -759,21 +766,21 @@ class ChildChunkAddApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) @@ -781,7 +788,7 @@ class ChildChunkAddApi(Resource): # validate args payload = ChildChunkBatchUpdatePayload.model_validate(console_ns.payload or {}) try: - child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, db.session) + child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkBatchUpdateResponse, {"data": child_chunks}), 200 @@ -811,31 +818,31 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) segment_ref, _ = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") try: - SegmentService.delete_child_chunk(child_chunk, dataset, db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session()) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) return "", 204 @@ -862,34 +869,34 @@ class ChildChunkUpdateApi(Resource): ): # check dataset dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if not dataset: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) # check document document_id_str = str(document_id) - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # The role of the current user in the ta table must be admin, owner, dataset_operator, or editor if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) segment_id_str = str(segment_id) segment_ref, segment = _get_segment_for_document(dataset, document, segment_id_str) child_chunk_id_str = str(child_chunk_id) - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") # validate args try: payload = ChildChunkUpdatePayload.model_validate(console_ns.payload or {}) child_chunk = SegmentService.update_child_chunk( - payload.content, child_chunk, segment, document, dataset, db.session + payload.content, child_chunk, segment, document, dataset, db.session() ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index 5b036641d4d..9cdca96f69a 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -299,7 +299,9 @@ class ExternalApiTemplateApi(Resource): if not (current_user.has_edit_permission or current_user.is_dataset_operator): raise Forbidden() - ExternalDatasetService.delete_external_knowledge_api(session, current_tenant_id, external_knowledge_api_id_str) + ExternalDatasetService.delete_external_knowledge_api( + current_tenant_id, external_knowledge_api_id_str, session=session + ) return "", 204 @@ -318,9 +320,7 @@ class ExternalApiUseCheckApi(Resource): external_knowledge_api_id_str = str(external_knowledge_api_id) external_knowledge_api_is_using, count = ExternalDatasetService.external_knowledge_api_use_check( - session, - external_knowledge_api_id_str, - current_tenant_id, + external_knowledge_api_id_str, current_tenant_id, session=session ) return {"is_using": external_knowledge_api_is_using, "count": count}, 200 @@ -366,6 +366,7 @@ class ExternalDatasetCreateApi(Resource): str(current_tenant_id), current_user.id, [dataset_id_str], + session=session, ) item["permission_keys"] = permission_keys_map.get(dataset_id_str, []) @@ -393,12 +394,12 @@ class ExternalKnowledgeHitTestingApi(Resource): @with_session def post(self, session: Session, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index cc02a990168..656a426c125 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -86,12 +86,12 @@ class DatasetsHitTestingBase: dataset_id: str, current_user: Account | None = None, current_tenant_id: str | None = None ) -> Dataset: current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/metadata.py b/api/controllers/console/datasets/metadata.py index 8802fcf2814..42ae4903673 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -61,13 +61,13 @@ class DatasetMetadataCreateApi(Resource): metadata_args = MetadataArgs.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata = MetadataService.create_metadata( - db.session(), dataset_id_str, metadata_args, current_user, current_tenant_id + dataset_id_str, metadata_args, current_user, current_tenant_id, session=db.session() ) return dump_response(DatasetMetadataResponse, metadata), 201 @@ -81,10 +81,10 @@ class DatasetMetadataCreateApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) def get(self, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) + metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -105,13 +105,13 @@ class DatasetMetadataApi(Resource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata = MetadataService.update_metadata_name( - db.session(), dataset_id_str, metadata_id_str, name, current_user, current_tenant_id + dataset_id_str, metadata_id_str, name, current_user, current_tenant_id, session=db.session() ) return dump_response(DatasetMetadataResponse, metadata), 200 @@ -125,12 +125,12 @@ class DatasetMetadataApi(Resource): def delete(self, current_user: Account, dataset_id: UUID, metadata_id: UUID): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -162,16 +162,16 @@ class DatasetMetadataBuiltInFieldActionApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable"]): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) match action: case "enable": - MetadataService.enable_built_in_field(db.session(), dataset) + MetadataService.enable_built_in_field(dataset, session=db.session()) case "disable": - MetadataService.disable_built_in_field(db.session(), dataset) + MetadataService.disable_built_in_field(dataset, session=db.session()) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -191,14 +191,14 @@ class DocumentMetadataEditApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) def post(self, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata_args = MetadataOperationData.model_validate(console_ns.payload or {}) - MetadataService.update_documents_metadata(db.session(), dataset, metadata_args, current_user) + MetadataService.update_documents_metadata(dataset, metadata_args, current_user, session=db.session()) # Frontend callers only await success and invalidate caches; no response body is consumed. return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py index 389515bd4f9..57d6b628d4b 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py @@ -23,6 +23,7 @@ from core.entities.provider_entities import ProviderConfig from core.plugin.entities.plugin_daemon import PluginOAuthAuthorizationUrlResponse from core.plugin.impl.oauth import OAuthHandler from core.tools.entities.common_entities import I18nObject +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.errors.validate import CredentialsValidateFailedError from libs.helper import dump_response @@ -309,6 +310,7 @@ class DatasourceAuth(Resource): provider=datasource_provider_id.provider_name, plugin_id=datasource_provider_id.plugin_id, user=user, + session=db.session(), ) return dump_response(DatasourceCredentialListResponse, {"result": datasources}), 200 @@ -335,6 +337,7 @@ class DatasourceAuthDeleteApi(Resource): auth_id=payload.credential_id, provider=provider_name, plugin_id=plugin_id, + session=db.session(), ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -380,7 +383,9 @@ class DatasourceAuthListApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str): datasource_provider_service = DatasourceProviderService() - datasources = datasource_provider_service.get_all_datasource_credentials(tenant_id=current_tenant_id) + datasources = datasource_provider_service.get_all_datasource_credentials( + tenant_id=current_tenant_id, session=db.session() + ) return dump_response(DatasourceProviderAuthListResponse, {"result": datasources}), 200 @@ -397,7 +402,9 @@ class DatasourceHardCodeAuthListApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str): datasource_provider_service = DatasourceProviderService() - datasources = datasource_provider_service.get_hard_code_datasource_credentials(tenant_id=current_tenant_id) + datasources = datasource_provider_service.get_hard_code_datasource_credentials( + tenant_id=current_tenant_id, session=db.session() + ) return dump_response(DatasourceProviderAuthListResponse, {"result": datasources}), 200 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py index 213337fedc9..873ba130064 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py @@ -9,6 +9,7 @@ from controllers.common.schema import register_schema_models from controllers.console import console_ns from controllers.console.datasets.wraps import get_rag_pipeline from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from extensions.ext_database import db from libs.login import login_required from models import Account from models.dataset import Pipeline @@ -41,7 +42,7 @@ class DataSourceContentPreviewApi(Resource): inputs = args.inputs datasource_type = args.datasource_type - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) preview_content = rag_pipeline_service.run_datasource_node_preview( pipeline=pipeline, node_id=node_id, diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py index 4027fa487a2..2d824afb6ef 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py @@ -108,7 +108,10 @@ class PipelineTemplateListApi(Resource): query = PipelineTemplateListQuery.model_validate(request.args.to_dict(flat=True)) # get pipeline templates pipeline_templates = RagPipelineService.get_pipeline_templates( - session, query.type, query.language, current_tenant_id + type=query.type, + language=query.language, + current_tenant_id=current_tenant_id, + session=session, ) return dump_response(PipelineTemplateListResponse, pipeline_templates), 200 @@ -124,8 +127,11 @@ class PipelineTemplateDetailApi(Resource): @with_session def get(self, session: Session, template_id: str) -> JsonResponseWithStatus: query = PipelineTemplateDetailQuery.model_validate(request.args.to_dict(flat=True)) - rag_pipeline_service = RagPipelineService() - pipeline_template = rag_pipeline_service.get_pipeline_template_detail(session, template_id, query.type) + pipeline_template = RagPipelineService.get_pipeline_template_detail( + template_id, + type=query.type, + session=session, + ) if pipeline_template is None: raise NotFound("Pipeline template not found from upstream service.") return dump_response(PipelineTemplateDetailResponse, pipeline_template), 200 @@ -145,7 +151,7 @@ class CustomizedPipelineTemplateApi(Resource): payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) pipeline_template_info = PipelineTemplateInfoEntity.model_validate(payload.model_dump()) RagPipelineService.update_customized_pipeline_template( - template_id, pipeline_template_info, current_user, current_tenant_id + template_id, pipeline_template_info, current_user, current_tenant_id, session=db.session() ) return "", 204 @@ -156,7 +162,7 @@ class CustomizedPipelineTemplateApi(Resource): @enterprise_license_required @with_current_tenant_id def delete(self, current_tenant_id: str, template_id: str) -> tuple[str, int]: - RagPipelineService.delete_customized_pipeline_template(template_id, current_tenant_id) + RagPipelineService.delete_customized_pipeline_template(template_id, current_tenant_id, session=db.session()) return "", 204 @setup_required @@ -188,8 +194,8 @@ class PublishCustomizedPipelineTemplateApi(Resource): @with_current_tenant_id def post(self, current_tenant_id: str, current_user: Account, pipeline_id: str) -> tuple[str, int]: payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) rag_pipeline_service.publish_customized_pipeline_template( - pipeline_id, payload.model_dump(), current_user, current_tenant_id + pipeline_id, payload.model_dump(), current_user, current_tenant_id, session=db.session() ) return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py index a373c8b1a41..5ad764871e4 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py @@ -65,7 +65,7 @@ class CreateRagPipelineDatasetApi(Resource): yaml_content=payload.yaml_content, ) try: - rag_pipeline_dsl_service = RagPipelineDslService(db.session) + rag_pipeline_dsl_service = RagPipelineDslService(db.session()) import_info = rag_pipeline_dsl_service.create_rag_pipeline_dataset( tenant_id=current_tenant_id, rag_pipeline_dataset_create_entity=rag_pipeline_dataset_create_entity, @@ -75,7 +75,7 @@ class CreateRagPipelineDatasetApi(Resource): current_tenant_id, import_info["dataset_id"], rag_pipeline_dataset_create_entity.partial_member_list, - db.session, + db.session(), ) db.session.commit() except services.errors.dataset.DatasetNameDuplicateError: @@ -110,6 +110,6 @@ class CreateEmptyRagPipelineDatasetApi(Resource): permission=DatasetPermissionEnum.ONLY_ME, partial_member_list=None, ), - session=db.session, + session=db.session(), ) return dump_response(DatasetDetailResponse, dataset), 201 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py index af417f24dfe..25628a67177 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py @@ -98,7 +98,7 @@ class RagPipelineVariableCollectionApi(Resource): query = PaginationQuery.model_validate(request.args.to_dict()) # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_exist = rag_pipeline_service.is_workflow_exist(pipeline=pipeline) if not workflow_exist: raise DraftWorkflowNotExist() @@ -290,7 +290,7 @@ class RagPipelineVariableResetApi(Resource): session=db.session(), ) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) draft_workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if draft_workflow is None: raise NotFoundError( @@ -347,7 +347,7 @@ class RagPipelineEnvironmentVariableCollectionApi(Resource): Get draft workflow """ # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if workflow is None: raise DraftWorkflowNotExist() diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index c52385f6cf2..a61fc2639db 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -197,7 +197,7 @@ class DraftRagPipelineApi(Resource): Get draft rag pipeline's workflow """ # fetch draft workflow by app_model - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if not workflow: @@ -231,7 +231,7 @@ class DraftRagPipelineApi(Resource): return {"message": "Invalid JSON data"}, 400 else: abort(415) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) try: environment_variables_list = Workflow.normalize_environment_variable_mappings( @@ -283,7 +283,7 @@ class RagPipelineDraftRunIterationNodeApi(Resource): try: response = PipelineGenerateService.generate_single_iteration( - pipeline=pipeline, user=current_user, node_id=node_id, args=args, streaming=True + pipeline=pipeline, user=current_user, node_id=node_id, args=args, session=db.session(), streaming=True ) return helper.compact_generate_response(response) @@ -318,7 +318,7 @@ class RagPipelineDraftRunLoopNodeApi(Resource): try: response = PipelineGenerateService.generate_single_loop( - pipeline=pipeline, user=current_user, node_id=node_id, args=args, streaming=True + pipeline=pipeline, user=current_user, node_id=node_id, args=args, session=db.session(), streaming=True ) return helper.compact_generate_response(response) @@ -343,8 +343,8 @@ class DraftRagPipelineRunApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user - @get_rag_pipeline @with_session + @get_rag_pipeline def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run draft workflow @@ -377,8 +377,8 @@ class PublishedRagPipelineRunApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user - @get_rag_pipeline @with_session + @get_rag_pipeline def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run published workflow @@ -419,7 +419,7 @@ class RagPipelinePublishedDatasourceNodeRunApi(Resource): """ payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( PipelineGenerator.convert_to_event_stream( rag_pipeline_service.run_datasource_workflow_node( @@ -452,7 +452,7 @@ class RagPipelineDraftDatasourceNodeRunApi(Resource): """ payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( PipelineGenerator.convert_to_event_stream( rag_pipeline_service.run_datasource_workflow_node( @@ -490,7 +490,7 @@ class RagPipelineDraftNodeRunApi(Resource): payload = NodeRunRequiredPayload.model_validate(console_ns.payload or {}) inputs = payload.inputs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.run_draft_workflow_node( pipeline=pipeline, node_id=node_id, user_inputs=inputs, account=current_user ) @@ -543,7 +543,7 @@ class PublishedRagPipelineApi(Resource): if not pipeline.is_published: return None # fetch published workflow by pipeline - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_published_workflow(pipeline=pipeline) # return workflow, if not found, return None @@ -564,9 +564,9 @@ class PublishedRagPipelineApi(Resource): """ Publish workflow """ - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.publish_workflow( - session=db.session, # type: ignore[reportArgumentType,arg-type] + session=db.session(), pipeline=pipeline, account=current_user, ) @@ -599,7 +599,7 @@ class DefaultRagPipelineBlockConfigsApi(Resource): Get default block config """ # Get default block configs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return rag_pipeline_service.get_default_block_configs() @@ -631,7 +631,7 @@ class DefaultRagPipelineBlockConfigApi(Resource): raise ValueError("Invalid filters") # Get default block configs - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) return rag_pipeline_service.get_default_block_config(node_type=block_type, filters=filters) @@ -666,7 +666,7 @@ class PublishedAllRagPipelineApi(Resource): if user_id != current_user.id: raise Forbidden() - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) with sessionmaker(db.engine).begin() as session: workflows, has_more = rag_pipeline_service.get_all_published_workflow( session=session, @@ -698,7 +698,7 @@ class RagPipelineDraftWorkflowRestoreApi(Resource): @with_current_user @get_rag_pipeline def post(self, current_user: Account, pipeline: Pipeline, workflow_id: str): - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) try: workflow = rag_pipeline_service.restore_published_workflow_to_draft( @@ -743,7 +743,7 @@ class RagPipelineByIdApi(Resource): if not update_data: return {"message": "No valid fields to update"}, 400 - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_ref = WorkflowRefService.create_pipeline_workflow_ref(pipeline, workflow_id) # Create a session and manage the transaction @@ -809,7 +809,7 @@ class PublishedRagPipelineSecondStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { "variables": variables, @@ -832,7 +832,7 @@ class PublishedRagPipelineFirstStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { "variables": variables, @@ -855,7 +855,7 @@ class DraftRagPipelineFirstStepApi(Resource): """ query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) return { "variables": variables, @@ -879,7 +879,7 @@ class DraftRagPipelineSecondStepApi(Resource): query = NodeIdQuery.model_validate(request.args.to_dict()) node_id = query.node_id - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) return { "variables": variables, @@ -913,7 +913,7 @@ class RagPipelineWorkflowRunListApi(Resource): "limit": query.limit, } - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) result = rag_pipeline_service.get_rag_pipeline_paginate_workflow_runs(pipeline=pipeline, args=args) return WorkflowRunPaginationResponse.model_validate(result, from_attributes=True).model_dump(mode="json") @@ -936,7 +936,7 @@ class RagPipelineWorkflowRunDetailApi(Resource): """ run_id_str = str(run_id) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_run = rag_pipeline_service.get_rag_pipeline_workflow_run(pipeline=pipeline, run_id=run_id_str) if workflow_run is None: raise NotFound("Workflow run not found") @@ -962,7 +962,7 @@ class RagPipelineWorkflowRunNodeExecutionListApi(Resource): """ run_id_str = str(run_id) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) user = cast("Account | EndUser", current_user) node_executions = rag_pipeline_service.get_rag_pipeline_workflow_run_node_executions( pipeline=pipeline, @@ -998,7 +998,7 @@ class RagPipelineWorkflowLastRunApi(Resource): @account_initialization_required @get_rag_pipeline def get(self, pipeline: Pipeline, node_id: str): - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) if not workflow: raise NotFound("Workflow not found") @@ -1051,7 +1051,7 @@ class RagPipelineDatasourceVariableApi(Resource): """ args = DatasourceVariablesPayload.model_validate(console_ns.payload or {}).model_dump() - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.set_datasource_variables( pipeline=pipeline, args=args, @@ -1074,6 +1074,6 @@ class RagPipelineRecommendedPluginApi(Resource): def get(self, current_tenant_id: str, current_user: Account): query = RagPipelineRecommendedPluginQuery.model_validate(request.args.to_dict()) - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) recommended_plugins = rag_pipeline_service.get_recommended_plugins(query.type, current_user, current_tenant_id) return recommended_plugins diff --git a/api/controllers/console/datasets/wraps.py b/api/controllers/console/datasets/wraps.py index b58a07029c8..b5a9cd753ff 100644 --- a/api/controllers/console/datasets/wraps.py +++ b/api/controllers/console/datasets/wraps.py @@ -2,6 +2,7 @@ from collections.abc import Callable from functools import wraps from sqlalchemy import select +from sqlalchemy.orm import Session from controllers.console.datasets.error import PipelineNotFoundError from extensions.ext_database import db @@ -22,9 +23,10 @@ def get_rag_pipeline[**P, R](view_func: Callable[P, R]) -> Callable[P, R]: del kwargs["pipeline_id"] - pipeline = db.session.scalar( - select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id).limit(1) - ) + stmt = select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == current_tenant_id).limit(1) + # Migrated handlers pass the request Session as args[1]; legacy handlers still use db.session. + session = args[1] if len(args) > 1 and isinstance(args[1], Session) else db.session + pipeline = session.scalar(stmt) if not pipeline: raise PipelineNotFoundError() diff --git a/api/controllers/console/explore/audio.py b/api/controllers/console/explore/audio.py index c0b86c19e43..e5f98f0f655 100644 --- a/api/controllers/console/explore/audio.py +++ b/api/controllers/console/explore/audio.py @@ -113,7 +113,7 @@ class ChatTextApi(InstalledAppResource): response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, message_ref=message_ref, diff --git a/api/controllers/console/explore/conversation.py b/api/controllers/console/explore/conversation.py index 2004e648f19..25239203d8d 100644 --- a/api/controllers/console/explore/conversation.py +++ b/api/controllers/console/explore/conversation.py @@ -111,7 +111,7 @@ class ConversationApi(InstalledAppResource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, current_user) + ConversationService.delete(app_model, conversation_id, current_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -140,7 +140,7 @@ class ConversationRenameApi(InstalledAppResource): try: conversation = ConversationService.rename( - app_model, conversation_id, current_user, payload.name, payload.auto_generate + app_model, conversation_id, current_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -169,7 +169,7 @@ class ConversationPinApi(InstalledAppResource): conversation_id = str(c_id) try: - WebConversationService.pin(app_model, conversation_id, current_user) + WebConversationService.pin(app_model, conversation_id, current_user, db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -192,6 +192,6 @@ class ConversationUnPinApi(InstalledAppResource): raise NotChatAppError() conversation_id = str(c_id) - WebConversationService.unpin(app_model, conversation_id, current_user) + WebConversationService.unpin(app_model, conversation_id, current_user, db.session()) return ResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/explore/installed_app.py b/api/controllers/console/explore/installed_app.py index 71cb03ce6a0..1fe1201bab7 100644 --- a/api/controllers/console/explore/installed_app.py +++ b/api/controllers/console/explore/installed_app.py @@ -181,7 +181,7 @@ class InstalledAppsListApi(Resource): if current_user.current_tenant is None: raise ValueError("current_user.current_tenant must not be None") - current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session) + current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session()) installed_app_list: list[dict[str, Any]] = [] for installed_app, app_model in installed_apps: installed_app_list.append( diff --git a/api/controllers/console/explore/message.py b/api/controllers/console/explore/message.py index 0e27e2db25b..7b316b0382d 100644 --- a/api/controllers/console/explore/message.py +++ b/api/controllers/console/explore/message.py @@ -27,6 +27,7 @@ from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import with_current_user from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError +from extensions.ext_database import db from fields.conversation_fields import ResultResponse from fields.message_fields import ( ExploreMessageInfiniteScrollPagination, @@ -91,6 +92,7 @@ class MessageListApi(InstalledAppResource): args.conversation_id, args.first_id or None, args.limit, + session=db.session(), ) adapter = TypeAdapter(ExploreMessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -129,6 +131,7 @@ class MessageFeedbackApi(InstalledAppResource): user=current_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -207,7 +210,11 @@ class MessageSuggestedQuestionApi(InstalledAppResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=current_user, message_id=message_id_str, invoke_from=InvokeFrom.EXPLORE + app_model=app_model, + user=current_user, + message_id=message_id_str, + invoke_from=InvokeFrom.EXPLORE, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") diff --git a/api/controllers/console/explore/parameter.py b/api/controllers/console/explore/parameter.py index 0bc6e032bf0..680885f9bd3 100644 --- a/api/controllers/console/explore/parameter.py +++ b/api/controllers/console/explore/parameter.py @@ -8,6 +8,7 @@ from controllers.console import console_ns from controllers.console.app.error import AppUnavailableError from controllers.console.explore.wraps import InstalledAppResource from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict +from extensions.ext_database import db from models.model import AppMode, InstalledApp from services.app_service import AppService @@ -64,4 +65,4 @@ class ExploreAppMetaApi(InstalledAppResource): app_model = installed_app.app if not app_model: raise ValueError("App not found") - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) diff --git a/api/controllers/console/explore/recommended_app.py b/api/controllers/console/explore/recommended_app.py index abe170bf90a..79eaa305d61 100644 --- a/api/controllers/console/explore/recommended_app.py +++ b/api/controllers/console/explore/recommended_app.py @@ -120,7 +120,7 @@ class RecommendedAppListApi(Resource): language_prefix = _resolve_language(args.language, current_user) return RecommendedAppListResponse.model_validate( - RecommendedAppService.get_recommended_apps_and_categories(db.session, language_prefix), + RecommendedAppService.get_recommended_apps_and_categories(language_prefix, session=db.session()), from_attributes=True, ).model_dump(mode="json") @@ -137,7 +137,7 @@ class LearnDifyAppListApi(Resource): language_prefix = _resolve_language(args.language, current_user) return LearnDifyAppListResponse.model_validate( - RecommendedAppService.get_learn_dify_apps(db.session, language_prefix), + RecommendedAppService.get_learn_dify_apps(language_prefix, session=db.session()), from_attributes=True, ).model_dump(mode="json") @@ -148,4 +148,4 @@ class RecommendedAppApi(Resource): @login_required @account_initialization_required def get(self, app_id: UUID): - return RecommendedAppService.get_recommend_app_detail(db.session, str(app_id)) + return RecommendedAppService.get_recommend_app_detail(str(app_id), session=db.session()) diff --git a/api/controllers/console/explore/saved_message.py b/api/controllers/console/explore/saved_message.py index ce43ff18c93..e3fd730a3cc 100644 --- a/api/controllers/console/explore/saved_message.py +++ b/api/controllers/console/explore/saved_message.py @@ -38,11 +38,7 @@ class SavedMessageListApi(InstalledAppResource): args = SavedMessageListQuery.model_validate(request.args.to_dict()) pagination = SavedMessageService.pagination_by_last_id( - db.session(), - app_model, - current_user, - str(args.last_id) if args.last_id else None, - args.limit, + app_model, current_user, str(args.last_id) if args.last_id else None, args.limit, session=db.session() ) adapter = TypeAdapter(SavedMessageItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -65,7 +61,7 @@ class SavedMessageListApi(InstalledAppResource): payload = SavedMessageCreatePayload.model_validate(console_ns.payload or {}) try: - SavedMessageService.save(db.session(), app_model, current_user, str(payload.message_id)) + SavedMessageService.save(app_model, current_user, str(payload.message_id), session=db.session()) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -88,6 +84,6 @@ class SavedMessageApi(InstalledAppResource): if app_model.mode != "completion": raise NotCompletionAppError() - SavedMessageService.delete(db.session(), app_model, current_user, message_id_str) + SavedMessageService.delete(app_model, current_user, message_id_str, session=db.session()) return "", 204 diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index b28116c9a2d..d01eb9c38b1 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -431,7 +431,7 @@ class TrialAppWorkflowRunApi(TrialAppResource): invoke_from=InvokeFrom.EXPLORE, streaming=True, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except ProviderTokenNotInitError as ex: @@ -511,7 +511,7 @@ class TrialChatApi(TrialAppResource): invoke_from=InvokeFrom.EXPLORE, streaming=True, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except services.errors.conversation.ConversationNotExistsError: @@ -551,7 +551,11 @@ class TrialMessageSuggestedQuestionApi(TrialAppResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=current_user, message_id=message_id, invoke_from=InvokeFrom.EXPLORE + app_model=app_model, + user=current_user, + message_id=message_id, + invoke_from=InvokeFrom.EXPLORE, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message not found") @@ -589,7 +593,7 @@ class TrialChatAudioApi(TrialAppResource): user_id = current_user.id response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=None) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=db.session()) return response except services.errors.app_model_config.AppModelConfigBrokenError: logger.exception("App model config broken.") @@ -645,12 +649,12 @@ class TrialChatTextApi(TrialAppResource): response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, message_ref=message_ref, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=db.session()) return response except services.errors.app_model_config.AppModelConfigBrokenError: logger.exception("App model config broken.") @@ -709,7 +713,7 @@ class TrialCompletionApi(TrialAppResource): streaming=streaming, ) - RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) + RecommendedAppService.add_trial_app_record(app_id, user_id, session=session) # response-contract:ignore compact_generate_response return helper.compact_generate_response(response) except services.errors.conversation.ConversationNotExistsError: diff --git a/api/controllers/console/extension.py b/api/controllers/console/extension.py index 4b149b9c08d..cc06204a905 100644 --- a/api/controllers/console/extension.py +++ b/api/controllers/console/extension.py @@ -112,7 +112,7 @@ class APIBasedExtensionAPI(Resource): def get(self, current_tenant_id: str): return dump_response( APIBasedExtensionListResponse, - APIBasedExtensionService.get_all_by_tenant_id(db.session(), current_tenant_id), + APIBasedExtensionService.get_all_by_tenant_id(current_tenant_id, session=db.session()), ) @console_ns.doc("create_api_based_extension") @@ -133,7 +133,7 @@ class APIBasedExtensionAPI(Resource): api_key=payload.api_key, ) - extension = APIBasedExtensionService.save(db.session(), extension_data) + extension = APIBasedExtensionService.save(extension_data, session=db.session()) return APIBasedExtensionResponse( id=extension.id, name=extension.name, @@ -158,7 +158,9 @@ class APIBasedExtensionDetailAPI(Resource): return dump_response( APIBasedExtensionResponse, - APIBasedExtensionService.get_with_tenant_id(db.session(), current_tenant_id, api_based_extension_id), + APIBasedExtensionService.get_with_tenant_id( + current_tenant_id, api_based_extension_id, session=db.session() + ), ) @console_ns.doc("update_api_based_extension") @@ -174,7 +176,7 @@ class APIBasedExtensionDetailAPI(Resource): api_based_extension_id = str(id) extension_data_from_db = APIBasedExtensionService.get_with_tenant_id( - db.session(), current_tenant_id, api_based_extension_id + current_tenant_id, api_based_extension_id, session=db.session() ) payload = APIBasedExtensionPayload.model_validate(console_ns.payload or {}) @@ -187,7 +189,7 @@ class APIBasedExtensionDetailAPI(Resource): extension_data_from_db.api_key = payload.api_key api_key_for_response = payload.api_key - APIBasedExtensionService.save(db.session(), extension_data_from_db) + APIBasedExtensionService.save(extension_data_from_db, session=db.session()) return APIBasedExtensionResponse( id=extension_data_from_db.id, name=extension_data_from_db.name, @@ -208,9 +210,9 @@ class APIBasedExtensionDetailAPI(Resource): api_based_extension_id = str(id) extension_data_from_db = APIBasedExtensionService.get_with_tenant_id( - db.session(), current_tenant_id, api_based_extension_id + current_tenant_id, api_based_extension_id, session=db.session() ) - APIBasedExtensionService.delete(db.session(), extension_data_from_db) + APIBasedExtensionService.delete(extension_data_from_db, session=db.session()) return "", 204 diff --git a/api/controllers/console/init_validate.py b/api/controllers/console/init_validate.py index 27f6bcc36dc..f155f222e19 100644 --- a/api/controllers/console/init_validate.py +++ b/api/controllers/console/init_validate.py @@ -50,7 +50,7 @@ def get_init_status() -> InitStatusResponse: @only_edition_self_hosted def validate_init_password(payload: InitValidatePayload) -> InitValidateResponse: """Validate initialization password.""" - tenant_count = TenantService.get_tenant_count(session=db.session) + tenant_count = TenantService.get_tenant_count(session=db.session()) if tenant_count > 0: raise AlreadySetupError() diff --git a/api/controllers/console/setup.py b/api/controllers/console/setup.py index 2b99693a9ca..e0a0fba3329 100644 --- a/api/controllers/console/setup.py +++ b/api/controllers/console/setup.py @@ -79,7 +79,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: if get_setup_status(): raise AlreadySetupError() - tenant_count = TenantService.get_tenant_count(session=db.session) + tenant_count = TenantService.get_tenant_count(session=db.session()) if tenant_count > 0: raise AlreadySetupError() @@ -94,7 +94,7 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: password=payload.password, ip_address=extract_remote_ip(request), language=payload.language, - session=db.session, + session=db.session(), ) mark_setup_completed() diff --git a/api/controllers/console/socketio/workflow.py b/api/controllers/console/socketio/workflow.py index 99e56df3cb8..db5a4144dd3 100644 --- a/api/controllers/console/socketio/workflow.py +++ b/api/controllers/console/socketio/workflow.py @@ -44,7 +44,7 @@ def socket_connect(sid, environ, auth): return False with sio.app.app_context(): - user = AccountService.load_logged_in_account(account_id=user_id, session=db.session) + user = AccountService.load_logged_in_account(account_id=user_id, session=db.session()) if not user: logging.warning("Socket connect rejected: user not found (user_id=%s, sid=%s)", user_id, sid) return False @@ -69,7 +69,7 @@ def handle_user_connect(sid, data): if not workflow_id: return {"msg": "workflow_id is required"}, 400 - result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid) + result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid, session=db.session()) if not result: return {"msg": "unauthorized"}, 401 diff --git a/api/controllers/console/tag/tags.py b/api/controllers/console/tag/tags.py index c4ec925c9a3..86c1ad9c54c 100644 --- a/api/controllers/console/tag/tags.py +++ b/api/controllers/console/tag/tags.py @@ -137,7 +137,7 @@ class TagListApi(Resource): def get(self, current_tenant_id: str): raw_args = request.args.to_dict() param = TagListQueryParam.model_validate(raw_args) - tags = TagService.get_tags(db.session(), param.type, current_tenant_id, param.keyword) + tags = TagService.get_tags(param.type, current_tenant_id, param.keyword, session=db.session()) return dump_response(TagListResponse, tags), 200 @@ -154,7 +154,7 @@ class TagListApi(Resource): payload = TagBasePayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_if_needed(payload.type) - tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=payload.type), db.session) + tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=payload.type), db.session()) return dump_response(TagResponse, {"id": tag.id, "name": tag.name, "type": tag.type, "binding_count": 0}), 200 @@ -175,9 +175,9 @@ class TagUpdateDeleteApi(Resource): payload = TagUpdateRequestPayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_by_tag_id(tag_id_str) - tag = TagService.update_tags(UpdateTagPayload(name=payload.name), tag_id_str, db.session) + tag = TagService.update_tags(UpdateTagPayload(name=payload.name), tag_id_str, db.session()) - binding_count = TagService.get_tag_binding_count(tag_id_str, db.session) + binding_count = TagService.get_tag_binding_count(tag_id_str, db.session()) return ( dump_response( @@ -196,7 +196,7 @@ class TagUpdateDeleteApi(Resource): tag_id_str = str(tag_id) _enforce_snippet_tag_rbac_by_tag_id(tag_id_str) - TagService.delete_tag(tag_id_str, db.session) + TagService.delete_tag(tag_id_str, db.session()) return "", 204 @@ -223,7 +223,7 @@ def _create_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: target_id=payload.target_id, type=payload.type, ), - db.session, + db.session(), ) return {"result": "success"}, 200 @@ -239,7 +239,7 @@ def _remove_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: target_id=payload.target_id, type=payload.type, ), - db.session, + db.session(), ) return {"result": "success"}, 200 diff --git a/api/controllers/console/workspace/account.py b/api/controllers/console/workspace/account.py index 2f4ef1c5b42..6a06ed2d3e6 100644 --- a/api/controllers/console/workspace/account.py +++ b/api/controllers/console/workspace/account.py @@ -317,7 +317,7 @@ class AccountNameApi(Resource): def post(self, current_user: Account): payload = console_ns.payload or {} args = AccountNamePayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, name=args.name) + updated_account = AccountService.update_account(current_user, session=db.session(), name=args.name) return dump_response(AccountResponse, updated_account) @@ -363,7 +363,7 @@ class AccountAvatarApi(Resource): payload = console_ns.payload or {} args = AccountAvatarPayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, avatar=args.avatar) + updated_account = AccountService.update_account(current_user, session=db.session(), avatar=args.avatar) return dump_response(AccountResponse, updated_account) @@ -381,7 +381,7 @@ class AccountInterfaceLanguageApi(Resource): args = AccountInterfaceLanguagePayload.model_validate(payload) updated_account = AccountService.update_account( - current_user, session=db.session, interface_language=args.interface_language + current_user, session=db.session(), interface_language=args.interface_language ) return dump_response(AccountResponse, updated_account) @@ -400,7 +400,7 @@ class AccountInterfaceThemeApi(Resource): args = AccountInterfaceThemePayload.model_validate(payload) updated_account = AccountService.update_account( - current_user, session=db.session, interface_theme=args.interface_theme + current_user, session=db.session(), interface_theme=args.interface_theme ) return dump_response(AccountResponse, updated_account) @@ -418,7 +418,7 @@ class AccountTimezoneApi(Resource): payload = console_ns.payload or {} args = AccountTimezonePayload.model_validate(payload) - updated_account = AccountService.update_account(current_user, session=db.session, timezone=args.timezone) + updated_account = AccountService.update_account(current_user, session=db.session(), timezone=args.timezone) return dump_response(AccountResponse, updated_account) @@ -437,7 +437,7 @@ class AccountPasswordApi(Resource): try: assert args.password is not None - AccountService.update_account_password(current_user, args.password, args.new_password, session=db.session) + AccountService.update_account_password(current_user, args.password, args.new_password, session=db.session()) except ServiceCurrentPasswordIncorrectError: raise CurrentPasswordIncorrectError() @@ -514,7 +514,7 @@ class AccountDeleteApi(Resource): if not AccountService.verify_account_deletion_code(args.token, args.code): raise InvalidAccountDeletionCodeError() - AccountService.delete_account(account) + AccountService.delete_account(account, session=db.session()) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -726,7 +726,7 @@ class ChangeEmailResetApi(Resource): if AccountService.is_account_in_freeze(normalized_new_email): raise AccountInFreezeError() - if not AccountService.check_email_unique(normalized_new_email, session=db.session): + if not AccountService.check_email_unique(normalized_new_email, session=db.session()): raise EmailAlreadyInUseError() reset_data = AccountService.get_change_email_data(args.token) @@ -751,7 +751,7 @@ class ChangeEmailResetApi(Resource): AccountService.revoke_change_email_token(args.token) updated_account = AccountService.update_account_email( - current_user, email=normalized_new_email, session=db.session + current_user, email=normalized_new_email, session=db.session() ) AccountService.send_change_email_completed_notify_email( @@ -772,6 +772,6 @@ class CheckEmailUnique(Resource): normalized_email = args.email.lower() if AccountService.is_account_in_freeze(normalized_email): raise AccountInFreezeError() - if not AccountService.check_email_unique(normalized_email, session=db.session): + if not AccountService.check_email_unique(normalized_email, session=db.session()): raise EmailAlreadyInUseError() return SimpleResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/workspace/load_balancing_config.py b/api/controllers/console/workspace/load_balancing_config.py index 5983a4e10be..abeb691be03 100644 --- a/api/controllers/console/workspace/load_balancing_config.py +++ b/api/controllers/console/workspace/load_balancing_config.py @@ -10,6 +10,7 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.errors.validate import CredentialsValidateFailedError @@ -69,6 +70,7 @@ class LoadBalancingCredentialsValidateApi(Resource): model=payload.model, model_type=payload.model_type, credentials=payload.credentials, + session=db.session(), ) except CredentialsValidateFailedError as ex: result = False @@ -118,6 +120,7 @@ class LoadBalancingConfigCredentialsValidateApi(Resource): model=payload.model, model_type=payload.model_type, credentials=payload.credentials, + session=db.session(), config_id=config_id, ) except CredentialsValidateFailedError as ex: diff --git a/api/controllers/console/workspace/members.py b/api/controllers/console/workspace/members.py index 7e44d511bcf..72330aba5f0 100644 --- a/api/controllers/console/workspace/members.py +++ b/api/controllers/console/workspace/members.py @@ -135,7 +135,7 @@ def _normalize_enum_value(value: object) -> str: def _count_new_member_invites(tenant_id: str, emails: list[str]) -> int: new_member_count = 0 for email in emails: - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if not account: new_member_count += 1 continue @@ -190,7 +190,7 @@ class MemberListApi(Resource): current_user, _ = current_account_with_tenant() if not current_user.current_tenant: raise ValueError("No current tenant") - members = TenantService.get_tenant_members(current_user.current_tenant, session=db.session) + members = TenantService.get_tenant_members(current_user.current_tenant, session=db.session()) if dify_config.RBAC_ENABLED: member_ids = [member.id for member in members] member_roles = enterprise_rbac_service.RBACService.MemberRoles.batch_get( @@ -275,7 +275,7 @@ class MemberInviteEmailApi(Resource): language=interface_language, role=invitee_role, inviter=inviter, - session=db.session, + session=db.session(), ) encoded_invitee_email = parse.quote(invitee_email) invitation_results.append( @@ -323,7 +323,7 @@ class MemberCancelInviteApi(Resource): else: try: TenantService.remove_member_from_tenant( - current_user.current_tenant, member, current_user, session=db.session + current_user.current_tenant, member, current_user, session=db.session() ) except services.errors.account.CannotOperateSelfError as e: return {"code": "cannot-operate-self", "message": str(e)}, HTTPStatus.BAD_REQUEST @@ -368,7 +368,7 @@ class MemberUpdateRoleApi(Resource): try: assert member is not None, "Member not found" TenantService.update_member_role( - current_user.current_tenant, member, new_role, current_user, session=db.session + current_user.current_tenant, member, new_role, current_user, session=db.session() ) except services.errors.account.CannotOperateSelfError as e: return {"code": "cannot-operate-self", "message": str(e)}, HTTPStatus.BAD_REQUEST @@ -396,7 +396,7 @@ class DatasetOperatorMemberListApi(Resource): def get(self, current_user: Account): if not current_user.current_tenant: raise ValueError("No current tenant") - members = TenantService.get_dataset_operator_members(current_user.current_tenant, session=db.session) + members = TenantService.get_dataset_operator_members(current_user.current_tenant, session=db.session()) return dump_response(AccountWithRoleListResponse, {"accounts": members}), HTTPStatus.OK @@ -420,7 +420,7 @@ class SendOwnerTransferEmailApi(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() if args.language is not None and args.language == "zh-Hans": @@ -455,7 +455,7 @@ class OwnerTransferCheckApi(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() user_email = current_user.email @@ -501,7 +501,7 @@ class OwnerTransfer(Resource): # check if the current user is the owner of the workspace if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session): + if not TenantService.is_owner(current_user, current_user.current_tenant, session=db.session()): raise NotOwnerError() if current_user.id == str(member_id): @@ -522,13 +522,13 @@ class OwnerTransfer(Resource): if not current_user.current_tenant: raise ValueError("No current tenant") - if not TenantService.is_member(member, current_user.current_tenant, session=db.session): + if not TenantService.is_member(member, current_user.current_tenant, session=db.session()): raise MemberNotInTenantError() try: assert member is not None, "Member not found" TenantService.update_member_role( - current_user.current_tenant, member, "owner", current_user, session=db.session + current_user.current_tenant, member, "owner", current_user, session=db.session() ) AccountService.send_new_owner_transfer_notify_email( diff --git a/api/controllers/console/workspace/model_providers.py b/api/controllers/console/workspace/model_providers.py index 779399f9055..3bafa2ab6a6 100644 --- a/api/controllers/console/workspace/model_providers.py +++ b/api/controllers/console/workspace/model_providers.py @@ -353,7 +353,7 @@ class ModelProviderPaymentCheckoutUrlApi(Resource): def get(self, current_tenant_id: str, current_user: Account, provider: str): if provider != "anthropic": raise ValueError(f"provider name {provider} is invalid") - BillingService.is_tenant_owner_or_admin(db.session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) data = BillingService.get_model_provider_payment_link( provider_name=provider, tenant_id=current_tenant_id, diff --git a/api/controllers/console/workspace/models.py b/api/controllers/console/workspace/models.py index 1da72ef4362..0f735a7479e 100644 --- a/api/controllers/console/workspace/models.py +++ b/api/controllers/console/workspace/models.py @@ -24,6 +24,7 @@ from controllers.console.wraps import ( with_current_user, ) from core.entities.provider_entities import CredentialConfiguration +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.entities.model_entities import ModelType, ParameterRule from graphon.model_runtime.errors.validate import CredentialsValidateFailedError @@ -297,6 +298,7 @@ class ModelProviderModelApi(Resource): model_type=args.model_type, configs=args.load_balancing.configs, config_from=args.config_from or "", + session=db.session(), ) if args.load_balancing.enabled: @@ -356,6 +358,7 @@ class ModelProviderModelCredentialApi(Resource): provider=provider, model=args.model, model_type=args.model_type, + session=db.session(), config_from=args.config_from or "", ) diff --git a/api/controllers/console/workspace/plugin.py b/api/controllers/console/workspace/plugin.py index 682aa5b6190..c7644af2b48 100644 --- a/api/controllers/console/workspace/plugin.py +++ b/api/controllers/console/workspace/plugin.py @@ -38,6 +38,7 @@ from core.tools.builtin_tool.providers._positions import BuiltinToolProviderSort from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolProviderType from core.tools.tool_manager import ToolManager +from extensions.ext_database import db from fields.base import ResponseModel from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.helper import dump_response @@ -973,7 +974,7 @@ class PluginChangePermissionApi(Resource): args = ParserPermissionChange.model_validate(console_ns.payload) set_permission_result = PluginPermissionService.change_permission( - tenant_id, args.install_permission, args.debug_permission + tenant_id, args.install_permission, args.debug_permission, session=db.session() ) if not set_permission_result: return jsonable_encoder({"success": False, "message": "Failed to set permission"}) @@ -989,7 +990,7 @@ class PluginFetchPermissionApi(Resource): @account_initialization_required @with_current_tenant_id def get(self, tenant_id: str): - permission = PluginPermissionService.get_permission(tenant_id) + permission = PluginPermissionService.get_permission(tenant_id, session=db.session()) if not permission: return jsonable_encoder( { @@ -1094,6 +1095,7 @@ class PluginChangeAutoUpgradeApi(Resource): auto_upgrade.exclude_plugins, auto_upgrade.include_plugins, category=args.category, + session=db.session(), ) if not set_auto_upgrade_strategy_result: return jsonable_encoder({"success": False, "message": "Failed to set auto upgrade strategy"}) @@ -1111,7 +1113,7 @@ class PluginFetchAutoUpgradeApi(Resource): @with_current_tenant_id def get(self, tenant_id: str): args = ParserAutoUpgradeFetch.model_validate(request.args.to_dict(flat=True)) - auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category) + auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category, session=db.session()) auto_upgrade_dict = ( _auto_upgrade_settings_to_dict(auto_upgrade) if auto_upgrade @@ -1140,7 +1142,11 @@ class PluginAutoUpgradeExcludePluginApi(Resource): args = ParserExcludePlugin.model_validate(console_ns.payload) return jsonable_encoder( - {"success": PluginAutoUpgradeService.exclude_plugin(tenant_id, args.plugin_id, args.category)} + { + "success": PluginAutoUpgradeService.exclude_plugin( + tenant_id, args.plugin_id, args.category, session=db.session() + ) + } ) diff --git a/api/controllers/console/workspace/rbac.py b/api/controllers/console/workspace/rbac.py index a155bb0cf0b..39ad12080f3 100644 --- a/api/controllers/console/workspace/rbac.py +++ b/api/controllers/console/workspace/rbac.py @@ -14,6 +14,7 @@ from controllers.console import console_ns from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required from core.db.session_factory import session_factory from core.rbac import RBACResourceWhitelistScope +from extensions.ext_database import db from libs.login import current_account_with_tenant, login_required from models import Account from services.enterprise import rbac_service as svc @@ -564,6 +565,7 @@ class RBACMyPermissionsApi(Resource): account_id, app_id=request.args.get("app_id") or None, dataset_id=request.args.get("dataset_id") or None, + session=db.session(), ) ) @@ -902,7 +904,7 @@ class RBACMemberRolesApi(Resource): @console_ns.response(200, "Success", console_ns.models[svc.MemberRolesResponse.__name__]) def get(self, member_id): tenant_id, account_id = _current_ids() - return _dump(svc.RBACService.MemberRoles.get(tenant_id, account_id, str(member_id))) + return _dump(svc.RBACService.MemberRoles.get(tenant_id, account_id, str(member_id), session=db.session())) @login_required @console_ns.expect(console_ns.models[_ReplaceMemberRolesRequest.__name__]) @@ -916,6 +918,7 @@ class RBACMemberRolesApi(Resource): account_id, str(member_id), role_ids=list(request.role_ids), + session=db.session(), ) ) diff --git a/api/controllers/console/workspace/snippets.py b/api/controllers/console/workspace/snippets.py index c849336401c..18407b62d52 100644 --- a/api/controllers/console/workspace/snippets.py +++ b/api/controllers/console/workspace/snippets.py @@ -188,7 +188,7 @@ class CustomizedSnippetsApi(Resource): snippet_service = _snippet_service() snippets, total, has_more = snippet_service.get_snippets( tenant_id=current_tenant_id, - session=db.session, + session=db.session(), page=query.page, limit=query.limit, keyword=query.keyword, diff --git a/api/controllers/console/workspace/tool_providers.py b/api/controllers/console/workspace/tool_providers.py index 7a3f158b0c3..30eeec2fcc5 100644 --- a/api/controllers/console/workspace/tool_providers.py +++ b/api/controllers/console/workspace/tool_providers.py @@ -459,6 +459,7 @@ class ToolBuiltinProviderGetCredentialsApi(Resource): BuiltinToolManageService.get_builtin_tool_provider_credentials( tenant_id=tenant_id, provider_name=provider, + session=db.session(), user=user, include_credential_ids=query.include_credential_ids or None, ) @@ -1064,6 +1065,7 @@ class ToolBuiltinProviderGetCredentialInfoApi(Resource): BuiltinToolManageService.get_builtin_tool_provider_credential_info( tenant_id=tenant_id, provider=provider, + session=db.session(), user=user, include_credential_ids=query.include_credential_ids or None, ) diff --git a/api/controllers/console/workspace/workspace.py b/api/controllers/console/workspace/workspace.py index 0630281de75..23ce116b349 100644 --- a/api/controllers/console/workspace/workspace.py +++ b/api/controllers/console/workspace/workspace.py @@ -223,7 +223,7 @@ class TenantListApi(Resource): def get(self, current_tenant_id: str, current_user: Account): tenant_rows: list[tuple[Tenant, TenantAccountJoin]] = [ (tenant, membership) - for tenant, membership in TenantService.get_workspaces_for_account(db.session, current_user.id) + for tenant, membership in TenantService.get_workspaces_for_account(current_user.id, session=db.session()) if tenant.status == TenantStatus.NORMAL ] tenants = [tenant for tenant, _ in tenant_rows] @@ -306,16 +306,19 @@ class TenantApi(Resource): raise ValueError("No current tenant") if tenant.status == TenantStatus.ARCHIVE: - tenants = TenantService.get_join_tenants(current_user, session=db.session) + tenants = TenantService.get_join_tenants(current_user, session=db.session()) # if there is any tenant, switch to the first one if len(tenants) > 0: - TenantService.switch_tenant(current_user, tenants[0].id, session=db.session) + TenantService.switch_tenant(current_user, tenants[0].id, session=db.session()) tenant = tenants[0] # else, raise Unauthorized else: raise Unauthorized("workspace is archived") - return dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant)), HTTPStatus.OK + return ( + dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant, session=db.session())), + HTTPStatus.OK, + ) @console_ns.route("/workspaces/switch") @@ -332,7 +335,7 @@ class SwitchWorkspaceApi(Resource): # Check whether the tenant_id belongs to the current account. try: - TenantService.switch_tenant(current_user, args.tenant_id, session=db.session) + TenantService.switch_tenant(current_user, args.tenant_id, session=db.session()) except Exception: raise AccountNotLinkTenantError("Account not link tenant") @@ -341,7 +344,7 @@ class SwitchWorkspaceApi(Resource): raise ValueError("Tenant not found") return SwitchWorkspaceResponse( - result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant) + result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant, session=db.session()) ).model_dump(mode="json") @@ -372,7 +375,7 @@ class CustomConfigWorkspaceApi(Resource): db.session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) ).model_dump(mode="json") @@ -438,7 +441,7 @@ class WorkspaceInfoApi(Resource): db.session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) ).model_dump(mode="json") diff --git a/api/controllers/files/agent_drive_archive.py b/api/controllers/files/agent_drive_archive.py index afa6ac79483..8ecec2e9a4c 100644 --- a/api/controllers/files/agent_drive_archive.py +++ b/api/controllers/files/agent_drive_archive.py @@ -8,6 +8,7 @@ from werkzeug.exceptions import Forbidden, NotFound from controllers.common.file_response import enforce_download_for_html from controllers.common.schema import register_schema_models from controllers.files import files_ns +from extensions.ext_database import db from models.agent import AgentDriveFileKind from services.agent_drive_service import AgentDriveError, AgentDriveService @@ -54,6 +55,7 @@ class AgentDriveArchiveMemberApi(Resource): archive_file_kind=args.archive_file_kind, archive_file_id=args.archive_file_id, member_path=args.member_path, + session=db.session(), ) except AgentDriveError as exc: raise NotFound(exc.message) from exc diff --git a/api/controllers/inner_api/app/dsl.py b/api/controllers/inner_api/app/dsl.py index 915a11dcddc..9fd111f86dc 100644 --- a/api/controllers/inner_api/app/dsl.py +++ b/api/controllers/inner_api/app/dsl.py @@ -98,6 +98,7 @@ class EnterpriseAppDSLExport(Resource): data = AppDslService.export_dsl( app_model=app_model, + session=db.session(), include_secret=include_secret, ) diff --git a/api/controllers/inner_api/plugin/agent_drive.py b/api/controllers/inner_api/plugin/agent_drive.py index 0cdb9dab35f..e06720a8e99 100644 --- a/api/controllers/inner_api/plugin/agent_drive.py +++ b/api/controllers/inner_api/plugin/agent_drive.py @@ -17,6 +17,7 @@ from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns from controllers.inner_api.plugin.wraps import get_user from controllers.inner_api.wraps import plugin_inner_api_only +from extensions.ext_database import db from services.agent_drive_service import ( AgentDriveError, AgentDriveService, @@ -53,6 +54,7 @@ class AgentDriveManifestApi(Resource): agent_id=agent_id, prefix=request.args.get("prefix", ""), include_download_url=include_download_url, + session=db.session(), ) except AgentDriveError as exc: return _error_response(exc) @@ -71,7 +73,7 @@ class AgentDriveSkillsApi(Resource): tenant_id = (request.args.get("tenant_id") or "").strip() if not tenant_id: raise AgentDriveError("missing_tenant_id", "tenant_id is required", status_code=400) - items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=agent_id) + items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=agent_id, session=db.session()) except AgentDriveError as exc: return _error_response(exc) return {"items": items} @@ -96,6 +98,7 @@ class AgentDriveCommitApi(Resource): user_id=user.id, agent_id=agent_id, items=body.items, + session=db.session(), ) except AgentDriveError as exc: return _error_response(exc) diff --git a/api/controllers/inner_api/workspace/workspace.py b/api/controllers/inner_api/workspace/workspace.py index 1f25eb576d3..b3a571112f6 100644 --- a/api/controllers/inner_api/workspace/workspace.py +++ b/api/controllers/inner_api/workspace/workspace.py @@ -47,8 +47,8 @@ class EnterpriseWorkspace(Resource): if account is None: return {"message": "owner account not found."}, 404 - tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session) - TenantService.create_tenant_member(tenant, account, db.session, role="owner") + tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session()) + TenantService.create_tenant_member(tenant, account, db.session(), role="owner") tenant_was_created.send(tenant) @@ -84,7 +84,7 @@ class EnterpriseWorkspaceNoOwnerEmail(Resource): def post(self): args = WorkspaceOwnerlessPayload.model_validate(inner_api_ns.payload or {}) - tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session) + tenant = TenantService.create_tenant(args.name, is_from_dashboard=True, session=db.session()) tenant_was_created.send(tenant) diff --git a/api/controllers/openapi/account.py b/api/controllers/openapi/account.py index 8ad0b02f4a0..b4786f2ae25 100644 --- a/api/controllers/openapi/account.py +++ b/api/controllers/openapi/account.py @@ -45,8 +45,10 @@ class AccountApi(Resource): enforce(LIMIT_ME_PER_ACCOUNT, key=f"account:{auth_data.account_id}") account_id_str = str(auth_data.account_id) if auth_data.account_id else None - account = AccountService.get_account_by_id(db.session, account_id_str) if account_id_str else None - memberships = TenantService.get_account_memberships(db.session, account_id_str) if account_id_str else [] + account = AccountService.get_account_by_id(account_id_str, session=db.session()) if account_id_str else None + memberships = ( + TenantService.get_account_memberships(account_id_str, session=db.session()) if account_id_str else [] + ) default_ws_id = _pick_default_workspace(memberships) return AccountResponse( @@ -63,7 +65,7 @@ class AccountSessionsSelfApi(Resource): @auth_router.guard(scope=Scope.FULL, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, RevokeResponse, description="Session revoked") def delete(self, *, auth_data: AuthData): - revoke_oauth_token(db.session, redis_client, str(auth_data.token_id)) + revoke_oauth_token(redis_client, str(auth_data.token_id), session=db.session()) return RevokeResponse(status="revoked") @@ -81,7 +83,7 @@ class AccountSessionsApi(Resource): page = query.page limit = query.limit - all_rows = list_active_sessions(db.session, ctx, now) + all_rows = list_active_sessions(ctx, now, session=db.session()) total = len(all_rows) sliced = all_rows[(page - 1) * limit : page * limit] @@ -117,10 +119,10 @@ class AccountSessionByIdApi(Resource): # 404 (not 403) on cross-subject so the endpoint doesn't leak # token IDs that belong to other subjects. - if not token_belongs_to_subject(db.session, session_id, ctx): + if not token_belongs_to_subject(session_id, ctx, session=db.session()): raise NotFound("session not found") - revoke_oauth_token(db.session, redis_client, session_id) + revoke_oauth_token(redis_client, session_id, session=db.session()) return RevokeResponse(status="revoked") diff --git a/api/controllers/openapi/app_dsl.py b/api/controllers/openapi/app_dsl.py index cea7127bd07..d06845dada4 100644 --- a/api/controllers/openapi/app_dsl.py +++ b/api/controllers/openapi/app_dsl.py @@ -145,6 +145,7 @@ class AppDslExportApi(Resource): try: data = AppDslService.export_dsl( app_model=app, + session=db.session(), include_secret=query.include_secret, workflow_id=query.workflow_id, ) diff --git a/api/controllers/openapi/apps.py b/api/controllers/openapi/apps.py index d4cb175ba5e..882b55b7041 100644 --- a/api/controllers/openapi/apps.py +++ b/api/controllers/openapi/apps.py @@ -66,13 +66,13 @@ class AppReadResource(Resource): if is_uuid: # ``str(parsed_uuid)`` normalises to the canonical dashed form. - app = AppService.get_visible_app_by_id(db.session, str(parsed_uuid)) + app = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) if app is None: raise NotFound("app not found") else: if not workspace_id: raise UnprocessableEntity("workspace_id is required for name-based lookup") - matches = AppService.find_visible_apps_by_name(db.session, name=app_id, tenant_id=workspace_id) + matches = AppService.find_visible_apps_by_name(name=app_id, tenant_id=workspace_id, session=db.session()) if len(matches) == 0: raise NotFound("app not found") if len(matches) > 1: @@ -177,7 +177,7 @@ class AppListApi(Resource): tenant_name: str | None = None if parsed_uuid is not None: - app: App | None = AppService.get_visible_app_by_id(db.session, str(parsed_uuid)) + app: App | None = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) if app is None or str(app.tenant_id) != workspace_id: return empty if not _is_listable(app): @@ -188,7 +188,7 @@ class AppListApi(Resource): str(app.id), str(app.maintainer) if app.maintainer else None, str(auth_data.account_id) ): return empty - tenant_name = TenantService.get_tenant_name(db.session, workspace_id) + tenant_name = TenantService.get_tenant_name(workspace_id, session=db.session()) item = AppListRow( id=str(app.id), name=app.name, @@ -215,13 +215,13 @@ class AppListApi(Resource): if apply_rbac_filter: access_filter.apply_to_params(params) - pagination = AppService().get_paginate_apps(str(auth_data.account_id), workspace_id, params, db.session) + pagination = AppService().get_paginate_apps(str(auth_data.account_id), workspace_id, params, db.session()) if pagination is None: return empty tenant_name = None if pagination.items: - tenant_name = TenantService.get_tenant_name(db.session, workspace_id) + tenant_name = TenantService.get_tenant_name(workspace_id, session=db.session()) items = [ AppListRow( diff --git a/api/controllers/openapi/apps_permitted_external.py b/api/controllers/openapi/apps_permitted_external.py index 5c6fdce5141..353a1ec1cb3 100644 --- a/api/controllers/openapi/apps_permitted_external.py +++ b/api/controllers/openapi/apps_permitted_external.py @@ -55,10 +55,10 @@ class PermittedExternalAppsListApi(Resource): return env apps_by_id: dict[str, App] = { - str(a.id): a for a in AppService.find_visible_apps_by_ids(db.session, page_result.app_ids) + str(a.id): a for a in AppService.find_visible_apps_by_ids(page_result.app_ids, session=db.session()) } tenant_ids = list({str(a.tenant_id) for a in apps_by_id.values()}) - tenants_by_id = {str(t.id): t for t in TenantService.get_tenants_by_ids(db.session, tenant_ids)} + tenants_by_id = {str(t.id): t for t in TenantService.get_tenants_by_ids(tenant_ids, session=db.session())} items: list[AppListRow] = [] for app_id in page_result.app_ids: diff --git a/api/controllers/openapi/auth/prepare.py b/api/controllers/openapi/auth/prepare.py index 6704b27decc..96cf9a8858f 100644 --- a/api/controllers/openapi/auth/prepare.py +++ b/api/controllers/openapi/auth/prepare.py @@ -23,7 +23,7 @@ def load_app(data: AuthData) -> None: uuid.UUID(app_id) except ValueError: raise NotFound("app not found") - app = AppService.get_app_by_id(db.session, app_id) + app = AppService.get_app_by_id(app_id, session=db.session()) if not app or app.status != AppStatus.NORMAL: raise NotFound("app not found") data.app = app @@ -34,7 +34,7 @@ def load_tenant(data: AuthData) -> None: return if data.app is None: raise InternalServerError("pipeline_invariant_violated: app not loaded before load_tenant") - tenant = TenantService.get_tenant_by_id(db.session, str(data.app.tenant_id)) + tenant = TenantService.get_tenant_by_id(str(data.app.tenant_id), session=db.session()) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise Forbidden("workspace unavailable") data.tenant = tenant @@ -50,7 +50,7 @@ def load_tenant_from_request(data: AuthData) -> None: uuid.UUID(workspace_id) except ValueError: raise NotFound("workspace not found") - tenant = TenantService.get_tenant_by_id(db.session, workspace_id) + tenant = TenantService.get_tenant_by_id(workspace_id, session=db.session()) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise NotFound("workspace not found") data.tenant = tenant @@ -59,7 +59,7 @@ def load_tenant_from_request(data: AuthData) -> None: def load_account(data: AuthData) -> None: if data.caller is not None: return - account = AccountService.get_account_by_id(db.session, str(data.account_id)) + account = AccountService.get_account_by_id(str(data.account_id), session=db.session()) if account is None: raise Unauthorized("account not found") if data.tenant: @@ -75,7 +75,7 @@ def load_workspace_role(data: AuthData) -> None: return if data.caller is not None and getattr(data.caller, "status", None) != AccountStatus.ACTIVE: return - role = TenantService.get_account_role_in_tenant(db.session, str(data.account_id), str(data.tenant.id)) + role = TenantService.get_account_role_in_tenant(str(data.account_id), str(data.tenant.id), session=db.session()) if role is None: return data.tenant_role = role diff --git a/api/controllers/openapi/auth/verify.py b/api/controllers/openapi/auth/verify.py index b5f10f66b34..b6ef95e3ea3 100644 --- a/api/controllers/openapi/auth/verify.py +++ b/api/controllers/openapi/auth/verify.py @@ -82,7 +82,7 @@ def check_app_api_enabled(data: AuthData) -> None: def check_app_access(data: AuthData) -> None: if data.tenant is None: return - if not TenantService.account_belongs_to_tenant(db.session, data.account_id, data.tenant.id): + if not TenantService.account_belongs_to_tenant(data.account_id, data.tenant.id, session=db.session()): raise Forbidden("subject_no_app_access") @@ -127,5 +127,5 @@ def _resolve_user_id(data: AuthData) -> str | None: return str(data.account_id) if data.account_id is not None else None if data.external_identity is None: return None - account = AccountService.get_account_by_email(db.session, data.external_identity.email) + account = AccountService.get_account_by_email(data.external_identity.email, session=db.session()) return str(account.id) if account is not None else None diff --git a/api/controllers/openapi/oauth_device.py b/api/controllers/openapi/oauth_device.py index cee187daaf3..3ba5f2ee207 100644 --- a/api/controllers/openapi/oauth_device.py +++ b/api/controllers/openapi/oauth_device.py @@ -247,7 +247,6 @@ class DeviceApproveApi(Resource): raise BadRequest(description=str(e)) from None ttl_days = oauth_ttl_days(tenant_id=tenant) mint = mint_oauth_token( - db.session, redis_client, subject_email=account.email, subject_issuer=ACCOUNT_ISSUER_SENTINEL, @@ -256,6 +255,7 @@ class DeviceApproveApi(Resource): device_label=state.device_label, prefix=profile.prefix, ttl_days=ttl_days, + session=db.session(), ) poll_payload = _build_account_poll_payload(account, tenant, mint) @@ -342,7 +342,7 @@ def _audit_cross_ip_if_needed(state) -> None: def _build_account_poll_payload(account, tenant, mint) -> PollPayload: - rows = TenantService.get_workspaces_for_account(db.session, str(account.id)) + rows = TenantService.get_workspaces_for_account(str(account.id), session=db.session()) workspaces = [WorkspacePayload(id=str(t.id), name=t.name, role=getattr(m, "role", "")) for t, m in rows] # Prefer active session tenant → DB-flagged current join → first membership. default_ws_id = None diff --git a/api/controllers/openapi/oauth_device_sso.py b/api/controllers/openapi/oauth_device_sso.py index 79538f48059..fbf7bfa6295 100644 --- a/api/controllers/openapi/oauth_device_sso.py +++ b/api/controllers/openapi/oauth_device_sso.py @@ -194,7 +194,7 @@ def _sso_complete_impl(): if state.status is not DeviceFlowStatus.PENDING: return _device_error_redirect("sso_failed", user_code) - if AccountService.has_active_account_with_email(db.session, claims.email): + if AccountService.has_active_account_with_email(claims.email, session=db.session()): _emit_external_rejection_audit( state, _RejectedClaims(subject_email=claims.email, subject_issuer=claims.issuer), @@ -274,7 +274,7 @@ def approve_external(): if state.status is not DeviceFlowStatus.PENDING: raise Conflict("user_code_not_pending") - if AccountService.has_active_account_with_email(db.session, claims.subject_email): + if AccountService.has_active_account_with_email(claims.subject_email, session=db.session()): _emit_external_rejection_audit(state, claims, reason="email_belongs_to_dify_account") raise Forbidden("email_belongs_to_dify_account") @@ -293,7 +293,6 @@ def approve_external(): ttl_days = oauth_ttl_days(tenant_id=None) mint = mint_oauth_token( - db.session, redis_client, subject_email=claims.subject_email, subject_issuer=claims.subject_issuer, @@ -302,6 +301,7 @@ def approve_external(): device_label=state.device_label, prefix=profile.prefix, ttl_days=ttl_days, + session=db.session(), ) # SSO branch of the shared PollPayload contract: account/workspace diff --git a/api/controllers/openapi/workspaces.py b/api/controllers/openapi/workspaces.py index c45c02e54d3..7f8eb0f7012 100644 --- a/api/controllers/openapi/workspaces.py +++ b/api/controllers/openapi/workspaces.py @@ -64,14 +64,14 @@ def _member_response(account: Account) -> MemberResponse: def _load_tenant(workspace_id: str) -> Tenant: - tenant = TenantService.get_tenant_by_id(db.session, workspace_id) + tenant = TenantService.get_tenant_by_id(workspace_id, session=db.session()) if tenant is None or tenant.status != TenantStatus.NORMAL: raise NotFound("workspace not found") return tenant def _load_account(account_id: object) -> Account: - account = AccountService.get_account_by_id(db.session, str(account_id)) if account_id else None + account = AccountService.get_account_by_id(str(account_id), session=db.session()) if account_id else None if account is None: raise RuntimeError("authenticated account_id has no Account row") return account @@ -94,7 +94,7 @@ class WorkspacesApi(Resource): @auth_router.guard(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, WorkspaceListResponse, description="Workspace list") def get(self, *, auth_data: AuthData): - rows = TenantService.get_workspaces_for_account(db.session, str(auth_data.account_id)) + rows = TenantService.get_workspaces_for_account(str(auth_data.account_id), session=db.session()) return WorkspaceListResponse(workspaces=list(starmap(_workspace_summary, rows))) @@ -104,7 +104,7 @@ class WorkspaceByIdApi(Resource): @auth_router.guard(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, WorkspaceDetailResponse, description="Workspace detail") def get(self, workspace_id: str, *, auth_data: AuthData): - row = TenantService.find_workspace_for_account(db.session, str(auth_data.account_id), workspace_id) + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=db.session()) # 404 (not 403) on non-member so workspace IDs don't leak across tenants. if row is None: raise NotFound("workspace not found") @@ -128,11 +128,11 @@ class WorkspaceSwitchApi(Resource): account = _load_account(auth_data.account_id) try: - TenantService.switch_tenant(account, workspace_id, session=db.session) + TenantService.switch_tenant(account, workspace_id, session=db.session()) except AccountNotLinkTenantError: raise NotFound("workspace not found") - row = TenantService.find_workspace_for_account(db.session, str(auth_data.account_id), workspace_id) + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=db.session()) if row is None: raise NotFound("workspace not found") tenant, membership = row @@ -152,7 +152,7 @@ class WorkspaceMembersApi(Resource): @accepts(query=MemberListQuery) def get(self, workspace_id: str, *, auth_data: AuthData, query: MemberListQuery): tenant = _load_tenant(workspace_id) - members = TenantService.get_tenant_members(tenant, session=db.session) + members = TenantService.get_tenant_members(tenant, session=db.session()) total = len(members) start = (query.page - 1) * query.limit page_items = members[start : start + query.limit] @@ -184,7 +184,7 @@ class WorkspaceMembersApi(Resource): language=None, role=body.role, inviter=inviter, - session=db.session, + session=db.session(), ) except AccountAlreadyInTenantError as exc: raise BadRequest(str(exc)) @@ -194,7 +194,7 @@ class WorkspaceMembersApi(Resource): raise BadRequest(str(exc)) normalized_email = body.email.lower() - member = AccountService.get_account_by_email_with_case_fallback(db.session, normalized_email) + member = AccountService.get_account_by_email_with_case_fallback(normalized_email, session=db.session()) if member is None: # invite_new_member just created or fetched this account. raise RuntimeError("invited member missing from DB after invite") @@ -229,12 +229,12 @@ class WorkspaceMemberApi(Resource): def delete(self, workspace_id: str, member_id: str, *, auth_data: AuthData): operator = _load_account(auth_data.account_id) tenant = _load_tenant(workspace_id) - member = AccountService.get_account_by_id(db.session, member_id) + member = AccountService.get_account_by_id(member_id, session=db.session()) if member is None: raise NotFound("member not found") try: - TenantService.remove_member_from_tenant(tenant, member, operator, session=db.session) + TenantService.remove_member_from_tenant(tenant, member, operator, session=db.session()) except CannotOperateSelfError as exc: raise BadRequest(str(exc)) except NoPermissionError as exc: @@ -254,12 +254,12 @@ class WorkspaceMemberApi(Resource): def patch(self, workspace_id: str, member_id: str, *, auth_data: AuthData, body: MemberRoleUpdatePayload): operator = _load_account(auth_data.account_id) tenant = _load_tenant(workspace_id) - member = AccountService.get_account_by_id(db.session, member_id) + member = AccountService.get_account_by_id(member_id, session=db.session()) if member is None: raise NotFound("member not found") try: - TenantService.update_member_role(tenant, member, body.role, operator, session=db.session) + TenantService.update_member_role(tenant, member, body.role, operator, session=db.session()) except CannotOperateSelfError as exc: raise BadRequest(str(exc)) except NoPermissionError as exc: diff --git a/api/controllers/service_api/app/annotation.py b/api/controllers/service_api/app/annotation.py index 0fbf8125ed9..126c67b5d61 100644 --- a/api/controllers/service_api/app/annotation.py +++ b/api/controllers/service_api/app/annotation.py @@ -201,7 +201,7 @@ class AnnotationListApi(Resource): query = AnnotationListQuery.model_validate(request.args.to_dict(flat=True)) annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app_model.id, query.page, query.limit, query.keyword + app_model.id, query.page, query.limit, query.keyword, session=db.session() ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) response = AnnotationList( @@ -243,7 +243,9 @@ class AnnotationListApi(Resource): """Create a new annotation.""" payload = AnnotationCreatePayload.model_validate(service_api_ns.payload or {}) insert_args: InsertAnnotationArgs = {"question": payload.question, "answer": payload.answer} - annotation = AppAnnotationService.insert_app_annotation_directly(insert_args, app_model.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + insert_args, app_model.id, session=db.session() + ) response = Annotation.model_validate(annotation, from_attributes=True) return response.model_dump(mode="json"), HTTPStatus.CREATED @@ -285,7 +287,7 @@ class AnnotationUpdateDeleteApi(Resource): update_args: UpdateAnnotationArgs = {"question": payload.question, "answer": payload.answer} app_ref = AppRefService.create_app_ref(app_model) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session) + annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, db.session()) response = Annotation.model_validate(annotation, from_attributes=True) return response.model_dump(mode="json") @@ -316,5 +318,5 @@ class AnnotationUpdateDeleteApi(Resource): """Delete an annotation.""" app_ref = AppRefService.create_app_ref(app_model) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) - AppAnnotationService.delete_app_annotation(annotation_ref, db.session) + AppAnnotationService.delete_app_annotation(annotation_ref, db.session()) return "", 204 diff --git a/api/controllers/service_api/app/app.py b/api/controllers/service_api/app/app.py index 3ac44b12c66..60f83d7d070 100644 --- a/api/controllers/service_api/app/app.py +++ b/api/controllers/service_api/app/app.py @@ -11,6 +11,7 @@ from controllers.service_api.app.error import AgentNotPublishedError, AppUnavail from controllers.service_api.wraps import validate_app_token from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError +from extensions.ext_database import db from fields.base import ResponseModel from models.model import App, AppMode from services.app_service import AppService @@ -122,7 +123,7 @@ class AppMetaApi(Resource): Returns metadata about the application including configuration and settings. """ - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) @service_api_ns.route("/info") diff --git a/api/controllers/service_api/app/audio.py b/api/controllers/service_api/app/audio.py index 53b31c8e6c4..68ab5f31ea5 100644 --- a/api/controllers/service_api/app/audio.py +++ b/api/controllers/service_api/app/audio.py @@ -188,7 +188,7 @@ class TextApi(Resource): ) response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, end_user=end_user.external_user_id, diff --git a/api/controllers/service_api/app/conversation.py b/api/controllers/service_api/app/conversation.py index 9b5533ea07a..a395dcb93fc 100644 --- a/api/controllers/service_api/app/conversation.py +++ b/api/controllers/service_api/app/conversation.py @@ -249,7 +249,7 @@ class ConversationDetailApi(Resource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, end_user) + ConversationService.delete(app_model, conversation_id, end_user, session=db.session()) except services.errors.conversation.ConversationNotExistsError: raise NotFound("Conversation Not Exists.") return "", 204 @@ -299,7 +299,7 @@ class ConversationRenameApi(Resource): try: conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -356,7 +356,13 @@ class ConversationVariablesApi(Resource): try: pagination = ConversationService.get_conversational_variable( - app_model, conversation_id, end_user, query_args.limit, last_id, query_args.variable_name + app_model, + conversation_id, + end_user, + query_args.limit, + last_id, + query_args.variable_name, + session=db.session(), ) return ConversationVariableInfiniteScrollPaginationResponse.model_validate( pagination, from_attributes=True @@ -417,7 +423,7 @@ class ConversationVariableDetailApi(Resource): try: variable = ConversationService.update_conversation_variable( - app_model, conversation_id, variable_id_str, end_user, payload.value + app_model, conversation_id, variable_id_str, end_user, payload.value, session=db.session() ) return ConversationVariableResponse.model_validate(variable, from_attributes=True).model_dump(mode="json") except services.errors.conversation.ConversationNotExistsError: diff --git a/api/controllers/service_api/app/message.py b/api/controllers/service_api/app/message.py index 18d1c5d3254..3acb2c74872 100644 --- a/api/controllers/service_api/app/message.py +++ b/api/controllers/service_api/app/message.py @@ -15,6 +15,7 @@ from controllers.service_api.app.error import NotChatAppError from controllers.service_api.schema import expect_with_user from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token from core.app.entities.app_invoke_entities import InvokeFrom +from extensions.ext_database import db from fields.base import ResponseModel from fields.conversation_fields import ResultResponse from fields.message_fields import MessageInfiniteScrollPagination, MessageListItem @@ -109,7 +110,7 @@ class MessageListApi(Resource): try: pagination = MessageService.pagination_by_first_id( - app_model, end_user, conversation_id, first_id, query_args.limit + app_model, end_user, conversation_id, first_id, query_args.limit, session=db.session() ) adapter = TypeAdapter(MessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -167,6 +168,7 @@ class MessageFeedbackApi(Resource): user=end_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -208,7 +210,9 @@ class AppGetFeedbacksApi(Resource): Returns paginated list of all feedback submitted for messages in this app. """ query_args = FeedbackListQuery.model_validate(request.args.to_dict()) - feedbacks = MessageService.get_all_messages_feedbacks(app_model, page=query_args.page, limit=query_args.limit) + feedbacks = MessageService.get_all_messages_feedbacks( + app_model, page=query_args.page, limit=query_args.limit, session=db.session() + ) return {"data": feedbacks} @@ -258,7 +262,11 @@ class MessageSuggestedApi(Resource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=end_user, message_id=message_id_str, invoke_from=InvokeFrom.SERVICE_API + app_model=app_model, + user=end_user, + message_id=message_id_str, + invoke_from=InvokeFrom.SERVICE_API, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index 56836f56895..66085ca0642 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -414,7 +414,7 @@ class DatasetListApi(DatasetApiResource): datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session, + db.session(), tenant_id, current_user, query.keyword, @@ -565,11 +565,11 @@ class DatasetApi(DatasetApiResource): ) def get(self, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) data = _dump_service_dataset_detail(dataset) @@ -601,7 +601,7 @@ class DatasetApi(DatasetApiResource): retrieval_model_dict["search_method"] = "keyword_search" if data.get("permission") == "partial_members": - part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + part_users_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) data.update({"partial_member_list": part_users_list}) return _dump_service_dataset_with_partial_members(data), 200 @@ -640,7 +640,7 @@ class DatasetApi(DatasetApiResource): @with_session def patch(self, session: Session, _, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") @@ -681,10 +681,10 @@ class DatasetApi(DatasetApiResource): dataset, str(payload.permission) if payload.permission else None, payload.partial_member_list, - db.session, + session=db.session(), ) - dataset = DatasetService.update_dataset(session, dataset_id_str, update_data, current_user) + dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user, session=session) if dataset is None: raise NotFound("Dataset not found.") @@ -695,13 +695,13 @@ class DatasetApi(DatasetApiResource): if payload.partial_member_list and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: DatasetPermissionService.update_partial_member_list( - tenant_id, dataset_id_str, payload.partial_member_list, db.session + tenant_id, dataset_id_str, payload.partial_member_list, db.session() ) # clear partial member list when permission is only_me or all_team_members elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) - partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session) + partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, db.session()) result_data.update({"partial_member_list": partial_member_list}) return _dump_service_dataset_with_partial_members(result_data), 200 @@ -754,8 +754,8 @@ class DatasetApi(DatasetApiResource): dataset_id_str = str(dataset_id) try: - if DatasetService.delete_dataset(dataset_id_str, current_user, db.session): - DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session) + if DatasetService.delete_dataset(dataset_id_str, current_user, db.session()): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, db.session()) return "", 204 else: raise NotFound("Dataset not found.") @@ -820,14 +820,14 @@ class DocumentStatusApi(DatasetApiResource): InvalidActionError: If the action is invalid or cannot be performed. """ dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") # Check user's permission try: - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -839,7 +839,7 @@ class DocumentStatusApi(DatasetApiResource): document_ids = data.get("document_ids", []) try: - DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session) + DocumentService.batch_update_document_status(dataset, document_ids, action, current_user, db.session()) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -876,7 +876,7 @@ class DatasetTagsApi(DatasetApiResource): assert isinstance(current_user, Account) cid = current_user.current_tenant_id assert cid is not None - tags = TagService.get_tags(db.session(), "knowledge", cid) + tags = TagService.get_tags("knowledge", cid, session=db.session()) return dump_response(KnowledgeTagListResponse, tags), 200 @service_api_ns.doc( @@ -909,7 +909,7 @@ class DatasetTagsApi(DatasetApiResource): raise Forbidden() payload = TagCreatePayload.model_validate(service_api_ns.payload or {}) - tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=TagType.KNOWLEDGE), db.session) + tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=TagType.KNOWLEDGE), db.session()) response = dump_response( KnowledgeTagResponse, @@ -948,10 +948,10 @@ class DatasetTagsApi(DatasetApiResource): payload = TagUpdatePayload.model_validate(service_api_ns.payload or {}) tag_id = payload.tag_id tag = TagService.update_tags( - UpdateTagServicePayload(name=payload.name), tag_id, db.session, tag_type=TagType.KNOWLEDGE + UpdateTagServicePayload(name=payload.name), tag_id, db.session(), tag_type=TagType.KNOWLEDGE ) - binding_count = TagService.get_tag_binding_count(tag_id, db.session, tag_type=TagType.KNOWLEDGE) + binding_count = TagService.get_tag_binding_count(tag_id, db.session(), tag_type=TagType.KNOWLEDGE) response = dump_response( KnowledgeTagResponse, @@ -981,7 +981,7 @@ class DatasetTagsApi(DatasetApiResource): def delete(self, _): """Delete a knowledge type tag.""" payload = TagDeletePayload.model_validate(service_api_ns.payload or {}) - TagService.delete_tag(payload.tag_id, db.session, tag_type=TagType.KNOWLEDGE) + TagService.delete_tag(payload.tag_id, db.session(), tag_type=TagType.KNOWLEDGE) return "", 204 @@ -1015,7 +1015,7 @@ class DatasetTagBindingApi(DatasetApiResource): payload = TagBindingPayload.model_validate(service_api_ns.payload or {}) TagService.save_tag_binding( TagBindingCreatePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE), - db.session, + db.session(), ) return "", 204 @@ -1050,7 +1050,7 @@ class DatasetTagUnbindingApi(DatasetApiResource): payload = TagUnbindingPayload.model_validate(service_api_ns.payload or {}) TagService.delete_tag_binding( TagBindingDeletePayload(tag_ids=payload.tag_ids, target_id=payload.target_id, type=TagType.KNOWLEDGE), - db.session, + db.session(), ) return "", 204 @@ -1086,7 +1086,7 @@ class DatasetTagsBindingStatusApi(DatasetApiResource): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None tags = TagService.get_tags_by_target_id( - "knowledge", current_user.current_tenant_id, str(dataset_id), db.session + "knowledge", current_user.current_tenant_id, str(dataset_id), db.session() ) tags_list = [{"id": tag.id, "name": tag.name} for tag in tags] return dump_response(DatasetBoundTagListResponse, {"data": tags_list, "total": len(tags)}), 200 diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 4c083d3d50f..5e5919a7048 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -401,7 +401,7 @@ def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Mapping[ account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -461,7 +461,7 @@ def _update_document_by_text(tenant_id: str, dataset_id: UUID, document_id: UUID account=current_user, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -759,7 +759,7 @@ class DocumentAddByFileApi(DatasetApiResource): account=dataset.created_by_account, dataset_process_rule=dataset_process_rule, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -836,7 +836,7 @@ def _update_document_by_file(tenant_id: str, dataset_id: UUID, document_id: UUID account=dataset.created_by_account, dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, created_from="api", - session=db.session, + session=db.session(), ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -955,6 +955,7 @@ class DocumentListApi(DatasetApiResource): documents=documents, dataset=dataset, tenant_id=tenant_id, + session=db.session(), ) response = { @@ -1007,7 +1008,7 @@ class DocumentBatchDownloadZipApi(DatasetApiResource): document_ids=[str(document_id) for document_id in payload.document_ids], tenant_id=str(tenant_id), current_user=current_user, - session=db.session, + session=db.session(), ) with ExitStack() as stack: @@ -1064,7 +1065,7 @@ class DocumentIndexingStatusApi(DatasetApiResource): if not dataset: raise NotFound("Dataset not found.") # get documents - documents = DocumentService.get_batch_documents(dataset_id_str, batch, db.session) + documents = DocumentService.get_batch_documents(dataset_id_str, batch, db.session()) if not documents: raise NotFound("Documents not found.") documents_status = [] @@ -1140,7 +1141,7 @@ class DocumentDownloadApi(DatasetApiResource): @cloud_edition_billing_rate_limit_check("knowledge", "dataset") def get(self, tenant_id, dataset_id: UUID, document_id: UUID): dataset = self.get_dataset(str(dataset_id), str(tenant_id)) - document = DocumentService.get_document(dataset.id, str(document_id), session=db.session) + document = DocumentService.get_document(dataset.id, str(document_id), session=db.session()) if not document: raise NotFound("Document not found.") @@ -1148,7 +1149,7 @@ class DocumentDownloadApi(DatasetApiResource): if document.tenant_id != str(tenant_id): raise Forbidden("No permission.") - return {"url": DocumentService.get_document_download_url(document, db.session)} + return {"url": DocumentService.get_document_download_url(document, db.session())} @service_api_ns.route("/datasets//documents/") @@ -1196,7 +1197,7 @@ class DocumentApi(DatasetApiResource): dataset = self.get_dataset(dataset_id_str, tenant_id) - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -1216,12 +1217,13 @@ class DocumentApi(DatasetApiResource): document_id=document_id_str, dataset_id=dataset_id_str, tenant_id=tenant_id, + session=db.session(), ) if metadata == "only": response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} elif metadata == "without": - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1256,7 +1258,7 @@ class DocumentApi(DatasetApiResource): "need_summary": document.need_summary if document.need_summary is not None else False, } else: - dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session) + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, db.session()) document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} data_source_info = document.data_source_detail_dict response = { @@ -1351,7 +1353,7 @@ class DocumentApi(DatasetApiResource): if not dataset: raise ValueError("Dataset does not exist.") - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) # 404 if document not found if document is None: @@ -1363,7 +1365,7 @@ class DocumentApi(DatasetApiResource): try: # delete document - DocumentService.delete_document(document, db.session) + DocumentService.delete_document(document, db.session()) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index aec3b06a91e..1d793583cc2 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -81,12 +81,12 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): metadata_args = MetadataArgs.model_validate(service_api_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - metadata = MetadataService.create_metadata(db.session(), dataset_id_str, metadata_args) + metadata = MetadataService.create_metadata(dataset_id_str, metadata_args, session=db.session()) return dump_response(DatasetMetadataResponse, metadata), 201 @service_api_ns.doc( @@ -116,10 +116,10 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): def get(self, tenant_id, dataset_id: UUID): """Get all metadata for a dataset.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(db.session(), dataset) + metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -154,12 +154,14 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - metadata = MetadataService.update_metadata_name(db.session(), dataset_id_str, metadata_id_str, payload.name) + metadata = MetadataService.update_metadata_name( + dataset_id_str, metadata_id_str, payload.name, session=db.session() + ) return dump_response(DatasetMetadataResponse, metadata), 200 @service_api_ns.doc( @@ -189,12 +191,12 @@ class DatasetMetadataServiceApi(DatasetApiResource): """Delete metadata.""" dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) - MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) return "", 204 @@ -257,16 +259,16 @@ class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID, action: Literal["enable", "disable"]): """Enable or disable built-in metadata field.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) match action: case "enable": - MetadataService.enable_built_in_field(db.session(), dataset) + MetadataService.enable_built_in_field(dataset, session=db.session()) case "disable": - MetadataService.disable_built_in_field(db.session(), dataset) + MetadataService.disable_built_in_field(dataset, session=db.session()) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 @@ -303,13 +305,13 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): def post(self, tenant_id, dataset_id: UUID): """Update metadata for multiple documents.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, db.session) + dataset = DatasetService.get_dataset(dataset_id_str, db.session()) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user, db.session) + DatasetService.check_dataset_permission(dataset, current_user, db.session()) metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {}) - MetadataService.update_documents_metadata(db.session(), dataset, metadata_args) + MetadataService.update_documents_metadata(dataset, metadata_args, session=db.session()) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 diff --git a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py index f0f953462c9..35f3a4c01a0 100644 --- a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py @@ -159,7 +159,7 @@ class DatasourcePluginsApi(DatasetApiResource): query = query_params_from_request(DatasourcePluginsQuery) - rag_pipeline_service: RagPipelineService = RagPipelineService() + rag_pipeline_service = RagPipelineService(db.session()) datasource_plugins: list[dict[Any, Any]] = rag_pipeline_service.get_datasource_plugins( tenant_id=tenant_id, dataset_id=dataset_id_str, is_published=query.is_published ) @@ -204,7 +204,7 @@ class DatasourceNodeRunApi(DatasetApiResource): payload = DatasourceNodeRunPayload.model_validate(service_api_ns.payload or {}) assert isinstance(current_user, Account) - rag_pipeline_service: RagPipelineService = RagPipelineService() + rag_pipeline_service: RagPipelineService = RagPipelineService(db.session()) pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) datasource_node_run_api_entity = DatasourceNodeRunApiEntity.model_validate( { @@ -272,7 +272,7 @@ class PipelineRunApi(DatasetApiResource): dataset_id_str = str(dataset_id) # Verify dataset ownership stmt = select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str) - dataset = db.session.scalar(stmt) + dataset = session.scalar(stmt) if not dataset: raise NotFound("Dataset not found.") @@ -281,8 +281,8 @@ class PipelineRunApi(DatasetApiResource): if not isinstance(current_user, Account): raise Forbidden() - rag_pipeline_service: RagPipelineService = RagPipelineService() - pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) + rag_pipeline_service = RagPipelineService(session) + pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) try: response: dict[Any, Any] | Generator[str, Any, None] = PipelineGenerateService.generate( session=session, diff --git a/api/controllers/service_api/dataset/segment.py b/api/controllers/service_api/dataset/segment.py index 41fbc709fdd..e911c454c9e 100644 --- a/api/controllers/service_api/dataset/segment.py +++ b/api/controllers/service_api/dataset/segment.py @@ -137,7 +137,7 @@ def _get_segment_for_document( raise NotFound("Document not found.") segment_ref = DatasetRefService.create_segment_ref(document_ref, segment_id) - segment = SegmentService.get_segment_by_ref(segment_ref) + segment = SegmentService.get_segment_by_ref(segment_ref, db.session()) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -191,7 +191,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if document.indexing_status != "completed": @@ -227,13 +227,13 @@ class SegmentApi(DatasetApiResource): for args_item in segment_items: SegmentService.segment_create_args_validate(args_item, document) segments = cast( - list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset, db.session) + list[DocumentSegment], SegmentService.multi_create_segment(segment_items, document, dataset, db.session()) ) segment_ids = [segment.id for segment in segments] summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} response = { @@ -285,7 +285,7 @@ class SegmentApi(DatasetApiResource): raise NotFound("Dataset not found.") document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") # check embedding model setting @@ -317,7 +317,7 @@ class SegmentApi(DatasetApiResource): summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str + segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} @@ -367,12 +367,12 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - SegmentService.delete_segment(segment, document, dataset, db.session) + SegmentService.delete_segment(segment, document, dataset, db.session()) return "", 204 @service_api_ns.doc( @@ -410,7 +410,7 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -434,8 +434,10 @@ class DatasetSegmentApi(DatasetApiResource): payload = SegmentUpdatePayload.model_validate(service_api_ns.payload or {}) - updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session) - summary = SummaryIndexService.get_segment_summary(segment_id=updated_segment.id, dataset_id=dataset_id_str) + updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session()) + summary = SummaryIndexService.get_segment_summary( + segment_id=updated_segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(updated_segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -481,13 +483,15 @@ class DatasetSegmentApi(DatasetApiResource): DatasetService.check_dataset_model_setting(dataset) document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) _, segment = _get_segment_for_document(dataset, document, segment_id_str) - summary = SummaryIndexService.get_segment_summary(segment_id=segment.id, dataset_id=dataset_id_str) + summary = SummaryIndexService.get_segment_summary( + segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + ) response = { "data": segment_response_with_summary(segment, summary.summary_content if summary else None), "doc_form": document.doc_form, @@ -542,7 +546,7 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -570,7 +574,7 @@ class ChildChunkApi(DatasetApiResource): payload = ChildChunkCreatePayload.model_validate(service_api_ns.payload or {}) try: - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session) + child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, db.session()) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) @@ -613,7 +617,7 @@ class ChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -680,7 +684,7 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # check document - document = DocumentService.get_document(dataset.id, document_id_str, session=db.session) + document = DocumentService.get_document(dataset.id, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -689,12 +693,12 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # check child chunk - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") try: - SegmentService.delete_child_chunk(child_chunk, dataset, db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, db.session()) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) @@ -741,7 +745,7 @@ class DatasetChildChunkApi(DatasetApiResource): document_id_str = str(document_id) # get document - document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session) + document = DocumentService.get_document(dataset_id_str, document_id_str, session=db.session()) if not document: raise NotFound("Document not found.") @@ -750,7 +754,7 @@ class DatasetChildChunkApi(DatasetApiResource): child_chunk_id_str = str(child_chunk_id) # get child chunk - child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, db.session()) if not child_chunk: raise NotFound("Child chunk not found.") @@ -759,7 +763,7 @@ class DatasetChildChunkApi(DatasetApiResource): try: child_chunk = SegmentService.update_child_chunk( - payload.content, child_chunk, segment, document, dataset, db.session + payload.content, child_chunk, segment, document, dataset, db.session() ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/web/app.py b/api/controllers/web/app.py index 17ff05f7137..6804d072ef0 100644 --- a/api/controllers/web/app.py +++ b/api/controllers/web/app.py @@ -12,6 +12,7 @@ from controllers.common.agent_app_parameters import get_published_agent_app_feat from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError +from extensions.ext_database import db from libs.passport import PassportService from libs.token import extract_webapp_passport from models.model import App, AppMode, EndUser @@ -122,7 +123,7 @@ class AppMeta(WebApiResource): @web_ns.response(200, "Success", web_ns.models[AppMetaResponse.__name__]) def get(self, app_model: App, end_user: EndUser): """Get app meta""" - return AppService().get_app_meta(app_model) + return AppService().get_app_meta(app_model, session=db.session()) @web_ns.route("/webapp/access-mode") @@ -148,7 +149,7 @@ class AppAccessMode(Resource): app_id = args.app_id if args.app_code: - app_id = AppService.get_app_id_by_code(args.app_code) + app_id = AppService.get_app_id_by_code(args.app_code, session=db.session()) if not app_id: raise ValueError("appId or appCode must be provided") @@ -179,7 +180,9 @@ class AppWebAuthPermission(Resource): if not app_id or not app_code: raise ValueError("appId must be provided") - require_permission_check = WebAppAuthService.is_app_require_permission_check(app_id=app_id) + require_permission_check = WebAppAuthService.is_app_require_permission_check( + app_id=app_id, session=db.session() + ) if not require_permission_check: return {"result": True} @@ -200,6 +203,6 @@ class AppWebAuthPermission(Resource): return {"result": True} res = True - if WebAppAuthService.is_app_require_permission_check(app_id=app_id): + if WebAppAuthService.is_app_require_permission_check(app_id=app_id, session=db.session()): res = EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(str(user_id), app_id) return {"result": res} diff --git a/api/controllers/web/audio.py b/api/controllers/web/audio.py index 47e72ff95a5..b7856f7dd90 100644 --- a/api/controllers/web/audio.py +++ b/api/controllers/web/audio.py @@ -141,7 +141,7 @@ class TextApi(WebApiResource): ) response = AudioService.transcript_tts( app_model=app_model, - session=db.session, + session=db.session(), text=text, voice=voice, end_user=end_user.external_user_id, diff --git a/api/controllers/web/completion.py b/api/controllers/web/completion.py index 343afd68f9a..c1a7d1f8d10 100644 --- a/api/controllers/web/completion.py +++ b/api/controllers/web/completion.py @@ -30,6 +30,7 @@ from core.errors.error import ( ProviderTokenNotInitError, QuotaExceededError, ) +from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs import helper from libs.helper import uuid_value @@ -219,7 +220,10 @@ class ChatApi(WebApiResource): # Eagerly validate conversation to avoid hanging on invalid conversation_id if payload.conversation_id: ConversationService.get_conversation( - app_model=app_model, conversation_id=payload.conversation_id, user=end_user + app_model=app_model, + conversation_id=payload.conversation_id, + user=end_user, + session=db.session(), ) response = AppGenerateService.generate( diff --git a/api/controllers/web/conversation.py b/api/controllers/web/conversation.py index 73461b1a294..09a3a508824 100644 --- a/api/controllers/web/conversation.py +++ b/api/controllers/web/conversation.py @@ -112,7 +112,7 @@ class ConversationApi(WebApiResource): conversation_id = str(c_id) try: - ConversationService.delete(app_model, conversation_id, end_user) + ConversationService.delete(app_model, conversation_id, end_user, session=db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") return "", 204 @@ -157,7 +157,7 @@ class ConversationRenameApi(WebApiResource): try: conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() ) return ( TypeAdapter(SimpleConversation) @@ -192,7 +192,7 @@ class ConversationPinApi(WebApiResource): conversation_id = str(c_id) try: - WebConversationService.pin(app_model, conversation_id, end_user) + WebConversationService.pin(app_model, conversation_id, end_user, db.session()) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -221,6 +221,6 @@ class ConversationUnPinApi(WebApiResource): raise NotChatAppError() conversation_id = str(c_id) - WebConversationService.unpin(app_model, conversation_id, end_user) + WebConversationService.unpin(app_model, conversation_id, end_user, db.session()) return ResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/web/forgot_password.py b/api/controllers/web/forgot_password.py index ecc91113c32..a9374555ed4 100644 --- a/api/controllers/web/forgot_password.py +++ b/api/controllers/web/forgot_password.py @@ -69,7 +69,7 @@ class ForgotPasswordSendEmailApi(Resource): else: language = "en-US" - account = AccountService.get_account_by_email_with_case_fallback(db.session, request_email) + account = AccountService.get_account_by_email_with_case_fallback(request_email, session=db.session()) if account is None: raise AuthenticationFailedError() else: @@ -168,7 +168,7 @@ class ForgotPasswordResetApi(Resource): email = reset_data.get("email", "") - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) if account: account = db.session.merge(account) diff --git a/api/controllers/web/login.py b/api/controllers/web/login.py index 011bb43b880..0aa42f43687 100644 --- a/api/controllers/web/login.py +++ b/api/controllers/web/login.py @@ -30,6 +30,7 @@ from controllers.console.wraps import ( ) from controllers.web import web_ns from controllers.web.wraps import decode_jwt_token +from extensions.ext_database import db from libs.helper import EmailStr, extract_remote_ip from libs.passport import PassportService from libs.password import valid_password @@ -104,7 +105,7 @@ class LoginApi(Resource): normalized_email = payload.email.lower() try: - account = WebAppAuthService.authenticate(payload.email, payload.password) + account = WebAppAuthService.authenticate(payload.email, payload.password, db.session()) except services.errors.account.AccountLoginError: _log_web_login_failure(email=normalized_email, reason=LoginFailureReason.ACCOUNT_BANNED) raise AccountBannedError() @@ -144,9 +145,9 @@ class LoginStatusApi(Resource): token = extract_webapp_access_token(request) if not app_code: return LoginStatusResponse(logged_in=bool(token), app_logged_in=False).model_dump(mode="json") - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) is_public = not dify_config.ENTERPRISE_ENABLED or not WebAppAuthService.is_app_require_permission_check( - app_id=app_id + app_id=app_id, session=db.session() ) user_logged_in = False @@ -211,7 +212,7 @@ class EmailCodeLoginSendEmailApi(Resource): else: language = "en-US" - account = WebAppAuthService.get_user_through_email(payload.email) + account = WebAppAuthService.get_user_through_email(payload.email, db.session()) if account is None: raise AuthenticationFailedError() token = WebAppAuthService.send_email_code_login_email(account=account, language=language) @@ -264,7 +265,7 @@ class EmailCodeLoginApi(Resource): WebAppAuthService.revoke_email_code_login_token(payload.token) try: - account = WebAppAuthService.get_user_through_email(token_email) + account = WebAppAuthService.get_user_through_email(token_email, db.session()) except Unauthorized as exc: _log_web_login_failure(email=user_email, reason=LoginFailureReason.ACCOUNT_BANNED) raise AccountBannedError() from exc diff --git a/api/controllers/web/message.py b/api/controllers/web/message.py index 691eba05491..45fea9a328e 100644 --- a/api/controllers/web/message.py +++ b/api/controllers/web/message.py @@ -25,6 +25,7 @@ from controllers.web.error import ( from controllers.web.wraps import WebApiResource from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError +from extensions.ext_database import db from fields.conversation_fields import ResultResponse from fields.message_fields import SuggestedQuestionsResponse, WebMessageInfiniteScrollPagination, WebMessageListItem from graphon.model_runtime.errors.invoke import InvokeError @@ -86,7 +87,7 @@ class MessageListApi(WebApiResource): try: pagination = MessageService.pagination_by_first_id( - app_model, end_user, query.conversation_id, query.first_id, query.limit + app_model, end_user, query.conversation_id, query.first_id, query.limit, session=db.session() ) adapter = TypeAdapter(WebMessageListItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -141,6 +142,7 @@ class MessageFeedbackApi(WebApiResource): user=end_user, rating=FeedbackRating(payload.rating) if payload.rating else None, content=payload.content, + session=db.session(), ) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -231,7 +233,11 @@ class MessageSuggestedQuestionApi(WebApiResource): try: questions = MessageService.get_suggested_questions_after_answer( - app_model=app_model, user=end_user, message_id=message_id_str, invoke_from=InvokeFrom.WEB_APP + app_model=app_model, + user=end_user, + message_id=message_id_str, + invoke_from=InvokeFrom.WEB_APP, + session=db.session(), ) # questions is a list of strings, not a list of Message objects except MessageNotExistsError: diff --git a/api/controllers/web/passport.py b/api/controllers/web/passport.py index c11ce824731..4b0b25fb971 100644 --- a/api/controllers/web/passport.py +++ b/api/controllers/web/passport.py @@ -62,7 +62,7 @@ class PassportResource(Resource): raise Unauthorized("X-App-Code header is missing.") if system_features.webapp_auth.enabled: enterprise_user_decoded = decode_enterprise_webapp_user_id(access_token) - app_auth_type = WebAppAuthService.get_app_auth_type(app_code=app_code) + app_auth_type = WebAppAuthService.get_app_auth_type(app_code=app_code, session=db.session()) if app_auth_type != WebAppAuthType.PUBLIC: if not enterprise_user_decoded: raise WebAppAuthRequiredError() diff --git a/api/controllers/web/saved_message.py b/api/controllers/web/saved_message.py index 6e59a85e2b0..d61ffd545c3 100644 --- a/api/controllers/web/saved_message.py +++ b/api/controllers/web/saved_message.py @@ -44,7 +44,7 @@ class SavedMessageListApi(WebApiResource): query = SavedMessageListQuery.model_validate(raw_args) pagination = SavedMessageService.pagination_by_last_id( - db.session(), app_model, end_user, query.last_id, query.limit + app_model, end_user, query.last_id, query.limit, session=db.session() ) adapter = TypeAdapter(SavedMessageItem) items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] @@ -80,7 +80,7 @@ class SavedMessageListApi(WebApiResource): payload = SavedMessageCreatePayload.model_validate(web_ns.payload or {}) try: - SavedMessageService.save(db.session(), app_model, end_user, payload.message_id) + SavedMessageService.save(app_model, end_user, payload.message_id, session=db.session()) except MessageNotExistsError: raise NotFound("Message Not Exists.") @@ -108,6 +108,6 @@ class SavedMessageApi(WebApiResource): if app_model.mode != "completion": raise NotCompletionAppError() - SavedMessageService.delete(db.session(), app_model, end_user, message_id_str) + SavedMessageService.delete(app_model, end_user, message_id_str, session=db.session()) return "", 204 diff --git a/api/controllers/web/wraps.py b/api/controllers/web/wraps.py index ccc9c0f8f60..eff4b70ff0f 100644 --- a/api/controllers/web/wraps.py +++ b/api/controllers/web/wraps.py @@ -70,7 +70,7 @@ def decode_jwt_token(app_code: str | None = None, user_id: str | None = None) -> app_web_auth_enabled = False webapp_settings = None if system_features.webapp_auth.enabled: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) webapp_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id) if not webapp_settings: raise NotFound("Web app settings not found.") @@ -86,7 +86,7 @@ def decode_jwt_token(app_code: str | None = None, user_id: str | None = None) -> if system_features.webapp_auth.enabled: if not app_code: raise Unauthorized("Please re-login to access the web app.") - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) app_web_auth_enabled = ( EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=app_id).access_mode != WebAppAccessMode.PUBLIC @@ -129,8 +129,10 @@ def _validate_user_accessibility( if not webapp_settings: raise WebAppAuthRequiredError("Web app settings not found.") - if WebAppAuthService.is_app_require_permission_check(access_mode=webapp_settings.access_mode): - app_id = AppService.get_app_id_by_code(app_code) + if WebAppAuthService.is_app_require_permission_check( + access_mode=webapp_settings.access_mode, session=db.session() + ): + app_id = AppService.get_app_id_by_code(app_code, session=db.session()) if not EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(user_id, app_id): raise WebAppAuthAccessDeniedError() diff --git a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py index 140d4e6a2a6..0108e7d7c72 100644 --- a/api/core/app/app_config/easy_ui_based_app/dataset/manager.py +++ b/api/core/app/app_config/easy_ui_based_app/dataset/manager.py @@ -257,7 +257,7 @@ class DatasetConfigManager: @classmethod def is_dataset_exists(cls, tenant_id: str, dataset_id: str) -> bool: # verify if the dataset ID exists - dataset = DatasetService.get_dataset(dataset_id, db.session) + dataset = DatasetService.get_dataset(dataset_id, db.session()) if not dataset: return False diff --git a/api/core/app/apps/advanced_chat/app_generator.py b/api/core/app/apps/advanced_chat/app_generator.py index f52fd1046f8..75ada4fe888 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -156,7 +156,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): if conversation_id: try: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) except ConversationNotExistsError: if invoke_from == InvokeFrom.SERVICE_API: diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index 9531f9092a4..f2f62496883 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -105,7 +105,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # Build the EasyUI-shaped config from the Agent Soul so the chat pipeline @@ -284,7 +284,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): out of scope here — the message is persisted and can be re-fetched. """ conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( app_model, diff --git a/api/core/app/apps/agent_chat/app_generator.py b/api/core/app/apps/agent_chat/app_generator.py index d640bcdc863..a3cc913abf3 100644 --- a/api/core/app/apps/agent_chat/app_generator.py +++ b/api/core/app/apps/agent_chat/app_generator.py @@ -108,7 +108,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # get app model config app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) diff --git a/api/core/app/apps/chat/app_generator.py b/api/core/app/apps/chat/app_generator.py index 4873168b885..678525e0f77 100644 --- a/api/core/app/apps/chat/app_generator.py +++ b/api/core/app/apps/chat/app_generator.py @@ -105,7 +105,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): conversation_id = args.get("conversation_id") if conversation_id: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() ) # get app model config app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) diff --git a/api/core/app/features/annotation_reply/annotation_reply.py b/api/core/app/features/annotation_reply/annotation_reply.py index 520ba7b85b3..9eff9747764 100644 --- a/api/core/app/features/annotation_reply/annotation_reply.py +++ b/api/core/app/features/annotation_reply/annotation_reply.py @@ -45,7 +45,7 @@ class AnnotationReplyFeature: embedding_model_name = collection_binding_detail.model_name dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( - embedding_provider_name, embedding_model_name, db.session, CollectionBindingType.ANNOTATION + embedding_provider_name, embedding_model_name, db.session(), CollectionBindingType.ANNOTATION ) dataset = Dataset( @@ -66,7 +66,7 @@ class AnnotationReplyFeature: if documents and documents[0].metadata: annotation_id = documents[0].metadata["annotation_id"] score = documents[0].metadata["score"] - annotation = AppAnnotationService.get_annotation_by_id(annotation_id) + annotation = AppAnnotationService.get_annotation_by_id(annotation_id, session=db.session()) if annotation: if invoke_from in {InvokeFrom.SERVICE_API, InvokeFrom.WEB_APP}: from_source = ConversationFromSource.API @@ -84,6 +84,7 @@ class AnnotationReplyFeature: message.id, from_source, score, + session=db.session(), ) return annotation diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index 5bf3334a7b2..d26d5d8a998 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -125,6 +125,7 @@ def _deduct_used_llm_quota(*, tenant_id: str, provider: str, provider_configurat CreditPoolService.deduct_credits_capped( tenant_id=tenant_id, credits_required=used_quota, + session=db.session(), ) case ProviderQuotaType.PAID: from services.credit_pool_service import CreditPoolService @@ -133,6 +134,7 @@ def _deduct_used_llm_quota(*, tenant_id: str, provider: str, provider_configurat tenant_id=tenant_id, credits_required=used_quota, pool_type="paid", + session=db.session(), ) case ProviderQuotaType.FREE: _deduct_free_llm_quota( diff --git a/api/core/app/task_pipeline/message_cycle_manager.py b/api/core/app/task_pipeline/message_cycle_manager.py index 6b6437adac3..5ada7d0ba2d 100644 --- a/api/core/app/task_pipeline/message_cycle_manager.py +++ b/api/core/app/task_pipeline/message_cycle_manager.py @@ -154,7 +154,7 @@ class MessageCycleManager: :param event: event :return: """ - annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id) + annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session=db.session()) if annotation: account = annotation.account self._task_state.metadata.annotation_reply = AnnotationReply( diff --git a/api/core/callback_handler/index_tool_callback_handler.py b/api/core/callback_handler/index_tool_callback_handler.py index 26dc1a12a2c..d2024454a68 100644 --- a/api/core/callback_handler/index_tool_callback_handler.py +++ b/api/core/callback_handler/index_tool_callback_handler.py @@ -2,7 +2,7 @@ import logging from collections.abc import Sequence from sqlalchemy import select, update -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.entities.app_invoke_entities import InvokeFrom @@ -30,7 +30,7 @@ class DatasetIndexToolCallbackHandler: self._user_id = user_id self._invoke_from = invoke_from - def on_query(self, query: str, dataset_id: str, session: scoped_session): + def on_query(self, query: str, dataset_id: str, session: Session): """ Handle query. """ @@ -52,7 +52,7 @@ class DatasetIndexToolCallbackHandler: with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as independent_session: independent_session.add(dataset_query) - def on_tool_end(self, documents: list[Document], session: scoped_session): + def on_tool_end(self, documents: list[Document], session: Session): """Handle tool end.""" # Use an independent session so hit-count updates do not # interfere with the caller's request-scoped session. diff --git a/api/core/llm_generator/llm_generator.py b/api/core/llm_generator/llm_generator.py index f97f9c38330..29a93fff815 100644 --- a/api/core/llm_generator/llm_generator.py +++ b/api/core/llm_generator/llm_generator.py @@ -6,6 +6,7 @@ from typing import Any, Literal, NotRequired, Protocol, TypedDict, cast import json_repair from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import ModelConfig from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload @@ -117,7 +118,9 @@ def _parse_string_list(text: str) -> list[str]: class WorkflowServiceInterface(Protocol): - def get_draft_workflow(self, app_model: App, workflow_id: str | None = None) -> Workflow | None: + def get_draft_workflow( + self, app_model: App, workflow_id: str | None = None, *, session: Session + ) -> Workflow | None: pass def get_node_last_run(self, app_model: App, workflow: Workflow, node_id: str) -> WorkflowNodeExecutionModel | None: @@ -758,7 +761,7 @@ class LLMGenerator: app: App | None = session.scalar(select(App).where(App.id == flow_id, App.tenant_id == tenant_id).limit(1)) if not app: raise ValueError("App not found.") - workflow = workflow_service.get_draft_workflow(app_model=app) + workflow = workflow_service.get_draft_workflow(app_model=app, session=session) if not workflow: raise ValueError("Workflow not found for the given app model.") last_run = workflow_service.get_node_last_run(app_model=app, workflow=workflow, node_id=node_id) diff --git a/api/core/mcp/server/streamable_http.py b/api/core/mcp/server/streamable_http.py index 964f3211db0..7fd03788c7e 100644 --- a/api/core/mcp/server/streamable_http.py +++ b/api/core/mcp/server/streamable_http.py @@ -207,11 +207,11 @@ def handle_call_tool( raise ValueError("End user not found") response = AppGenerateService.generate( - session, - app, - end_user, - args, - InvokeFrom.SERVICE_API, + session=session, + app_model=app, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, streaming=app.mode == AppMode.AGENT_CHAT, ) diff --git a/api/core/provider_manager.py b/api/core/provider_manager.py index e2c710923b5..ebfe77e8f30 100644 --- a/api/core/provider_manager.py +++ b/api/core/provider_manager.py @@ -1544,10 +1544,12 @@ class ProviderManager: trail_pool = CreditPoolService.get_pool( tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, + session=db.session(), ) paid_pool = CreditPoolService.get_pool( tenant_id=tenant_id, pool_type=ProviderQuotaType.PAID, + session=db.session(), ) else: trail_pool = None diff --git a/api/core/rag/datasource/retrieval_service.py b/api/core/rag/datasource/retrieval_service.py index 50381f5e75c..3b20f8bc530 100644 --- a/api/core/rag/datasource/retrieval_service.py +++ b/api/core/rag/datasource/retrieval_service.py @@ -199,7 +199,7 @@ class RetrievalService: metadata_filtering_conditions: dict[str, Any] | None = None, ): stmt = select(Dataset).where(Dataset.id == dataset_id) - dataset = db.session.scalar(stmt) + dataset = session.scalar(stmt) if not dataset: return [] metadata_condition = ( @@ -208,12 +208,12 @@ class RetrievalService: else None ) all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( - session, - dataset.tenant_id, - dataset_id, - query, - external_retrieval_model or {}, + tenant_id=dataset.tenant_id, + dataset_id=dataset_id, + query=query, + external_retrieval_parameters=external_retrieval_model or {}, metadata_condition=metadata_condition, + session=session, ) return all_documents diff --git a/api/core/rag/index_processor/processor/paragraph_index_processor.py b/api/core/rag/index_processor/processor/paragraph_index_processor.py index dd173207b09..b31c1bb634b 100644 --- a/api/core/rag/index_processor/processor/paragraph_index_processor.py +++ b/api/core/rag/index_processor/processor/paragraph_index_processor.py @@ -5,7 +5,7 @@ import re import uuid from typing import Any, TypedDict, cast, override -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session logger = logging.getLogger(__name__) @@ -162,10 +162,10 @@ class ParagraphIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: vector = Vector(dataset) @@ -226,7 +226,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): all_multimodal_documents.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session) + account = AccountService.load_user(document.created_by, db.session()) if not account: raise ValueError("Invalid account") doc.attachments = self._get_content_files(doc, current_user=account) @@ -414,12 +414,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor): # First, try to get images from SegmentAttachmentBinding (preferred method) if segment_id: image_files = ParagraphIndexProcessor._extract_images_from_segment_attachments( - tenant_id, segment_id, db.session + tenant_id, segment_id, db.session() ) # If no images from attachments, fall back to extracting from text if not image_files: - image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, db.session) + image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, db.session()) # Build prompt messages prompt_messages = [] @@ -473,7 +473,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): return summary_content, usage @staticmethod - def _extract_images_from_text(tenant_id: str, text: str, session: scoped_session) -> list[File]: + def _extract_images_from_text(tenant_id: str, text: str, session: Session) -> list[File]: """ Extract images from markdown text and convert them to File objects. @@ -553,9 +553,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): return file_objects @staticmethod - def _extract_images_from_segment_attachments( - tenant_id: str, segment_id: str, session: scoped_session - ) -> list[File]: + def _extract_images_from_segment_attachments(tenant_id: str, segment_id: str, session: Session) -> list[File]: """ Extract images from SegmentAttachmentBinding table (preferred method). This matches how DatasetRetrieval gets segment attachments. diff --git a/api/core/rag/index_processor/processor/parent_child_index_processor.py b/api/core/rag/index_processor/processor/parent_child_index_processor.py index 78d8b7dcd53..aecb4154d6f 100644 --- a/api/core/rag/index_processor/processor/parent_child_index_processor.py +++ b/api/core/rag/index_processor/processor/parent_child_index_processor.py @@ -169,10 +169,10 @@ class ParentChildIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: delete_child_chunks = kwargs.get("delete_child_chunks") or False @@ -291,7 +291,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): attachments.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session) + account = AccountService.load_user(document.created_by, db.session()) if not account: raise ValueError("Invalid account") doc.attachments = self._get_content_files(doc, current_user=account) diff --git a/api/core/rag/index_processor/processor/qa_index_processor.py b/api/core/rag/index_processor/processor/qa_index_processor.py index 253acebc2c6..7b7443a621f 100644 --- a/api/core/rag/index_processor/processor/qa_index_processor.py +++ b/api/core/rag/index_processor/processor/qa_index_processor.py @@ -173,10 +173,10 @@ class QAIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset, None) + SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) vector = Vector(dataset) if node_ids: diff --git a/api/core/rag/summary_index/summary_index.py b/api/core/rag/summary_index/summary_index.py index bff5f85decb..d9ce3879890 100644 --- a/api/core/rag/summary_index/summary_index.py +++ b/api/core/rag/summary_index/summary_index.py @@ -74,11 +74,16 @@ class SummaryIndex: def process_segment(segment_id: str) -> None: """Process a single segment in a thread with a fresh DB session.""" with session_factory.create_session() as session: + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if dataset is None: + return segment = session.scalar(select(DocumentSegment).where(DocumentSegment.id == segment_id).limit(1)) if segment is None: return try: - SummaryIndexService.generate_and_vectorize_summary(segment, dataset, summary_index_setting) + SummaryIndexService.generate_and_vectorize_summary( + segment, dataset, summary_index_setting, session=session + ) except Exception: logger.exception( "Failed to generate summary for segment %s", diff --git a/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py b/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py index a3afe659563..c26523b9be5 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_multi_retriever_tool.py @@ -80,7 +80,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): all_documents = rerank_runner.run(query, all_documents, self.score_threshold, self.top_k) for hit_callback in self.hit_callbacks: - hit_callback.on_tool_end(all_documents, db.session) + hit_callback.on_tool_end(all_documents, db.session()) document_score_list = {} for item in all_documents: @@ -167,7 +167,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): return [] for hit_callback in hit_callbacks: - hit_callback.on_query(query, dataset.id, db.session) + hit_callback.on_query(query, dataset.id, db.session()) # get retrieval model , if the model is not setting , using default retrieval_model = dataset.retrieval_model or default_retrieval_model diff --git a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py index 247bd0705fc..d7e390ca877 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py @@ -65,7 +65,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): if not dataset: return "" for hit_callback in self.hit_callbacks: - hit_callback.on_query(query, dataset.id, db.session) + hit_callback.on_query(query, dataset.id, db.session()) dataset_retrieval = DatasetRetrieval() metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition( session, @@ -162,7 +162,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): else: documents = [] for hit_callback in self.hit_callbacks: - hit_callback.on_tool_end(documents, db.session) + hit_callback.on_tool_end(documents, db.session()) document_score_list = {} if dataset.indexing_technique != IndexTechniqueType.ECONOMY: for item in documents: diff --git a/api/core/workflow/nodes/agent_v2/dify_tools_builder.py b/api/core/workflow/nodes/agent_v2/dify_tools_builder.py index fc2719a6204..0e6fb3d1830 100644 --- a/api/core/workflow/nodes/agent_v2/dify_tools_builder.py +++ b/api/core/workflow/nodes/agent_v2/dify_tools_builder.py @@ -15,7 +15,6 @@ from dify_agent.layers.dify_plugin import ( DifyPluginToolsLayerConfig, ) from sqlalchemy import select -from sqlalchemy.orm import Session from core.agent.entities import AgentToolEntity from core.app.entities.app_invoke_entities import InvokeFrom @@ -132,7 +131,7 @@ def _list_provider_tool_names( def _resolve_mcp_provider_id(*, tenant_id: str, provider_id: str) -> str: """Normalize MCP provider ids to the runtime-facing server identifier.""" - service = MCPToolManageService(session=cast(Session, db.session)) + service = MCPToolManageService(session=db.session()) try: return service.get_provider_entity(provider_id, tenant_id, by_server_id=True).provider_id except ValueError: diff --git a/api/events/event_handlers/update_provider_when_message_created.py b/api/events/event_handlers/update_provider_when_message_created.py index 8dec5876a9b..15b40afdbf2 100644 --- a/api/events/event_handlers/update_provider_when_message_created.py +++ b/api/events/event_handlers/update_provider_when_message_created.py @@ -204,6 +204,7 @@ def _deduct_credit_pool_quota_capped(*, tenant_id: str, credits_required: int, p tenant_id=tenant_id, credits_required=credits_required, pool_type=pool_type, + session=db.session(), ) if deducted_credits < credits_required: logger.warning( diff --git a/api/extensions/ext_login.py b/api/extensions/ext_login.py index f6496c70a78..6515b22eb36 100644 --- a/api/extensions/ext_login.py +++ b/api/extensions/ext_login.py @@ -84,7 +84,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non if not user_id: raise Unauthorized("Invalid Authorization token.") - logged_in_account = AccountService.load_logged_in_account(account_id=user_id, session=db.session) + logged_in_account = AccountService.load_logged_in_account(account_id=user_id, session=db.session()) return logged_in_account elif request.blueprint == "openapi": # Account-branch device-flow approval routes (approve / deny / @@ -103,7 +103,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non source = decoded.get("token_source") if source or not user_id: return None - return AccountService.load_logged_in_account(account_id=user_id, session=db.session) + return AccountService.load_logged_in_account(account_id=user_id, session=db.session()) elif request.blueprint == "web": app_code = request.headers.get(HEADER_NAME_APP_CODE) webapp_token = extract_webapp_passport(app_code, request) if app_code else None diff --git a/api/services/account_service.py b/api/services/account_service.py index 1b9fd724a71..b5439467a23 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -16,7 +16,7 @@ from typing import Any, NotRequired, TypedDict, cast from pydantic import BaseModel, TypeAdapter, ValidationError from sqlalchemy import Row, delete, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Unauthorized from configs import dify_config @@ -188,12 +188,12 @@ class AccountService: raise ValueError(f"Builtin RBAC role not found for {role.value} in tenant {tenant_id}") @staticmethod - def get_workspace_permission_keys(tenant_id: str, account_id: str) -> set[str]: - permissions = RBACService.MyPermissions.get(tenant_id, account_id) + def get_workspace_permission_keys(tenant_id: str, account_id: str, *, session: Session) -> set[str]: + permissions = RBACService.MyPermissions.get(tenant_id, account_id, session=session) return set(getattr(getattr(permissions, "workspace", None), "permission_keys", []) or []) @staticmethod - def get_rbac_workspace_owner_account_id(tenant_id: str, actor_account_id: str) -> str: + def get_rbac_workspace_owner_account_id(tenant_id: str, actor_account_id: str, *, session: Session) -> str: """Return the account id bound to the workspace owner RBAC role.""" owner_role_id = AccountService._resolve_legacy_role_id( tenant_id=tenant_id, @@ -211,11 +211,14 @@ class AccountService: return owner_members[0].account_id @staticmethod - def is_rbac_workspace_owner(tenant_id: str, actor_account_id: str, member_account_id: str) -> bool: + def is_rbac_workspace_owner( + tenant_id: str, actor_account_id: str, member_account_id: str, *, session: Session + ) -> bool: roles = RBACService.MemberRoles.get( tenant_id=tenant_id, account_id=actor_account_id, member_account_id=member_account_id, + session=session, ).roles return any( role.is_builtin and role.category == "global_system_default" and role.role_tag == "owner" for role in roles @@ -246,7 +249,7 @@ class AccountService: ) @staticmethod - def _refresh_account_last_active(account: Account, session: scoped_session | Session) -> None: + def _refresh_account_last_active(account: Account, session: Session) -> None: now = naive_utc_now() refresh_before = now - ACCOUNT_LAST_ACTIVE_REFRESH_INTERVAL @@ -276,7 +279,7 @@ class AccountService: redis_client.delete(AccountService._get_account_refresh_token_key(account_id)) @staticmethod - def get_account_by_email(session: Session | scoped_session, email: str) -> Account | None: + def get_account_by_email(email: str, *, session: Session) -> Account | None: """Plain ``Account`` getter keyed by email. Case-sensitive — use :meth:`has_active_account_with_email` for the case-insensitive existence check that backs the SSO collision rule. @@ -284,7 +287,7 @@ class AccountService: return session.execute(select(Account).where(Account.email == email)).scalar_one_or_none() @staticmethod - def has_active_account_with_email(session: Session | scoped_session, email: str) -> bool: + def has_active_account_with_email(email: str, *, session: Session) -> bool: if not email: return False normalized = email.strip().lower() @@ -299,7 +302,7 @@ class AccountService: return row is not None @staticmethod - def get_account_by_id(session: Session | scoped_session, account_id: str) -> Account | None: + def get_account_by_id(account_id: str, *, session: Session) -> Account | None: """Plain ``Account`` getter — no banned check, no tenant rotation, no ``last_active_at`` write. Use this from read-only identity endpoints (``/openapi/v1/account``) where ``load_user``'s @@ -311,7 +314,7 @@ class AccountService: return session.get(Account, account_id) @staticmethod - def load_user(user_id: str, session: scoped_session | Session) -> None | Account: + def load_user(user_id: str, session: Session) -> None | Account: account = session.get(Account, user_id) if not account: return None @@ -363,9 +366,7 @@ class AccountService: return token @staticmethod - def authenticate( - email: str, password: str, invite_token: str | None = None, *, session: scoped_session | Session - ) -> Account: + def authenticate(email: str, password: str, invite_token: str | None = None, *, session: Session) -> Account: """authenticate account with email and password""" account = session.scalar(select(Account).where(Account.email == email).limit(1)) @@ -396,9 +397,7 @@ class AccountService: return account @staticmethod - def update_account_password( - account: Account, password: str, new_password: str, *, session: scoped_session | Session - ): + def update_account_password(account: Account, password: str, new_password: str, *, session: Session): """update account password""" if account.password and not compare_password(password, account.password, account.password_salt): raise CurrentPasswordIncorrectError("Current password is incorrect.") @@ -429,7 +428,7 @@ class AccountService: is_setup: bool | None = False, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Create an account, preferring explicit user timezone over language-derived defaults.""" if not FeatureService.get_system_features().is_allow_register and not is_setup: @@ -487,7 +486,7 @@ class AccountService: password: str | None = None, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Create an account and owner workspace.""" account = AccountService.create_account( @@ -544,12 +543,12 @@ class AccountService: return True @staticmethod - def delete_account(account: Account): + def delete_account(account: Account, *, session: Session): """Delete account. This method only adds a task to the queue for deletion.""" # Queue account deletion sync tasks for all workspaces BEFORE account deletion (enterprise only) from services.enterprise.account_deletion_sync import sync_account_deletion - sync_success = sync_account_deletion(account_id=account.id, source="account_deleted") + sync_success = sync_account_deletion(account_id=account.id, source="account_deleted", session=session) if not sync_success: logger.warning( "Enterprise account deletion sync failed for account %s; proceeding with local deletion.", @@ -560,7 +559,7 @@ class AccountService: delete_account_task.delay(account.id) @staticmethod - def link_account_integrate(provider: str, open_id: str, account: Account, *, session: scoped_session | Session): + def link_account_integrate(provider: str, open_id: str, account: Account, *, session: Session): """Link account integrate""" try: # Query whether there is an existing binding record for the same provider @@ -589,13 +588,13 @@ class AccountService: raise LinkAccountIntegrateError("Failed to link account.") from e @staticmethod - def close_account(account: Account, *, session: scoped_session | Session): + def close_account(account: Account, *, session: Session): """Close account""" account.status = AccountStatus.CLOSED session.commit() @staticmethod - def update_account(account: Account, *, session: scoped_session | Session, **kwargs): + def update_account(account: Account, *, session: Session, **kwargs): """Update account fields""" account = session.merge(account) for field, value in kwargs.items(): @@ -608,7 +607,7 @@ class AccountService: return account @staticmethod - def update_account_email(account: Account, email: str, session: scoped_session | Session) -> Account: + def update_account_email(account: Account, email: str, session: Session) -> Account: """Update account email""" account.email = email account_integrate = session.scalar( @@ -621,7 +620,7 @@ class AccountService: return account @staticmethod - def update_login_info(account: Account, session: scoped_session | Session, *, ip_address: str): + def update_login_info(account: Account, session: Session, *, ip_address: str): """Update last login time and ip""" account.last_login_at = naive_utc_now() account.last_login_ip = ip_address @@ -629,7 +628,7 @@ class AccountService: session.commit() @staticmethod - def login(account: Account, *, session: scoped_session | Session, ip_address: str | None = None) -> TokenPair: + def login(account: Account, *, session: Session, ip_address: str | None = None) -> TokenPair: if ip_address: AccountService.update_login_info(account=account, session=session, ip_address=ip_address) @@ -652,7 +651,7 @@ class AccountService: AccountService._delete_refresh_token(refresh_token.decode("utf-8"), account.id) @staticmethod - def refresh_token(refresh_token: str, *, session: scoped_session | Session) -> TokenPair: + def refresh_token(refresh_token: str, *, session: Session) -> TokenPair: # Verify the refresh token account_id = redis_client.get(AccountService._get_refresh_token_key(refresh_token)) if not account_id: @@ -673,7 +672,7 @@ class AccountService: return TokenPair(access_token=new_access_token, refresh_token=new_refresh_token, csrf_token=csrf_token) @staticmethod - def load_logged_in_account(*, account_id: str, session: scoped_session | Session): + def load_logged_in_account(*, account_id: str, session: Session): return AccountService.load_user(account_id, session) @classmethod @@ -1004,7 +1003,7 @@ class AccountService: return token @staticmethod - def get_account_by_email_with_case_fallback(session: Session | scoped_session, email: str) -> Account | None: + def get_account_by_email_with_case_fallback(email: str, *, session: Session) -> Account | None: """ Retrieve an account by email and fall back to the lowercase email if the original lookup fails. @@ -1026,7 +1025,7 @@ class AccountService: TokenManager.revoke_token(token, "email_code_login") @classmethod - def get_user_through_email(cls, email: str, *, session: scoped_session | Session): + def get_user_through_email(cls, email: str, *, session: Session): if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email): raise AccountRegisterError( description=( @@ -1234,7 +1233,7 @@ class AccountService: return False @staticmethod - def check_email_unique(email: str, *, session: scoped_session | Session) -> bool: + def check_email_unique(email: str, *, session: Session) -> bool: return session.scalar(select(Account).where(Account.email == email).limit(1)) is None @@ -1245,7 +1244,7 @@ class TenantService: is_setup: bool | None = False, is_from_dashboard: bool | None = False, *, - session: scoped_session | Session, + session: Session, ) -> Tenant: """Create tenant""" if ( @@ -1279,13 +1278,13 @@ class TenantService: from services.credit_pool_service import CreditPoolService - CreditPoolService.create_default_pool(tenant.id) + CreditPoolService.create_default_pool(tenant.id, session=session) return tenant @staticmethod def create_owner_tenant_if_not_exist( - account: Account, name: str | None = None, is_setup: bool | None = False, *, session: scoped_session | Session + account: Account, name: str | None = None, is_setup: bool | None = False, *, session: Session ): """Check if user have a workspace or not""" available_ta = session.scalar( @@ -1318,6 +1317,7 @@ class TenantService: account_id=account.id, member_account_id=account.id, role_ids=[owner_role_id], + session=session, ) account.current_tenant = tenant session.commit() @@ -1325,7 +1325,7 @@ class TenantService: @staticmethod def create_tenant_member( - tenant: Tenant, account: Account, session: scoped_session | Session, role: str = "normal" + tenant: Tenant, account: Account, session: Session, role: str = "normal" ) -> TenantAccountJoin: """Create tenant member""" if role == TenantAccountRole.OWNER: @@ -1350,7 +1350,7 @@ class TenantService: return ta @staticmethod - def get_join_tenants(account: Account, *, session: scoped_session | Session) -> list[Tenant]: + def get_join_tenants(account: Account, *, session: Session) -> list[Tenant]: """Get account join tenants""" return list( session.scalars( @@ -1361,10 +1361,7 @@ class TenantService: ) @staticmethod - def get_account_memberships( - session: Session | scoped_session, - account_id: str, - ) -> list[Row[tuple[TenantAccountJoin, Tenant]]]: + def get_account_memberships(account_id: str, *, session: Session) -> list[Row[tuple[TenantAccountJoin, Tenant]]]: """Return ``(TenantAccountJoin, Tenant)`` rows for every workspace the account belongs to. Unlike :meth:`get_join_tenants` this keeps the join row so callers can read ``role``/``current`` alongside the @@ -1385,10 +1382,7 @@ class TenantService: ) @staticmethod - def get_workspaces_for_account( - session: Session | scoped_session, - account_id: str, - ) -> list[Row[tuple[Tenant, TenantAccountJoin]]]: + def get_workspaces_for_account(account_id: str, *, session: Session) -> list[Row[tuple[Tenant, TenantAccountJoin]]]: """``(Tenant, TenantAccountJoin)`` rows for every workspace the account belongs to, ordered by ``Tenant.created_at`` ASC — the canonical ordering for ``/openapi/v1/workspaces``. @@ -1407,11 +1401,7 @@ class TenantService: ) @staticmethod - def account_belongs_to_tenant( - session: Session | scoped_session, - account_id: uuid.UUID | str | None, - tenant_id: str, - ) -> bool: + def account_belongs_to_tenant(account_id: uuid.UUID | str | None, tenant_id: str, *, session: Session) -> bool: """Existence check for ``TenantAccountJoin(account_id, tenant_id)``. Backs the CE-deployment membership fallback in ``controllers.openapi.auth.strategies.MembershipStrategy``. @@ -1431,9 +1421,7 @@ class TenantService: @staticmethod def get_account_role_in_tenant( - session: Session | scoped_session, - account_id: uuid.UUID | str | None, - tenant_id: str, + account_id: uuid.UUID | str | None, tenant_id: str, *, session: Session ) -> TenantAccountRole | None: """Return the caller's role in ``tenant_id``, or ``None`` if not a member. @@ -1459,7 +1447,7 @@ class TenantService: return TenantAccountRole(role) if role is not None else None @staticmethod - def get_tenant_by_id(session: Session | scoped_session, tenant_id: str) -> Tenant | None: + def get_tenant_by_id(tenant_id: str, *, session: Session) -> Tenant | None: """Plain ``session.get(Tenant, tenant_id)`` — no status filter. Callers map ``status == ARCHIVE`` to their own error code (the openapi auth pipeline raises 403 ``workspace unavailable``). @@ -1467,10 +1455,7 @@ class TenantService: return session.get(Tenant, tenant_id) @staticmethod - def get_tenants_by_ids( - session: Session | scoped_session, - tenant_ids: list[str], - ) -> list[Tenant]: + def get_tenants_by_ids(tenant_ids: list[str], *, session: Session) -> list[Tenant]: """Bulk ``Tenant`` fetch by primary-key list. Order is unspecified — callers index by ``tenant.id`` (e.g. for cross-tenant denorm in ``/openapi/v1/permitted-external-apps``). @@ -1483,7 +1468,7 @@ class TenantService: return list(session.execute(select(Tenant).where(Tenant.id.in_(tenant_ids))).scalars().all()) @staticmethod - def get_tenant_name(session: Session | scoped_session, tenant_id: str) -> str | None: + def get_tenant_name(tenant_id: str, *, session: Session) -> str | None: """Single-column tenant name read. Used by openapi list endpoints to denormalize ``workspace_name`` onto each row without dragging the full ``Tenant`` ORM entity through. @@ -1492,9 +1477,7 @@ class TenantService: @staticmethod def find_workspace_for_account( - session: Session | scoped_session, - account_id: str, - workspace_id: str, + account_id: str, workspace_id: str, *, session: Session ) -> Row[tuple[Tenant, TenantAccountJoin]] | None: """Single ``(Tenant, TenantAccountJoin)`` row scoped to the account's membership in ``workspace_id``. ``None`` on non-member @@ -1511,7 +1494,7 @@ class TenantService: ).first() @staticmethod - def get_current_tenant_by_account(account: Account, *, session: scoped_session | Session): + def get_current_tenant_by_account(account: Account, *, session: Session): """Get tenant by account and add the role""" tenant = account.current_tenant if not tenant: @@ -1529,7 +1512,7 @@ class TenantService: return tenant @staticmethod - def switch_tenant(account: Account, tenant_id: str | None = None, *, session: scoped_session | Session): + def switch_tenant(account: Account, tenant_id: str | None = None, *, session: Session): """Switch the current workspace for the account""" # Ensure tenant_id is provided @@ -1562,7 +1545,7 @@ class TenantService: session.commit() @staticmethod - def get_tenant_members(tenant: Tenant, *, session: scoped_session | Session) -> list[Account]: + def get_tenant_members(tenant: Tenant, *, session: Session) -> list[Account]: """Get tenant members""" stmt = ( select(Account, TenantAccountJoin.role) @@ -1581,7 +1564,7 @@ class TenantService: return updated_accounts @staticmethod - def get_dataset_operator_members(tenant: Tenant, *, session: scoped_session | Session) -> list[Account]: + def get_dataset_operator_members(tenant: Tenant, *, session: Session) -> list[Account]: """Get dataset admin members""" stmt = ( select(Account, TenantAccountJoin.role) @@ -1601,7 +1584,7 @@ class TenantService: return updated_accounts @staticmethod - def has_roles(tenant: Tenant, roles: list[TenantAccountRole], *, session: scoped_session | Session) -> bool: + def has_roles(tenant: Tenant, roles: list[TenantAccountRole], *, session: Session) -> bool: """Check if user has any of the given roles for a tenant""" if not all(isinstance(role, TenantAccountRole) for role in roles): raise ValueError("all roles must be TenantAccountRole") @@ -1619,9 +1602,7 @@ class TenantService: ) @staticmethod - def get_user_role( - account: Account, tenant: Tenant, *, session: scoped_session | Session - ) -> TenantAccountRole | None: + def get_user_role(account: Account, tenant: Tenant, *, session: Session) -> TenantAccountRole | None: """Get the role of the current account for a given tenant""" join = session.scalar( select(TenantAccountJoin) @@ -1631,13 +1612,13 @@ class TenantService: return TenantAccountRole(join.role) if join else None @staticmethod - def get_tenant_count(*, session: scoped_session | Session) -> int: + def get_tenant_count(*, session: Session) -> int: """Get tenant count""" return cast(int, session.scalar(select(func.count(Tenant.id)))) @staticmethod def check_member_permission( - tenant: Tenant, operator: Account, member: Account | None, action: str, *, session: scoped_session | Session + tenant: Tenant, operator: Account, member: Account | None, action: str, *, session: Session ): """Check member permission""" if action not in {"add", "remove", "update"}: @@ -1651,6 +1632,7 @@ class TenantService: workspace_permission_keys = AccountService.get_workspace_permission_keys( str(tenant.id), str(operator.id), + session=session, ) required_permission_key = ( "workspace.member.manage" if action in {"add", "remove"} else "workspace.role.manage" @@ -1661,7 +1643,9 @@ class TenantService: if ( action == "remove" and member - and AccountService.is_rbac_workspace_owner(str(tenant.id), str(operator.id), str(member.id)) + and AccountService.is_rbac_workspace_owner( + str(tenant.id), str(operator.id), str(member.id), session=session + ) ): raise NoPermissionError(f"No permission to {action} member.") return @@ -1691,9 +1675,7 @@ class TenantService: raise NoPermissionError(f"No permission to {action} member.") @staticmethod - def remove_member_from_tenant( - tenant: Tenant, account: Account, operator: Account, *, session: scoped_session | Session - ): + def remove_member_from_tenant(tenant: Tenant, account: Account, operator: Account, *, session: Session): """Remove member from tenant. Apps and datasets maintained by the removed member are reassigned to @@ -1722,7 +1704,9 @@ class TenantService: owner_id: str | None if dify_config.RBAC_ENABLED: - owner_id = AccountService.get_rbac_workspace_owner_account_id(str(tenant.id), str(operator.id)) + owner_id = AccountService.get_rbac_workspace_owner_account_id( + str(tenant.id), str(operator.id), session=session + ) else: owner_id = session.scalar( select(TenantAccountJoin.account_id) @@ -1796,9 +1780,7 @@ class TenantService: RBACService.MemberRoles.delete_rbac_bindings(tenant_id=tenant.id, account_id=account_id) @staticmethod - def update_member_role( - tenant: Tenant, member: Account, new_role: str, operator: Account, *, session: scoped_session | Session - ): + def update_member_role(tenant: Tenant, member: Account, new_role: str, operator: Account, *, session: Session): """Update member role""" TenantService.check_member_permission(tenant, operator, member, "update", session=session) new_tenant_role = TenantAccountRole(new_role) @@ -1841,6 +1823,7 @@ class TenantService: account_id=operator.id, member_account_id=str(current_owner_join.account_id), role_ids=[admin_role_id], + session=session, ) # Update the role of the target member @@ -1855,6 +1838,7 @@ class TenantService: account_id=operator.id, member_account_id=member.id, role_ids=[resolved_role_id], + session=session, ) else: target_member_join.role = new_tenant_role @@ -1867,11 +1851,11 @@ class TenantService: return tenant.custom_config_dict @staticmethod - def is_owner(account: Account, tenant: Tenant, *, session: scoped_session | Session) -> bool: + def is_owner(account: Account, tenant: Tenant, *, session: Session) -> bool: return TenantService.get_user_role(account, tenant, session=session) == TenantAccountRole.OWNER @staticmethod - def is_member(account: Account, tenant: Tenant, *, session: scoped_session | Session) -> bool: + def is_member(account: Account, tenant: Tenant, *, session: Session) -> bool: """Check if the account is a member of the tenant""" return TenantService.get_user_role(account, tenant, session=session) is not None @@ -1890,7 +1874,7 @@ class RegisterService: ip_address: str, language: str | None, *, - session: scoped_session | Session, + session: Session, ): """ Setup dify @@ -1943,7 +1927,7 @@ class RegisterService: create_workspace_required: bool | None = True, timezone: str | None = None, *, - session: scoped_session | Session, + session: Session, ) -> Account: """Register account""" session.begin_nested() @@ -2005,7 +1989,7 @@ class RegisterService: role: str = "normal", inviter: Account | None = None, *, - session: scoped_session | Session, + session: Session, ) -> str: if not inviter: raise ValueError("Inviter is required") @@ -2019,7 +2003,7 @@ class RegisterService: check_workspace_member_invite_permission(tenant.id) - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) requires_setup = False if not account: @@ -2057,6 +2041,7 @@ class RegisterService: account_id=inviter.id, member_account_id=account.id, role_ids=[role], + session=session, ) if ta or dify_config.RBAC_ENABLED: raise AccountAlreadyInTenantError("Account already in tenant.") @@ -2068,6 +2053,7 @@ class RegisterService: account_id=inviter.id, member_account_id=account.id, role_ids=[role], + session=session, ) token = cls.generate_invite_token(tenant, account, role, requires_setup=requires_setup) @@ -2116,7 +2102,7 @@ class RegisterService: @classmethod def get_invitation_if_token_valid( - cls, workspace_id: str | None, email: str | None, token: str, *, session: scoped_session | Session + cls, workspace_id: str | None, email: str | None, token: str, *, session: Session ) -> InvitationDetailDict | None: invitation_data = cls.get_invitation_by_token(token, workspace_id, email) if not invitation_data: @@ -2169,7 +2155,7 @@ class RegisterService: @classmethod def get_invitation_with_case_fallback( - cls, workspace_id: str | None, email: str | None, token: str, *, session: scoped_session | Session + cls, workspace_id: str | None, email: str | None, token: str, *, session: Session ) -> InvitationDetailDict | None: invitation = cls.get_invitation_if_token_valid(workspace_id, email, token, session=session) if invitation or not email or email == email.lower(): diff --git a/api/services/agent/composer_service.py b/api/services/agent/composer_service.py index 28b86916b1d..4839012dd3a 100644 --- a/api/services/agent/composer_service.py +++ b/api/services/agent/composer_service.py @@ -4,6 +4,7 @@ from typing import Any from sqlalchemy import func, or_, select from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session from sqlalchemy.sql.elements import ColumnElement from extensions.ext_database import db @@ -104,23 +105,35 @@ def _agent_soul_config_json(agent_soul: AgentSoulConfig | dict[str, Any]) -> dic class AgentComposerService: @classmethod def load_workflow_composer( - cls, *, tenant_id: str, app_id: str, node_id: str, account_id: str | None = None, snapshot_id: str | None = None + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + account_id: str | None = None, + snapshot_id: str | None = None, + session: Session, ) -> dict[str, Any]: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) if not binding: if snapshot_id: raise AgentVersionNotFoundError() return cls._empty_workflow_state(app_id=app_id, workflow_id=workflow.id, node_id=node_id) - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._workflow_composer_version( tenant_id=tenant_id, binding=binding, agent=agent, snapshot_id=snapshot_id, + session=session, + ) + return cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - return cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) @classmethod def _workflow_composer_version( @@ -130,6 +143,7 @@ class AgentComposerService: binding: WorkflowAgentNodeBinding, agent: Agent | None, snapshot_id: str | None, + session: Session, ) -> AgentConfigSnapshot | None: if snapshot_id: if agent is None: @@ -147,7 +161,7 @@ class AgentComposerService: raise AgentVersionNotFoundError() else: raise AgentVersionNotFoundError() - return cls._require_version(tenant_id=tenant_id, agent_id=agent.id, version_id=snapshot_id) + return cls._require_version(tenant_id=tenant_id, agent_id=agent.id, version_id=snapshot_id, session=session) version_id = ( agent.active_config_snapshot_id @@ -158,11 +172,19 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=version_id, + session=session, ) @classmethod def save_workflow_composer( - cls, *, tenant_id: str, app_id: str, node_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.WORKFLOW: raise ValueError("Workflow composer endpoint only accepts workflow variant") @@ -171,8 +193,10 @@ class AgentComposerService: _validate_composer_payload_for_strategy(payload) if payload.save_strategy in _PUBLISH_SAVE_STRATEGIES: cls.validate_knowledge_datasets(tenant_id=tenant_id, agent_soul=payload.agent_soul) - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) match payload.save_strategy: case ComposerSaveStrategy.NODE_JOB_ONLY: @@ -184,14 +208,15 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) case ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION: binding = cls._save_to_current_version( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) case ComposerSaveStrategy.SAVE_AS_NEW_VERSION: binding = cls._save_as_new_version( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) case ComposerSaveStrategy.SAVE_AS_NEW_AGENT: binding = cls._save_as_new_agent( @@ -202,14 +227,15 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) case ComposerSaveStrategy.SAVE_TO_ROSTER: binding = cls._save_to_roster( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session ) - db.session.commit() - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + session.commit() + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version_id = ( agent.active_config_snapshot_id if agent and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT @@ -219,12 +245,16 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=version_id, + session=session, + ) + state = cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - state = cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) state["validation"] = cls.collect_validation_findings( tenant_id=tenant_id, payload=payload, agent_id=binding.agent_id, + session=session, ) return state @@ -239,33 +269,38 @@ class AgentComposerService: source_agent_id: str, source_snapshot_id: str | None = None, idempotency_key: str | None = None, + session: Session, ) -> dict[str, Any]: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) binding = cls._require_binding( - cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session) ) if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and idempotency_key: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._get_version_if_present( tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=binding.current_snapshot_id, + session=session, + ) + return cls._serialize_workflow_state( + binding=binding, agent=agent, version=version, account_id=account_id, session=session ) - return cls._serialize_workflow_state(binding=binding, agent=agent, version=version, account_id=account_id) if binding.binding_type != WorkflowAgentBindingType.ROSTER_AGENT: raise InvalidComposerConfigError("Workflow agent node must be bound to a roster agent.") if binding.agent_id != source_agent_id: raise InvalidComposerConfigError("Source agent does not match the current workflow node binding.") - source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=source_agent_id) + source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=source_agent_id, session=session) if source_agent.scope != AgentScope.ROSTER or source_agent.status != AgentStatus.ACTIVE: raise InvalidComposerConfigError("Source agent must be an active roster agent.") source_version = cls._require_version( tenant_id=tenant_id, agent_id=source_agent.id, version_id=source_agent.active_config_snapshot_id, + session=session, ) if source_snapshot_id and source_snapshot_id != source_version.id: raise AgentVersionConflictError() @@ -284,6 +319,7 @@ class AgentComposerService: icon_type=source_agent.icon_type, icon=source_agent.icon, icon_background=source_agent.icon_background, + session=session, ) cls._copy_agent_drive_rows( tenant_id=tenant_id, @@ -292,45 +328,48 @@ class AgentComposerService: account_id=account_id, agent_soul=agent_soul, node_job=WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict), + session=session, ) binding.binding_type = WorkflowAgentBindingType.INLINE_AGENT binding.agent_id = inline_agent.id binding.current_snapshot_id = inline_agent.active_config_snapshot_id binding.updated_by = account_id - db.session.flush() - db.session.commit() + session.flush() + session.commit() version = cls._require_version( tenant_id=tenant_id, agent_id=inline_agent.id, version_id=inline_agent.active_config_snapshot_id, + session=session, ) return cls._serialize_workflow_state( - binding=binding, agent=inline_agent, version=version, account_id=account_id + binding=binding, agent=inline_agent, version=version, account_id=account_id, session=session ) @classmethod - def load_agent_app_composer(cls, *, tenant_id: str, app_id: str) -> dict[str, Any]: - agent = cls._require_agent_app_agent(tenant_id=tenant_id, app_id=app_id) - return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent) + def load_agent_app_composer(cls, *, tenant_id: str, app_id: str, session: Session) -> dict[str, Any]: + agent = cls._require_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) + return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent, session=session) @classmethod - def load_agent_composer(cls, *, tenant_id: str, agent_id: str) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) - return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent) + def load_agent_composer(cls, *, tenant_id: str, agent_id: str, session: Session) -> dict[str, Any]: + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) + return cls._load_agent_composer_for_agent(tenant_id=tenant_id, agent=agent, session=session) @classmethod - def _load_agent_composer_for_agent(cls, *, tenant_id: str, agent: Agent) -> dict[str, Any]: + def _load_agent_composer_for_agent(cls, *, tenant_id: str, agent: Agent, session: Session) -> dict[str, Any]: draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=agent.updated_by or agent.created_by, + session=session, ) version = cls._get_version_if_present( - tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id + tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, session=session ) return { "variant": ComposerVariant.AGENT_APP.value, @@ -347,7 +386,13 @@ class AgentComposerService: @classmethod def save_agent_app_composer( - cls, *, tenant_id: str, app_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + app_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent App composer endpoint only accepts agent_app variant") @@ -360,7 +405,7 @@ class AgentComposerService: _backfill_cli_tool_ids(payload.agent_soul) _validate_composer_payload_for_strategy(payload) - agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id) + agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) if not agent: agent = Agent( tenant_id=tenant_id, @@ -375,22 +420,29 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(agent) + session.add(agent) try: - db.session.flush() + session.flush() except IntegrityError as exc: - db.session.rollback() + session.rollback() raise AgentNameConflictError() from exc return cls._save_agent_composer_for_agent( tenant_id=tenant_id, agent=agent, account_id=account_id, payload=payload, + session=session, ) @classmethod def save_agent_composer( - cls, *, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent composer endpoint only accepts agent_app variant") @@ -402,17 +454,24 @@ class AgentComposerService: raise ValueError("agent_soul is required") _backfill_cli_tool_ids(payload.agent_soul) _validate_composer_payload_for_strategy(payload) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) return cls._save_agent_composer_for_agent( tenant_id=tenant_id, agent=agent, account_id=account_id, payload=payload, + session=session, ) @classmethod def _save_agent_composer_for_agent( - cls, *, tenant_id: str, agent: Agent, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent: Agent, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") @@ -423,20 +482,23 @@ class AgentComposerService: account_id=None, agent_soul=payload.agent_soul, account_id_for_audit=account_id, + session=session, ) agent.updated_by = account_id agent.active_config_is_published = cls._agent_soul_matches_active_config( tenant_id=tenant_id, agent=agent, agent_soul=payload.agent_soul, + session=session, ) - db.session.commit() - state = cls.load_agent_composer(tenant_id=tenant_id, agent_id=agent.id) + session.commit() + state = cls.load_agent_composer(tenant_id=tenant_id, agent_id=agent.id, session=session) state["validation"] = cls.collect_validation_findings( tenant_id=tenant_id, payload=payload, agent_id=agent.id, + session=session, ) return state @@ -447,6 +509,7 @@ class AgentComposerService: tenant_id: str, agent: Agent, agent_soul: AgentSoulConfig, + session: Session, ) -> bool: if not agent.active_config_snapshot_id: return False @@ -455,6 +518,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) if not active_version: return False @@ -490,9 +554,15 @@ class AgentComposerService: @classmethod def publish_agent_app_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, version_note: str | None = None + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + version_note: str | None = None, + session: Session, ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) if agent.scope != AgentScope.ROSTER or agent.source != AgentSource.AGENT_APP: raise AgentNotFoundError() draft = cls._get_or_create_agent_draft( @@ -501,6 +571,7 @@ class AgentComposerService: draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=account_id, + session=session, ) agent_soul = AgentSoulConfig.model_validate(draft.config_snapshot_dict) ComposerConfigValidator.validate_publish_payload( @@ -522,6 +593,7 @@ class AgentComposerService: operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, version_note=version_note, previous_snapshot_id=agent.active_config_snapshot_id, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -529,7 +601,7 @@ class AgentComposerService: agent.updated_by = account_id draft.base_snapshot_id = version.id draft.updated_by = account_id - db.session.commit() + session.commit() return { "result": "success", "active_config_snapshot_id": version.id, @@ -539,21 +611,29 @@ class AgentComposerService: @classmethod def checkout_agent_app_build_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, force: bool = False + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + force: bool = False, + session: Session, ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) normal_draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=account_id, + session=session, ) build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is not None and not force: return cls._serialize_build_draft_state(build_draft) @@ -566,20 +646,23 @@ class AgentComposerService: draft_owner_key=account_id, created_by=account_id, ) - db.session.add(build_draft) + session.add(build_draft) build_draft.base_snapshot_id = normal_draft.base_snapshot_id build_draft.config_snapshot = AgentSoulConfig.model_validate(normal_draft.config_snapshot_dict) build_draft.updated_by = account_id - db.session.commit() + session.commit() return cls._serialize_build_draft_state(build_draft) @classmethod - def load_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: + def load_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent_id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is None: raise AgentVersionNotFoundError() @@ -587,13 +670,19 @@ class AgentComposerService: @classmethod def save_agent_app_build_draft( - cls, *, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + tenant_id: str, + agent_id: str, + account_id: str, + payload: ComposerSavePayload, + session: Session, ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") _backfill_cli_tool_ids(payload.agent_soul) ComposerConfigValidator.validate_draft_save_payload(payload) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) build_draft = cls._save_agent_draft( tenant_id=tenant_id, agent=agent, @@ -601,18 +690,22 @@ class AgentComposerService: account_id=account_id, agent_soul=payload.agent_soul, account_id_for_audit=account_id, + session=session, ) - db.session.commit() + session.commit() return cls._serialize_build_draft_state(build_draft) @classmethod - def apply_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id) + def apply_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: + agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is None: raise AgentVersionNotFoundError() @@ -625,28 +718,33 @@ class AgentComposerService: agent_soul=applied_agent_soul, account_id_for_audit=account_id, base_snapshot_id=build_draft.base_snapshot_id, + session=session, ) agent.active_config_is_published = cls._agent_soul_matches_active_config( tenant_id=tenant_id, agent=agent, agent_soul=applied_agent_soul, + session=session, ) agent.updated_by = account_id - db.session.delete(build_draft) - db.session.commit() + session.delete(build_draft) + session.commit() return {"result": "success", "draft": cls._serialize_draft(normal_draft)} @classmethod - def discard_agent_app_build_draft(cls, *, tenant_id: str, agent_id: str, account_id: str) -> dict[str, Any]: + def discard_agent_app_build_draft( + cls, *, tenant_id: str, agent_id: str, account_id: str, session: Session + ) -> dict[str, Any]: build_draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent_id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, + session=session, ) if build_draft is not None: - db.session.delete(build_draft) - db.session.commit() + session.delete(build_draft) + session.commit() return {"result": "success"} @classmethod @@ -656,6 +754,7 @@ class AgentComposerService: tenant_id: str, payload: ComposerSavePayload, agent_id: str | None = None, + session: Session, ) -> dict[str, Any]: """ENG-617 soft findings, with DB-backed dataset and drive mention checks.""" existing_knowledge_set_ids = ( @@ -673,6 +772,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent_id, prompt=payload.agent_soul.prompt.system_prompt, + session=session, ) ) return findings @@ -696,9 +796,9 @@ class AgentComposerService: ) @classmethod - def resolve_bound_agent_id(cls, *, tenant_id: str, app_id: str) -> str | None: + def resolve_bound_agent_id(cls, *, tenant_id: str, app_id: str, session: Session) -> str | None: """The Agent App's bound roster agent id, if any (validate-endpoint context).""" - return db.session.scalar( + return session.scalar( select(Agent.id) .where( Agent.tenant_id == tenant_id, @@ -711,13 +811,17 @@ class AgentComposerService: ) @classmethod - def resolve_workflow_node_agent_id(cls, *, tenant_id: str, app_id: str, node_id: str) -> str | None: + def resolve_workflow_node_agent_id( + cls, *, tenant_id: str, app_id: str, node_id: str, session: Session + ) -> str | None: """The draft workflow node binding's agent id, if any (validate-endpoint context).""" try: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) except ValueError: return None - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) return binding.agent_id if binding else None @classmethod @@ -727,6 +831,7 @@ class AgentComposerService: tenant_id: str, agent_id: str, prompt: str, + session: Session, ) -> list[dict[str, str | None]]: """Soft warnings for missing drive-backed prompt mentions.""" from services.agent.prompt_mentions import MentionKind, parse_prompt_mentions @@ -744,7 +849,7 @@ class AgentComposerService: return [] existing_keys = set( - db.session.scalars( + session.scalars( select(AgentDriveFile.key).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id, @@ -768,22 +873,32 @@ class AgentComposerService: return findings @classmethod - def get_workflow_candidates(cls, *, tenant_id: str, app_id: str, node_id: str, user_id: str) -> dict[str, Any]: + def get_workflow_candidates( + cls, + *, + tenant_id: str, + app_id: str, + node_id: str, + user_id: str, + session: Session, + ) -> dict[str, Any]: """Slash-menu data source for the workflow Agent node composer (ENG-615).""" from services.agent.composer_candidates import previous_node_output_candidates, soul_candidates try: - workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id) + workflow = cls._get_draft_workflow(tenant_id=tenant_id, app_id=app_id, session=session) except ValueError: workflow = None node_job: WorkflowNodeJobConfig | None = None agent_soul: AgentSoulConfig | None = None if workflow is not None: - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + ) if binding is not None: node_job = cls._parse_node_job(binding) - agent_soul = cls._load_binding_soul(tenant_id=tenant_id, binding=binding) + agent_soul = cls._load_binding_soul(tenant_id=tenant_id, binding=binding, session=session) truncated = False previous_outputs: list[dict[str, Any]] = [] @@ -794,7 +909,7 @@ class AgentComposerService: graph=workflow.graph_dict, node_id=node_id, declared_outputs_loader=lambda nid: cls._binding_declared_outputs( - tenant_id=tenant_id, workflow_id=workflow.id, node_id=nid + tenant_id=tenant_id, workflow_id=workflow.id, node_id=nid, session=session ), draft_variables_loader=lambda nid: cls._draft_node_variables( session=draft_variable_session, app_id=app_id, node_id=nid, user_id=user_id @@ -829,11 +944,13 @@ class AgentComposerService: return response.model_dump(mode="json") @classmethod - def get_agent_app_candidates(cls, *, tenant_id: str, agent_id: str, user_id: str) -> dict[str, Any]: + def get_agent_app_candidates( + cls, *, tenant_id: str, agent_id: str, user_id: str, session: Session + ) -> dict[str, Any]: """Slash-menu data source for the Agent App (Console) composer (ENG-615).""" from services.agent.composer_candidates import soul_candidates - agent_soul = cls._load_agent_soul(tenant_id=tenant_id, agent_id=agent_id) + agent_soul = cls._load_agent_soul(tenant_id=tenant_id, agent_id=agent_id, session=session) soul_lists, truncated = soul_candidates( agent_soul=agent_soul, dataset_lookup=lambda ids: get_tenant_knowledge_dataset_rows(tenant_id=tenant_id, dataset_ids=ids), @@ -858,18 +975,21 @@ class AgentComposerService: return None @classmethod - def _load_binding_soul(cls, *, tenant_id: str, binding: WorkflowAgentNodeBinding) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id) + def _load_binding_soul( + cls, *, tenant_id: str, binding: WorkflowAgentNodeBinding, session: Session + ) -> AgentSoulConfig | None: + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) version = cls._get_version_if_present( tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=binding.current_snapshot_id, + session=session, ) return cls._parse_soul_snapshot(version) @classmethod - def _load_agent_soul(cls, *, tenant_id: str, agent_id: str) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=agent_id) + def _load_agent_soul(cls, *, tenant_id: str, agent_id: str, session: Session) -> AgentSoulConfig | None: + agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=agent_id, session=session) if agent is None: return None draft = cls._get_or_create_agent_draft( @@ -878,6 +998,7 @@ class AgentComposerService: draft_type=AgentConfigDraftType.DRAFT, account_id=None, created_by=agent.updated_by or agent.created_by, + session=session, ) return AgentSoulConfig.model_validate(draft.config_snapshot_dict) @@ -893,9 +1014,11 @@ class AgentComposerService: @classmethod def _binding_declared_outputs( - cls, *, tenant_id: str, workflow_id: str, node_id: str + cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session ) -> list[DeclaredOutputConfig] | None: - binding = cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id) + binding = cls._get_workflow_binding( + tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id, session=session + ) if binding is None: return None node_job = cls._parse_node_job(binding) @@ -970,8 +1093,8 @@ class AgentComposerService: return tools @classmethod - def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str) -> dict[str, Any]: - snapshot = db.session.scalar( + def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str, session: Session) -> dict[str, Any]: + snapshot = session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -987,7 +1110,7 @@ class AgentComposerService: & (WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT) ) bindings = list( - db.session.scalars( + session.scalars( select(WorkflowAgentNodeBinding).where( WorkflowAgentNodeBinding.tenant_id == tenant_id, or_(*predicates), @@ -1018,6 +1141,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: node_job = payload.node_job or WorkflowNodeJobConfig() if binding: @@ -1030,6 +1154,7 @@ class AgentComposerService: account_id=account_id, binding=binding, payload=payload, + session=session, ) binding.node_job_config = node_job if payload.agent_soul is not None and binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT: @@ -1037,6 +1162,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1044,8 +1170,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) if agent.scope != AgentScope.WORKFLOW_ONLY: raise ValueError("Inline workflow agent binding must point to a workflow-only agent") agent.active_config_snapshot_id = version.id @@ -1064,6 +1191,7 @@ class AgentComposerService: node_id=node_id, account_id=account_id, agent_soul=agent_soul, + session=session, ) binding = WorkflowAgentNodeBinding( tenant_id=tenant_id, @@ -1078,8 +1206,8 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(binding) - db.session.flush() + session.add(binding) + session.flush() return binding @classmethod @@ -1101,6 +1229,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: if payload.binding and (payload.binding.agent_id or payload.binding.current_snapshot_id): raise ValueError("Start from Scratch must not provide an existing inline agent binding.") @@ -1113,13 +1242,14 @@ class AgentComposerService: node_id=node_id, account_id=account_id, agent_soul=agent_soul, + session=session, ) binding.binding_type = WorkflowAgentBindingType.INLINE_AGENT binding.agent_id = agent.id binding.current_snapshot_id = agent.active_config_snapshot_id binding.node_job_config = payload.node_job or binding.node_job_config binding.updated_by = account_id - db.session.flush() + session.flush() return binding @classmethod @@ -1130,6 +1260,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) if payload.agent_soul is None: @@ -1138,6 +1269,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1145,8 +1277,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(payload.agent_soul) agent.active_config_is_published = True @@ -1165,6 +1298,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) if not binding.agent_id or payload.agent_soul is None: @@ -1176,8 +1310,9 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note=payload.version_note, + session=session, ) - agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(payload.agent_soul) agent.active_config_is_published = True @@ -1199,6 +1334,7 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: if payload.agent_soul is None: raise ValueError("agent_soul is required") @@ -1215,6 +1351,7 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_AGENT, version_note=payload.version_note, + session=session, ) node_job = payload.node_job or WorkflowNodeJobConfig() if not binding: @@ -1226,13 +1363,13 @@ class AgentComposerService: node_id=node_id, created_by=account_id, ) - db.session.add(binding) + session.add(binding) binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT binding.agent_id = agent.id binding.current_snapshot_id = agent.active_config_snapshot_id binding.node_job_config = node_job binding.updated_by = account_id - db.session.flush() + session.flush() return binding @classmethod @@ -1243,13 +1380,15 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, + session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) - source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id) + source_agent = cls._require_agent(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) source_version = cls._require_version( tenant_id=tenant_id, agent_id=source_agent.id, version_id=binding.current_snapshot_id, + session=session, ) agent_soul = payload.agent_soul or AgentSoulConfig.model_validate(source_version.config_snapshot_dict) agent_name = payload.new_agent_name or source_agent.name @@ -1267,6 +1406,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_TO_ROSTER, version_note=payload.version_note, + session=session, ) cls._copy_agent_drive_rows( tenant_id=tenant_id, @@ -1275,6 +1415,7 @@ class AgentComposerService: account_id=account_id, agent_soul=agent_soul, node_job=payload.node_job or WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict), + session=session, ) binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT binding.agent_id = roster_agent.id @@ -1300,8 +1441,9 @@ class AgentComposerService: icon_type: Any | None = None, icon: str | None = None, icon_background: str | None = None, + session: Session, ) -> Agent: - backing_app = AgentRosterService(db.session).create_hidden_backing_app_for_workflow_agent( + backing_app = AgentRosterService(session).create_hidden_backing_app_for_workflow_agent( tenant_id=tenant_id, account_id=account_id, name=name or f"Workflow Agent {node_id}", @@ -1329,8 +1471,8 @@ class AgentComposerService: created_by=account_id, updated_by=account_id, ) - db.session.add(agent) - db.session.flush() + session.add(agent) + session.flush() version = cls._create_config_version( tenant_id=tenant_id, agent_id=agent.id, @@ -1338,6 +1480,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -1354,6 +1497,7 @@ class AgentComposerService: account_id: str, agent_soul: AgentSoulConfig, node_job: WorkflowNodeJobConfig | None = None, + session: Session, ) -> None: exact_keys, prefixes = cls._drive_copy_scopes_from_agent_configs(agent_soul=agent_soul, node_job=node_job) predicates: list[ColumnElement[bool]] = [] @@ -1364,7 +1508,7 @@ class AgentComposerService: return source_rows = list( - db.session.scalars( + session.scalars( select(AgentDriveFile).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == source_agent_id, @@ -1376,7 +1520,7 @@ class AgentComposerService: return existing_target_keys = set( - db.session.scalars( + session.scalars( select(AgentDriveFile.key).where( AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == target_agent_id, @@ -1387,7 +1531,7 @@ class AgentComposerService: for row in source_rows: if row.key in existing_target_keys: continue - db.session.add( + session.add( AgentDriveFile( tenant_id=tenant_id, agent_id=target_agent_id, @@ -1451,8 +1595,9 @@ class AgentComposerService: icon_type: AgentIconType | None = None, icon: str | None = None, icon_background: str | None = None, + session: Session, ) -> Agent: - account = cls._require_account(account_id=account_id) + account = cls._require_account(account_id=account_id, session=session) try: app = AppService().create_app( tenant_id, @@ -1466,12 +1611,13 @@ class AgentComposerService: icon_background=icon_background, ), account, + session=session, ) except IntegrityError as exc: - db.session.rollback() + session.rollback() raise AgentNameConflictError() from exc - agent = AgentRosterService(db.session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id) + agent = AgentRosterService(session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id) if agent is None: raise AgentNotFoundError() @@ -1479,6 +1625,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) version = cls._update_current_version( current_snapshot=current_snapshot, @@ -1486,6 +1633,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=operation, version_note=version_note, + session=session, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -1504,9 +1652,10 @@ class AgentComposerService: operation: AgentConfigRevisionOperation, version_note: str | None, previous_snapshot_id: str | None = None, + session: Session, ) -> AgentConfigSnapshot: next_version = ( - db.session.scalar( + session.scalar( select(func.max(AgentConfigSnapshot.version)).where( AgentConfigSnapshot.tenant_id == tenant_id, AgentConfigSnapshot.agent_id == agent_id, @@ -1522,20 +1671,20 @@ class AgentComposerService: version_note=version_note, created_by=account_id, ) - db.session.add(version) - db.session.flush() + session.add(version) + session.flush() revision = AgentConfigRevision( tenant_id=tenant_id, agent_id=agent_id, previous_snapshot_id=previous_snapshot_id, current_snapshot_id=version.id, - revision=cls._next_revision(tenant_id=tenant_id, agent_id=agent_id), + revision=cls._next_revision(tenant_id=tenant_id, agent_id=agent_id, session=session), operation=operation, version_note=version_note, created_by=account_id, ) - db.session.add(revision) - db.session.flush() + session.add(revision) + session.flush() return version @classmethod @@ -1547,6 +1696,7 @@ class AgentComposerService: agent_soul: AgentSoulConfig, operation: AgentConfigRevisionOperation, version_note: str | None, + session: Session, ) -> AgentConfigSnapshot: return cls._create_config_version( tenant_id=current_snapshot.tenant_id, @@ -1556,12 +1706,13 @@ class AgentComposerService: operation=operation, version_note=version_note, previous_snapshot_id=current_snapshot.id, + session=session, ) @classmethod - def _next_revision(cls, *, tenant_id: str, agent_id: str) -> int: + def _next_revision(cls, *, tenant_id: str, agent_id: str, session: Session) -> int: return ( - db.session.scalar( + session.scalar( select(func.max(AgentConfigRevision.revision)).where( AgentConfigRevision.tenant_id == tenant_id, AgentConfigRevision.agent_id == agent_id, @@ -1571,8 +1722,8 @@ class AgentComposerService: ) + 1 @classmethod - def _get_agent_app_agent(cls, *, tenant_id: str, app_id: str) -> Agent | None: - return db.session.scalar( + def _get_agent_app_agent(cls, *, tenant_id: str, app_id: str, session: Session) -> Agent | None: + return session.scalar( select(Agent) .where( Agent.tenant_id == tenant_id, @@ -1586,8 +1737,8 @@ class AgentComposerService: ) @classmethod - def _require_agent_app_agent(cls, *, tenant_id: str, app_id: str) -> Agent: - agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id) + def _require_agent_app_agent(cls, *, tenant_id: str, app_id: str, session: Session) -> Agent: + agent = cls._get_agent_app_agent(tenant_id=tenant_id, app_id=app_id, session=session) if agent is None: raise AgentNotFoundError() return agent @@ -1600,6 +1751,7 @@ class AgentComposerService: agent_id: str, draft_type: AgentConfigDraftType, account_id: str | None, + session: Session, ) -> AgentConfigDraft | None: stmt = select(AgentConfigDraft).where( AgentConfigDraft.tenant_id == tenant_id, @@ -1610,7 +1762,7 @@ class AgentComposerService: stmt = stmt.where(AgentConfigDraft.account_id == account_id) else: stmt = stmt.where(AgentConfigDraft.account_id.is_(None)) - return db.session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) + return session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) @classmethod def _get_or_create_agent_draft( @@ -1621,12 +1773,14 @@ class AgentComposerService: draft_type: AgentConfigDraftType, account_id: str | None, created_by: str | None, + session: Session, ) -> AgentConfigDraft: draft = cls._get_agent_draft( tenant_id=tenant_id, agent_id=agent.id, draft_type=draft_type, account_id=account_id, + session=session, ) if draft is not None: return draft @@ -1634,6 +1788,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, + session=session, ) agent_soul = ( AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict) @@ -1651,8 +1806,8 @@ class AgentComposerService: created_by=created_by, updated_by=created_by, ) - db.session.add(draft) - db.session.flush() + session.add(draft) + session.flush() return draft @classmethod @@ -1666,6 +1821,7 @@ class AgentComposerService: agent_soul: AgentSoulConfig, account_id_for_audit: str, base_snapshot_id: str | None = None, + session: Session, ) -> AgentConfigDraft: draft = cls._get_or_create_agent_draft( tenant_id=tenant_id, @@ -1673,6 +1829,7 @@ class AgentComposerService: draft_type=draft_type, account_id=account_id, created_by=account_id_for_audit, + session=session, ) draft.config_snapshot = agent_soul if base_snapshot_id is not None: @@ -1682,7 +1839,7 @@ class AgentComposerService: draft.updated_by = account_id_for_audit if draft_type == AgentConfigDraftType.DRAFT and account_id is None: agent.active_config_is_published = False - db.session.flush() + session.flush() return draft @classmethod @@ -1710,8 +1867,8 @@ class AgentComposerService: } @classmethod - def _get_draft_workflow(cls, *, tenant_id: str, app_id: str) -> Workflow: - workflow = db.session.scalar( + def _get_draft_workflow(cls, *, tenant_id: str, app_id: str, session: Session) -> Workflow: + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == tenant_id, @@ -1726,13 +1883,13 @@ class AgentComposerService: @classmethod def _get_workflow_binding( - cls, *, tenant_id: str, workflow_id: str, node_id: str + cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session ) -> WorkflowAgentNodeBinding | None: # Composer always operates against the draft workflow row, so this lookup # is scoped to ``workflow_version="draft"``. Published bindings are # materialized by WorkflowAgentPublishService.copy_agent_node_bindings_to_published # and are not edited through the Composer. - return db.session.scalar( + return session.scalar( select(WorkflowAgentNodeBinding) .where( WorkflowAgentNodeBinding.tenant_id == tenant_id, @@ -1750,32 +1907,34 @@ class AgentComposerService: return binding @classmethod - def _require_agent(cls, *, tenant_id: str, agent_id: str | None) -> Agent: + def _require_agent(cls, *, tenant_id: str, agent_id: str | None, session: Session) -> Agent: if not agent_id: raise AgentNotFoundError() - agent = db.session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) + agent = session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) if not agent: raise AgentNotFoundError() return agent @classmethod - def _require_account(cls, *, account_id: str) -> Account: - account = db.session.get(Account, account_id) + def _require_account(cls, *, account_id: str, session: Session) -> Account: + account = session.get(Account, account_id) if not account: raise ValueError("Account not found") return account @classmethod - def _get_agent_if_present(cls, *, tenant_id: str, agent_id: str | None) -> Agent | None: + def _get_agent_if_present(cls, *, tenant_id: str, agent_id: str | None, session: Session) -> Agent | None: if not agent_id: return None - return db.session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) + return session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) @classmethod - def _require_version(cls, *, tenant_id: str, agent_id: str | None, version_id: str | None) -> AgentConfigSnapshot: + def _require_version( + cls, *, tenant_id: str, agent_id: str | None, version_id: str | None, session: Session + ) -> AgentConfigSnapshot: if not agent_id or not version_id: raise AgentVersionNotFoundError() - version = db.session.scalar( + version = session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -1790,11 +1949,11 @@ class AgentComposerService: @classmethod def _get_version_if_present( - cls, *, tenant_id: str, agent_id: str | None, version_id: str | None + cls, *, tenant_id: str, agent_id: str | None, version_id: str | None, session: Session ) -> AgentConfigSnapshot | None: if not agent_id or not version_id: return None - return db.session.scalar( + return session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == tenant_id, @@ -1853,6 +2012,7 @@ class AgentComposerService: agent: Agent | None, version: AgentConfigSnapshot | None, account_id: str | None = None, + session: Session, ) -> dict[str, Any]: locked = bool(agent and agent.scope == AgentScope.ROSTER) save_options = [ComposerSaveStrategy.NODE_JOB_ONLY.value] @@ -1871,9 +2031,10 @@ class AgentComposerService: binding=binding, agent=agent, account_id=account_id, + session=session, ) debug_conversation_message_count = ( - AgentRosterService(db.session).count_agent_app_debug_conversation_messages( + AgentRosterService(session).count_agent_app_debug_conversation_messages( conversation_id=debug_conversation_id ) if debug_conversation_id @@ -1906,7 +2067,9 @@ class AgentComposerService: # this is the same list (so callers don't need to special-case). "effective_declared_outputs": cls._serialize_effective_outputs(cls._declared_outputs_from_binding(binding)), "save_options": save_options, - "impact_summary": cls.calculate_impact(tenant_id=binding.tenant_id, current_snapshot_id=version.id) + "impact_summary": cls.calculate_impact( + tenant_id=binding.tenant_id, current_snapshot_id=version.id, session=session + ) if version else None, "app_id": binding.app_id, @@ -1927,6 +2090,7 @@ class AgentComposerService: binding: WorkflowAgentNodeBinding, agent: Agent | None, account_id: str | None, + session: Session, ) -> str | None: if ( not account_id @@ -1938,7 +2102,7 @@ class AgentComposerService: from services.agent.roster_service import AgentRosterService - return AgentRosterService(db.session).get_or_create_agent_app_debug_conversation_id( + return AgentRosterService(session).get_or_create_agent_app_debug_conversation_id( tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, diff --git a/api/services/agent/roster_service.py b/api/services/agent/roster_service.py index a34fca67105..de76b3c4eb2 100644 --- a/api/services/agent/roster_service.py +++ b/api/services/agent/roster_service.py @@ -826,6 +826,7 @@ class AgentRosterService: max_active_requests=source_app.max_active_requests, ), account, + session=self._session, ) target_app.enable_site = source_app.enable_site diff --git a/api/services/agent/skill_standardize_service.py b/api/services/agent/skill_standardize_service.py index cc2ba4b9bdc..2639f7a9a18 100644 --- a/api/services/agent/skill_standardize_service.py +++ b/api/services/agent/skill_standardize_service.py @@ -18,6 +18,8 @@ from __future__ import annotations import re from typing import Any +from sqlalchemy.orm import Session + from core.tools.tool_file_manager import ToolFileManager from services.agent.skill_package_service import SkillPackageService from services.agent_drive_service import AgentDriveService, DriveCommitItem, DriveFileRef, DriveSkillMetadata @@ -59,6 +61,7 @@ class SkillStandardizeService: tenant_id: str, user_id: str, agent_id: str, + session: Session, ) -> dict[str, Any]: """Create two ToolFiles, commit two drive-owned keys, and return skill metadata. @@ -113,6 +116,7 @@ class SkillStandardizeService: value_owned_by_drive=True, ), ], + session=session, ) self.last_committed_items = committed_items diff --git a/api/services/agent/skill_tool_inference_service.py b/api/services/agent/skill_tool_inference_service.py index a6d5e6b2de9..7ce53dd4666 100644 --- a/api/services/agent/skill_tool_inference_service.py +++ b/api/services/agent/skill_tool_inference_service.py @@ -19,6 +19,7 @@ from typing import Any import json_repair from pydantic import BaseModel, Field, ValidationError +from sqlalchemy.orm import Session from core.errors.error import ProviderTokenNotInitError from core.model_manager import ModelManager @@ -91,8 +92,8 @@ class SkillToolInferenceService: def __init__(self, *, drive_service: AgentDriveService | None = None) -> None: self._drive = drive_service or AgentDriveService() - def infer(self, *, tenant_id: str, agent_id: str, slug: str) -> dict[str, Any]: - skill_md = self._load_skill_md(tenant_id=tenant_id, agent_id=agent_id, slug=slug) + def infer(self, *, tenant_id: str, agent_id: str, slug: str, session: Session) -> dict[str, Any]: + skill_md = self._load_skill_md(tenant_id=tenant_id, agent_id=agent_id, slug=slug, session=session) user_prompt = f"SKILL.md of skill '{slug}':\n\n{skill_md}" @@ -115,9 +116,11 @@ class SkillToolInferenceService: tool.inferred_from = slug return result.model_dump(mode="json") - def _load_skill_md(self, *, tenant_id: str, agent_id: str, slug: str) -> str: + def _load_skill_md(self, *, tenant_id: str, agent_id: str, slug: str, session: Session) -> str: try: - preview = self._drive.preview(tenant_id=tenant_id, agent_id=agent_id, key=f"{slug}/SKILL.md") + preview = self._drive.preview( + tenant_id=tenant_id, agent_id=agent_id, key=f"{slug}/SKILL.md", session=session + ) except AgentDriveError as exc: if exc.code == "drive_key_not_found": raise SkillToolInferenceError( diff --git a/api/services/agent_app_feature_service.py b/api/services/agent_app_feature_service.py index 5fd794bb10f..d336cdf29de 100644 --- a/api/services/agent_app_feature_service.py +++ b/api/services/agent_app_feature_service.py @@ -13,7 +13,7 @@ from __future__ import annotations from typing import Any, cast -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager from core.app.app_config.features.opening_statement.manager import OpeningStatementConfigManager @@ -74,7 +74,7 @@ class AgentAppFeatureConfigService: app_model: App, account: Account, config: dict[str, Any], - session: scoped_session, + session: Session, ) -> AppModelConfig: """Persist the presentation features as a new app_model_config version. diff --git a/api/services/agent_app_sandbox_service.py b/api/services/agent_app_sandbox_service.py index b1652d628e9..3f5a0bf41b2 100644 --- a/api/services/agent_app_sandbox_service.py +++ b/api/services/agent_app_sandbox_service.py @@ -19,12 +19,12 @@ from dify_agent.client import Client from dify_agent.protocol import RuntimeLayerSpec, SandboxLocator, build_sandbox_locator_from_layer_specs from pydantic import BaseModel, TypeAdapter from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore from core.app.file_access import DatabaseFileAccessController from core.app.workflow.file_runtime import DifyWorkflowFileRuntime -from core.db.session_factory import session_factory from factories import file_factory from models.agent import AgentRuntimeSessionOwnerType, WorkflowAgentRuntimeSession, WorkflowAgentRuntimeSessionStatus @@ -134,6 +134,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ): locator = self._resolve_locator( tenant_id=tenant_id, @@ -141,6 +142,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) return self._client_factory().list_sandbox_files_sync(locator, path) @@ -153,6 +155,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ): locator = self._resolve_locator( tenant_id=tenant_id, @@ -160,6 +163,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) return self._client_factory().read_sandbox_file_sync(locator, path) @@ -172,6 +176,7 @@ class WorkflowAgentSandboxService: node_id: str, node_execution_id: str | None, path: str, + session: Session, ) -> AgentSandboxUploadDownload: locator = self._resolve_locator( tenant_id=tenant_id, @@ -179,6 +184,7 @@ class WorkflowAgentSandboxService: workflow_run_id=workflow_run_id, node_id=node_id, node_execution_id=node_execution_id, + session=session, ) uploaded = self._client_factory().upload_sandbox_file_sync(locator, path) return _upload_download_response( @@ -194,6 +200,7 @@ class WorkflowAgentSandboxService: workflow_run_id: str, node_id: str, node_execution_id: str | None, + session: Session, ) -> SandboxLocator: """Resolve one workflow Agent sandbox from product-facing identifiers. @@ -216,8 +223,7 @@ class WorkflowAgentSandboxService: stmt = stmt.where(WorkflowAgentRuntimeSession.node_execution_id == node_execution_id) stmt = stmt.order_by(WorkflowAgentRuntimeSession.updated_at.desc()).limit(1) - with session_factory.create_session() as session: - row = session.scalar(stmt) + row = session.scalar(stmt) if row is None: raise AgentSandboxInspectorError( diff --git a/api/services/agent_drive_service.py b/api/services/agent_drive_service.py index d79aa1abe9c..eb375f997d1 100644 --- a/api/services/agent_drive_service.py +++ b/api/services/agent_drive_service.py @@ -41,7 +41,6 @@ from sqlalchemy.orm import Session from configs import dify_config from core.app.file_access.controller import DatabaseFileAccessController -from core.db.session_factory import session_factory from extensions.ext_storage import storage from factories import file_factory from libs.uuid_utils import uuidv7 @@ -195,38 +194,38 @@ class AgentDriveService: *, tenant_id: str, agent_id: str, + session: Session, prefix: str = "", include_download_url: bool = False, ) -> list[dict[str, Any]]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - stmt = ( - select(AgentDriveFile) - .where(AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id) - .order_by(AgentDriveFile.key) - ) - if prefix: - stmt = stmt.where(AgentDriveFile.key.startswith(prefix)) - rows = list(session.scalars(stmt)) - items: list[dict[str, Any]] = [] - for row in rows: - item: dict[str, Any] = { - "key": row.key, - "size": row.size, - "hash": row.hash, - "mime_type": row.mime_type, - "file_kind": row.file_kind.value, - "file_id": row.file_id, - "is_skill": row.is_skill, - "skill_metadata": row.skill_metadata, - "created_at": int(row.created_at.timestamp()) if row.created_at else None, - } - if include_download_url: - item["download_url"] = self._resolve_download_url( - tenant_id=tenant_id, file_kind=row.file_kind, file_id=row.file_id - ) - items.append(item) - return items + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + stmt = ( + select(AgentDriveFile) + .where(AgentDriveFile.tenant_id == tenant_id, AgentDriveFile.agent_id == agent_id) + .order_by(AgentDriveFile.key) + ) + if prefix: + stmt = stmt.where(AgentDriveFile.key.startswith(prefix)) + rows = list(session.scalars(stmt)) + items: list[dict[str, Any]] = [] + for row in rows: + item: dict[str, Any] = { + "key": row.key, + "size": row.size, + "hash": row.hash, + "mime_type": row.mime_type, + "file_kind": row.file_kind.value, + "file_id": row.file_id, + "is_skill": row.is_skill, + "skill_metadata": row.skill_metadata, + "created_at": int(row.created_at.timestamp()) if row.created_at else None, + } + if include_download_url: + item["download_url"] = self._resolve_download_url( + tenant_id=tenant_id, file_kind=row.file_kind, file_id=row.file_id + ) + items.append(item) + return items def commit( self, @@ -235,25 +234,25 @@ class AgentDriveService: user_id: str, agent_id: str, items: list[DriveCommitItem], + session: Session, ) -> list[dict[str, Any]]: if not items: raise AgentDriveError("empty_commit", "commit requires at least one item", status_code=400) committed: list[dict[str, Any]] = [] pending_storage_deletes: list[str] = [] - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - for item in items: - committed.append( - self._commit_one( - session, - tenant_id=tenant_id, - user_id=user_id, - agent_id=agent_id, - item=item, - pending_storage_deletes=pending_storage_deletes, - ) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + for item in items: + committed.append( + self._commit_one( + session, + tenant_id=tenant_id, + user_id=user_id, + agent_id=agent_id, + item=item, + pending_storage_deletes=pending_storage_deletes, ) - session.commit() + ) + session.commit() for storage_key in pending_storage_deletes: self._delete_storage(storage_key) return committed @@ -263,6 +262,7 @@ class AgentDriveService: *, tenant_id: str, agent_id: str, + session: Session, prefix: str | None = None, key: str | None = None, ) -> list[str]: @@ -276,59 +276,57 @@ class AgentDriveService: raise AgentDriveError("invalid_delete_scope", "delete requires exactly one of prefix or key") removed_keys: list[str] = [] pending_storage_deletes: list[str] = [] - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - stmt = select(AgentDriveFile).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - ) - if key is not None: - stmt = stmt.where(AgentDriveFile.key == normalize_drive_key(key)) - else: - stmt = stmt.where(AgentDriveFile.key.startswith(normalize_drive_key(prefix or ""))) - rows = list(session.scalars(stmt)) - for row in rows: - if row.value_owned_by_drive: - self._cleanup_value( - session, - tenant_id=tenant_id, - file_kind=row.file_kind, - file_id=row.file_id, - exclude_row_id=row.id, - pending_storage_deletes=pending_storage_deletes, - ) - removed_keys.append(row.key) - session.delete(row) - session.commit() + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + stmt = select(AgentDriveFile).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + ) + if key is not None: + stmt = stmt.where(AgentDriveFile.key == normalize_drive_key(key)) + else: + stmt = stmt.where(AgentDriveFile.key.startswith(normalize_drive_key(prefix or ""))) + rows = list(session.scalars(stmt)) + for row in rows: + if row.value_owned_by_drive: + self._cleanup_value( + session, + tenant_id=tenant_id, + file_kind=row.file_kind, + file_id=row.file_id, + exclude_row_id=row.id, + pending_storage_deletes=pending_storage_deletes, + ) + removed_keys.append(row.key) + session.delete(row) + session.commit() for storage_key in pending_storage_deletes: self._delete_storage(storage_key) return removed_keys - def list_skills(self, *, tenant_id: str, agent_id: str) -> list[AgentDriveSkillInfo]: + def list_skills(self, *, tenant_id: str, agent_id: str, session: Session) -> list[AgentDriveSkillInfo]: """Return the drive-backed skill catalog derived from canonical ``SKILL.md`` rows.""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - skill_rows = list( - session.scalars( - select(AgentDriveFile) - .where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.is_skill.is_(True), - ) - .order_by(AgentDriveFile.key) - ) - ) - archive_keys = set( - session.scalars( - select(AgentDriveFile.key).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.key.in_([self._skill_archive_key(row.key) for row in skill_rows]), - ) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + skill_rows = list( + session.scalars( + select(AgentDriveFile) + .where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.is_skill.is_(True), + ) + .order_by(AgentDriveFile.key) + ) + ) + archive_keys = set( + session.scalars( + select(AgentDriveFile.key).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.key.in_([self._skill_archive_key(row.key) for row in skill_rows]), ) ) + ) skills: list[AgentDriveSkillInfo] = [] for row in skill_rows: @@ -349,14 +347,20 @@ class AgentDriveService: ) return skills - def inspect_skill(self, *, tenant_id: str, agent_id: str, skill_path: str) -> AgentDriveSkillInspectInfo: + def inspect_skill( + self, *, tenant_id: str, agent_id: str, skill_path: str, session: Session + ) -> AgentDriveSkillInspectInfo: """Return the UI-facing skill inspect view for slash-menu hover/detail.""" skill_path = normalize_drive_key(skill_path) skill_md_key = skill_path if skill_path.endswith(_SKILL_MD_SUFFIX) else f"{skill_path}{_SKILL_MD_SUFFIX}" skill_path = self._skill_path_from_key(skill_md_key) catalog = next( - (item for item in self.list_skills(tenant_id=tenant_id, agent_id=agent_id) if item["path"] == skill_path), + ( + item + for item in self.list_skills(tenant_id=tenant_id, agent_id=agent_id, session=session) + if item["path"] == skill_path + ), None, ) if catalog is None: @@ -366,10 +370,11 @@ class AgentDriveService: tenant_id=tenant_id, agent_id=agent_id, skill_md_key=skill_md_key, + session=session, ) - drive_items = self.manifest(tenant_id=tenant_id, agent_id=agent_id, prefix=f"{skill_path}/") + drive_items = self.manifest(tenant_id=tenant_id, agent_id=agent_id, prefix=f"{skill_path}/", session=session) drive_keys = {item["key"] for item in drive_items} - preview = self.preview(tenant_id=tenant_id, agent_id=agent_id, key=skill_md_key) + preview = self.preview(tenant_id=tenant_id, agent_id=agent_id, key=skill_md_key, session=session) files, warnings = self._skill_file_entries( skill_path=skill_path, skill_md_key=skill_md_key, @@ -582,23 +587,24 @@ class AgentDriveService: ) from exc @staticmethod - def _manifest_files_from_skill_metadata(*, tenant_id: str, agent_id: str, skill_md_key: str) -> list[str] | None: - with session_factory.create_session() as session: - row = session.scalar( - select(AgentDriveFile).where( - AgentDriveFile.tenant_id == tenant_id, - AgentDriveFile.agent_id == agent_id, - AgentDriveFile.key == skill_md_key, - AgentDriveFile.is_skill.is_(True), - ) + def _manifest_files_from_skill_metadata( + *, tenant_id: str, agent_id: str, skill_md_key: str, session: Session + ) -> list[str] | None: + row = session.scalar( + select(AgentDriveFile).where( + AgentDriveFile.tenant_id == tenant_id, + AgentDriveFile.agent_id == agent_id, + AgentDriveFile.key == skill_md_key, + AgentDriveFile.is_skill.is_(True), ) - if row is None: - return None - try: - metadata = AgentDriveService._parse_skill_metadata(row.key, row.skill_metadata) - except Exception: - logger.warning("drive skill inspect: malformed skill metadata for %s", skill_md_key, exc_info=True) - return None + ) + if row is None: + return None + try: + metadata = AgentDriveService._parse_skill_metadata(row.key, row.skill_metadata) + except Exception: + logger.warning("drive skill inspect: malformed skill metadata for %s", skill_md_key, exc_info=True) + return None return [str(item) for item in (metadata.manifest_files or []) if str(item).strip()] or None @classmethod @@ -932,15 +938,15 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> bytes: member_path = normalize_drive_key(member_path) - with session_factory.create_session() as session: - storage_key = self._storage_key_for_ref( - session, - tenant_id=tenant_id, - file_kind=archive_file_kind, - file_id=archive_file_id, - ) + storage_key = self._storage_key_for_ref( + session, + tenant_id=tenant_id, + file_kind=archive_file_kind, + file_id=archive_file_id, + ) archive_bytes = b"".join(storage.load_stream(storage_key)) try: with zipfile.ZipFile(io.BytesIO(archive_bytes)) as archive: @@ -978,26 +984,25 @@ class AgentDriveService: return {"key": key, "size": size, "truncated": truncated, "binary": True, "text": None} return {"key": key, "size": size, "truncated": truncated, "binary": False, "text": text} - def preview(self, *, tenant_id: str, agent_id: str, key: str) -> dict[str, Any]: + def preview(self, *, tenant_id: str, agent_id: str, key: str, session: Session) -> dict[str, Any]: """Truncated text preview of one drive value (binary-safe, never 500s on size).""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - try: - row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) - storage_key = self._storage_key_for_row(session, tenant_id=tenant_id, row=row) - size = row.size - response_key = row.key - archive_ref: tuple[AgentDriveFile, str] | None = None - except AgentDriveError: - archive_ref = self._archive_member_for_key( - session, - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - ) - storage_key = None - size = None - response_key = normalize_drive_key(key) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + try: + row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) + storage_key = self._storage_key_for_row(session, tenant_id=tenant_id, row=row) + size = row.size + response_key = row.key + archive_ref: tuple[AgentDriveFile, str] | None = None + except AgentDriveError: + archive_ref = self._archive_member_for_key( + session, + tenant_id=tenant_id, + agent_id=agent_id, + key=key, + ) + storage_key = None + size = None + response_key = normalize_drive_key(key) if archive_ref is not None: archive_row, member_path = archive_ref @@ -1006,6 +1011,7 @@ class AgentDriveService: archive_file_kind=archive_row.file_kind, archive_file_id=archive_row.file_id, member_path=member_path, + session=session, ) return self._preview_bytes(key=response_key, size=len(payload), payload=payload) @@ -1026,47 +1032,47 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> dict[str, Any]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) payload = self._load_archive_member_bytes( tenant_id=tenant_id, archive_file_kind=archive_file_kind, archive_file_id=archive_file_id, member_path=member_path, + session=session, ) return self._preview_bytes(key=normalize_drive_key(key), size=len(payload), payload=payload) - def download_url(self, *, tenant_id: str, agent_id: str, key: str) -> str: + def download_url(self, *, tenant_id: str, agent_id: str, key: str, session: Session) -> str: """External signed URL for a browser download of one drive value.""" - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) - try: - row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) - except AgentDriveError: - archive_row, member_path = self._archive_member_for_key( - session, - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - ) - return self.sign_archive_member_url( - tenant_id=tenant_id, - agent_id=agent_id, - key=key, - archive_file_kind=archive_row.file_kind, - archive_file_id=archive_row.file_id, - member_path=member_path, - for_external=True, - as_attachment=True, - ) - url = self._resolve_download_url( + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + try: + row = self._require_row(session, tenant_id=tenant_id, agent_id=agent_id, key=key) + except AgentDriveError: + archive_row, member_path = self._archive_member_for_key( + session, tenant_id=tenant_id, - file_kind=row.file_kind, - file_id=row.file_id, + agent_id=agent_id, + key=key, + ) + return self.sign_archive_member_url( + tenant_id=tenant_id, + agent_id=agent_id, + key=key, + archive_file_kind=archive_row.file_kind, + archive_file_id=archive_row.file_id, + member_path=member_path, for_external=True, as_attachment=True, ) + url = self._resolve_download_url( + tenant_id=tenant_id, + file_kind=row.file_kind, + file_id=row.file_id, + for_external=True, + as_attachment=True, + ) if url is None: raise AgentDriveError("drive_key_not_found", "drive value cannot be resolved", status_code=404) return url @@ -1080,10 +1086,10 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, for_external: bool = True, ) -> str: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) return self.sign_archive_member_url( tenant_id=tenant_id, agent_id=agent_id, @@ -1211,14 +1217,15 @@ class AgentDriveService: archive_file_kind: AgentDriveFileKind, archive_file_id: str, member_path: str, + session: Session, ) -> tuple[bytes, str, str]: - with session_factory.create_session() as session: - self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) + self._assert_agent_belongs_to_tenant(session, tenant_id=tenant_id, agent_id=agent_id) payload = self._load_archive_member_bytes( tenant_id=tenant_id, archive_file_kind=archive_file_kind, archive_file_id=archive_file_id, member_path=member_path, + session=session, ) mime_type = mimetypes.guess_type(member_path)[0] or "application/octet-stream" filename = normalize_drive_key(key).rsplit("/", 1)[-1] diff --git a/api/services/agent_service.py b/api/services/agent_service.py index d8f4e11e758..a201eeb0485 100644 --- a/api/services/agent_service.py +++ b/api/services/agent_service.py @@ -3,13 +3,13 @@ from typing import Any import pytz from sqlalchemy import select +from sqlalchemy.orm import Session import contexts from core.app.app_config.easy_ui_based_app.agent.manager import AgentConfigManager from core.plugin.impl.agent import PluginAgentClient from core.plugin.impl.exc import PluginDaemonClientSideError from core.tools.tool_manager import ToolManager -from extensions.ext_database import db from libs.login import current_user from models import Account from models.model import App, Conversation, EndUser, Message @@ -17,14 +17,14 @@ from models.model import App, Conversation, EndUser, Message class AgentService: @classmethod - def get_agent_logs(cls, app_model: App, conversation_id: str, message_id: str): + def get_agent_logs(cls, app_model: App, conversation_id: str, message_id: str, session: Session): """ Service to get agent logs """ contexts.plugin_tool_providers.set({}) contexts.plugin_tool_providers_lock.set(threading.Lock()) - conversation: Conversation | None = db.session.scalar( + conversation: Conversation | None = session.scalar( select(Conversation) .where( Conversation.id == conversation_id, @@ -36,7 +36,7 @@ class AgentService: if not conversation: raise ValueError(f"Conversation not found: {conversation_id}") - message: Message | None = db.session.scalar( + message: Message | None = session.scalar( select(Message) .where( Message.id == message_id, @@ -52,9 +52,9 @@ class AgentService: if conversation.from_end_user_id: # only select name field - executor_name = db.session.scalar(select(EndUser.name).where(EndUser.id == conversation.from_end_user_id)) + executor_name = session.scalar(select(EndUser.name).where(EndUser.id == conversation.from_end_user_id)) else: - executor_name = db.session.scalar(select(Account.name).where(Account.id == conversation.from_account_id)) + executor_name = session.scalar(select(Account.name).where(Account.id == conversation.from_account_id)) executor = executor_name or "Unknown" assert isinstance(current_user, Account) diff --git a/api/services/agent_tool_inner_service.py b/api/services/agent_tool_inner_service.py index 4420f1b66b0..633ca893007 100644 --- a/api/services/agent_tool_inner_service.py +++ b/api/services/agent_tool_inner_service.py @@ -41,7 +41,7 @@ from services.errors.agent_tool_inner import AgentToolInnerServiceError class AgentToolInnerService: """Invoke one API-owned Agent tool declaration, including explicit plugin-via-core calls.""" - def invoke(self, session: Session, request: AgentToolInvokeRequest) -> AgentToolInvokeResponse: + def invoke(self, request: AgentToolInvokeRequest, *, session: Session) -> AgentToolInvokeResponse: app = session.get(App, request.caller.app_id) if app is None: raise AgentToolInnerServiceError( diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 03e445a938b..ccca621aab5 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -4,12 +4,12 @@ from typing import TypedDict import pandas as pd from sqlalchemy import delete, or_, select, update -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from werkzeug.exceptions import NotFound from core.helper.csv_sanitizer import CSVSanitizer -from extensions.ext_database import db +from extensions.ext_database import db # noqa: F401 from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now from libs.login import current_account_with_tenant @@ -91,7 +91,7 @@ class UpdateAnnotationSettingArgs(TypedDict): class AppAnnotationService: @staticmethod - def _get_annotation_by_ref(annotation_ref: AnnotationRef, session: scoped_session) -> MessageAnnotation | None: + def _get_annotation_by_ref(annotation_ref: AnnotationRef, session: Session) -> MessageAnnotation | None: return session.scalar( select(MessageAnnotation) .where( @@ -102,10 +102,12 @@ class AppAnnotationService: ) @classmethod - def up_insert_app_annotation_from_message(cls, args: UpsertAnnotationArgs, app_id: str) -> MessageAnnotation: + def up_insert_app_annotation_from_message( + cls, args: UpsertAnnotationArgs, app_id: str, *, session: Session + ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -119,9 +121,7 @@ class AppAnnotationService: raw_message_id = args.get("message_id") if raw_message_id: message_id = str(raw_message_id) - message = db.session.scalar( - select(Message).where(Message.id == message_id, Message.app_id == app.id).limit(1) - ) + message = session.scalar(select(Message).where(Message.id == message_id, Message.app_id == app.id).limit(1)) if not message: raise NotFound("Message Not Exists.") @@ -155,10 +155,10 @@ class AppAnnotationService: question=question, account_id=current_user.id, ) - db.session.add(annotation) - db.session.commit() + session.add(annotation) + session.commit() - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) assert current_tenant_id is not None @@ -213,10 +213,10 @@ class AppAnnotationService: return {"job_id": job_id, "job_status": "waiting"} @classmethod - def get_annotation_list_by_app_id(cls, app_id: str, page: int, limit: int, keyword: str): + def get_annotation_list_by_app_id(cls, app_id: str, page: int, limit: int, keyword: str, *, session: Session): # get app info _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -247,7 +247,7 @@ class AppAnnotationService: return annotations.items, annotations.total or 0 @classmethod - def export_annotation_list_by_app_id(cls, app_id: str): + def export_annotation_list_by_app_id(cls, app_id: str, *, session: Session): """ Export all annotations for an app with CSV injection protection. @@ -256,13 +256,13 @@ class AppAnnotationService: """ # get app info _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotations = db.session.scalars( + annotations = session.scalars( select(MessageAnnotation) .where(MessageAnnotation.app_id == app_id) .order_by(MessageAnnotation.created_at.desc()) @@ -280,10 +280,12 @@ class AppAnnotationService: return annotations @classmethod - def insert_app_annotation_directly(cls, args: InsertAnnotationArgs, app_id: str) -> MessageAnnotation: + def insert_app_annotation_directly( + cls, args: InsertAnnotationArgs, app_id: str, *, session: Session + ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -297,10 +299,10 @@ class AppAnnotationService: annotation = MessageAnnotation( app_id=app.id, content=args["answer"], question=question, account_id=current_user.id ) - db.session.add(annotation) - db.session.commit() + session.add(annotation) + session.commit() # if annotation reply is enabled , add annotation to index - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) if annotation_setting: @@ -315,7 +317,7 @@ class AppAnnotationService: @classmethod def update_app_annotation_directly( - cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: scoped_session + cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: Session ): annotation = cls._get_annotation_by_ref(annotation_ref, session) @@ -351,7 +353,7 @@ class AppAnnotationService: return annotation @classmethod - def delete_app_annotation(cls, annotation_ref: AnnotationRef, session: scoped_session): + def delete_app_annotation(cls, annotation_ref: AnnotationRef, session: Session): annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: @@ -384,9 +386,9 @@ class AppAnnotationService: ) @classmethod - def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str]): + def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str], *, session: Session): # Fetch annotations and their settings in a single query - annotations_to_delete = db.session.execute( + annotations_to_delete = session.execute( select(MessageAnnotation, AppAnnotationSetting) .outerjoin(AppAnnotationSetting, MessageAnnotation.app_id == AppAnnotationSetting.app_id) .where(MessageAnnotation.id.in_(annotation_ids), MessageAnnotation.app_id == app_ref.app_id) @@ -399,7 +401,7 @@ class AppAnnotationService: annotation_ids_to_delete = [annotation.id for annotation, _ in annotations_to_delete] # Step 2: Bulk delete hit histories in a single query - db.session.execute( + session.execute( delete(AppAnnotationHitHistory).where( AppAnnotationHitHistory.app_id == app_ref.app_id, AppAnnotationHitHistory.annotation_id.in_(annotation_ids_to_delete), @@ -414,7 +416,7 @@ class AppAnnotationService: ) # Step 4: Bulk delete annotations in a single query - delete_result = db.session.execute( + delete_result = session.execute( delete(MessageAnnotation).where( MessageAnnotation.id.in_(annotation_ids_to_delete), MessageAnnotation.app_id == app_ref.app_id, @@ -422,11 +424,11 @@ class AppAnnotationService: ) deleted_count = getattr(delete_result, "rowcount", 0) - db.session.commit() + session.commit() return {"deleted_count": deleted_count} @classmethod - def batch_import_app_annotations(cls, app_id: str, file: FileStorage): + def batch_import_app_annotations(cls, app_id: str, file: FileStorage, *, session: Session): """ Batch import annotations from CSV file with enhanced security checks. @@ -441,7 +443,7 @@ class AppAnnotationService: # get app info current_user, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -560,8 +562,8 @@ class AppAnnotationService: return {"job_id": job_id, "job_status": "waiting", "record_count": len(result)} @classmethod - def get_annotation_hit_histories(cls, annotation_ref: AnnotationRef, page, limit): - annotation = cls._get_annotation_by_ref(annotation_ref, db.session) + def get_annotation_hit_histories(cls, annotation_ref: AnnotationRef, page, limit, *, session: Session): + annotation = cls._get_annotation_by_ref(annotation_ref, session) if not annotation: raise NotFound("Annotation not found") @@ -578,8 +580,8 @@ class AppAnnotationService: return annotation_hit_histories.items, annotation_hit_histories.total or 0 @classmethod - def get_annotation_by_id(cls, annotation_id: str) -> MessageAnnotation | None: - annotation = db.session.get(MessageAnnotation, annotation_id) + def get_annotation_by_id(cls, annotation_id: str, *, session: Session) -> MessageAnnotation | None: + annotation = session.get(MessageAnnotation, annotation_id) if not annotation: return None @@ -597,9 +599,11 @@ class AppAnnotationService: message_id: str, from_source: str, score: float, - ): + *, + session: Session, + ) -> None: # add hit count to annotation - db.session.execute( + session.execute( update(MessageAnnotation) .where(MessageAnnotation.id == annotation_id) .values(hit_count=MessageAnnotation.hit_count + 1) @@ -616,21 +620,23 @@ class AppAnnotationService: annotation_question=annotation_question, annotation_content=annotation_content, ) - db.session.add(annotation_hit_history) - db.session.commit() + session.add(annotation_hit_history) + session.commit() @classmethod - def get_app_annotation_setting_by_app_id(cls, app_id: str) -> AnnotationSettingDict | AnnotationSettingDisabledDict: + def get_app_annotation_setting_by_app_id( + cls, app_id: str, *, session: Session + ) -> AnnotationSettingDict | AnnotationSettingDisabledDict: _, current_tenant_id = current_account_with_tenant() # get app info - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) if annotation_setting: @@ -656,18 +662,18 @@ class AppAnnotationService: @classmethod def update_app_annotation_setting( - cls, app_id: str, annotation_setting_id: str, args: UpdateAnnotationSettingArgs + cls, app_id: str, annotation_setting_id: str, args: UpdateAnnotationSettingArgs, *, session: Session ) -> AnnotationSettingDict: current_user, current_tenant_id = current_account_with_tenant() # get app info - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) if not app: raise NotFound("App not found") - annotation_setting = db.session.scalar( + annotation_setting = session.scalar( select(AppAnnotationSetting) .where( AppAnnotationSetting.app_id == app_id, @@ -680,8 +686,8 @@ class AppAnnotationService: annotation_setting.score_threshold = args["score_threshold"] annotation_setting.updated_user_id = current_user.id annotation_setting.updated_at = naive_utc_now() - db.session.add(annotation_setting) - db.session.commit() + session.add(annotation_setting) + session.commit() collection_binding_detail = annotation_setting.collection_binding_detail @@ -704,9 +710,9 @@ class AppAnnotationService: } @classmethod - def clear_all_annotations(cls, app_id: str): + def clear_all_annotations(cls, app_id: str, *, session: Session): _, current_tenant_id = current_account_with_tenant() - app = db.session.scalar( + app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) ) @@ -714,19 +720,19 @@ class AppAnnotationService: raise NotFound("App not found") # if annotation reply is enabled, delete annotation index - app_annotation_setting = db.session.scalar( + app_annotation_setting = session.scalar( select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) - annotations_iter = db.session.scalars( + annotations_iter = session.scalars( select(MessageAnnotation).where(MessageAnnotation.app_id == app_id) ).yield_per(100) for annotation in annotations_iter: - hit_histories_iter = db.session.scalars( + hit_histories_iter = session.scalars( select(AppAnnotationHitHistory).where(AppAnnotationHitHistory.annotation_id == annotation.id) ).yield_per(100) for annotation_hit_history in hit_histories_iter: - db.session.delete(annotation_hit_history) + session.delete(annotation_hit_history) # if annotation reply is enabled, delete annotation index if app_annotation_setting: @@ -734,7 +740,7 @@ class AppAnnotationService: annotation.id, app_id, current_tenant_id, app_annotation_setting.collection_binding_id ) - db.session.delete(annotation) + session.delete(annotation) - db.session.commit() + session.commit() return {"result": "success"} diff --git a/api/services/api_based_extension_service.py b/api/services/api_based_extension_service.py index 25f554b6bdc..e855780d6a1 100644 --- a/api/services/api_based_extension_service.py +++ b/api/services/api_based_extension_service.py @@ -8,7 +8,7 @@ from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint class APIBasedExtensionService: @staticmethod - def get_all_by_tenant_id(session: Session, tenant_id: str) -> list[APIBasedExtension]: + def get_all_by_tenant_id(tenant_id: str, *, session: Session) -> list[APIBasedExtension]: extension_list = list( session.scalars( select(APIBasedExtension) @@ -23,7 +23,7 @@ class APIBasedExtensionService: return extension_list @classmethod - def save(cls, session: Session, extension_data: APIBasedExtension) -> APIBasedExtension: + def save(cls, extension_data: APIBasedExtension, *, session: Session) -> APIBasedExtension: cls._validation(session, extension_data) extension_data.api_key = encrypt_token(extension_data.tenant_id, extension_data.api_key) @@ -33,12 +33,12 @@ class APIBasedExtensionService: return extension_data @staticmethod - def delete(session: Session, extension_data: APIBasedExtension): + def delete(extension_data: APIBasedExtension, *, session: Session): session.delete(extension_data) session.commit() @staticmethod - def get_with_tenant_id(session: Session, tenant_id: str, api_based_extension_id: str) -> APIBasedExtension: + def get_with_tenant_id(tenant_id: str, api_based_extension_id: str, *, session: Session) -> APIBasedExtension: extension = session.scalar( select(APIBasedExtension) .where(APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id) diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index d042ad69f88..e8c12586856 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -474,7 +474,7 @@ class AppDslService: ] workflow_service = WorkflowService() - current_draft_workflow = workflow_service.get_draft_workflow(app_model=app) + current_draft_workflow = workflow_service.get_draft_workflow(app_model=app, session=self._session) if current_draft_workflow: unique_hash = current_draft_workflow.unique_hash else: @@ -500,6 +500,7 @@ class AppDslService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=self._session, ) case AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION: # Initialize model config @@ -521,7 +522,14 @@ class AppDslService: return app @classmethod - def export_dsl(cls, app_model: App, include_secret: bool = False, workflow_id: str | None = None) -> str: + def export_dsl( + cls, + app_model: App, + *, + session: Session, + include_secret: bool = False, + workflow_id: str | None = None, + ) -> str: """ Export app :param app_model: App instance @@ -548,7 +556,11 @@ class AppDslService: if app_mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: cls._append_workflow_export_data( - export_data=export_data, app_model=app_model, include_secret=include_secret, workflow_id=workflow_id + export_data=export_data, + app_model=app_model, + include_secret=include_secret, + workflow_id=workflow_id, + session=session, ) else: cls._append_model_config_export_data(export_data, app_model) @@ -557,7 +569,13 @@ class AppDslService: @classmethod def _append_workflow_export_data( - cls, *, export_data: dict[str, Any], app_model: App, include_secret: bool, workflow_id: str | None = None + cls, + *, + export_data: dict[str, Any], + app_model: App, + include_secret: bool, + session: Session, + workflow_id: str | None = None, ): """ Append workflow export data @@ -565,7 +583,7 @@ class AppDslService: :param app_model: App instance """ workflow_service = WorkflowService() - workflow = workflow_service.get_draft_workflow(app_model, workflow_id) + workflow = workflow_service.get_draft_workflow(app_model, workflow_id, session=session) if not workflow: raise WorkflowNotFoundError("Missing draft workflow configuration, please check.") diff --git a/api/services/app_generate_service.py b/api/services/app_generate_service.py index 3e2c3c96403..940cab5f678 100644 --- a/api/services/app_generate_service.py +++ b/api/services/app_generate_service.py @@ -120,11 +120,12 @@ class AppGenerateService: @trace_span(AppGenerateHandler) def generate( cls, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, + *, + session: Session, streaming: bool = True, root_node_id: str | None = None, ): @@ -141,13 +142,13 @@ class AppGenerateService: app_model=app_model, streaming=streaming, action=lambda rate_limit, request_id: cls._dispatch_generate( - session=session, app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=streaming, root_node_id=root_node_id, + session=session, rate_limit=rate_limit, request_id=request_id, ), @@ -189,13 +190,13 @@ class AppGenerateService: def _dispatch_generate( cls, *, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool, root_node_id: str | None, + session: Session, rate_limit: RateLimit, request_id: str, ): @@ -251,7 +252,7 @@ class AppGenerateService: ) case AppMode.ADVANCED_CHAT: workflow_id = args.get("workflow_id") - workflow = cls._get_workflow(app_model, invoke_from, workflow_id) + workflow = cls._get_workflow(app_model, invoke_from, workflow_id, session=session) if streaming: # Streaming mode: subscribe to SSE and enqueue the execution on first subscriber @@ -308,7 +309,7 @@ class AppGenerateService: ) case AppMode.WORKFLOW: workflow_id = args.get("workflow_id") - workflow = cls._get_workflow(app_model, invoke_from, workflow_id) + workflow = cls._get_workflow(app_model, invoke_from, workflow_id, session=session) if streaming: with rate_limit_context(rate_limit, request_id): payload = AppExecutionParams.new( @@ -384,12 +385,21 @@ class AppGenerateService: return min(limits) if limits else 0 @classmethod - def generate_single_iteration(cls, app_model: App, user: Account, node_id: str, args: Any, streaming: bool = True): + def generate_single_iteration( + cls, + app_model: App, + user: Account, + node_id: str, + args: Any, + *, + session: Session, + streaming: bool = True, + ): match app_model.mode: case AppMode.COMPLETION | AppMode.CHAT | AppMode.AGENT_CHAT: raise ValueError(f"Invalid app mode {app_model.mode}") case AppMode.ADVANCED_CHAT: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( AdvancedChatAppGenerator().single_iteration_generate( app_model=app_model, @@ -401,7 +411,7 @@ class AppGenerateService: ) ) case AppMode.WORKFLOW: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( WorkflowAppGenerator().single_iteration_generate( app_model=app_model, @@ -419,13 +429,20 @@ class AppGenerateService: @classmethod def generate_single_loop( - cls, app_model: App, user: Account, node_id: str, args: LoopNodeRunPayload, streaming: bool = True + cls, + app_model: App, + user: Account, + node_id: str, + args: LoopNodeRunPayload, + *, + session: Session, + streaming: bool = True, ): match app_model.mode: case AppMode.COMPLETION | AppMode.CHAT | AppMode.AGENT_CHAT: raise ValueError(f"Invalid app mode {app_model.mode}") case AppMode.ADVANCED_CHAT: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( AdvancedChatAppGenerator().single_loop_generate( app_model=app_model, @@ -437,7 +454,7 @@ class AppGenerateService: ) ) case AppMode.WORKFLOW: - workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(app_model, InvokeFrom.DEBUGGER, session=session) return AdvancedChatAppGenerator.convert_to_event_stream( WorkflowAppGenerator().single_loop_generate( app_model=app_model, @@ -456,11 +473,12 @@ class AppGenerateService: @classmethod def generate_more_like_this( cls, - session: Session, app_model: App, user: Account | EndUser, message_id: str, invoke_from: InvokeFrom, + *, + session: Session, streaming: bool = True, ) -> Mapping | Generator: """ @@ -482,7 +500,14 @@ class AppGenerateService: ) @classmethod - def _get_workflow(cls, app_model: App, invoke_from: InvokeFrom, workflow_id: str | None = None) -> Workflow: + def _get_workflow( + cls, + app_model: App, + invoke_from: InvokeFrom, + workflow_id: str | None = None, + *, + session: Session, + ) -> Workflow: """ Get workflow :param app_model: app model @@ -498,20 +523,22 @@ class AppGenerateService: _ = uuid.UUID(workflow_id) except ValueError: raise WorkflowIdFormatError(f"Invalid workflow_id format: '{workflow_id}'. ") - workflow = workflow_service.get_published_workflow_by_id(app_model=app_model, workflow_id=workflow_id) + workflow = workflow_service.get_published_workflow_by_id( + app_model=app_model, workflow_id=workflow_id, session=session + ) if not workflow: raise WorkflowNotFoundError(f"Workflow not found with id: {workflow_id}") return workflow if invoke_from == InvokeFrom.DEBUGGER: # fetch draft workflow by app_model - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("Workflow not initialized") else: # fetch published workflow by app_model - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("Workflow not published") diff --git a/api/services/app_service.py b/api/services/app_service.py index 08cd30974e3..139513e87ee 100644 --- a/api/services/app_service.py +++ b/api/services/app_service.py @@ -8,7 +8,7 @@ import sqlalchemy as sa from pydantic import BaseModel, Field from sqlalchemy import ColumnElement, select from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from configs import dify_config from constants.model_template import default_app_templates @@ -18,7 +18,7 @@ from core.model_manager import ModelManager from core.tools.tool_manager import ToolManager from core.tools.utils.configuration import ToolParameterConfigurationManager from events.app_event import app_was_created, app_was_deleted, app_was_updated -from extensions.ext_database import db +from extensions.ext_database import db # noqa: F401 from graphon.model_runtime.entities.model_entities import ModelPropertyKey, ModelType from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel from libs.datetime_utils import naive_utc_now @@ -80,7 +80,7 @@ class CreateAppParams(BaseModel): class AppService: @staticmethod def _build_app_list_filters( - user_id: str, tenant_id: str, params: AppListBaseParams, session: scoped_session + user_id: str, tenant_id: str, params: AppListBaseParams, session: Session ) -> list[sa.ColumnElement[bool]]: filters = [App.tenant_id == tenant_id, App.is_universal == False] @@ -153,13 +153,7 @@ class AppService: }[sort_by] @staticmethod - def get_starred_app_ids( - session: Session | scoped_session, - *, - tenant_id: str, - account_id: str, - app_ids: Sequence[str], - ) -> set[str]: + def get_starred_app_ids(*, tenant_id: str, account_id: str, app_ids: Sequence[str], session: Session) -> set[str]: """Return app IDs starred by this account within the tenant.""" if not app_ids: return set() @@ -174,38 +168,24 @@ class AppService: return set(starred_app_ids) @staticmethod - def get_app_by_id( - session: Session | scoped_session, - app_id: str, - ) -> App | None: + def get_app_by_id(app_id: str, *, session: Session) -> App | None: return session.get(App, app_id) @staticmethod - def get_visible_app_by_id( - session: Session | scoped_session, - app_id: str, - ) -> App | None: + def get_visible_app_by_id(app_id: str, *, session: Session) -> App | None: app = session.get(App, app_id) if not app or app.status != "normal" or not is_openapi_visible(app): return None return app @staticmethod - def find_visible_apps_by_ids( - session: Session | scoped_session, - app_ids: Sequence[str], - ) -> list[App]: + def find_visible_apps_by_ids(app_ids: Sequence[str], *, session: Session) -> list[App]: if not app_ids: return [] return list(session.execute(apply_openapi_gate(select(App).where(App.id.in_(list(app_ids))))).scalars().all()) @staticmethod - def find_visible_apps_by_name( - session: Session | scoped_session, - *, - name: str, - tenant_id: str, - ) -> list[App]: + def find_visible_apps_by_name(*, name: str, tenant_id: str, session: Session) -> list[App]: return list( session.execute( apply_openapi_gate( @@ -219,7 +199,7 @@ class AppService: ) def get_paginate_apps( - self, user_id: str, tenant_id: str, params: AppListParams, session: scoped_session + self, user_id: str, tenant_id: str, params: AppListParams, session: Session ) -> PaginatedResult | None: """ Get app list with pagination, filters, and explicit sort order. @@ -238,14 +218,12 @@ class AppService: sa.select(App).where(*filters).order_by(order_by), page=params.page, per_page=params.limit, + session=session, ) app_ids = [str(app.id) for app in app_models.items] starred_app_ids = self.get_starred_app_ids( - db.session, - tenant_id=tenant_id, - account_id=user_id, - app_ids=app_ids, + tenant_id=tenant_id, account_id=user_id, app_ids=app_ids, session=session ) for app in app_models.items: app.is_starred = str(app.id) in starred_app_ids @@ -253,7 +231,7 @@ class AppService: return app_models def get_paginate_starred_apps( - self, user_id: str, tenant_id: str, params: StarredAppListParams, session: scoped_session + self, user_id: str, tenant_id: str, params: StarredAppListParams, session: Session ) -> PaginatedResult | None: """ Get apps starred by the current account with pagination, filters, and explicit sort order. @@ -277,6 +255,7 @@ class AppService: .order_by(order_by), page=params.page, per_page=params.limit, + session=session, ) for app in app_models.items: @@ -285,7 +264,7 @@ class AppService: return app_models @staticmethod - def star_app(session: Session, *, app: App, account_id: str) -> None: + def star_app(*, app: App, account_id: str, session: Session) -> None: """Create the account's app star if it does not already exist.""" existing_star = session.scalar( select(AppStar) @@ -302,7 +281,7 @@ class AppService: session.add(AppStar(tenant_id=app.tenant_id, app_id=app.id, account_id=account_id)) @staticmethod - def unstar_app(session: Session, *, app: App, account_id: str) -> None: + def unstar_app(*, app: App, account_id: str, session: Session) -> None: """Remove the account's app star if present.""" existing_star = session.scalar( select(AppStar) @@ -318,7 +297,7 @@ class AppService: session.delete(existing_star) - def create_app(self, tenant_id: str, params: CreateAppParams, account: Account) -> App: + def create_app(self, tenant_id: str, params: CreateAppParams, account: Account, *, session: Session) -> App: """ Create app :param tenant_id: tenant id @@ -397,15 +376,15 @@ class AppService: app.maintainer = account.id app.updated_by = account.id - db.session.add(app) - db.session.flush() + session.add(app) + session.flush() if default_model_config: app_model_config = AppModelConfig( **default_model_config, app_id=app.id, created_by=account.id, updated_by=account.id ) - db.session.add(app_model_config) - db.session.flush() + session.add(app_model_config) + session.flush() app.app_model_config_id = app_model_config.id elif app_mode == AppMode.AGENT: @@ -418,8 +397,8 @@ class AppService: # left unset so App.is_agent stays False (this is the new Agent App # type, not a legacy function-call/react agent). agent_app_model_config = AppModelConfig(app_id=app.id, created_by=account.id, updated_by=account.id) - db.session.add(agent_app_model_config) - db.session.flush() + session.add(agent_app_model_config) + session.flush() app.app_model_config_id = agent_app_model_config.id @@ -431,7 +410,7 @@ class AppService: from services.agent.roster_service import AgentRosterService icon_type = AgentIconType(params.icon_type) if params.icon_type else None - AgentRosterService(db.session).create_backing_agent_for_app( + AgentRosterService(session).create_backing_agent_for_app( tenant_id=tenant_id, account_id=account.id, app_id=app.id, @@ -443,7 +422,7 @@ class AppService: icon_background=params.icon_background, ) - db.session.commit() + session.commit() app_was_created.send(app, account=account) enterprise_rbac_service.try_sync_creator_access_policy_member_bindings( @@ -542,10 +521,10 @@ class AppService: role: NotRequired[str | None] @staticmethod - def _get_backing_agent_for_update(app: App) -> Agent | None: + def _get_backing_agent_for_update(app: App, *, session: Session) -> Agent | None: if app.mode != AppMode.AGENT: return None - return db.session.scalar( + return session.scalar( select(Agent).where( Agent.tenant_id == app.tenant_id, Agent.app_id == app.id, @@ -574,6 +553,7 @@ class AppService: icon_background: str | None = None, account_id: str | None = None, updated_at: datetime | None = None, + session: Session, ) -> None: """Keep the Roster identity aligned with its Agent App shell. @@ -584,7 +564,7 @@ class AppService: Role omission is intentional: ``role=None`` preserves the backing Agent's current role, while ``role=""`` explicitly clears it. """ - agent = self._get_backing_agent_for_update(app) + agent = self._get_backing_agent_for_update(app, session=session) if agent is None: return @@ -605,16 +585,16 @@ class AppService: agent.updated_at = updated_at @staticmethod - def _commit_app_identity_update(app: App) -> None: + def _commit_app_identity_update(app: App, *, session: Session) -> None: try: - db.session.commit() + session.commit() except IntegrityError as exc: - db.session.rollback() + session.rollback() if app.mode == AppMode.AGENT: raise AgentNameConflictError() from exc raise - def update_app(self, app: App, args: ArgsDict) -> App: + def update_app(self, app: App, args: ArgsDict, *, session: Session) -> App: """ Update app :param app: App instance @@ -649,14 +629,15 @@ class AppService: icon_background=app.icon_background, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - self._commit_app_identity_update(app) + self._commit_app_identity_update(app, session=session) app_was_updated.send(app) return app - def update_app_name(self, app: App, name: str) -> App: + def update_app_name(self, app: App, name: str, *, session: Session) -> App: """ Update app name :param app: App instance @@ -672,15 +653,22 @@ class AppService: name=app.name, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - self._commit_app_identity_update(app) + self._commit_app_identity_update(app, session=session) app_was_updated.send(app) return app def update_app_icon( - self, app: App, icon: str, icon_background: str, icon_type: IconType | str | None = None + self, + app: App, + icon: str, + icon_background: str, + icon_type: IconType | str | None = None, + *, + session: Session, ) -> App: """ Update app icon @@ -704,14 +692,15 @@ class AppService: icon_background=app.icon_background, account_id=current_user.id, updated_at=app.updated_at, + session=session, ) - db.session.commit() + session.commit() app_was_updated.send(app) return app - def update_app_site_status(self, app: App, enable_site: bool) -> App: + def update_app_site_status(self, app: App, enable_site: bool, *, session: Session) -> App: """ Update app site status :param app: App instance @@ -724,13 +713,13 @@ class AppService: app.enable_site = enable_site app.updated_by = current_user.id app.updated_at = naive_utc_now() - db.session.commit() + session.commit() app_was_updated.send(app) return app - def update_app_api_status(self, app: App, enable_api: bool) -> App: + def update_app_api_status(self, app: App, enable_api: bool, *, session: Session) -> App: """ Update app api status :param app: App instance @@ -744,20 +733,20 @@ class AppService: app.enable_api = enable_api app.updated_by = current_user.id app.updated_at = naive_utc_now() - db.session.commit() + session.commit() app_was_updated.send(app) return app - def delete_app(self, app: App): + def delete_app(self, app: App, *, session: Session) -> None: """ Delete app :param app: App instance """ app_was_deleted.send(app) - backing_agent = self._get_backing_agent_for_update(app) + backing_agent = self._get_backing_agent_for_update(app, session=session) if backing_agent is not None: now = naive_utc_now() account_id = getattr(current_user, "id", None) @@ -767,8 +756,8 @@ class AppService: backing_agent.updated_by = account_id backing_agent.updated_at = now - db.session.delete(app) - db.session.commit() + session.delete(app) + session.commit() # clean up web app settings if FeatureService.get_system_features().webapp_auth.enabled: @@ -780,7 +769,7 @@ class AppService: # Trigger asynchronous deletion of app and related data remove_app_and_related_data_task.delay(tenant_id=app.tenant_id, app_id=app.id) - def get_app_meta(self, app_model: App): + def get_app_meta(self, app_model: App, *, session: Session): """ Get app meta info :param app_model: app model @@ -833,7 +822,7 @@ class AppService: meta["tool_icons"][tool_name] = url_prefix + provider_id + "/icon" elif provider_type == "api": try: - provider: ApiToolProvider | None = db.session.get(ApiToolProvider, provider_id) + provider: ApiToolProvider | None = session.get(ApiToolProvider, provider_id) if provider is None: raise ValueError(f"provider not found for tool {tool_name}") meta["tool_icons"][tool_name] = json.loads(provider.icon) @@ -843,25 +832,25 @@ class AppService: return meta @staticmethod - def get_app_code_by_id(app_id: str) -> str: + def get_app_code_by_id(app_id: str, *, session: Session) -> str: """ Get app code by app id :param app_id: app id :return: app code """ - site = db.session.scalar(select(Site).where(Site.app_id == app_id).limit(1)) + site = session.scalar(select(Site).where(Site.app_id == app_id).limit(1)) if not site: raise ValueError(f"App with id {app_id} not found") return str(site.code) @staticmethod - def get_app_id_by_code(app_code: str) -> str: + def get_app_id_by_code(app_code: str, *, session: Session) -> str: """ Get app id by app code :param app_code: app code :return: app id """ - site = db.session.scalar(select(Site).where(Site.code == app_code).limit(1)) + site = session.scalar(select(Site).where(Site.code == app_code).limit(1)) if not site: raise ValueError(f"App with code {app_code} not found") return str(site.app_id) diff --git a/api/services/async_workflow_service.py b/api/services/async_workflow_service.py index ceda30e950f..601cad7557a 100644 --- a/api/services/async_workflow_service.py +++ b/api/services/async_workflow_service.py @@ -51,7 +51,7 @@ class AsyncWorkflowService: @classmethod def trigger_workflow_async( - cls, session: Session, user: Account | EndUser, trigger_data: TriggerData + cls, user: Account | EndUser, trigger_data: TriggerData, *, session: Session ) -> AsyncTriggerResponse: """ Universal entry point for async workflow execution - THIS METHOD WILL NOT BLOCK @@ -187,7 +187,7 @@ class AsyncWorkflowService: @classmethod def reinvoke_trigger( - cls, session: Session, user: Account | EndUser, workflow_trigger_log_id: str + cls, user: Account | EndUser, workflow_trigger_log_id: str, *, session: Session ) -> AsyncTriggerResponse: """ Re-invoke a previously failed or rate-limited trigger - THIS METHOD WILL NOT BLOCK @@ -231,7 +231,7 @@ class AsyncWorkflowService: session.commit() # Re-trigger workflow (this will create a new trigger log) - return cls.trigger_workflow_async(session, user, trigger_data) + return cls.trigger_workflow_async(user, trigger_data, session=session) @classmethod def get_trigger_log( @@ -309,7 +309,8 @@ class AsyncWorkflowService: workflow_service: WorkflowService, app_model: App, workflow_id: str | None = None, - session: Session | None = None, + *, + session: Session, ) -> Workflow: """ Get workflow for the app @@ -317,9 +318,7 @@ class AsyncWorkflowService: Args: app_model: App model instance workflow_id: Optional specific workflow ID - session: Reuse this SQLAlchemy session for the lookup when provided, - so the caller's explicit session bears the connection cost - instead of Flask's request-scoped ``db.session``. + session: SQLAlchemy session used for the workflow lookup. Returns: Workflow instance diff --git a/api/services/audio_service.py b/api/services/audio_service.py index 86c56e60a13..52c71edd576 100644 --- a/api/services/audio_service.py +++ b/api/services/audio_service.py @@ -6,7 +6,7 @@ from typing import cast from flask import Response, stream_with_context from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from constants import AUDIO_EXTENSIONS @@ -32,7 +32,7 @@ logger = logging.getLogger(__name__) class AudioService: @staticmethod - def _get_message_by_ref(session: Session | scoped_session, message_ref: MessageRef) -> Message | None: + def _get_message_by_ref(session: Session, message_ref: MessageRef) -> Message | None: stmt = select(Message).where(Message.id == message_ref.message_id, Message.app_id == message_ref.app_id) if message_ref.end_user_id is not None: stmt = stmt.where(Message.from_end_user_id == message_ref.end_user_id) @@ -89,7 +89,7 @@ class AudioService: cls, app_model: App, *, - session: Session | scoped_session, + session: Session, text: str | None = None, voice: str | None = None, end_user: str | None = None, diff --git a/api/services/auth/api_key_auth_service.py b/api/services/auth/api_key_auth_service.py index 42f1d4d8d40..f9ad7cf27b0 100644 --- a/api/services/auth/api_key_auth_service.py +++ b/api/services/auth/api_key_auth_service.py @@ -11,7 +11,7 @@ from services.auth.api_key_auth_factory import ApiKeyAuthFactory class ApiKeyAuthService: @staticmethod - def get_provider_auth_list(session: Session, tenant_id: str): + def get_provider_auth_list(tenant_id: str, *, session: Session): data_source_api_key_bindings = session.scalars( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, DataSourceApiKeyAuthBinding.disabled.is_(False) @@ -20,7 +20,7 @@ class ApiKeyAuthService: return data_source_api_key_bindings @staticmethod - def create_provider_auth(session: Session, tenant_id: str, args: dict[str, Any]): + def create_provider_auth(tenant_id: str, args: dict[str, Any], *, session: Session): auth_result = ApiKeyAuthFactory(args["provider"], args["credentials"]).validate_credentials() if auth_result: # Encrypt the api key @@ -35,7 +35,7 @@ class ApiKeyAuthService: session.commit() @staticmethod - def get_auth_credentials(session: Session, tenant_id: str, category: str, provider: str): + def get_auth_credentials(tenant_id: str, category: str, provider: str, *, session: Session): data_source_api_key_bindings = session.scalar( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, @@ -52,7 +52,7 @@ class ApiKeyAuthService: return credentials @staticmethod - def delete_provider_auth(session: Session, tenant_id: str, binding_id: str): + def delete_provider_auth(tenant_id: str, binding_id: str, *, session: Session): data_source_api_key_binding = session.scalar( select(DataSourceApiKeyAuthBinding).where( DataSourceApiKeyAuthBinding.tenant_id == tenant_id, diff --git a/api/services/billing_service.py b/api/services/billing_service.py index ec391e51676..2ee7179f432 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -7,7 +7,7 @@ from typing import Any, Literal, NotRequired, TypedDict import httpx from pydantic import TypeAdapter from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from tenacity import retry, retry_if_exception_type, stop_before_delay, wait_fixed from werkzeug.exceptions import InternalServerError @@ -363,7 +363,7 @@ class BillingService: return response.json() @staticmethod - def is_tenant_owner_or_admin(session: Session | scoped_session, current_user: Account): + def is_tenant_owner_or_admin(current_user: Account, *, session: Session): tenant_id = current_user.current_tenant_id join: TenantAccountJoin | None = session.scalar( diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index 557ae8e89f3..7c3b8d451c5 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -8,16 +8,13 @@ from sqlalchemy.orm import Session from configs import dify_config from core.app.entities.app_invoke_entities import InvokeFrom -from core.db.session_factory import session_factory from core.llm_generator.llm_generator import LLMGenerator -from extensions.ext_database import db from factories import variable_factory from graphon.variables.types import SegmentType from libs.datetime_utils import naive_utc_now from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account, ConversationVariable from models.model import App, Conversation, EndUser, Message -from services.conversation_variable_updater import ConversationVariableUpdater from services.errors.conversation import ( ConversationNotExistsError, ConversationVariableNotExistsError, @@ -122,24 +119,26 @@ class ConversationService: user: Account | EndUser | None, name: str | None, auto_generate: bool, + *, + session: Session, ): - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) if auto_generate: - return cls.auto_generate_name(app_model, conversation) + return cls.auto_generate_name(app_model, conversation, session=session) else: if name is None: raise ValueError("name is required when auto_generate is false") conversation.name = name conversation.updated_at = naive_utc_now() - db.session.commit() + session.commit() return conversation @classmethod - def auto_generate_name(cls, app_model: App, conversation: Conversation): + def auto_generate_name(cls, app_model: App, conversation: Conversation, *, session: Session): # get conversation first message - message = db.session.scalar( + message = session.scalar( select(Message) .where(Message.app_id == app_model.id, Message.conversation_id == conversation.id) .order_by(Message.created_at.asc()) @@ -156,13 +155,15 @@ class ConversationService: ) conversation.name = name - db.session.commit() + session.commit() return conversation @classmethod - def get_conversation(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): - conversation = db.session.scalar( + def get_conversation( + cls, app_model: App, conversation_id: str, user: Account | EndUser | None, *, session: Session + ): + conversation = session.scalar( select(Conversation) .where( Conversation.id == conversation_id, @@ -181,14 +182,14 @@ class ConversationService: return conversation @classmethod - def delete(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def delete(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, *, session: Session): """ Delete a conversation only if it belongs to the given user and app context. Raises: ConversationNotExistsError: When the conversation is not visible to the current user. """ - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) try: logger.info( @@ -197,13 +198,13 @@ class ConversationService: conversation_id, ) - db.session.delete(conversation) - db.session.commit() + session.delete(conversation) + session.commit() delete_conversation_related_data.delay(conversation.id) except Exception as e: - db.session.rollback() + session.rollback() raise e @classmethod @@ -215,8 +216,10 @@ class ConversationService: limit: int, last_id: str | None, variable_name: str | None = None, + *, + session: Session, ) -> InfiniteScrollPagination: - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) stmt = ( select(ConversationVariable) @@ -245,18 +248,17 @@ class ConversationService: ) ) - with session_factory.create_session() as session: - if last_id: - last_variable = session.scalar(stmt.where(ConversationVariable.id == last_id)) - if not last_variable: - raise ConversationVariableNotExistsError() + if last_id: + last_variable = session.scalar(stmt.where(ConversationVariable.id == last_id)) + if not last_variable: + raise ConversationVariableNotExistsError() - # Filter for variables created after the last_id - stmt = stmt.where(ConversationVariable.created_at > last_variable.created_at) + # Filter for variables created after the last_id + stmt = stmt.where(ConversationVariable.created_at > last_variable.created_at) - # Apply limit to query: fetch one extra row to determine has_more - query_stmt = stmt.limit(limit + 1) - rows = session.scalars(query_stmt).all() + # Apply limit to query: fetch one extra row to determine has_more + query_stmt = stmt.limit(limit + 1) + rows = session.scalars(query_stmt).all() has_more = False if len(rows) > limit: @@ -282,6 +284,8 @@ class ConversationService: variable_id: str, user: Account | EndUser | None, new_value: Any, + *, + session: Session, ): """ Update a conversation variable's value. @@ -302,7 +306,7 @@ class ConversationService: ConversationVariableTypeMismatchError: If the new value type doesn't match the variable's expected type """ # Verify conversation exists and user has access - conversation = cls.get_conversation(app_model, conversation_id, user) + conversation = cls.get_conversation(app_model, conversation_id, user, session=session) # Get the existing conversation variable stmt = ( @@ -312,48 +316,43 @@ class ConversationService: .where(ConversationVariable.id == variable_id) ) - with session_factory.create_session() as session: - existing_variable = session.scalar(stmt) - if not existing_variable: - raise ConversationVariableNotExistsError() + existing_variable = session.scalar(stmt) + if not existing_variable: + raise ConversationVariableNotExistsError() - # Convert existing variable to Variable object - current_variable = existing_variable.to_variable() + # Convert existing variable to Variable object + current_variable = existing_variable.to_variable() - # Validate that the new value type matches the expected variable type - expected_type = SegmentType(current_variable.value_type) + # Validate that the new value type matches the expected variable type + expected_type = SegmentType(current_variable.value_type) - # There is showing number in web ui but int in db - if expected_type == SegmentType.INTEGER: - expected_type = SegmentType.NUMBER + # There is showing number in web ui but int in db + if expected_type == SegmentType.INTEGER: + expected_type = SegmentType.NUMBER - if not expected_type.is_valid(new_value): - inferred_type = SegmentType.infer_segment_type(new_value) - raise ConversationVariableTypeMismatchError( - f"Type mismatch: variable '{current_variable.name}' expects {expected_type.value}, " - f"but got {inferred_type.value if inferred_type else 'unknown'} type" - ) + if not expected_type.is_valid(new_value): + inferred_type = SegmentType.infer_segment_type(new_value) + raise ConversationVariableTypeMismatchError( + f"Type mismatch: variable '{current_variable.name}' expects {expected_type.value}, " + f"but got {inferred_type.value if inferred_type else 'unknown'} type" + ) - # Create updated variable with new value only, preserving everything else - updated_variable_dict = { - "id": current_variable.id, - "name": current_variable.name, - "description": current_variable.description, - "value_type": current_variable.value_type, - "value": new_value, - "selector": current_variable.selector, - } + # Create updated variable with new value only, preserving everything else + updated_variable_dict = { + "id": current_variable.id, + "name": current_variable.name, + "description": current_variable.description, + "value_type": current_variable.value_type, + "value": new_value, + "selector": current_variable.selector, + } - updated_variable = variable_factory.build_conversation_variable_from_mapping(updated_variable_dict) + updated_variable = variable_factory.build_conversation_variable_from_mapping(updated_variable_dict) + existing_variable.data = updated_variable.model_dump_json() + session.commit() - # Use the conversation variable updater to persist the changes - updater = ConversationVariableUpdater(session_factory.get_session_maker()) - updater.update(conversation_id, updated_variable) - updater.flush() - - # Return the updated variable data - return { - "created_at": existing_variable.created_at, - "updated_at": naive_utc_now(), # Update timestamp - **updated_variable.model_dump(), - } + return { + "created_at": existing_variable.created_at, + "updated_at": naive_utc_now(), # Update timestamp + **updated_variable.model_dump(), + } diff --git a/api/services/credential_permission_service.py b/api/services/credential_permission_service.py index 2b1082d132b..d9ce5e7c502 100644 --- a/api/services/credential_permission_service.py +++ b/api/services/credential_permission_service.py @@ -1,7 +1,7 @@ from collections.abc import Sequence from sqlalchemy import or_, select -from sqlalchemy.orm import InstrumentedAttribute, Session, scoped_session +from sqlalchemy.orm import InstrumentedAttribute, Session from models.account import Account from models.credential_permission import CredentialPermission @@ -16,9 +16,7 @@ class CredentialPermissionService: """ @classmethod - def get_partial_member_list( - cls, session: Session | scoped_session, credential_id: str, credential_type: str - ) -> Sequence[str]: + def get_partial_member_list(cls, credential_id: str, credential_type: str, *, session: Session) -> Sequence[str]: """Return account_ids that have partial-member access to a credential.""" return session.scalars( select(CredentialPermission.account_id).where( diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index 94515309e79..afc49181185 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -12,9 +12,7 @@ from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config -from core.db.session_factory import session_factory from core.errors.error import QuotaExceededError -from extensions.ext_database import db from extensions.ext_redis import redis_client from models import TenantCreditPool from models.enums import ProviderQuotaType @@ -66,7 +64,7 @@ class CreditPoolService: ) @classmethod - def create_default_pool(cls, tenant_id: str) -> TenantCreditPool: + def create_default_pool(cls, tenant_id: str, session: Session) -> TenantCreditPool: """create default credit pool for new tenant""" credit_pool = TenantCreditPool( tenant_id=tenant_id, @@ -74,22 +72,21 @@ class CreditPoolService: quota_used=0, pool_type=ProviderQuotaType.TRIAL, ) - db.session.add(credit_pool) - db.session.commit() + session.add(credit_pool) + session.commit() return credit_pool @classmethod - def get_pool(cls, tenant_id: str, pool_type: str = "trial") -> TenantCreditPool | None: + def get_pool(cls, tenant_id: str, pool_type: str = "trial", *, session: Session) -> TenantCreditPool | None: """get tenant credit pool""" - with session_factory.get_session_maker().begin() as session: - return session.scalar( - select(TenantCreditPool) - .where( - TenantCreditPool.tenant_id == tenant_id, - TenantCreditPool.pool_type == pool_type, - ) - .limit(1) + return session.scalar( + select(TenantCreditPool) + .where( + TenantCreditPool.tenant_id == tenant_id, + TenantCreditPool.pool_type == pool_type, ) + .limit(1) + ) @classmethod def check_credits_available( @@ -97,9 +94,11 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> bool: """check if credits are available without deducting""" - pool = cls.get_pool(tenant_id, pool_type) + pool = cls.get_pool(tenant_id, pool_type, session=session) if not pool: return False return pool.remaining_credits >= credits_required @@ -110,25 +109,27 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> int: """Deduct exactly the requested credits or raise without mutating the pool.""" if credits_required <= 0: return 0 def deduct() -> int: - with session_factory.get_session_maker().begin() as session: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) - if not pool: - raise QuotaExceededError("Credit pool not found") + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") - remaining_credits = pool.remaining_credits - if remaining_credits <= 0: - raise QuotaExceededError("No credits remaining") - if remaining_credits < credits_required: - raise QuotaExceededError("Insufficient credits remaining") + remaining_credits = pool.remaining_credits + if remaining_credits <= 0: + raise QuotaExceededError("No credits remaining") + if remaining_credits < credits_required: + raise QuotaExceededError("Insufficient credits remaining") - pool.quota_used += credits_required - return credits_required + pool.quota_used += credits_required + session.commit() + return credits_required try: return cls._deduct_with_tenant_lock(tenant_id, deduct) @@ -144,24 +145,26 @@ class CreditPoolService: tenant_id: str, credits_required: int, pool_type: str = "trial", + *, + session: Session, ) -> int: """Deduct up to the available balance and return the actual deducted credits.""" if credits_required <= 0: return 0 def deduct() -> int: - with session_factory.get_session_maker().begin() as session: - pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) - if not pool: - logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type) - return 0 + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + logger.warning("Credit pool not found, tenant_id=%s, pool_type=%s", tenant_id, pool_type) + return 0 - deducted_credits = min(credits_required, pool.remaining_credits) - if deducted_credits <= 0: - return 0 + deducted_credits = min(credits_required, pool.remaining_credits) + if deducted_credits <= 0: + return 0 - pool.quota_used += deducted_credits - return deducted_credits + pool.quota_used += deducted_credits + session.commit() + return deducted_credits try: return cls._deduct_with_tenant_lock(tenant_id, deduct) diff --git a/api/services/data_migration/export_service.py b/api/services/data_migration/export_service.py index f5d214d230b..c0233006690 100644 --- a/api/services/data_migration/export_service.py +++ b/api/services/data_migration/export_service.py @@ -120,8 +120,8 @@ class MigrationExportService: self.package_service = package_service or MigrationPackageService() self.dependency_discovery_service = dependency_discovery_service or DependencyDiscoveryService() - def export(self, session: Session, selection: ExportSelection) -> ExportResult: - tenant = self._get_tenant(session, selection) + def export(self, selection: ExportSelection, *, session: Session) -> ExportResult: + tenant = self._get_tenant(selection, session=session) package = self.package_service.build_empty_package( source_tenant_id=tenant.id, source_tenant_name=tenant.name, @@ -131,10 +131,12 @@ class MigrationExportService: report_items: list[ResourceReportItem] = [] discovered_dependencies: list[DiscoveredDependency] = [] - apps = self._selected_apps(session, tenant.id, selection) + apps = self._selected_apps(tenant.id, selection, session=session) exported_app_ids = {app.id for app in apps} for app in apps: - dsl_content = AppDslService.export_dsl(app_model=app, include_secret=selection.include_secrets) + dsl_content = AppDslService.export_dsl( + app_model=app, session=session, include_secret=selection.include_secrets + ) package.workflows.append( { "id": app.id, @@ -157,7 +159,6 @@ class MigrationExportService: report_items=report_items, ) self._export_workflow_tools( - session, tenant, self._provider_ids( selection.additional_workflow_tools, discovered_dependencies, DependencyKind.WORKFLOW_TOOL @@ -166,9 +167,9 @@ class MigrationExportService: exported_workflow_tools=package.workflow_tools, dependencies=package.dependencies, report_items=report_items, + session=session, ) self._export_mcp_tools( - session, tenant_id=tenant.id, provider_ids=self._provider_ids( selection.additional_mcp_tools, @@ -179,6 +180,7 @@ class MigrationExportService: exported_mcp_tools=package.mcp_tools, dependencies=package.dependencies, report_items=report_items, + session=session, ) self._record_dependency_metadata( self._dependencies_by_kind(discovered_dependencies, DependencyKind.BUILTIN_OR_PLUGIN_TOOL), @@ -195,7 +197,7 @@ class MigrationExportService: ), ) - def _get_tenant(self, session: Session, selection: ExportSelection) -> Tenant: + def _get_tenant(self, selection: ExportSelection, *, session: Session) -> Tenant: if selection.source_tenant_id: tenant = session.get(Tenant, selection.source_tenant_id) if tenant is None: @@ -214,7 +216,7 @@ class MigrationExportService: ) return tenants[0] - def _selected_apps(self, session: Session, tenant_id: str, selection: ExportSelection) -> list[App]: + def _selected_apps(self, tenant_id: str, selection: ExportSelection, *, session: Session) -> list[App]: query = sa.select(App).where(App.tenant_id == tenant_id, App.mode.in_(SUPPORTED_APP_MODES)) if not selection.export_all_apps: if not selection.app_ids: @@ -267,7 +269,6 @@ class MigrationExportService: def _export_workflow_tools( self, - session: Session, tenant: Tenant, provider_ids: Iterable[str], *, @@ -275,11 +276,12 @@ class MigrationExportService: exported_workflow_tools: list[dict[str, Any]], dependencies: list[dict[str, Any]], report_items: list[ResourceReportItem], + session: Session, ) -> None: provider_ids = self._dedupe(provider_ids) if not provider_ids: return - owner = self._get_tenant_owner(session, tenant.id) + owner = self._get_tenant_owner(tenant.id, session=session) if owner is None: for provider_id in provider_ids: report_items.append( @@ -330,7 +332,7 @@ class MigrationExportService: ResourceReportItem(ResourceType.WORKFLOW_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_tenant_owner(self, session: Session, tenant_id: str) -> Account | None: + def _get_tenant_owner(self, tenant_id: str, *, session: Session) -> Account | None: return session.scalar( sa.select(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) @@ -341,7 +343,6 @@ class MigrationExportService: def _export_mcp_tools( self, - session: Session, *, tenant_id: str, provider_ids: Iterable[str], @@ -349,6 +350,7 @@ class MigrationExportService: exported_mcp_tools: list[dict[str, Any]], dependencies: list[dict[str, Any]], report_items: list[ResourceReportItem], + session: Session, ) -> None: for provider_id in self._dedupe(provider_ids): if not include_secrets: @@ -359,7 +361,7 @@ class MigrationExportService: ) continue try: - provider = self._get_mcp_provider(session, tenant_id, provider_id) + provider = self._get_mcp_provider(tenant_id, provider_id, session=session) exported_mcp_tools.append(self._serialize_mcp_provider(provider)) report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider.name, "exported")) except Exception as exc: @@ -367,7 +369,7 @@ class MigrationExportService: ResourceReportItem(ResourceType.MCP_TOOL, provider_id, provider_id, "unresolved", str(exc)) ) - def _get_mcp_provider(self, session: Session, tenant_id: str, provider_id: str) -> MCPToolProvider: + def _get_mcp_provider(self, tenant_id: str, provider_id: str, *, session: Session) -> MCPToolProvider: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): predicates.append(MCPToolProvider.id == provider_id) diff --git a/api/services/data_migration/import_service.py b/api/services/data_migration/import_service.py index 3eb251bbaef..b3354413ba1 100644 --- a/api/services/data_migration/import_service.py +++ b/api/services/data_migration/import_service.py @@ -82,11 +82,11 @@ class ImportTargetResolver: "Target tenant must be provided by --target-tenant, import config, or package metadata." ) - def resolve(self, session: Session, request: ImportRequest) -> ImportTarget: + def resolve(self, request: ImportRequest, *, session: Session) -> ImportTarget: target_tenant_name = self.select_target_tenant_name(request) package_target = request.package.metadata.target_tenant or {} if request.cli_target_tenant or request.config_target_tenant: - tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(target_tenant_name, session=session) elif package_target.get("id") and self._is_uuid(package_target["id"]): tenant = session.get(Tenant, package_target["id"]) if tenant is not None and package_target.get("name") and tenant.name != package_target.get("name"): @@ -94,7 +94,7 @@ class ImportTargetResolver: f"Target tenant id/name mismatch: {package_target['id']} / {package_target['name']}" ) else: - tenant = self._resolve_tenant_by_id_or_name(session, target_tenant_name) + tenant = self._resolve_tenant_by_id_or_name(target_tenant_name, session=session) if tenant is None: raise MigrationDataError(f"Target tenant not found: {target_tenant_name}") @@ -123,7 +123,7 @@ class ImportTargetResolver: operator_email=account.email, ) - def _resolve_tenant_by_id_or_name(self, session: Session, value: str) -> Tenant | None: + def _resolve_tenant_by_id_or_name(self, value: str, *, session: Session) -> Tenant | None: if self._is_uuid(value): tenant = session.get(Tenant, value) if tenant is not None: @@ -149,8 +149,8 @@ class MigrationImportService: def __init__(self, *, target_resolver: ImportTargetResolver | None = None) -> None: self.target_resolver = target_resolver or ImportTargetResolver() - def import_package(self, session: Session, request: ImportRequest) -> ImportResult: - target = self.target_resolver.resolve(session, request) + def import_package(self, request: ImportRequest, *, session: Session) -> ImportResult: + target = self.target_resolver.resolve(request, session=session) options = request.options_override or request.package.metadata.import_options report_items = [ ResourceReportItem( @@ -165,7 +165,6 @@ class MigrationImportService: id_mapping_details: list[ResourceIdMapping] = [] self._import_api_tools( - session, request.package, target, options, @@ -173,14 +172,16 @@ class MigrationImportService: id_mapping, id_mapping_details, self._source_api_provider_ids_by_name(request.package), + session=session, ) - self._import_mcp_tools(session, request.package, target, options, report_items, id_mapping, id_mapping_details) - self._preflight_dependency_only_mcp(session, request.package, target, report_items) + self._import_mcp_tools( + request.package, target, options, report_items, id_mapping, id_mapping_details, session=session + ) + self._preflight_dependency_only_mcp(request.package, target, report_items, session=session) workflow_tool_app_ids = self._workflow_tool_source_app_ids(request.package) imported_workflow_ids: set[str] = set() if workflow_tool_app_ids: self._import_workflows( - session, request.package, target, options, @@ -189,12 +190,12 @@ class MigrationImportService: id_mapping_details=id_mapping_details, imported_workflow_ids=imported_workflow_ids, only_app_ids=workflow_tool_app_ids, + session=session, ) self._import_workflow_tools( - session, request.package, target, options, id_mapping, id_mapping_details, report_items + request.package, target, options, id_mapping, id_mapping_details, report_items, session=session ) self._import_workflows( - session, request.package, target, options, @@ -203,6 +204,7 @@ class MigrationImportService: id_mapping_details=id_mapping_details, imported_workflow_ids=imported_workflow_ids, skip_app_ids=imported_workflow_ids, + session=session, ) return ImportResult( report_items=report_items, @@ -218,7 +220,6 @@ class MigrationImportService: def _import_workflows( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -228,6 +229,8 @@ class MigrationImportService: imported_workflow_ids: set[str] | None = None, only_app_ids: set[str] | None = None, skip_app_ids: set[str] | None = None, + *, + session: Session, ) -> None: account = session.get(Account, target.operator_id) tenant = session.get(Tenant, target.tenant_id) @@ -248,7 +251,7 @@ class MigrationImportService: id_mapping, ) existing_app = ( - self._find_existing_app(session, app_id, target.tenant_id) + self._find_existing_app(app_id, target.tenant_id, session=session) if options.id_strategy == IdStrategy.PRESERVE_ID else None ) @@ -270,13 +273,13 @@ class MigrationImportService: continue imported_app_id = self._import_workflow_app( - session=session, account=account, workflow_data=workflow_data, dsl_content=dsl_content, app_id=app_id, existing_app=existing_app, options=options, + session=session, ) if app_id: self._record_id_mappings( @@ -290,7 +293,7 @@ class MigrationImportService: if imported_workflow_ids is not None: imported_workflow_ids.add(app_id) if options.create_app_api_token_on_import: - self._create_or_reuse_app_api_token(session, imported_app_id, target.tenant_id) + self._create_or_reuse_app_api_token(imported_app_id, target.tenant_id, session=session) report_items.append( ResourceReportItem( ResourceType.WORKFLOW, @@ -311,15 +314,15 @@ class MigrationImportService: def _import_workflow_app( self, *, - session: Session, account: Account, workflow_data: dict[str, object], dsl_content: str, app_id: str | None, existing_app: App | None, options: ImportOptions, + session: Session, ) -> str: - import_service = AppDslService(session) + import_service = AppDslService(cast(Session, session)) if existing_app is not None: import_result = import_service.import_app( account=account, @@ -408,12 +411,12 @@ class MigrationImportService: def _should_preserve_source_app_id(self, options: ImportOptions) -> bool: return options.id_strategy == IdStrategy.PRESERVE_ID - def _find_existing_app(self, session: Session, app_id: str | None, tenant_id: str) -> App | None: + def _find_existing_app(self, app_id: str | None, tenant_id: str, *, session: Session) -> App | None: if not self._is_uuid_string(app_id): return None return session.scalar(sa.select(App).where(App.id == app_id, App.tenant_id == tenant_id)) - def _create_or_reuse_app_api_token(self, session: Session, app_id: str, tenant_id: str) -> None: + def _create_or_reuse_app_api_token(self, app_id: str, tenant_id: str, *, session: Session) -> None: existing = session.scalar( sa.select(ApiToken).where( ApiToken.type == ApiTokenType.APP, @@ -433,7 +436,6 @@ class MigrationImportService: def _import_api_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -441,6 +443,8 @@ class MigrationImportService: id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], source_provider_ids_by_name: dict[str, set[str]], + *, + session: Session, ) -> None: for tool_data in package.tools: provider_name = self._required_string(tool_data, "provider_name", "api_tool") @@ -510,7 +514,7 @@ class MigrationImportService: icon=icon, ) status = "created" - target_provider = self._find_api_tool_provider(session, target.tenant_id, provider_name) + target_provider = self._find_api_tool_provider(target.tenant_id, provider_name, session=session) if target_provider is not None: self._record_id_mappings( id_mapping, @@ -522,7 +526,9 @@ class MigrationImportService: ) report_items.append(ResourceReportItem(ResourceType.API_TOOL, provider_name, provider_name, status)) - def _find_api_tool_provider(self, session: Session, tenant_id: str, provider_name: str) -> ApiToolProvider | None: + def _find_api_tool_provider( + self, tenant_id: str, provider_name: str, *, session: Session + ) -> ApiToolProvider | None: return session.scalar( sa.select(ApiToolProvider).where( ApiToolProvider.tenant_id == tenant_id, @@ -558,13 +564,14 @@ class MigrationImportService: def _import_workflow_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], report_items: list[ResourceReportItem], + *, + session: Session, ) -> None: if not package.workflow_tools: return @@ -574,7 +581,10 @@ class MigrationImportService: for workflow_tool_data in package.workflow_tools: app_id = self._optional_string(workflow_tool_data.get("app_id")) resolved_app_id = id_mapping.get(app_id or "", app_id) - if not resolved_app_id or self._find_existing_app(session, resolved_app_id, target.tenant_id) is None: + if ( + not resolved_app_id + or self._find_existing_app(resolved_app_id, target.tenant_id, session=session) is None + ): report_items.append( ResourceReportItem( ResourceType.WORKFLOW_TOOL, @@ -586,7 +596,7 @@ class MigrationImportService: ) continue try: - self._ensure_workflow_app_is_published(session, target, account, resolved_app_id) + self._ensure_workflow_app_is_published(target, account, resolved_app_id, session=session) except Exception as exc: report_items.append( ResourceReportItem( @@ -602,7 +612,7 @@ class MigrationImportService: tool_name = self._required_string(workflow_tool_data, "name", "workflow_tool") lookup_workflow_tool_id = workflow_tool_id if options.id_strategy == IdStrategy.PRESERVE_ID else None existing = self._find_existing_workflow_tool( - session, target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id + target.tenant_id, lookup_workflow_tool_id, tool_name, resolved_app_id, session=session ) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"Workflow tool already exists and conflict_strategy=fail: {tool_name}") @@ -669,7 +679,7 @@ class MigrationImportService: ) status = "created" target_provider = self._find_existing_workflow_tool( - session, target.tenant_id, import_id or None, tool_name, resolved_app_id + target.tenant_id, import_id or None, tool_name, resolved_app_id, session=session ) if target_provider is None: raise MigrationDataError(f"Workflow tool was not created: {tool_name}") @@ -686,9 +696,9 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.WORKFLOW_TOOL, identifier, tool_name, status)) def _ensure_workflow_app_is_published( - self, session: Session, target: ImportTarget, account: Account, app_id: str + self, target: ImportTarget, account: Account, app_id: str, *, session: Session ) -> None: - app = self._find_existing_app(session, app_id, target.tenant_id) + app = self._find_existing_app(app_id, target.tenant_id, session=session) if app is None: raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}") if app.workflow_id: @@ -714,20 +724,23 @@ class MigrationImportService: def _import_mcp_tools( self, - session: Session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, report_items: list[ResourceReportItem], id_mapping: dict[str, str], id_mapping_details: list[ResourceIdMapping], + *, + session: Session, ) -> None: for mcp_data in package.mcp_tools: name = self._required_string(mcp_data, "name", "mcp_tool") server_identifier = self._required_string(mcp_data, "server_identifier", "mcp_tool") provider_id = self._optional_string(mcp_data.get("id")) lookup_provider_id = provider_id if options.id_strategy == IdStrategy.PRESERVE_ID else None - existing = self._find_existing_mcp_tool(session, target.tenant_id, lookup_provider_id, server_identifier) + existing = self._find_existing_mcp_tool( + target.tenant_id, lookup_provider_id, server_identifier, session=session + ) if existing is not None and options.conflict_strategy == ConflictStrategy.FAIL: raise MigrationDataError(f"MCP tool already exists and conflict_strategy=fail: {name}") if existing is not None and options.conflict_strategy == ConflictStrategy.SKIP: @@ -743,7 +756,7 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, existing.id, name, "skipped")) continue - service = MCPToolManageService(session=session) + service = MCPToolManageService(session=cast(Session, session)) configuration = MCPConfiguration.model_validate(mcp_data.get("configuration") or {}) authentication = ( MCPAuthentication.model_validate(mcp_data["authentication"]) if mcp_data.get("authentication") else None @@ -784,7 +797,7 @@ class MigrationImportService: authentication=authentication, ) created_provider = self._find_existing_mcp_tool( - session, target.tenant_id, lookup_provider_id, server_identifier + target.tenant_id, lookup_provider_id, server_identifier, session=session ) if created_provider is None: raise MigrationDataError(f"MCP provider was not created: {name}") @@ -812,7 +825,12 @@ class MigrationImportService: provider.authed = True def _find_existing_mcp_tool( - self, session: Session, tenant_id: str, provider_id: str | None, server_identifier: str + self, + tenant_id: str, + provider_id: str | None, + server_identifier: str, + *, + session: Session, ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == server_identifier] if self._is_uuid_string(provider_id): @@ -831,7 +849,13 @@ class MigrationImportService: return True def _find_existing_workflow_tool( - self, session: Session, tenant_id: str, workflow_tool_id: str | None, tool_name: str, app_id: str + self, + tenant_id: str, + workflow_tool_id: str | None, + tool_name: str, + app_id: str, + *, + session: Session, ) -> WorkflowToolProvider | None: predicates = [WorkflowToolProvider.name == tool_name, WorkflowToolProvider.app_id == app_id] if self._is_uuid_string(workflow_tool_id): @@ -843,14 +867,21 @@ class MigrationImportService: ) def _preflight_dependency_only_mcp( - self, session: Session, package: MigrationPackage, target: ImportTarget, report_items: list[ResourceReportItem] + self, + package: MigrationPackage, + target: ImportTarget, + report_items: list[ResourceReportItem], + *, + session: Session, ) -> None: for dependency in package.dependencies: if dependency.get("kind") != DependencyKind.MCP_TOOL.value: continue provider_id = str(dependency.get("provider_id", dependency.get("id", ""))) provider_name = self._optional_string(dependency.get("provider_name") or dependency.get("name")) - existing = self._find_dependency_only_mcp_provider(session, target.tenant_id, provider_id, provider_name) + existing = self._find_dependency_only_mcp_provider( + target.tenant_id, provider_id, provider_name, session=session + ) report_name = f"mcp_tool {provider_name or getattr(existing, 'name', None) or provider_id}" if existing is not None: report_items.append( @@ -879,7 +910,12 @@ class MigrationImportService: ) def _find_dependency_only_mcp_provider( - self, session: Session, tenant_id: str, provider_id: str, provider_name: str | None + self, + tenant_id: str, + provider_id: str, + provider_name: str | None, + *, + session: Session, ) -> MCPToolProvider | None: predicates = [MCPToolProvider.server_identifier == provider_id] if self._is_uuid_string(provider_id): diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index b36926a32c5..dda5440f772 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -13,7 +13,7 @@ import sqlalchemy as sa from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator from redis.exceptions import LockNotOwnedError from sqlalchemy import ColumnElement, delete, exists, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound from configs import dify_config @@ -110,13 +110,6 @@ from tasks.sync_website_document_indexing_task import sync_website_document_inde logger = logging.getLogger(__name__) -def _session_for_helpers(session: scoped_session | Session) -> Session: - """Return a concrete SQLAlchemy session for helpers that do not accept scoped_session.""" - if isinstance(session, scoped_session): - return session() - return session - - class ProcessRulesDict(TypedDict): mode: ProcessRuleMode rules: dict[str, Any] @@ -244,11 +237,11 @@ class _EstimateArgs(BaseModel): class DatasetService: @staticmethod - def _can_manage_all_datasets(tenant_id: str, account_id: str) -> bool: + def _can_manage_all_datasets(tenant_id: str, account_id: str, *, session: Session) -> bool: if not dify_config.RBAC_ENABLED: return False - permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id) + permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id, session=session) workspace_permission_keys = getattr(getattr(permissions, "workspace", None), "permission_keys", []) or [] return "dataset.create_and_management" in workspace_permission_keys @@ -256,7 +249,7 @@ class DatasetService: def get_datasets( page, per_page, - session: scoped_session | Session, + session: Session, tenant_id=None, user=None, search=None, @@ -291,7 +284,9 @@ class DatasetService: return [], 0 else: if dify_config.RBAC_ENABLED: - can_manage_all_datasets = DatasetService._can_manage_all_datasets(str(tenant_id), str(user.id)) + can_manage_all_datasets = DatasetService._can_manage_all_datasets( + str(tenant_id), str(user.id), session=session + ) should_show_all_datasets = include_all and can_manage_all_datasets else: should_show_all_datasets = user.current_role == TenantAccountRole.OWNER and include_all @@ -361,7 +356,7 @@ class DatasetService: return datasets.items, datasets.total @staticmethod - def get_process_rules(dataset_id, session: scoped_session | Session) -> ProcessRulesDict: + def get_process_rules(dataset_id, session: Session) -> ProcessRulesDict: # get the latest process rule dataset_process_rule = session.execute( select(DatasetProcessRule) @@ -406,7 +401,6 @@ class DatasetService: @staticmethod def create_empty_dataset( - session: Session, tenant_id: str, name: str, description: str | None, @@ -420,6 +414,8 @@ class DatasetService: embedding_model_name: str | None = None, retrieval_model: RetrievalModel | None = None, summary_index_setting: dict[str, Any] | None = None, + *, + session: Session, ): # check if dataset name already exists if session.scalar(select(Dataset).where(Dataset.name == name, Dataset.tenant_id == tenant_id).limit(1)): @@ -473,7 +469,7 @@ class DatasetService: if provider == "external" and external_knowledge_api_id: external_knowledge_api = ExternalDatasetService.get_external_knowledge_api( - session, external_knowledge_api_id, tenant_id + external_knowledge_api_id, tenant_id, session=session ) if not external_knowledge_api: raise ValueError("External API template not found.") @@ -501,7 +497,7 @@ class DatasetService: def create_empty_rag_pipeline_dataset( tenant_id: str, rag_pipeline_dataset_create_entity: RagPipelineDatasetCreateEntity, - session: scoped_session | Session, + session: Session, ): if rag_pipeline_dataset_create_entity.name: # check if dataset name already exists @@ -549,7 +545,7 @@ class DatasetService: return dataset @staticmethod - def get_dataset(dataset_id, session: scoped_session | Session) -> Dataset | None: + def get_dataset(dataset_id, session: Session) -> Dataset | None: dataset: Dataset | None = session.get(Dataset, dataset_id) return dataset @@ -632,7 +628,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def update_dataset(session: Session, dataset_id, data, user): + def update_dataset(dataset_id, data, user, *, session: Session): """ Update dataset configuration and settings. @@ -672,7 +668,7 @@ class DatasetService: return DatasetService._update_internal_dataset(dataset, data, user, session) @staticmethod - def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str, session: scoped_session | Session): + def _has_dataset_same_name(tenant_id: str, dataset_id: str, name: str, session: Session): dataset = session.scalar( select(Dataset) .where( @@ -725,7 +721,7 @@ class DatasetService: if not external_knowledge_api_id: raise ValueError("External knowledge api id is required.") # Ensure the referenced external API template exists and belongs to the dataset tenant. - ExternalDatasetService.get_external_knowledge_api(session, external_knowledge_api_id, dataset.tenant_id) + ExternalDatasetService.get_external_knowledge_api(external_knowledge_api_id, dataset.tenant_id, session=session) # Update metadata fields dataset.updated_by = user.id if user else None dataset.updated_at = naive_utc_now() @@ -743,7 +739,7 @@ class DatasetService: @staticmethod def _update_external_knowledge_binding( - dataset_id, external_knowledge_id, external_knowledge_api_id, session: scoped_session | Session + dataset_id, external_knowledge_id, external_knowledge_api_id, session: Session ): """ Update external knowledge binding configuration. @@ -770,7 +766,7 @@ class DatasetService: session.add(external_knowledge_binding) @staticmethod - def _update_internal_dataset(dataset, data, user, session: scoped_session | Session): + def _update_internal_dataset(dataset, data, user, session: Session): """ Update internal dataset configuration. @@ -836,9 +832,7 @@ class DatasetService: return dataset @staticmethod - def _update_pipeline_knowledge_base_node_data( - dataset: Dataset, updata_user_id: str, session: scoped_session | Session - ): + def _update_pipeline_knowledge_base_node_data(dataset: Dataset, updata_user_id: str, session: Session): """ Update pipeline knowledge base node data. """ @@ -850,7 +844,7 @@ class DatasetService: return try: - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(session) published_workflow = rag_pipeline_service.get_published_workflow(pipeline) draft_workflow = rag_pipeline_service.get_draft_workflow(pipeline) @@ -921,7 +915,7 @@ class DatasetService: raise @staticmethod - def _handle_indexing_technique_change(dataset, data, filtered_data, session: scoped_session | Session): + def _handle_indexing_technique_change(dataset, data, filtered_data, session: Session): """ Handle changes in indexing technique and configure embedding models accordingly. @@ -955,7 +949,7 @@ class DatasetService: return None @staticmethod - def _configure_embedding_model_for_high_quality(data, filtered_data, session: scoped_session | Session): + def _configure_embedding_model_for_high_quality(data, filtered_data, session: Session): """ Configure embedding model settings for high quality indexing. @@ -992,9 +986,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def _handle_embedding_model_update_when_technique_unchanged( - dataset, data, filtered_data, session: scoped_session | Session - ): + def _handle_embedding_model_update_when_technique_unchanged(dataset, data, filtered_data, session: Session): """ Handle embedding model updates when indexing technique remains the same. @@ -1043,7 +1035,7 @@ class DatasetService: del filtered_data["embedding_model"] @staticmethod - def _update_embedding_model_settings(dataset, data, filtered_data, session: scoped_session | Session): + def _update_embedding_model_settings(dataset, data, filtered_data, session: Session): """ Update embedding model settings with new values. @@ -1078,7 +1070,7 @@ class DatasetService: return None @staticmethod - def _apply_new_embedding_settings(dataset, data, filtered_data, session: scoped_session | Session): + def _apply_new_embedding_settings(dataset, data, filtered_data, session: Session): """ Apply new embedding model settings to the dataset. @@ -1176,7 +1168,11 @@ class DatasetService: @staticmethod def update_rag_pipeline_dataset_settings( - session: Session, dataset: Dataset, knowledge_configuration: KnowledgeConfiguration, has_published: bool = False + dataset: Dataset, + knowledge_configuration: KnowledgeConfiguration, + has_published: bool = False, + *, + session: Session, ): if not current_user or not current_user.current_tenant_id: raise ValueError("Current user or current tenant not found") @@ -1335,7 +1331,7 @@ class DatasetService: deal_dataset_index_update_task.delay(dataset.id, action) @staticmethod - def delete_dataset(dataset_id, user, session: scoped_session | Session): + def delete_dataset(dataset_id, user, session: Session): dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: @@ -1350,12 +1346,12 @@ class DatasetService: return True @staticmethod - def dataset_use_check(dataset_id, session: scoped_session | Session) -> bool: + def dataset_use_check(dataset_id, session: Session) -> bool: stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_id)) return session.execute(stmt).scalar_one() @staticmethod - def check_dataset_permission(dataset, user, session: scoped_session | Session): + def check_dataset_permission(dataset, user, session: Session): """Validate dataset access for a user, using the injected session for partial-member lookups.""" if dataset.tenant_id != user.current_tenant_id: logger.debug("User %s does not have permission to access dataset %s", user.id, dataset.id) @@ -1378,7 +1374,7 @@ class DatasetService: @staticmethod def check_dataset_operator_permission( - user: Account | None = None, dataset: Dataset | None = None, *, session: scoped_session | Session + user: Account | None = None, dataset: Dataset | None = None, *, session: Session ): if not dataset: raise ValueError("Dataset not found") @@ -1409,7 +1405,7 @@ class DatasetService: return dataset_queries.items, dataset_queries.total @staticmethod - def get_related_apps(dataset_id: str, session: scoped_session | Session): + def get_related_apps(dataset_id: str, session: Session): return session.scalars( select(AppDatasetJoin) .where(AppDatasetJoin.dataset_id == dataset_id) @@ -1417,7 +1413,7 @@ class DatasetService: ).all() @staticmethod - def update_dataset_api_status(dataset_id: str, status: bool, session: scoped_session | Session): + def update_dataset_api_status(dataset_id: str, status: bool, session: Session): dataset = DatasetService.get_dataset(dataset_id, session) if dataset is None: raise NotFound("Dataset not found.") @@ -1429,7 +1425,7 @@ class DatasetService: session.commit() @staticmethod - def get_dataset_auto_disable_logs(dataset_id: str, session: scoped_session | Session) -> AutoDisableLogsDict: + def get_dataset_auto_disable_logs(dataset_id: str, session: Session) -> AutoDisableLogsDict: assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None features = FeatureService.get_features(current_user.current_tenant_id, exclude_vector_space=True) @@ -1628,9 +1624,7 @@ class DocumentService: } @staticmethod - def get_document( - dataset_id: str, document_id: str | None = None, *, session: scoped_session | Session - ) -> Document | None: + def get_document(dataset_id: str, document_id: str | None = None, *, session: Session) -> Document | None: """Fetch a document by id within a dataset using the caller-provided session.""" if document_id: document = session.scalar( @@ -1641,9 +1635,7 @@ class DocumentService: return None @staticmethod - def get_documents_by_ids( - dataset_id: str, document_ids: Sequence[str], session: scoped_session | Session - ) -> Sequence[Document]: + def get_documents_by_ids(dataset_id: str, document_ids: Sequence[str], session: Session) -> Sequence[Document]: """Fetch documents for a dataset in a single batch query.""" if not document_ids: return [] @@ -1661,7 +1653,7 @@ class DocumentService: def update_documents_need_summary( dataset_id: str, document_ids: Sequence[str], - session: scoped_session | Session, + session: Session, need_summary: bool = True, ) -> int: """ @@ -1705,7 +1697,7 @@ class DocumentService: return updated_count @staticmethod - def get_document_download_url(document: Document, session: scoped_session | Session) -> str: + def get_document_download_url(document: Document, session: Session) -> str: """ Return a signed download URL for an upload-file document. """ @@ -1717,6 +1709,7 @@ class DocumentService: documents: Sequence[Document], dataset: Dataset, tenant_id: str, + session: Session, ) -> None: """ Enrich documents with summary_index_status based on dataset summary index settings. @@ -1728,6 +1721,7 @@ class DocumentService: documents: List of Document instances to enrich dataset: Dataset instance containing summary_index_setting tenant_id: Tenant ID for summary status lookup + session: SQLAlchemy session used to read summary status records """ # Check if dataset has summary index enabled has_summary_index = dataset.summary_index_setting and dataset.summary_index_setting.get("enable") is True @@ -1745,6 +1739,7 @@ class DocumentService: document_ids=document_ids_need_summary, dataset_id=dataset.id, tenant_id=tenant_id, + session=session, ) # Add summary_index_status to each document @@ -1763,7 +1758,7 @@ class DocumentService: document_ids: Sequence[str], tenant_id: str, current_user: Account, - session: scoped_session | Session, + session: Session, ) -> tuple[list[UploadFile], str]: """ Resolve upload files for batch ZIP downloads and generate a client-visible filename. @@ -1814,7 +1809,7 @@ class DocumentService: return str(upload_file_id) @staticmethod - def _get_upload_file_for_upload_file_document(document: Document, session: scoped_session | Session) -> UploadFile: + def _get_upload_file_for_upload_file_document(document: Document, session: Session) -> UploadFile: """ Load the `UploadFile` row for an upload-file document. """ @@ -1823,9 +1818,7 @@ class DocumentService: invalid_source_message="Document does not have an uploaded file to download.", missing_file_message="Uploaded file not found.", ) - upload_files_by_id = FileService.get_upload_files_by_ids( - _session_for_helpers(session), document.tenant_id, [upload_file_id] - ) + upload_files_by_id = FileService.get_upload_files_by_ids(document.tenant_id, [upload_file_id], session=session) upload_file = upload_files_by_id.get(upload_file_id) if not upload_file: raise NotFound("Uploaded file not found.") @@ -1837,7 +1830,7 @@ class DocumentService: dataset_id: str, document_ids: Sequence[str], tenant_id: str, - session: scoped_session | Session, + session: Session, ) -> dict[str, UploadFile]: """ Batch load upload files keyed by document id for ZIP downloads. @@ -1865,9 +1858,7 @@ class DocumentService: upload_file_ids.append(upload_file_id) upload_file_ids_by_document_id[document_id] = upload_file_id - upload_files_by_id = FileService.get_upload_files_by_ids( - _session_for_helpers(session), tenant_id, upload_file_ids - ) + upload_files_by_id = FileService.get_upload_files_by_ids(tenant_id, upload_file_ids, session=session) missing_upload_file_ids: set[str] = set(upload_file_ids) - set(upload_files_by_id.keys()) if missing_upload_file_ids: raise NotFound("Only uploaded-file documents can be downloaded as ZIP.") @@ -1878,13 +1869,13 @@ class DocumentService: } @staticmethod - def get_document_by_id(document_id: str, session: scoped_session | Session) -> Document | None: + def get_document_by_id(document_id: str, session: Session) -> Document | None: document = session.get(Document, document_id) return document @staticmethod - def get_document_by_ids(document_ids: list[str], session: scoped_session | Session) -> Sequence[Document]: + def get_document_by_ids(document_ids: list[str], session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.id.in_(document_ids), @@ -1896,7 +1887,7 @@ class DocumentService: return documents @staticmethod - def get_document_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_document_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1907,7 +1898,7 @@ class DocumentService: return documents @staticmethod - def get_working_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_working_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1920,7 +1911,7 @@ class DocumentService: return documents @staticmethod - def get_error_documents_by_dataset_id(dataset_id: str, session: scoped_session | Session) -> Sequence[Document]: + def get_error_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( Document.dataset_id == dataset_id, @@ -1930,7 +1921,7 @@ class DocumentService: return documents @staticmethod - def get_batch_documents(dataset_id: str, batch: str, session: scoped_session | Session) -> Sequence[Document]: + def get_batch_documents(dataset_id: str, batch: str, session: Session) -> Sequence[Document]: assert isinstance(current_user, Account) documents = session.scalars( select(Document).where( @@ -1943,7 +1934,7 @@ class DocumentService: return documents @staticmethod - def get_document_file_detail(file_id: str, session: scoped_session | Session): + def get_document_file_detail(file_id: str, session: Session): file_detail = session.get(UploadFile, file_id) return file_detail @@ -1955,7 +1946,7 @@ class DocumentService: return False @staticmethod - def delete_document(document, session: scoped_session | Session): + def delete_document(document, session: Session): # trigger document_was_deleted signal file_id = None if document.data_source_type == DataSourceType.UPLOAD_FILE: @@ -1975,7 +1966,7 @@ class DocumentService: dataset_ref: DatasetRef, document_ids: list[str], doc_form: str | None, - session: scoped_session | Session, + session: Session, ): # Check if document_ids is not empty to avoid WHERE false condition if not document_ids or len(document_ids) == 0: @@ -2006,7 +1997,7 @@ class DocumentService: batch_clean_document_task.delay(deleted_document_ids, dataset_ref.dataset_id, doc_form, file_ids) @staticmethod - def rename_document(dataset_id: str, document_id: str, name: str, session: scoped_session | Session) -> Document: + def rename_document(dataset_id: str, document_id: str, name: str, session: Session) -> Document: assert isinstance(current_user, Account) dataset = DatasetService.get_dataset(dataset_id, session) @@ -2041,7 +2032,7 @@ class DocumentService: return document @staticmethod - def pause_document(document, session: scoped_session | Session): + def pause_document(document, session: Session): if document.indexing_status not in { IndexingStatus.WAITING, IndexingStatus.PARSING, @@ -2063,7 +2054,7 @@ class DocumentService: redis_client.setnx(indexing_cache_key, "True") @staticmethod - def recover_document(document, session: scoped_session | Session): + def recover_document(document, session: Session): if not document.is_paused: raise DocumentIndexingError() # update document to be recover @@ -2080,7 +2071,7 @@ class DocumentService: recover_document_indexing_task.delay(document.dataset_id, document.id) @staticmethod - def retry_document(dataset_id: str, documents: list[Document], session: scoped_session | Session): + def retry_document(dataset_id: str, documents: list[Document], session: Session): for document in documents: # add retry flag retry_indexing_cache_key = f"document_{document.id}_is_retried" @@ -2100,7 +2091,7 @@ class DocumentService: retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id) @staticmethod - def sync_website_document(dataset_id: str, document: Document, session: scoped_session | Session): + def sync_website_document(dataset_id: str, document: Document, session: Session): # add sync flag sync_indexing_cache_key = f"document_{document.id}_is_sync" cache_result = redis_client.get(sync_indexing_cache_key) @@ -2120,7 +2111,7 @@ class DocumentService: sync_website_document_indexing_task.delay(dataset_id, document.id) @staticmethod - def get_documents_position(dataset_id, session: scoped_session | Session): + def get_documents_position(dataset_id, session: Session): document = session.scalar( select(Document).where(Document.dataset_id == dataset_id).order_by(Document.position.desc()).limit(1) ) @@ -2137,7 +2128,7 @@ class DocumentService: dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, *, - session: scoped_session | Session, + session: Session, ) -> tuple[list[Document], str]: # check doc_form DatasetService.check_doc_form(dataset, knowledge_config.doc_form) @@ -2793,7 +2784,7 @@ class DocumentService: return document @staticmethod - def get_tenant_documents_count(session: scoped_session | Session): + def get_tenant_documents_count(*, session: Session): assert isinstance(current_user, Account) documents_count = ( @@ -2817,7 +2808,7 @@ class DocumentService: dataset_process_rule: DatasetProcessRule | None = None, created_from: str = DocumentCreatedFrom.WEB, *, - session: scoped_session | Session, + session: Session, ): assert isinstance(current_user, Account) @@ -2944,7 +2935,7 @@ class DocumentService: @staticmethod def save_document_without_dataset_id( - tenant_id: str, knowledge_config: KnowledgeConfig, account: Account, session: scoped_session | Session + tenant_id: str, knowledge_config: KnowledgeConfig, account: Account, session: Session ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3133,7 +3124,7 @@ class DocumentService: document_ids: list[str], action: Literal["enable", "disable", "archive", "un_archive"], user, - session: scoped_session | Session, + session: Session, ): """ Batch update document status. @@ -3216,7 +3207,7 @@ class DocumentService: document = update_info["document"] indexing_cache_key = f"document_{document.id}_indexing" redis_client.setex(indexing_cache_key, 600, 1) - except Exception as e: + except Exception: # Log the error but do not rollback the transaction logger.exception("Error setting cache for document %s", update_info["document"].id) # Raise any propagation error after all updates @@ -3340,9 +3331,7 @@ class SegmentService: raise ValueError(f"Exceeded maximum attachment limit of {single_chunk_attachment_limit}") @classmethod - def create_segment( - cls, args: dict[str, Any], document: Document, dataset: Dataset, session: scoped_session | Session - ): + def create_segment(cls, args: dict[str, Any], document: Document, dataset: Dataset, session: Session): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3408,7 +3397,13 @@ class SegmentService: try: keywords = args.get("keywords") keywords_list = [keywords] if keywords is not None else None - VectorService.create_segments_vector(keywords_list, [segment_document], dataset, document.doc_form) + VectorService.create_segments_vector( + keywords_list, + [segment_document], + dataset, + document.doc_form, + session, + ) except Exception as e: logger.exception("create segment index failed") segment_document.enabled = False @@ -3422,9 +3417,7 @@ class SegmentService: pass @classmethod - def multi_create_segment( - cls, segments: list, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def multi_create_segment(cls, segments: list, document: Document, dataset: Dataset, session: Session): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3498,7 +3491,11 @@ class SegmentService: try: # save vector index VectorService.create_segments_vector( - keywords_list, pre_segment_data_list, dataset, document.doc_form + keywords_list, + pre_segment_data_list, + dataset, + document.doc_form, + session, ) except Exception as e: logger.exception("create segment index failed") @@ -3519,7 +3516,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ): assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -3597,7 +3594,13 @@ class SegmentService: processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( - segment, document, dataset, embedding_model_instance, processing_rule, True + segment, + document, + dataset, + embedding_model_instance, + processing_rule, + session, + True, ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): if args.enabled or keyword_changed: @@ -3628,7 +3631,12 @@ class SegmentService: from services.summary_index_service import SummaryIndexService try: - SummaryIndexService.update_summary_for_segment(segment, dataset, args.summary) + SummaryIndexService.update_summary_for_segment( + segment, + dataset, + args.summary, + session=session, + ) except Exception: logger.exception("Failed to update summary for segment %s", segment.id) # Don't fail the entire update if summary update fails @@ -3697,7 +3705,13 @@ class SegmentService: processing_rule = session.get(DatasetProcessRule, document.dataset_process_rule_id) if processing_rule: VectorService.generate_child_chunks( - segment, document, dataset, embedding_model_instance, processing_rule, True + segment, + document, + dataset, + embedding_model_instance, + processing_rule, + session, + True, ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): # update segment vector index @@ -3728,7 +3742,10 @@ class SegmentService: try: SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, dataset.summary_index_setting + segment, + dataset, + dataset.summary_index_setting, + session=session, ) logger.info("Auto-regenerated summary for segment %s after content change", segment.id) except Exception: @@ -3743,7 +3760,12 @@ class SegmentService: from services.summary_index_service import SummaryIndexService try: - SummaryIndexService.update_summary_for_segment(segment, dataset, args.summary) + SummaryIndexService.update_summary_for_segment( + segment, + dataset, + args.summary, + session=session, + ) logger.info("Updated summary for segment %s with user-provided content", segment.id) except Exception: logger.exception("Failed to update summary for segment %s", segment.id) @@ -3760,7 +3782,10 @@ class SegmentService: try: SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, dataset.summary_index_setting + segment, + dataset, + dataset.summary_index_setting, + session=session, ) logger.info( "Regenerated summary for segment %s after content change (summary unchanged)", @@ -3770,7 +3795,7 @@ class SegmentService: logger.exception("Failed to regenerate summary for segment %s", segment.id) # Don't fail the entire update if summary regeneration fails # update multimodel vector index - VectorService.update_multimodel_vector(segment, args.attachment_ids or [], dataset) + VectorService.update_multimodel_vector(segment, args.attachment_ids or [], dataset, session) except Exception as e: logger.exception("update segment index failed") segment.enabled = False @@ -3784,9 +3809,7 @@ class SegmentService: return new_segment @classmethod - def delete_segment( - cls, segment: DocumentSegment, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def delete_segment(cls, segment: DocumentSegment, document: Document, dataset: Dataset, session: Session): indexing_cache_key = f"segment_{segment.id}_delete_indexing" cache_result = redis_client.get(indexing_cache_key) if cache_result is not None: @@ -3821,9 +3844,7 @@ class SegmentService: session.commit() @classmethod - def delete_segments( - cls, segment_ids: list, document: Document, dataset: Dataset, session: scoped_session | Session - ): + def delete_segments(cls, segment_ids: list, document: Document, dataset: Dataset, session: Session): assert current_user is not None # Check if segment_ids is not empty to avoid WHERE false condition if not segment_ids or len(segment_ids) == 0: @@ -3882,7 +3903,7 @@ class SegmentService: action: Literal["enable", "disable"], dataset: Dataset, document: Document, - session: scoped_session | Session, + session: Session, ): assert current_user is not None @@ -3948,7 +3969,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> ChildChunk: assert isinstance(current_user, Account) @@ -3997,7 +4018,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> list[ChildChunk]: assert isinstance(current_user, Account) child_chunks = session.scalars( @@ -4072,7 +4093,7 @@ class SegmentService: segment: DocumentSegment, document: Document, dataset: Dataset, - session: scoped_session | Session, + session: Session, ) -> ChildChunk: assert current_user is not None @@ -4092,7 +4113,7 @@ class SegmentService: return child_chunk @classmethod - def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset, session: scoped_session | Session): + def delete_child_chunk(cls, child_chunk: ChildChunk, dataset: Dataset, session: Session): session.delete(child_chunk) try: VectorService.delete_child_chunk_vector(child_chunk, dataset) @@ -4124,9 +4145,7 @@ class SegmentService: return paginate_query(query, page=page, per_page=limit, max_per_page=100) @classmethod - def get_child_chunk_by_id( - cls, child_chunk_id: str, tenant_id: str, session: scoped_session | Session - ) -> ChildChunk | None: + def get_child_chunk_by_id(cls, child_chunk_id: str, tenant_id: str, session: Session) -> ChildChunk | None: """Get a child chunk by its ID.""" result = session.scalar( select(ChildChunk).where(ChildChunk.id == child_chunk_id, ChildChunk.tenant_id == tenant_id).limit(1) @@ -4134,9 +4153,11 @@ class SegmentService: return result if isinstance(result, ChildChunk) else None @classmethod - def get_child_chunk_by_segment_ref(cls, child_chunk_id: str, segment_ref: SegmentRef) -> ChildChunk | None: + def get_child_chunk_by_segment_ref( + cls, child_chunk_id: str, segment_ref: SegmentRef, session: Session + ) -> ChildChunk | None: """Get a child chunk through the full tenant/dataset/document/segment chain.""" - result = db.session.scalar( + result = session.scalar( select(ChildChunk) .where( ChildChunk.id == child_chunk_id, @@ -4178,9 +4199,7 @@ class SegmentService: return paginated_segments.items, paginated_segments.total @classmethod - def get_segment_by_id( - cls, segment_id: str, tenant_id: str, session: scoped_session | Session - ) -> DocumentSegment | None: + def get_segment_by_id(cls, segment_id: str, tenant_id: str, session: Session) -> DocumentSegment | None: """Get a segment by its ID.""" result = session.scalar( select(DocumentSegment) @@ -4190,9 +4209,9 @@ class SegmentService: return result if isinstance(result, DocumentSegment) else None @classmethod - def get_segment_by_ref(cls, segment_ref: SegmentRef) -> DocumentSegment | None: + def get_segment_by_ref(cls, segment_ref: SegmentRef, session: Session) -> DocumentSegment | None: """Get a segment through the full tenant/dataset/document ownership chain.""" - result = db.session.scalar( + result = session.scalar( select(DocumentSegment) .where( DocumentSegment.id == segment_ref.segment_id, @@ -4209,7 +4228,7 @@ class SegmentService: cls, document_id: str, dataset_id: str, - session: scoped_session | Session, + session: Session, status: str | None = None, enabled: bool | None = None, ) -> Sequence[DocumentSegment]: @@ -4242,7 +4261,7 @@ class SegmentService: class DatasetCollectionBindingService: @classmethod def get_dataset_collection_binding( - cls, provider_name: str, model_name: str, session: scoped_session | Session, collection_type: str = "dataset" + cls, provider_name: str, model_name: str, session: Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) @@ -4268,7 +4287,7 @@ class DatasetCollectionBindingService: @classmethod def get_dataset_collection_binding_by_id_and_type( - cls, collection_binding_id: str, session: scoped_session | Session, collection_type: str = "dataset" + cls, collection_binding_id: str, session: Session, collection_type: str = "dataset" ) -> DatasetCollectionBinding: dataset_collection_binding = session.scalar( select(DatasetCollectionBinding) @@ -4286,7 +4305,7 @@ class DatasetCollectionBindingService: class DatasetPermissionService: @classmethod - def get_dataset_partial_member_list(cls, dataset_id, session: scoped_session | Session): + def get_dataset_partial_member_list(cls, dataset_id, session: Session): user_list_query = session.scalars( select( DatasetPermission.account_id, @@ -4296,7 +4315,7 @@ class DatasetPermissionService: return user_list_query @classmethod - def update_partial_member_list(cls, tenant_id, dataset_id, user_list, session: scoped_session | Session): + def update_partial_member_list(cls, tenant_id, dataset_id, user_list, session: Session): try: session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) permissions = [] @@ -4315,7 +4334,7 @@ class DatasetPermissionService: raise e @classmethod - def check_permission(cls, session: Session, user, dataset, requested_permission, requested_partial_member_list): + def check_permission(cls, user, dataset, requested_permission, requested_partial_member_list, *, session: Session): if not user.is_dataset_editor: raise NoPermissionError("User does not have permission to edit this dataset.") @@ -4332,7 +4351,7 @@ class DatasetPermissionService: raise ValueError("Dataset operators cannot change the dataset permissions.") @classmethod - def clear_partial_member_list(cls, dataset_id, session: scoped_session | Session): + def clear_partial_member_list(cls, dataset_id, session: Session): try: session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) session.commit() diff --git a/api/services/datasource_provider_service.py b/api/services/datasource_provider_service.py index 12807a41f04..5de194dd262 100644 --- a/api/services/datasource_provider_service.py +++ b/api/services/datasource_provider_service.py @@ -446,12 +446,14 @@ class DatasourceProviderService: is not None ) - def is_tenant_oauth_params_enabled(self, tenant_id: str, datasource_provider_id: DatasourceProviderID) -> bool: + def is_tenant_oauth_params_enabled( + self, tenant_id: str, datasource_provider_id: DatasourceProviderID, *, session: Session + ) -> bool: """ check if tenant oauth params is enabled """ return ( - db.session.scalar( + session.scalar( select(func.count(DatasourceOauthTenantParamConfig.id)).where( DatasourceOauthTenantParamConfig.tenant_id == tenant_id, DatasourceOauthTenantParamConfig.provider == datasource_provider_id.provider_name, @@ -463,12 +465,17 @@ class DatasourceProviderService: ) > 0 def get_tenant_oauth_client( - self, tenant_id: str, datasource_provider_id: DatasourceProviderID, mask: bool = False + self, + tenant_id: str, + datasource_provider_id: DatasourceProviderID, + mask: bool = False, + *, + session: Session, ) -> Mapping[str, Any] | None: """ get tenant oauth client """ - tenant_oauth_client_params = db.session.scalar( + tenant_oauth_client_params = session.scalar( select(DatasourceOauthTenantParamConfig) .where( DatasourceOauthTenantParamConfig.tenant_id == tenant_id, @@ -547,7 +554,7 @@ class DatasourceProviderService: @staticmethod def generate_next_datasource_provider_name( - session: Session, tenant_id: str, provider_id: DatasourceProviderID, credential_type: CredentialType + tenant_id: str, provider_id: DatasourceProviderID, credential_type: CredentialType, *, session: Session ) -> str: db_providers = session.scalars( select(DatasourceProvider).where( @@ -800,6 +807,8 @@ class DatasourceProviderService: provider: str, plugin_id: str, user: "Account | None" = None, + *, + session: Session, ) -> list[dict]: """ list datasource credentials with obfuscated sensitive fields, @@ -829,11 +838,11 @@ class DatasourceProviderService: credential_type=CredPermType.DATASOURCE_PROVIDER, user=user, ) - datasource_providers: list[DatasourceProvider] = list(db.session.scalars(query).all()) + datasource_providers: list[DatasourceProvider] = list(session.scalars(query).all()) if not datasource_providers: return [] copy_credentials_list = [] - default_provider = db.session.execute( + default_provider = session.execute( select(DatasourceProvider.id) .where( DatasourceProvider.tenant_id == tenant_id, @@ -870,7 +879,7 @@ class DatasourceProviderService: return copy_credentials_list - def get_all_datasource_credentials(self, tenant_id: str) -> list[dict]: + def get_all_datasource_credentials(self, tenant_id: str, *, session: Session) -> list[dict]: """ get datasource credentials. @@ -883,7 +892,10 @@ class DatasourceProviderService: for datasource in datasources: datasource_provider_id = DatasourceProviderID(f"{datasource.plugin_id}/{datasource.provider}") credentials = self.list_datasource_credentials( - tenant_id=tenant_id, provider=datasource.provider, plugin_id=datasource.plugin_id + tenant_id=tenant_id, + provider=datasource.provider, + plugin_id=datasource.plugin_id, + session=session, ) redirect_uri = ( f"{dify_config.CONSOLE_API_URL}/console/api/oauth/plugin/{datasource_provider_id}/datasource/callback" @@ -912,10 +924,10 @@ class DatasourceProviderService: for credential_schema in datasource.declaration.oauth_schema.credentials_schema ], "oauth_custom_client_params": self.get_tenant_oauth_client( - tenant_id, datasource_provider_id, mask=True + tenant_id, datasource_provider_id, mask=True, session=session ), "is_oauth_custom_client_enabled": self.is_tenant_oauth_params_enabled( - tenant_id, datasource_provider_id + tenant_id, datasource_provider_id, session=session ), "is_system_oauth_params_exists": self.is_system_oauth_params_exist(datasource_provider_id), "redirect_uri": redirect_uri, @@ -926,7 +938,7 @@ class DatasourceProviderService: ) return datasource_credentials - def get_hard_code_datasource_credentials(self, tenant_id: str) -> list[dict]: + def get_hard_code_datasource_credentials(self, tenant_id: str, *, session: Session) -> list[dict]: """ get hard code datasource credentials. @@ -945,7 +957,10 @@ class DatasourceProviderService: ]: datasource_provider_id = DatasourceProviderID(f"{datasource.plugin_id}/{datasource.provider}") credentials = self.list_datasource_credentials( - tenant_id=tenant_id, provider=datasource.provider, plugin_id=datasource.plugin_id + tenant_id=tenant_id, + provider=datasource.provider, + plugin_id=datasource.plugin_id, + session=session, ) redirect_uri = "{}/console/api/oauth/plugin/{}/datasource/callback".format( dify_config.CONSOLE_API_URL, datasource_provider_id @@ -974,10 +989,10 @@ class DatasourceProviderService: for credential_schema in datasource.declaration.oauth_schema.credentials_schema ], "oauth_custom_client_params": self.get_tenant_oauth_client( - tenant_id, datasource_provider_id, mask=True + tenant_id, datasource_provider_id, mask=True, session=session ), "is_oauth_custom_client_enabled": self.is_tenant_oauth_params_enabled( - tenant_id, datasource_provider_id + tenant_id, datasource_provider_id, session=session ), "is_system_oauth_params_exists": self.is_system_oauth_params_exist(datasource_provider_id), "redirect_uri": redirect_uri, @@ -988,7 +1003,9 @@ class DatasourceProviderService: ) return datasource_credentials - def get_real_datasource_credentials(self, tenant_id: str, provider: str, plugin_id: str) -> list[dict]: + def get_real_datasource_credentials( + self, tenant_id: str, provider: str, plugin_id: str, *, session: Session + ) -> list[dict]: """ get datasource credentials. @@ -998,7 +1015,7 @@ class DatasourceProviderService: """ # Get all provider configurations of the current workspace datasource_providers: list[DatasourceProvider] = list( - db.session.scalars( + session.scalars( select(DatasourceProvider).where( DatasourceProvider.tenant_id == tenant_id, DatasourceProvider.provider == provider, @@ -1110,7 +1127,9 @@ class DatasourceProviderService: datasource_provider.encrypted_credentials = encrypted_credentials - def remove_datasource_credentials(self, tenant_id: str, auth_id: str, provider: str, plugin_id: str) -> None: + def remove_datasource_credentials( + self, tenant_id: str, auth_id: str, provider: str, plugin_id: str, *, session: Session + ) -> None: """ remove datasource credentials. @@ -1119,7 +1138,7 @@ class DatasourceProviderService: :param plugin_id: plugin id :return: """ - datasource_provider = db.session.scalar( + datasource_provider = session.scalar( select(DatasourceProvider) .where( DatasourceProvider.tenant_id == tenant_id, @@ -1130,5 +1149,5 @@ class DatasourceProviderService: .limit(1) ) if datasource_provider: - db.session.delete(datasource_provider) - db.session.commit() + session.delete(datasource_provider) + session.commit() diff --git a/api/services/enterprise/account_deletion_sync.py b/api/services/enterprise/account_deletion_sync.py index b5107fb0f66..89c4b80e670 100644 --- a/api/services/enterprise/account_deletion_sync.py +++ b/api/services/enterprise/account_deletion_sync.py @@ -5,9 +5,9 @@ from datetime import UTC, datetime from redis import RedisError from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config -from extensions.ext_database import db from extensions.ext_redis import redis_client from models.account import TenantAccountJoin @@ -87,7 +87,7 @@ def sync_workspace_member_removal(workspace_id: str, member_id: str, *, source: return _queue_task(workspace_id=workspace_id, member_id=member_id, source=source) -def sync_account_deletion(account_id: str, *, source: str) -> bool: +def sync_account_deletion(account_id: str, *, source: str, session: Session) -> bool: """ Sync full account deletion across all workspaces (enterprise only). @@ -97,6 +97,7 @@ def sync_account_deletion(account_id: str, *, source: str) -> bool: Args: account_id: The account ID being deleted source: Source of the sync request (e.g., "account_deleted") + session: SQLAlchemy session used to fetch workspace memberships Returns: bool: True if all tasks were queued (or skipped in community), False if any queueing failed @@ -105,9 +106,7 @@ def sync_account_deletion(account_id: str, *, source: str) -> bool: return True # Fetch all workspaces the account belongs to - workspace_joins = db.session.scalars( - select(TenantAccountJoin).where(TenantAccountJoin.account_id == account_id) - ).all() + workspace_joins = session.scalars(select(TenantAccountJoin).where(TenantAccountJoin.account_id == account_id)).all() # Queue sync task for each workspace success = True diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index b2e77156d3d..47ac8d5aeae 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -9,9 +9,9 @@ from flask import has_request_context, request from pydantic import AliasChoices, BaseModel, ConfigDict, Field, field_validator from sqlalchemy import select from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.orm import Session from configs import dify_config -from core.db.session_factory import session_factory from core.rbac import RBACResourceWhitelistScope from models import TenantAccountJoin, TenantAccountRole from services.enterprise.base import EnterpriseRequest @@ -565,25 +565,24 @@ def _legacy_member_roles_response( ) -def _legacy_my_permissions(tenant_id: str, account_id: str | None) -> MyPermissionsResponse: +def _legacy_my_permissions(tenant_id: str, account_id: str | None, *, session: Session) -> MyPermissionsResponse: if not account_id: return MyPermissionsResponse() try: - with session_factory.create_session() as session: - role = session.scalar( - select(TenantAccountJoin.role).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == account_id, - ) + role = session.scalar( + select(TenantAccountJoin.role).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == account_id, ) - if not role: - return MyPermissionsResponse() + ) + if not role: + return MyPermissionsResponse() - try: - tenant_role = TenantAccountRole(role) - except ValueError: - return MyPermissionsResponse() + try: + tenant_role = TenantAccountRole(role) + except ValueError: + return MyPermissionsResponse() except SQLAlchemyError: return MyPermissionsResponse() @@ -600,8 +599,10 @@ def _legacy_resource_permission_keys_batch( account_id: str | None, resource_ids: list[str], resource_type: RBACResourceType, + *, + session: Session, ) -> dict[str, list[str]]: - snapshot = _legacy_my_permissions(tenant_id, account_id) + snapshot = _legacy_my_permissions(tenant_id, account_id, session=session) if resource_type == RBACResourceType.APP: permission_keys = snapshot.app.default_permission_keys else: @@ -1597,7 +1598,9 @@ class RBACService: class MemberRoles: @staticmethod - def get(tenant_id: str, account_id: str | None, member_account_id: str) -> MemberRolesResponse: + def get( + tenant_id: str, account_id: str | None, member_account_id: str, *, session: Session + ) -> MemberRolesResponse: if dify_config.RBAC_ENABLED: data = _inner_call( "GET", @@ -1609,14 +1612,13 @@ class RBACService: rst = MemberRolesResponse.model_validate(data or {}) return rst else: - with session_factory.create_session() as session: - role = session.scalar( - select(TenantAccountJoin.role).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == member_account_id, - ) + role = session.scalar( + select(TenantAccountJoin.role).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == member_account_id, ) - return _legacy_member_roles_response(tenant_id, member_account_id, role) + ) + return _legacy_member_roles_response(tenant_id, member_account_id, role) @staticmethod def batch_get( @@ -1646,34 +1648,35 @@ class RBACService: account_id: str | None, member_account_id: str, role_ids: list[str], + *, + session: Session, ) -> MemberRolesResponse: if not dify_config.RBAC_ENABLED: if len(role_ids) != 1: raise ValueError("Legacy workspace member role update requires exactly one role.") tenant_role = TenantAccountRole(role_ids[0]) - with session_factory.create_session() as session: - target_member_join = session.scalar( + target_member_join = session.scalar( + select(TenantAccountJoin).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == member_account_id, + ) + ) + if not target_member_join: + raise ValueError("Member not in tenant.") + + if tenant_role == TenantAccountRole.OWNER: + current_owner_join = session.scalar( select(TenantAccountJoin).where( TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.account_id == member_account_id, + TenantAccountJoin.role == TenantAccountRole.OWNER, ) ) - if not target_member_join: - raise ValueError("Member not in tenant.") + if current_owner_join and current_owner_join.account_id != member_account_id: + current_owner_join.role = TenantAccountRole.ADMIN - if tenant_role == TenantAccountRole.OWNER: - current_owner_join = session.scalar( - select(TenantAccountJoin).where( - TenantAccountJoin.tenant_id == tenant_id, - TenantAccountJoin.role == TenantAccountRole.OWNER, - ) - ) - if current_owner_join and current_owner_join.account_id != member_account_id: - current_owner_join.role = TenantAccountRole.ADMIN - - target_member_join.role = tenant_role - session.commit() + target_member_join.role = tenant_role + session.commit() return _legacy_member_roles_response(tenant_id, member_account_id, tenant_role) @@ -1739,11 +1742,15 @@ class RBACService: tenant_id: str, account_id: str | None, app_ids: list[str], + *, + session: Session, ) -> dict[str, list[str]]: if not app_ids: return {} if not dify_config.RBAC_ENABLED: - return _legacy_resource_permission_keys_batch(tenant_id, account_id, app_ids, RBACResourceType.APP) + return _legacy_resource_permission_keys_batch( + tenant_id, account_id, app_ids, RBACResourceType.APP, session=session + ) data = _inner_call( "POST", f"{_INNER_PREFIX}/apps/permission-keys/batch", @@ -1759,12 +1766,14 @@ class RBACService: tenant_id: str, account_id: str | None, dataset_ids: list[str], + *, + session: Session, ) -> dict[str, list[str]]: if not dataset_ids: return {} if not dify_config.RBAC_ENABLED: return _legacy_resource_permission_keys_batch( - tenant_id, account_id, dataset_ids, RBACResourceType.DATASET + tenant_id, account_id, dataset_ids, RBACResourceType.DATASET, session=session ) data = _inner_call( "POST", @@ -1783,9 +1792,10 @@ class RBACService: *, app_id: str | None = None, dataset_id: str | None = None, + session: Session, ) -> MyPermissionsResponse: if not dify_config.RBAC_ENABLED: - return _legacy_my_permissions(tenant_id, account_id) + return _legacy_my_permissions(tenant_id, account_id, session=session) data = _inner_call( "GET", diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index 42e7eca29d7..cdd6c48342e 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.helper import ssrf_proxy from core.rag.entities import MetadataFilteringCondition +from extensions.ext_database import db # noqa: F401 from graphon.nodes.http_request.exc import InvalidHttpMethodError from libs.datetime_utils import naive_utc_now from libs.pagination import paginate_query @@ -57,7 +58,7 @@ class ExternalDatasetService: @staticmethod def create_external_knowledge_api( - tenant_id: str, user_id: str, args: dict[str, Any], session: Session + tenant_id: str, user_id: str, args: dict[str, Any], *, session: Session ) -> ExternalKnowledgeApis: settings = args.get("settings") if settings is None: @@ -105,7 +106,7 @@ class ExternalDatasetService: @staticmethod def get_external_knowledge_api( - session: Session, external_knowledge_api_id: str, tenant_id: str + external_knowledge_api_id: str, tenant_id: str, *, session: Session ) -> ExternalKnowledgeApis: external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) @@ -118,7 +119,12 @@ class ExternalDatasetService: @staticmethod def update_external_knowledge_api( - session: Session, tenant_id: str, user_id: str, external_knowledge_api_id: str, args + tenant_id: str, + user_id: str, + external_knowledge_api_id: str, + args: dict[str, Any], + *, + session: Session, ) -> ExternalKnowledgeApis: external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) @@ -131,9 +137,9 @@ class ExternalDatasetService: if settings and settings.get("api_key") == HIDDEN_VALUE and external_knowledge_api.settings_dict: settings["api_key"] = external_knowledge_api.settings_dict.get("api_key") - external_knowledge_api.name = args.get("name") - external_knowledge_api.description = args.get("description", "") - external_knowledge_api.settings = json.dumps(args.get("settings"), ensure_ascii=False) + external_knowledge_api.name = str(args.get("name")) + external_knowledge_api.description = str(args.get("description", "")) + external_knowledge_api.settings = json.dumps(settings, ensure_ascii=False) external_knowledge_api.updated_by = user_id external_knowledge_api.updated_at = naive_utc_now() session.commit() @@ -141,7 +147,7 @@ class ExternalDatasetService: return external_knowledge_api @staticmethod - def delete_external_knowledge_api(session: Session, tenant_id: str, external_knowledge_api_id: str): + def delete_external_knowledge_api(tenant_id: str, external_knowledge_api_id: str, *, session: Session) -> None: external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) @@ -155,7 +161,7 @@ class ExternalDatasetService: @staticmethod def external_knowledge_api_use_check( - session: Session, external_knowledge_api_id: str, tenant_id: str + external_knowledge_api_id: str, tenant_id: str, *, session: Session ) -> tuple[bool, int]: """ Return usage for an external knowledge API within a single tenant. @@ -176,7 +182,7 @@ class ExternalDatasetService: @staticmethod def get_external_knowledge_binding_with_dataset_id( - session: Session, tenant_id: str, dataset_id: str + tenant_id: str, dataset_id: str, *, session: Session ) -> ExternalKnowledgeBindings: external_knowledge_binding: ExternalKnowledgeBindings | None = session.scalar( select(ExternalKnowledgeBindings) @@ -189,8 +195,12 @@ class ExternalDatasetService: @staticmethod def document_create_args_validate( - session: Session, tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any] - ): + tenant_id: str, + external_knowledge_api_id: str, + process_parameter: dict[str, Any], + *, + session: Session, + ) -> None: external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) @@ -264,7 +274,7 @@ class ExternalDatasetService: return ExternalKnowledgeApiSetting.model_validate(settings) @staticmethod - def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], session: Session) -> Dataset: + def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], *, session: Session) -> Dataset: # check if dataset name already exists if session.scalar( select(Dataset).where(Dataset.name == args.get("name"), Dataset.tenant_id == tenant_id).limit(1) @@ -314,12 +324,13 @@ class ExternalDatasetService: @staticmethod def fetch_external_knowledge_retrieval( - session: Session, tenant_id: str, dataset_id: str, query: str, external_retrieval_parameters: dict[str, Any], metadata_condition: MetadataFilteringCondition | None = None, + *, + session: Session, ): """Fetch retrieval records from an external knowledge provider. diff --git a/api/services/file_service.py b/api/services/file_service.py index e41d74ad3eb..ec69af4e80c 100644 --- a/api/services/file_service.py +++ b/api/services/file_service.py @@ -268,7 +268,7 @@ class FileService: @staticmethod def get_upload_files_by_ids( - session: Session, tenant_id: str, upload_file_ids: Sequence[str] + tenant_id: str, upload_file_ids: Sequence[str], *, session: Session ) -> dict[str, UploadFile]: """ Fetch `UploadFile` rows for a tenant in a single batch query. diff --git a/api/services/hit_testing_service.py b/api/services/hit_testing_service.py index 1b51a2d279b..1bfa4025fa0 100644 --- a/api/services/hit_testing_service.py +++ b/api/services/hit_testing_service.py @@ -4,7 +4,7 @@ import time from typing import Any, TypedDict, cast from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.app.app_config.entities import ModelConfig from core.rag.datasource.retrieval_service import DefaultRetrievalModelDict, RetrievalService @@ -56,9 +56,7 @@ class HitTestingService: } @classmethod - def _dump_retrieval_records( - cls, session: Session | scoped_session, records: list[RetrievalSegments] - ) -> list[dict[str, Any]]: + def _dump_retrieval_records(cls, session: Session, records: list[RetrievalSegments]) -> list[dict[str, Any]]: document_ids = { document_id for record in records @@ -105,7 +103,6 @@ class HitTestingService: @classmethod def retrieve( cls, - session: Session, dataset: Dataset, query: str, account: Account, @@ -113,6 +110,8 @@ class HitTestingService: external_retrieval_model: dict[str, Any], attachment_ids: list | None = None, limit: int = 10, + *, + session: Session, ): start = time.perf_counter() @@ -144,7 +143,7 @@ class HitTestingService: if metadata_filter_document_ids: document_ids_filter = metadata_filter_document_ids.get(dataset.id, []) if metadata_condition and not document_ids_filter: - return cls.compact_retrieve_response(session, query, []) + return cls.compact_retrieve_response(query, [], session=session) all_documents = RetrievalService.retrieve( retrieval_method=RetrievalMethod( resolved_retrieval_model.get("search_method", RetrievalMethod.SEMANTIC_SEARCH) @@ -186,17 +185,18 @@ class HitTestingService: session.add(dataset_query) session.commit() - return cls.compact_retrieve_response(session, query, all_documents) + return cls.compact_retrieve_response(query, all_documents, session=session) @classmethod def external_retrieve( cls, - session: Session, dataset: Dataset, query: str, account: Account, external_retrieval_model: dict[str, Any] | None = None, metadata_filtering_conditions: dict[str, Any] | None = None, + *, + session: Session, ): if dataset.provider != "external": return { @@ -233,7 +233,7 @@ class HitTestingService: @classmethod def compact_retrieve_response( - cls, session: Session | scoped_session, query: str, documents: list[Document] + cls, query: str, documents: list[Document], *, session: Session ) -> RetrieveResponseDict: records = RetrievalService.format_retrieval_documents(documents) diff --git a/api/services/message_service.py b/api/services/message_service.py index e8d1b6232bc..4fbeb61e1f7 100644 --- a/api/services/message_service.py +++ b/api/services/message_service.py @@ -3,7 +3,7 @@ from collections.abc import Sequence from typing import cast from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager from core.app.entities.app_invoke_entities import InvokeFrom @@ -70,6 +70,8 @@ class MessageService: first_id: str | None, limit: int, order: str = "asc", + *, + session: Session, ) -> InfiniteScrollPagination: if not user: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) @@ -78,20 +80,20 @@ class MessageService: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) conversation = ConversationService.get_conversation( - app_model=app_model, user=user, conversation_id=conversation_id + app_model=app_model, user=user, conversation_id=conversation_id, session=session ) fetch_limit = limit + 1 if first_id: - first_message = db.session.scalar( + first_message = session.scalar( select(Message).where(Message.conversation_id == conversation.id, Message.id == first_id).limit(1) ) if not first_message: raise FirstMessageNotExistsError() - history_messages = db.session.scalars( + history_messages = session.scalars( select(Message) .where( Message.conversation_id == conversation.id, @@ -102,7 +104,7 @@ class MessageService: .limit(fetch_limit) ).all() else: - history_messages = db.session.scalars( + history_messages = session.scalars( select(Message) .where(Message.conversation_id == conversation.id) .order_by(Message.created_at.desc()) @@ -130,6 +132,8 @@ class MessageService: limit: int, conversation_id: str | None = None, include_ids: list | None = None, + *, + session: Session, ) -> InfiniteScrollPagination: if not user: return InfiniteScrollPagination(data=[], limit=limit, has_more=False) @@ -140,7 +144,7 @@ class MessageService: if conversation_id is not None: conversation = ConversationService.get_conversation( - app_model=app_model, user=user, conversation_id=conversation_id + app_model=app_model, user=user, conversation_id=conversation_id, session=session ) stmt = stmt.where(Message.conversation_id == conversation.id) @@ -152,18 +156,18 @@ class MessageService: stmt = stmt.where(Message.id.in_(include_ids)) if last_id: - last_message = db.session.scalar(stmt.where(Message.id == last_id).limit(1)) + last_message = session.scalar(stmt.where(Message.id == last_id).limit(1)) if not last_message: raise LastMessageNotExistsError() - history_messages = db.session.scalars( + history_messages = session.scalars( stmt.where(Message.created_at < last_message.created_at, Message.id != last_message.id) .order_by(Message.created_at.desc()) .limit(fetch_limit) ).all() else: - history_messages = db.session.scalars(stmt.order_by(Message.created_at.desc()).limit(fetch_limit)).all() + history_messages = session.scalars(stmt.order_by(Message.created_at.desc()).limit(fetch_limit)).all() has_more = False if len(history_messages) > limit: @@ -181,16 +185,17 @@ class MessageService: user: Account | EndUser | None, rating: FeedbackRating | None, content: str | None, + session: Session, ): if not user: raise ValueError("user cannot be None") - message = cls.get_message(app_model=app_model, user=user, message_id=message_id) + message = cls.get_message(app_model=app_model, user=user, message_id=message_id, session=session) feedback = message.user_feedback if isinstance(user, EndUser) else message.admin_feedback if not rating and feedback: - db.session.delete(feedback) + session.delete(feedback) elif rating and feedback: feedback.rating = rating feedback.content = content @@ -208,17 +213,17 @@ class MessageService: from_end_user_id=(user.id if isinstance(user, EndUser) else None), from_account_id=(user.id if isinstance(user, Account) else None), ) - db.session.add(feedback) + session.add(feedback) - db.session.commit() + session.commit() return feedback @classmethod - def get_all_messages_feedbacks(cls, app_model: App, page: int, limit: int): + def get_all_messages_feedbacks(cls, app_model: App, page: int, limit: int, *, session: Session): """Get all feedbacks of an app""" offset = (page - 1) * limit - feedbacks = db.session.scalars( + feedbacks = session.scalars( select(MessageFeedback) .where(MessageFeedback.app_id == app_model.id) .order_by(MessageFeedback.created_at.desc(), MessageFeedback.id.desc()) @@ -229,8 +234,8 @@ class MessageService: return [record.to_dict() for record in feedbacks] @classmethod - def get_message(cls, app_model: App, user: Account | EndUser | None, message_id: str): - message = db.session.scalar( + def get_message(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): + message = session.scalar( select(Message) .where( Message.id == message_id, @@ -249,15 +254,21 @@ class MessageService: @classmethod def get_suggested_questions_after_answer( - cls, app_model: App, user: Account | EndUser | None, message_id: str, invoke_from: InvokeFrom + cls, + app_model: App, + user: Account | EndUser | None, + message_id: str, + invoke_from: InvokeFrom, + *, + session: Session, ) -> list[str]: if not user: raise ValueError("user cannot be None") - message = cls.get_message(app_model=app_model, user=user, message_id=message_id) + message = cls.get_message(app_model=app_model, user=user, message_id=message_id, session=session) conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=message.conversation_id, user=user + app_model=app_model, conversation_id=message.conversation_id, user=user, session=session ) model_manager = ModelManager.for_tenant(tenant_id=app_model.tenant_id) @@ -266,9 +277,9 @@ class MessageService: if app_model.mode == AppMode.ADVANCED_CHAT: workflow_service = WorkflowService() if invoke_from == InvokeFrom.DEBUGGER: - workflow = workflow_service.get_draft_workflow(app_model=app_model) + workflow = workflow_service.get_draft_workflow(app_model=app_model, session=session) else: - workflow = workflow_service.get_published_workflow(app_model=app_model) + workflow = workflow_service.get_published_workflow(app_model=app_model, session=session) if workflow is None: return [] @@ -288,7 +299,7 @@ class MessageService: ) else: if not conversation.override_model_configs: - app_model_config = db.session.scalar( + app_model_config = session.scalar( select(AppModelConfig) .where(AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id) .limit(1) diff --git a/api/services/metadata_service.py b/api/services/metadata_service.py index 4e83858ea0e..481eb3b2e29 100644 --- a/api/services/metadata_service.py +++ b/api/services/metadata_service.py @@ -23,11 +23,12 @@ logger = logging.getLogger(__name__) class MetadataService: @staticmethod def create_metadata( - session: Session, dataset_id: str, metadata_args: MetadataArgs, current_user: Account | None = None, # TODO: the service_api is not migrated yet current_tenant_id: str | None = None, + *, + session: Session, ) -> DatasetMetadata: # check if metadata name is too long if len(metadata_args.name) > 255: @@ -60,12 +61,13 @@ class MetadataService: @staticmethod def update_metadata_name( - session: Session, dataset_id: str, metadata_id: str, name: str, current_user: Account | None = None, current_tenant_id: str | None = None, # TODO: the service_api is not migrated yet + *, + session: Session, ) -> DatasetMetadata | None: # check if metadata name is too long if len(name) > 255: @@ -126,7 +128,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def delete_metadata(session: Session, dataset_id: str, metadata_id: str): + def delete_metadata(dataset_id: str, metadata_id: str, *, session: Session): lock_key = f"dataset_metadata_lock_{dataset_id}" try: MetadataService.knowledge_base_metadata_lock_check(dataset_id, None) @@ -172,7 +174,7 @@ class MetadataService: ] @staticmethod - def enable_built_in_field(session: Session, dataset: Dataset): + def enable_built_in_field(dataset: Dataset, *, session: Session): if dataset.built_in_field_enabled: return lock_key = f"dataset_metadata_lock_{dataset.id}" @@ -201,7 +203,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def disable_built_in_field(session: Session, dataset: Dataset): + def disable_built_in_field(dataset: Dataset, *, session: Session): if not dataset.built_in_field_enabled: return lock_key = f"dataset_metadata_lock_{dataset.id}" @@ -233,11 +235,12 @@ class MetadataService: @staticmethod def update_documents_metadata( - session: Session, dataset: Dataset, metadata_args: MetadataOperationData, current_user: Account | None = None, # TODO: the service_api is not migrated yet current_tenant_id: str | None = None, + *, + session: Session, ): current_user, current_tenant_id = resolve_account_fallback( current_user, current_tenant_id, fallback_tenant_id=dataset.tenant_id @@ -316,7 +319,7 @@ class MetadataService: redis_client.set(lock_key, 1, ex=3600) @staticmethod - def get_dataset_metadatas(session: Session, dataset: Dataset): + def get_dataset_metadatas(dataset: Dataset, *, session: Session): return { "doc_metadata": [ { diff --git a/api/services/model_load_balancing_service.py b/api/services/model_load_balancing_service.py index 2a9094a35f2..6eab1ffbe3e 100644 --- a/api/services/model_load_balancing_service.py +++ b/api/services/model_load_balancing_service.py @@ -3,6 +3,7 @@ import logging from typing import Any, TypedDict, cast from sqlalchemy import or_, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.entities.provider_configuration import ProviderConfiguration @@ -14,7 +15,6 @@ from core.helper.model_provider_cache import ( from core.model_manager import LBModelManager from core.plugin.impl.model_runtime_factory import create_plugin_model_assembly, create_plugin_provider_manager from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager -from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ( ModelCredentialSchema, @@ -93,7 +93,13 @@ class ModelLoadBalancingService: provider_configuration.disable_model_load_balancing(model=model, model_type=ModelType(model_type)) def get_load_balancing_configs( - self, tenant_id: str, provider: str, model: str, model_type: str, config_from: str = "" + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + session: Session, + config_from: str = "", ) -> tuple[bool, list[LoadBalancingConfigSummaryDict]]: """ Get load balancing configurations. @@ -131,7 +137,7 @@ class ModelLoadBalancingService: # Get load balancing configurations load_balancing_configs = list( - db.session.scalars( + session.scalars( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, @@ -158,7 +164,7 @@ class ModelLoadBalancingService: if not inherit_config_exists: # Initialize the inherit configuration - inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type_enum) + inherit_config = self._init_inherit_config(tenant_id, provider, model, model_type_enum, session=session) # prepend the inherit configuration load_balancing_configs.insert(0, inherit_config) @@ -233,7 +239,13 @@ class ModelLoadBalancingService: return is_load_balancing_enabled, datas def get_load_balancing_config( - self, tenant_id: str, provider: str, model: str, model_type: str, config_id: str + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + config_id: str, + session: Session, ) -> LoadBalancingConfigDetailDict | None: """ Get load balancing configuration. @@ -256,7 +268,7 @@ class ModelLoadBalancingService: model_type_enum = ModelType(model_type) # Get load balancing configurations - load_balancing_model_config = db.session.scalar( + load_balancing_model_config = session.scalar( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, @@ -296,7 +308,12 @@ class ModelLoadBalancingService: return result def _init_inherit_config( - self, tenant_id: str, provider: str, model: str, model_type: ModelType + self, + tenant_id: str, + provider: str, + model: str, + model_type: ModelType, + session: Session, ) -> LoadBalancingModelConfig: """ Initialize the inherit configuration. @@ -314,8 +331,8 @@ class ModelLoadBalancingService: model_name=model, name="__inherit__", ) - db.session.add(inherit_config) - db.session.commit() + session.add(inherit_config) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -324,7 +341,14 @@ class ModelLoadBalancingService: return inherit_config def update_load_balancing_configs( - self, tenant_id: str, provider: str, model: str, model_type: str, configs: list[dict], config_from: str + self, + tenant_id: str, + provider: str, + model: str, + model_type: str, + configs: list[dict], + config_from: str, + session: Session, ): """ Update load balancing configurations. @@ -350,7 +374,7 @@ class ModelLoadBalancingService: if not isinstance(configs, list): raise ValueError("Invalid load balancing configs") - current_load_balancing_configs = db.session.scalars( + current_load_balancing_configs = session.scalars( select(LoadBalancingModelConfig).where( LoadBalancingModelConfig.tenant_id == tenant_id, LoadBalancingModelConfig.provider_name == provider_configuration.provider.provider, @@ -377,7 +401,7 @@ class ModelLoadBalancingService: if credential_id: if config_from == "predefined-model": - credential_record = db.session.scalar( + credential_record = session.scalar( select(ProviderCredential) .where( ProviderCredential.id == credential_id, @@ -387,7 +411,7 @@ class ModelLoadBalancingService: .limit(1) ) else: - credential_record = db.session.scalar( + credential_record = session.scalar( select(ProviderModelCredential) .where( ProviderModelCredential.id == credential_id, @@ -440,7 +464,7 @@ class ModelLoadBalancingService: load_balancing_config.name = name load_balancing_config.enabled = enabled load_balancing_config.updated_at = naive_utc_now() - db.session.commit() + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -496,8 +520,8 @@ class ModelLoadBalancingService: encrypted_config=json.dumps(credentials), ) - db.session.add(load_balancing_model_config) - db.session.commit() + session.add(load_balancing_model_config) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -506,8 +530,8 @@ class ModelLoadBalancingService: # get deleted config ids deleted_config_ids = set(current_load_balancing_configs_dict.keys()) - updated_config_ids for config_id in deleted_config_ids: - db.session.delete(current_load_balancing_configs_dict[config_id]) - db.session.commit() + session.delete(current_load_balancing_configs_dict[config_id]) + session.commit() ProviderManager.invalidate_configurations_cache( tenant_id, sources=(ProviderConfigurationCacheSource.PROVIDER_LOAD_BALANCING_CONFIGS,), @@ -522,6 +546,7 @@ class ModelLoadBalancingService: model: str, model_type: str, credentials: dict[str, Any], + session: Session, config_id: str | None = None, ): """ @@ -548,7 +573,7 @@ class ModelLoadBalancingService: load_balancing_model_config = None if config_id: # Get load balancing config - load_balancing_model_config = db.session.scalar( + load_balancing_model_config = session.scalar( select(LoadBalancingModelConfig) .where( LoadBalancingModelConfig.tenant_id == tenant_id, diff --git a/api/services/oauth_device_flow.py b/api/services/oauth_device_flow.py index 9ec5711890b..9e59b8c326a 100644 --- a/api/services/oauth_device_flow.py +++ b/api/services/oauth_device_flow.py @@ -13,7 +13,7 @@ from enum import StrEnum from typing import Any, NotRequired, TypedDict from sqlalchemy import and_, func, select, update -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from libs.oauth_bearer import TOKEN_CACHE_KEY_FMT, AuthContext, SubjectType from models.oauth import OAuthAccessToken @@ -335,9 +335,6 @@ def sha256_hex(token: str) -> str: def mint_oauth_token( - # Accept either Session or Flask-SQLAlchemy's request-scoped wrapper — - # the wrapper proxies the same execute/commit surface. - session: Session | scoped_session, redis_client, *, subject_email: str, @@ -347,6 +344,7 @@ def mint_oauth_token( device_label: str, prefix: str, ttl_days: int, + session: Session, ) -> MintResult: """Live row rotates in place via partial unique index ``uq_oauth_active_per_device``; hard-expired rows are excluded by the @@ -390,7 +388,7 @@ def mint_oauth_token( def _upsert( - session: Session | scoped_session, + session: Session, *, subject_email: str, subject_issuer: str | None, @@ -501,11 +499,7 @@ def subject_match_clauses(ctx: AuthContext) -> tuple[Any, ...]: ) -def list_active_sessions( - session: Session | scoped_session, - ctx: AuthContext, - now: datetime, -) -> list[OAuthAccessToken]: +def list_active_sessions(ctx: AuthContext, now: datetime, *, session: Session) -> list[OAuthAccessToken]: return list( session.execute( select(OAuthAccessToken) @@ -524,11 +518,7 @@ def list_active_sessions( ) -def token_belongs_to_subject( - session: Session | scoped_session, - token_id: str, - ctx: AuthContext, -) -> bool: +def token_belongs_to_subject(token_id: str, ctx: AuthContext, *, session: Session) -> bool: row = session.execute( select(OAuthAccessToken.id).where( and_( @@ -540,11 +530,7 @@ def token_belongs_to_subject( return row is not None -def revoke_oauth_token( - session: Session | scoped_session, - redis_client: Any, - token_id: str, -) -> None: +def revoke_oauth_token(redis_client: Any, token_id: str, *, session: Session) -> None: row = ( session.query(OAuthAccessToken.token_hash) .filter( diff --git a/api/services/oauth_server.py b/api/services/oauth_server.py index 5f3277c9525..47aa1bc99bf 100644 --- a/api/services/oauth_server.py +++ b/api/services/oauth_server.py @@ -2,7 +2,7 @@ import enum import uuid from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest from extensions.ext_database import db @@ -83,7 +83,7 @@ class OAuthServerService: return token @staticmethod - def validate_oauth_access_token(client_id: str, token: str) -> Account | None: + def validate_oauth_access_token(client_id: str, token: str, session: Session) -> Account | None: redis_key = OAUTH_ACCESS_TOKEN_REDIS_KEY.format(client_id=client_id, token=token) user_account_id = redis_client.get(redis_key) if not user_account_id: @@ -91,4 +91,4 @@ class OAuthServerService: user_id_str = user_account_id.decode("utf-8") - return AccountService.load_user(user_id_str, db.session) + return AccountService.load_user(user_id_str, session) diff --git a/api/services/ops_service.py b/api/services/ops_service.py index 3ad42faf249..b6f17168b3c 100644 --- a/api/services/ops_service.py +++ b/api/services/ops_service.py @@ -1,23 +1,23 @@ from typing import Any from sqlalchemy import select +from sqlalchemy.orm import Session from core.ops.entities.config_entity import BaseTracingConfig from core.ops.ops_trace_manager import OpsTraceManager, TracingProviderConfigEntry, provider_config_map -from extensions.ext_database import db from models.model import App, TraceAppConfig class OpsService: @classmethod - def get_tracing_app_config(cls, app_id: str, tracing_provider: str): + def get_tracing_app_config(cls, app_id: str, tracing_provider: str, session: Session): """ Get tracing app config :param app_id: app id :param tracing_provider: tracing provider :return: """ - trace_config_data: TraceAppConfig | None = db.session.scalar( + trace_config_data: TraceAppConfig | None = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -27,7 +27,7 @@ class OpsService: return None # decrypt_token and obfuscated_token - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -137,7 +137,9 @@ class OpsService: return trace_config_data.to_dict() @classmethod - def create_tracing_app_config(cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any]): + def create_tracing_app_config( + cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any], session: Session + ): """ Create tracing app config :param app_id: app id @@ -184,7 +186,7 @@ class OpsService: project_url = None # check if trace config already exists - trace_config_data: TraceAppConfig | None = db.session.scalar( + trace_config_data: TraceAppConfig | None = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -194,7 +196,7 @@ class OpsService: return None # get tenant id - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -206,13 +208,15 @@ class OpsService: tracing_provider=tracing_provider, tracing_config=tracing_config, ) - db.session.add(trace_config_data) - db.session.commit() + session.add(trace_config_data) + session.commit() return {"result": "success"} @classmethod - def update_tracing_app_config(cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any]): + def update_tracing_app_config( + cls, app_id: str, tracing_provider: str, tracing_config: dict[str, Any], session: Session + ): """ Update tracing app config :param app_id: app id @@ -226,7 +230,7 @@ class OpsService: raise ValueError(f"Invalid tracing provider: {tracing_provider}") # check if trace config already exists - current_trace_config = db.session.scalar( + current_trace_config = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -236,7 +240,7 @@ class OpsService: return None # get tenant id - app = db.session.get(App, app_id) + app = session.get(App, app_id) if not app: return None tenant_id = app.tenant_id @@ -251,19 +255,19 @@ class OpsService: raise ValueError("Invalid Credentials") current_trace_config.tracing_config = tracing_config - db.session.commit() + session.commit() return current_trace_config.to_dict() @classmethod - def delete_tracing_app_config(cls, app_id: str, tracing_provider: str): + def delete_tracing_app_config(cls, app_id: str, tracing_provider: str, session: Session): """ Delete tracing app config :param app_id: app id :param tracing_provider: tracing provider :return: """ - trace_config = db.session.scalar( + trace_config = session.scalar( select(TraceAppConfig) .where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider) .limit(1) @@ -272,7 +276,7 @@ class OpsService: if not trace_config: return None - db.session.delete(trace_config) - db.session.commit() + session.delete(trace_config) + session.commit() return True diff --git a/api/services/plugin/plugin_auto_upgrade_service.py b/api/services/plugin/plugin_auto_upgrade_service.py index f1e1918bdd2..79770063016 100644 --- a/api/services/plugin/plugin_auto_upgrade_service.py +++ b/api/services/plugin/plugin_auto_upgrade_service.py @@ -12,7 +12,6 @@ from hashlib import sha256 from sqlalchemy import select from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from core.plugin.impl.plugin import PluginInstaller from models.account import ( TenantPluginAutoUpgradeCategory, @@ -141,6 +140,8 @@ class PluginAutoUpgradeService: @staticmethod def backfill_strategy_categories( tenant_id: str, + *, + session: Session, ) -> PluginAutoUpgradeBackfillResult: """Create missing category strategies and split include/exclude lists when needed. @@ -148,89 +149,85 @@ class PluginAutoUpgradeService: New category rows copy it first, then plugin lists are narrowed by real plugin category when the source strategy contains include/exclude IDs. """ - with session_factory.create_session() as session, session.begin(): - strategies = list( - session.scalars( - select(TenantPluginAutoUpgradeStrategy).where( - TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id - ) - ).all() + strategies = list( + session.scalars( + select(TenantPluginAutoUpgradeStrategy).where(TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id) + ).all() + ) + if not strategies: + return PluginAutoUpgradeBackfillResult(created_count=0, normalized=False) + + # Schema migration marks the historical workspace-level row as tool. + source_strategy = next( + (strategy for strategy in strategies if strategy.category == PluginCategory.TOOL), + strategies[0], + ) + source_has_default_strategy = PluginAutoUpgradeService._has_default_strategy(source_strategy) + strategies_by_category = {strategy.category: strategy for strategy in strategies} + exclude_plugins = source_strategy.exclude_plugins + include_plugins = source_strategy.include_plugins + should_split_plugin_lists = bool(exclude_plugins or include_plugins) + # Query daemon only for tenants that actually customized plugin lists. + plugin_categories = ( + PluginAutoUpgradeService._get_installed_plugin_categories(tenant_id) if should_split_plugin_lists else {} + ) + if should_split_plugin_lists: + PluginAutoUpgradeService._log_unknown_plugin_ids( + tenant_id, + "exclude_plugins", + exclude_plugins, + plugin_categories, ) - if not strategies: - return PluginAutoUpgradeBackfillResult(created_count=0, normalized=False) - - # Schema migration marks the historical workspace-level row as tool. - source_strategy = next( - (strategy for strategy in strategies if strategy.category == PluginCategory.TOOL), - strategies[0], + PluginAutoUpgradeService._log_unknown_plugin_ids( + tenant_id, + "include_plugins", + include_plugins, + plugin_categories, ) - source_has_default_strategy = PluginAutoUpgradeService._has_default_strategy(source_strategy) - strategies_by_category = {strategy.category: strategy for strategy in strategies} - exclude_plugins = source_strategy.exclude_plugins - include_plugins = source_strategy.include_plugins - should_split_plugin_lists = bool(exclude_plugins or include_plugins) - # Query daemon only for tenants that actually customized plugin lists. - plugin_categories = ( - PluginAutoUpgradeService._get_installed_plugin_categories(tenant_id) - if should_split_plugin_lists - else {} + + created_count = 0 + for category in PLUGIN_CATEGORIES: + strategy = strategies_by_category.get(category) + if strategy is None: + # Start from the legacy workspace-level behavior before narrowing lists. + strategy = TenantPluginAutoUpgradeStrategy( + tenant_id=tenant_id, + category=category, + strategy_setting=PluginAutoUpgradeService._strategy_setting_for_category( + source_strategy, category, source_has_default_strategy + ), + upgrade_time_of_day=PluginAutoUpgradeService._upgrade_time_of_day_for_category( + tenant_id, source_strategy, source_has_default_strategy + ), + upgrade_mode=source_strategy.upgrade_mode, + exclude_plugins=source_strategy.exclude_plugins.copy(), + include_plugins=source_strategy.include_plugins.copy(), + ) + session.add(strategy) + created_count += 1 + elif source_has_default_strategy: + strategy.strategy_setting = PluginAutoUpgradeService.default_strategy_setting_for_category( + strategy.category + ) + strategy.upgrade_time_of_day = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id) + + if not should_split_plugin_lists: + continue + + # Narrow include/exclude lists to the current category after all rows exist. + strategy.exclude_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( + exclude_plugins, + strategy.category, + plugin_categories, + ) + strategy.include_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( + include_plugins, + strategy.category, + plugin_categories, ) - if should_split_plugin_lists: - PluginAutoUpgradeService._log_unknown_plugin_ids( - tenant_id, - "exclude_plugins", - exclude_plugins, - plugin_categories, - ) - PluginAutoUpgradeService._log_unknown_plugin_ids( - tenant_id, - "include_plugins", - include_plugins, - plugin_categories, - ) - created_count = 0 - for category in PLUGIN_CATEGORIES: - strategy = strategies_by_category.get(category) - if strategy is None: - # Start from the legacy workspace-level behavior before narrowing lists. - strategy = TenantPluginAutoUpgradeStrategy( - tenant_id=tenant_id, - category=category, - strategy_setting=PluginAutoUpgradeService._strategy_setting_for_category( - source_strategy, category, source_has_default_strategy - ), - upgrade_time_of_day=PluginAutoUpgradeService._upgrade_time_of_day_for_category( - tenant_id, source_strategy, source_has_default_strategy - ), - upgrade_mode=source_strategy.upgrade_mode, - exclude_plugins=source_strategy.exclude_plugins.copy(), - include_plugins=source_strategy.include_plugins.copy(), - ) - session.add(strategy) - created_count += 1 - elif source_has_default_strategy: - strategy.strategy_setting = PluginAutoUpgradeService.default_strategy_setting_for_category( - strategy.category - ) - strategy.upgrade_time_of_day = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id) - - if not should_split_plugin_lists: - continue - - # Narrow include/exclude lists to the current category after all rows exist. - strategy.exclude_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( - exclude_plugins, - strategy.category, - plugin_categories, - ) - strategy.include_plugins = PluginAutoUpgradeService._filter_plugin_ids_for_category( - include_plugins, - strategy.category, - plugin_categories, - ) - - return PluginAutoUpgradeBackfillResult(created_count=created_count, normalized=should_split_plugin_lists) + session.commit() + return PluginAutoUpgradeBackfillResult(created_count=created_count, normalized=should_split_plugin_lists) @staticmethod def _get_strategy( @@ -251,20 +248,18 @@ class PluginAutoUpgradeService: def get_strategy( tenant_id: str, category: PluginCategory, + *, + session: Session, ) -> TenantPluginAutoUpgradeStrategy | None: - with session_factory.create_session() as session: - return PluginAutoUpgradeService._get_strategy(session, tenant_id, category) + return PluginAutoUpgradeService._get_strategy(session, tenant_id, category) @staticmethod - def get_strategies(tenant_id: str) -> list[TenantPluginAutoUpgradeStrategy]: - with session_factory.create_session() as session: - return list( - session.scalars( - select(TenantPluginAutoUpgradeStrategy).where( - TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id - ) - ).all() - ) + def get_strategies(tenant_id: str, *, session: Session) -> list[TenantPluginAutoUpgradeStrategy]: + return list( + session.scalars( + select(TenantPluginAutoUpgradeStrategy).where(TenantPluginAutoUpgradeStrategy.tenant_id == tenant_id) + ).all() + ) @staticmethod def _change_strategy( @@ -305,20 +300,22 @@ class PluginAutoUpgradeService: exclude_plugins: list[str], include_plugins: list[str], category: PluginCategory, + *, + session: Session, ) -> bool: - with session_factory.create_session() as session, session.begin(): - PluginAutoUpgradeService._change_strategy( - session, - tenant_id=tenant_id, - category=category, - strategy_setting=strategy_setting, - upgrade_time_of_day=upgrade_time_of_day, - upgrade_mode=upgrade_mode, - exclude_plugins=exclude_plugins, - include_plugins=include_plugins, - ) + PluginAutoUpgradeService._change_strategy( + session, + tenant_id=tenant_id, + category=category, + strategy_setting=strategy_setting, + upgrade_time_of_day=upgrade_time_of_day, + upgrade_mode=upgrade_mode, + exclude_plugins=exclude_plugins, + include_plugins=include_plugins, + ) - return True + session.commit() + return True @staticmethod def _exclude_plugin( @@ -363,13 +360,15 @@ class PluginAutoUpgradeService: tenant_id: str, plugin_id: str, category: PluginCategory, + *, + session: Session, ) -> bool: - with session_factory.create_session() as session, session.begin(): - PluginAutoUpgradeService._exclude_plugin( - session, - tenant_id, - category, - plugin_id, - ) + PluginAutoUpgradeService._exclude_plugin( + session, + tenant_id, + category, + plugin_id, + ) - return True + session.commit() + return True diff --git a/api/services/plugin/plugin_permission_service.py b/api/services/plugin/plugin_permission_service.py index 19f3de2e52c..339a6ccb89b 100644 --- a/api/services/plugin/plugin_permission_service.py +++ b/api/services/plugin/plugin_permission_service.py @@ -1,35 +1,36 @@ from sqlalchemy import select +from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from models.account import TenantPluginDebugPermission, TenantPluginInstallPermission, TenantPluginPermission class PluginPermissionService: @staticmethod - def get_permission(tenant_id: str) -> TenantPluginPermission | None: - with session_factory.create_session() as session: - return session.scalar( - select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) - ) + def get_permission(tenant_id: str, *, session: Session) -> TenantPluginPermission | None: + return session.scalar( + select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + ) @staticmethod def change_permission( tenant_id: str, install_permission: TenantPluginInstallPermission, debug_permission: TenantPluginDebugPermission, - ): - with session_factory.create_session() as session, session.begin(): - permission = session.scalar( - select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + *, + session: Session, + ) -> bool: + permission = session.scalar( + select(TenantPluginPermission).where(TenantPluginPermission.tenant_id == tenant_id).limit(1) + ) + if not permission: + permission = TenantPluginPermission( + tenant_id=tenant_id, install_permission=install_permission, debug_permission=debug_permission ) - if not permission: - permission = TenantPluginPermission( - tenant_id=tenant_id, install_permission=install_permission, debug_permission=debug_permission - ) - session.add(permission) - else: - permission.install_permission = install_permission - permission.debug_permission = debug_permission + session.add(permission) + else: + permission.install_permission = install_permission + permission.debug_permission = debug_permission - return True + session.commit() + return True diff --git a/api/services/rag_pipeline/pipeline_generate_service.py b/api/services/rag_pipeline/pipeline_generate_service.py index e77ff9687ed..276bfaea158 100644 --- a/api/services/rag_pipeline/pipeline_generate_service.py +++ b/api/services/rag_pipeline/pipeline_generate_service.py @@ -17,12 +17,13 @@ class PipelineGenerateService: @classmethod def generate( cls, - session: Session, pipeline: Pipeline, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = True, + *, + session: Session, ): """ Pipeline Content Generate @@ -34,10 +35,10 @@ class PipelineGenerateService: :return: """ try: - workflow = cls._get_workflow(pipeline, invoke_from) + workflow = cls._get_workflow(pipeline, invoke_from, session) if original_document_id := args.get("original_document_id"): # update document status to waiting - cls.update_document_status(original_document_id, session) + cls.update_document_status(original_document_id, session=session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().generate( pipeline=pipeline, @@ -64,9 +65,9 @@ class PipelineGenerateService: @classmethod def generate_single_iteration( - cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, streaming: bool = True + cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True ): - workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER) + workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().single_iteration_generate( pipeline=pipeline, workflow=workflow, node_id=node_id, user=user, args=args, streaming=streaming @@ -74,8 +75,10 @@ class PipelineGenerateService: ) @classmethod - def generate_single_loop(cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, streaming: bool = True): - workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER) + def generate_single_loop( + cls, pipeline: Pipeline, user: Account, node_id: str, args: Any, session: Session, streaming: bool = True + ): + workflow = cls._get_workflow(pipeline, InvokeFrom.DEBUGGER, session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().single_loop_generate( pipeline=pipeline, workflow=workflow, node_id=node_id, user=user, args=args, streaming=streaming @@ -83,14 +86,14 @@ class PipelineGenerateService: ) @classmethod - def _get_workflow(cls, pipeline: Pipeline, invoke_from: InvokeFrom) -> Workflow: + def _get_workflow(cls, pipeline: Pipeline, invoke_from: InvokeFrom, session: Session) -> Workflow: """ Get workflow :param pipeline: pipeline :param invoke_from: invoke from :return: """ - rag_pipeline_service = RagPipelineService() + rag_pipeline_service = RagPipelineService(session) if invoke_from == InvokeFrom.DEBUGGER: # fetch draft workflow by app_model workflow = rag_pipeline_service.get_draft_workflow(pipeline=pipeline) @@ -107,7 +110,7 @@ class PipelineGenerateService: return workflow @classmethod - def update_document_status(cls, document_id: str, session: Session): + def update_document_status(cls, document_id: str, *, session: Session): """ Update document status to waiting :param document_id: document id diff --git a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py index 6de0be33a4a..d56c239ace2 100644 --- a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py @@ -23,14 +23,15 @@ class BuiltInPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: - del current_tenant_id + del current_tenant_id, session result = self.fetch_pipeline_templates_from_builtin(language) return result @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + del session result = self.fetch_pipeline_template_detail_from_builtin(template_id) return result diff --git a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py index 4faaf342f66..3d6baefcc46 100644 --- a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py @@ -41,16 +41,16 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) return self.fetch_pipeline_templates_from_customized( - session=session, tenant_id=current_tenant_id, language=language + tenant_id=current_tenant_id, language=language, session=session ) @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(session, template_id) + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(template_id, session=session) @override def get_type(self) -> str: @@ -58,7 +58,7 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @classmethod def fetch_pipeline_templates_from_customized( - cls, session: Session, tenant_id: str, language: str + cls, tenant_id: str, language: str, *, session: Session ) -> dict[str, Any]: """ Fetch pipeline templates from db. @@ -89,7 +89,7 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, template_id: str, *, session: Session) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param template_id: Template ID diff --git a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py index f6d2731e21a..d5c31ff74b2 100644 --- a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py @@ -41,21 +41,21 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: del current_tenant_id - return self.fetch_pipeline_templates_from_db(session, language) + return self.fetch_pipeline_templates_from_db(language, session=session) @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(session, template_id) + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(template_id, session=session) @override def get_type(self) -> str: return PipelineTemplateType.DATABASE @classmethod - def fetch_pipeline_templates_from_db(cls, session: Session, language: str) -> dict[str, Any]: + def fetch_pipeline_templates_from_db(cls, language: str, *, session: Session) -> dict[str, Any]: """ Fetch pipeline templates from db. :param language: language @@ -83,7 +83,7 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, template_id: str, *, session: Session) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param pipeline_id: Pipeline ID diff --git a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py index ff53dc1f79e..c61ac6d60f2 100644 --- a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py +++ b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py @@ -7,9 +7,9 @@ class PipelineTemplateRetrievalBase(Protocol): """Interface for pipeline template retrieval.""" def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: ... - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: ... + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: ... def get_type(self) -> str: ... diff --git a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py index 7f9fe1b56ea..29acbd198b6 100644 --- a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py @@ -18,23 +18,25 @@ class RemotePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: try: return self.fetch_pipeline_template_detail_from_dify_official(template_id) except Exception as e: logger.warning("fetch recommended app detail from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db(session, template_id) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db( + template_id, session=session + ) @override def get_pipeline_templates( - self, session: Session, language: str, current_tenant_id: str | None = None + self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: del current_tenant_id try: return self.fetch_pipeline_templates_from_dify_official(language) except Exception as e: logger.warning("fetch pipeline templates from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(session, language) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(language, session=session) @override def get_type(self) -> str: diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index 9e17a05be16..8bd3918eb15 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -27,7 +27,6 @@ from core.datasource.entities.datasource_entities import ( from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin from core.datasource.online_drive.online_drive_plugin import OnlineDriveDatasourcePlugin from core.datasource.website_crawl.website_crawl_plugin import WebsiteCrawlDatasourcePlugin -from core.db.session_factory import session_factory from core.helper import marketplace from core.rag.entities import DatasourceCompletedEvent, DatasourceErrorEvent, DatasourceProcessingEvent from core.repositories.factory import DifyCoreRepositoryFactory, OrderConfig @@ -96,11 +95,13 @@ def _build_seeded_variable_pool(variables: Sequence[Variable]) -> VariablePool: class RagPipelineService: - def __init__(self, session_maker: sessionmaker | None = None): + _session: Session + + def __init__(self, session: Session, session_maker: sessionmaker | None = None): """Initialize RagPipelineService with repository dependencies.""" + self._session = session if session_maker is None: - session_maker = session_factory.get_session_maker() - self._session_maker = session_maker + session_maker = sessionmaker(bind=db.engine, expire_on_commit=False) self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( session_maker ) @@ -109,15 +110,16 @@ class RagPipelineService: @classmethod def get_pipeline_templates( cls, - session: Session, type: str = "built-in", language: str = "en-US", current_tenant_id: str | None = None, + *, + session: Session, ) -> dict[str, Any]: if type == "built-in": mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(language, current_tenant_id, session=session) if not result.get("pipeline_templates") and language != "en-US": template_retrieval = PipelineTemplateRetrievalFactory.get_built_in_pipeline_template_retrieval() result = template_retrieval.fetch_pipeline_templates_from_builtin("en-US") @@ -125,12 +127,12 @@ class RagPipelineService: else: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(language, current_tenant_id, session=session) return result @classmethod def get_pipeline_template_detail( - cls, session: Session, template_id: str, type: str = "built-in" + cls, template_id: str, type: str = "built-in", *, session: Session ) -> dict[str, Any] | None: """ Get pipeline template detail. @@ -143,7 +145,7 @@ class RagPipelineService: mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() built_in_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - session, template_id + template_id, session=session ) if built_in_result is None: logger.warning( @@ -156,7 +158,7 @@ class RagPipelineService: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() customized_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - session, template_id + template_id, session=session ) return customized_result @@ -167,7 +169,8 @@ class RagPipelineService: template_info: PipelineTemplateInfoEntity, current_user: Account | None = None, current_tenant_id: str | None = None, - session: Session | None = None, + *, + session: Session, ): """ Update pipeline template. @@ -175,16 +178,6 @@ class RagPipelineService: :param template_info: template info """ current_user, current_tenant_id = resolve_account_fallback(current_user, current_tenant_id) - if session is None: - with session_factory.get_session_maker().begin() as new_session: - return cls.update_customized_pipeline_template( - template_id, - template_info, - current_user, - current_tenant_id, - session=new_session, - ) - customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( @@ -213,21 +206,17 @@ class RagPipelineService: customized_template.description = template_info.description customized_template.icon = template_info.icon_info.model_dump() customized_template.updated_by = current_user.id + session.commit() return customized_template @classmethod def delete_customized_pipeline_template( - cls, template_id: str, current_tenant_id: str | None = None, session: Session | None = None + cls, template_id: str, current_tenant_id: str | None = None, *, session: Session ): """ Delete customized pipeline template. """ current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) - if session is None: - with session_factory.get_session_maker().begin() as new_session: - cls.delete_customized_pipeline_template(template_id, current_tenant_id, session=new_session) - return - customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( @@ -239,22 +228,22 @@ class RagPipelineService: if not customized_template: raise ValueError("Customized pipeline template not found.") session.delete(customized_template) + session.commit() def get_draft_workflow(self, pipeline: Pipeline) -> Workflow | None: """ Get draft workflow """ # fetch draft workflow by rag pipeline - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == "draft", - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == "draft", ) + .limit(1) + ) # return draft workflow return workflow @@ -268,31 +257,29 @@ class RagPipelineService: return None # fetch published workflow by workflow_id - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == pipeline.workflow_id, ) + .limit(1) + ) return workflow def get_published_workflow_by_id(self, pipeline: Pipeline, workflow_id: str) -> Workflow | None: """Fetch a published workflow snapshot by ID for restore operations.""" - with self._session_maker() as session: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == workflow_id, - ) - .limit(1) + workflow = self._session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == workflow_id, ) + .limit(1) + ) if workflow and workflow.version == Workflow.VERSION_DRAFT: raise IsDraftWorkflowError("source workflow must be published") return workflow @@ -350,51 +337,39 @@ class RagPipelineService: Sync draft workflow :raises WorkflowHashNotEqualError """ - with self._session_maker.begin() as session: - managed_pipeline = session.get(Pipeline, pipeline.id) - if not managed_pipeline: - raise ValueError("Pipeline not found") + # fetch draft workflow by app_model + workflow = self.get_draft_workflow(pipeline=pipeline) - # fetch draft workflow by app_model - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.version == "draft", - ) - .limit(1) + if workflow and workflow.unique_hash != unique_hash: + raise WorkflowHashNotEqualError() + + # create draft workflow if not found + if not workflow: + workflow = Workflow( + tenant_id=pipeline.tenant_id, + app_id=pipeline.id, + features="{}", + type=WorkflowType.RAG_PIPELINE.value, + version="draft", + graph=json.dumps(graph), + created_by=account.id, + environment_variables=environment_variables, + conversation_variables=conversation_variables, + rag_pipeline_variables=rag_pipeline_variables, ) - - if workflow and workflow.unique_hash != unique_hash: - raise WorkflowHashNotEqualError() - - # create draft workflow if not found - if not workflow: - workflow = Workflow( - tenant_id=managed_pipeline.tenant_id, - app_id=managed_pipeline.id, - features="{}", - type=WorkflowType.RAG_PIPELINE.value, - version="draft", - graph=json.dumps(graph), - created_by=account.id, - environment_variables=environment_variables, - conversation_variables=conversation_variables, - rag_pipeline_variables=rag_pipeline_variables, - ) - session.add(workflow) - session.flush() - managed_pipeline.workflow_id = workflow.id - pipeline.workflow_id = workflow.id - # update draft workflow if found - else: - workflow.graph = json.dumps(graph) - workflow.updated_by = account.id - workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) - workflow.environment_variables = environment_variables - workflow.conversation_variables = conversation_variables - workflow.rag_pipeline_variables = rag_pipeline_variables + self._session.add(workflow) + self._session.flush() + pipeline.workflow_id = workflow.id + # update draft workflow if found + else: + workflow.graph = json.dumps(graph) + workflow.updated_by = account.id + workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) + workflow.environment_variables = environment_variables + workflow.conversation_variables = conversation_variables + workflow.rag_pipeline_variables = rag_pipeline_variables + # commit db session changes + self._session.commit() # trigger workflow events TODO # app_draft_workflow_was_synced.send(pipeline, synced_draft_workflow=workflow) @@ -415,48 +390,26 @@ class RagPipelineService: the pipeline-specific flush/link step that wires a newly created draft back onto ``pipeline.workflow_id``. """ - with self._session_maker.begin() as session: - managed_pipeline = session.get(Pipeline, pipeline.id) - if not managed_pipeline: - raise ValueError("Pipeline not found") + source_workflow = self.get_published_workflow_by_id(pipeline=pipeline, workflow_id=workflow_id) + if not source_workflow: + raise WorkflowNotFoundError("Workflow not found.") - source_workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.id == workflow_id, - ) - .limit(1) - ) - if source_workflow and source_workflow.version == Workflow.VERSION_DRAFT: - raise IsDraftWorkflowError("source workflow must be published") - if not source_workflow: - raise WorkflowNotFoundError("Workflow not found.") + draft_workflow = self.get_draft_workflow(pipeline=pipeline) + draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( + tenant_id=pipeline.tenant_id, + app_id=pipeline.id, + source_workflow=source_workflow, + draft_workflow=draft_workflow, + account=account, + updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), + ) - draft_workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == managed_pipeline.tenant_id, - Workflow.app_id == managed_pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) - .limit(1) - ) - draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( - tenant_id=managed_pipeline.tenant_id, - app_id=managed_pipeline.id, - source_workflow=source_workflow, - draft_workflow=draft_workflow, - account=account, - updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), - ) + if is_new_draft: + self._session.add(draft_workflow) + self._session.flush() + pipeline.workflow_id = draft_workflow.id - if is_new_draft: - session.add(draft_workflow) - session.flush() - managed_pipeline.workflow_id = draft_workflow.id - pipeline.workflow_id = draft_workflow.id + self._session.commit() return draft_workflow @@ -633,7 +586,7 @@ class RagPipelineService: workflow_node_execution.id ) - with self._session_maker.begin() as session: + with sessionmaker(bind=db.engine).begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -1050,22 +1003,23 @@ class RagPipelineService: dataset_id = get_system_segment(variable_pool, SystemVariableKey.DATASET_ID) pipeline_id = get_system_segment(variable_pool, SystemVariableKey.APP_ID) if document_id and dataset_id and pipeline_id: - with self._session_maker.begin() as session: - document = session.scalar( - select(Document) - .join(Dataset, Dataset.id == Document.dataset_id) - .where( - Document.id == document_id.value, - Document.tenant_id == tenant_id, - Document.dataset_id == dataset_id.value, - Dataset.tenant_id == tenant_id, - Dataset.pipeline_id == pipeline_id.value, - ) - .limit(1) + document = self._session.scalar( + select(Document) + .join(Dataset, Dataset.id == Document.dataset_id) + .where( + Document.id == document_id.value, + Document.tenant_id == tenant_id, + Document.dataset_id == dataset_id.value, + Dataset.tenant_id == tenant_id, + Dataset.pipeline_id == pipeline_id.value, ) - if document: - document.indexing_status = IndexingStatus.ERROR - document.error = error + .limit(1) + ) + if document: + document.indexing_status = IndexingStatus.ERROR + document.error = error + self._session.add(document) + self._session.commit() return workflow_node_execution @@ -1276,89 +1230,86 @@ class RagPipelineService: args: dict[str, Any], current_user: Account | None = None, current_tenant_id: str | None = None, + *, + session: Session, ): """ Publish customized pipeline template """ current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - with session_factory.get_session_maker().begin() as session: - pipeline = session.get(Pipeline, pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - if not pipeline.workflow_id: - raise ValueError("Pipeline workflow not found") - workflow = session.get(Workflow, pipeline.workflow_id) - if not workflow: - raise ValueError("Workflow not found") - dataset = pipeline.retrieve_dataset(session=session) - if not dataset: - raise ValueError("Dataset not found") + pipeline = session.get(Pipeline, pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + if not pipeline.workflow_id: + raise ValueError("Pipeline workflow not found") + workflow = session.get(Workflow, pipeline.workflow_id) + if not workflow: + raise ValueError("Workflow not found") + dataset = pipeline.retrieve_dataset(session=session) + if not dataset: + raise ValueError("Dataset not found") - # check template name is exist - template_name = args.get("name") - if template_name: - template = session.scalar( - select(PipelineCustomizedTemplate) - .where( - PipelineCustomizedTemplate.name == template_name, - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, - ) - .limit(1) - ) - if template: - raise ValueError("Template name is already exists") - - max_position = session.scalar( - select(func.max(PipelineCustomizedTemplate.position)).where( - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id + # check template name is exist + template_name = args.get("name") + if template_name: + template = session.scalar( + select(PipelineCustomizedTemplate) + .where( + PipelineCustomizedTemplate.name == template_name, + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, ) + .limit(1) ) + if template: + raise ValueError("Template name is already exists") - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService - - rag_pipeline_dsl_service = RagPipelineDslService(session) - dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) - if args.get("icon_info") is None: - args["icon_info"] = {} - if args.get("description") is None: - raise ValueError("Description is required") - if args.get("name") is None: - raise ValueError("Name is required") - pipeline_customized_template = PipelineCustomizedTemplate( - name=args.get("name") or "", - description=args.get("description") or "", - icon=args.get("icon_info") or {}, - tenant_id=pipeline.tenant_id, - yaml_content=dsl, - install_count=0, - position=max_position + 1 if max_position else 1, - chunk_structure=dataset.chunk_structure, - language="en-US", - created_by=current_user.id, + max_position = session.scalar( + select(func.max(PipelineCustomizedTemplate.position)).where( + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id ) - session.add(pipeline_customized_template) + ) + + from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService + + rag_pipeline_dsl_service = RagPipelineDslService(session) + dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) + if args.get("icon_info") is None: + args["icon_info"] = {} + if args.get("description") is None: + raise ValueError("Description is required") + if args.get("name") is None: + raise ValueError("Name is required") + pipeline_customized_template = PipelineCustomizedTemplate( + name=args.get("name") or "", + description=args.get("description") or "", + icon=args.get("icon_info") or {}, + tenant_id=pipeline.tenant_id, + yaml_content=dsl, + install_count=0, + position=max_position + 1 if max_position else 1, + chunk_structure=dataset.chunk_structure, + language="en-US", + created_by=current_user.id, + ) + session.add(pipeline_customized_template) + session.commit() def is_workflow_exist(self, pipeline: Pipeline) -> bool: - with self._session_maker() as session: - return ( - session.scalar( - select(func.count(Workflow.id)).where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) + return ( + self._session.scalar( + select(func.count(Workflow.id)).where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, ) - or 0 - ) > 0 + ) + or 0 + ) > 0 def get_node_last_run( self, pipeline: Pipeline, workflow: Workflow, node_id: str ) -> WorkflowNodeExecutionModel | None: - node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( - self._session_maker - ) - - node_exec = node_execution_service_repo.get_node_last_execution( + node_exec = self._node_execution_service_repo.get_node_last_execution( tenant_id=pipeline.tenant_id, app_id=pipeline.id, workflow_id=workflow.id, @@ -1431,7 +1382,7 @@ class RagPipelineService: # Convert node_execution to WorkflowNodeExecution after save workflow_node_execution_db_model = repository._to_db_model(workflow_node_execution) # type: ignore - with self._session_maker.begin() as session: + with sessionmaker(bind=db.engine).begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -1465,10 +1416,9 @@ class RagPipelineService: if type and type != "all": stmt = stmt.where(PipelineRecommendedPlugin.type == type) - with self._session_maker() as session: - pipeline_recommended_plugins = session.scalars( - stmt.order_by(PipelineRecommendedPlugin.position.asc()) - ).all() + pipeline_recommended_plugins = self._session.scalars( + stmt.order_by(PipelineRecommendedPlugin.position.asc()) + ).all() if not pipeline_recommended_plugins: return { @@ -1507,173 +1457,41 @@ class RagPipelineService: """ Retry error document """ - with self._session_maker() as session: - document_pipeline_execution_log = session.scalar( - select(DocumentPipelineExecutionLog) - .where(DocumentPipelineExecutionLog.document_id == document.id) - .limit(1) - ) - if not document_pipeline_execution_log: - raise ValueError("Document pipeline execution log not found") - pipeline = session.get(Pipeline, document_pipeline_execution_log.pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - # convert to app config - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) - ) - if not workflow: - raise ValueError("Workflow not found") - PipelineGenerator().generate( - pipeline=pipeline, - workflow=workflow, - user=user, - args={ - "inputs": document_pipeline_execution_log.input_data, - "start_node_id": document_pipeline_execution_log.datasource_node_id, - "datasource_type": document_pipeline_execution_log.datasource_type, - "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], - "original_document_id": document.id, - }, - invoke_from=InvokeFrom.PUBLISHED_PIPELINE, - streaming=False, - call_depth=0, - workflow_thread_pool_id=None, - is_retry=True, - ) + document_pipeline_execution_log = self._session.scalar( + select(DocumentPipelineExecutionLog).where(DocumentPipelineExecutionLog.document_id == document.id).limit(1) + ) + if not document_pipeline_execution_log: + raise ValueError("Document pipeline execution log not found") + pipeline = self._session.get(Pipeline, document_pipeline_execution_log.pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + # convert to app config + workflow = self.get_published_workflow(pipeline) + if not workflow: + raise ValueError("Workflow not found") + PipelineGenerator().generate( + pipeline=pipeline, + workflow=workflow, + user=user, + args={ + "inputs": document_pipeline_execution_log.input_data, + "start_node_id": document_pipeline_execution_log.datasource_node_id, + "datasource_type": document_pipeline_execution_log.datasource_type, + "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], + "original_document_id": document.id, + }, + invoke_from=InvokeFrom.PUBLISHED_PIPELINE, + streaming=False, + call_depth=0, + workflow_thread_pool_id=None, + is_retry=True, + ) def get_datasource_plugins(self, tenant_id: str, dataset_id: str, is_published: bool) -> list[dict]: """ Get datasource plugins """ - with self._session_maker() as session: - dataset: Dataset | None = session.scalar( - select(Dataset) - .where( - Dataset.id == dataset_id, - Dataset.tenant_id == tenant_id, - ) - .limit(1) - ) - if not dataset: - raise ValueError("Dataset not found") - pipeline: Pipeline | None = session.scalar( - select(Pipeline) - .where( - Pipeline.id == dataset.pipeline_id, - Pipeline.tenant_id == tenant_id, - ) - .limit(1) - ) - if not pipeline: - raise ValueError("Pipeline not found") - - if is_published: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, - ) - .limit(1) - ) - else: - workflow = session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, - ) - .limit(1) - ) - if not pipeline or not workflow: - raise ValueError("Pipeline or workflow not found") - - datasource_nodes = workflow.graph_dict.get("nodes", []) - datasource_plugins = [] - for datasource_node in datasource_nodes: - if datasource_node.get("data", {}).get("type") == "datasource": - datasource_node_data = datasource_node["data"] - if not datasource_node_data: - continue - - variables = workflow.rag_pipeline_variables - if variables: - variables_map = {item["variable"]: item for item in variables} - else: - variables_map = {} - - datasource_parameters = datasource_node_data.get("datasource_parameters", {}) - user_input_variables_keys = [] - user_input_variables = [] - - for _, value in datasource_parameters.items(): - if value.get("value") and isinstance(value.get("value"), str): - pattern = ( - r"\{\{#([a-zA-Z0-9_]{1,50}" - r"(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" - ) - match = re.match(pattern, value["value"]) - if match: - full_path = match.group(1) - last_part = full_path.split(".")[-1] - user_input_variables_keys.append(last_part) - elif value.get("value") and isinstance(value.get("value"), list): - last_part = value.get("value")[-1] - user_input_variables_keys.append(last_part) - for key, value in variables_map.items(): - if key in user_input_variables_keys: - user_input_variables.append(value) - - # get credentials - datasource_provider_service: DatasourceProviderService = DatasourceProviderService() - credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( - tenant_id=tenant_id, - provider=datasource_node_data.get("provider_name"), - plugin_id=datasource_node_data.get("plugin_id"), - ) - credential_info_list: list[Any] = [] - for credential in credentials: - credential_info_list.append( - { - "id": credential.get("id"), - "name": credential.get("name"), - "type": credential.get("type"), - "is_default": credential.get("is_default"), - } - ) - - datasource_plugins.append( - { - "node_id": datasource_node.get("id"), - "plugin_id": datasource_node_data.get("plugin_id"), - "provider_name": datasource_node_data.get("provider_name"), - "datasource_type": datasource_node_data.get("provider_type"), - "title": datasource_node_data.get("title"), - "user_input_variables": user_input_variables, - "credentials": credential_info_list, - } - ) - - return datasource_plugins - - def get_pipeline(self, tenant_id: str, dataset_id: str, session: Session | None = None) -> Pipeline: - """ - Get pipeline - """ - if session is None: - with self._session_maker() as new_session: - return self.get_pipeline(tenant_id, dataset_id, session=new_session) - - dataset: Dataset | None = session.scalar( + dataset: Dataset | None = self._session.scalar( select(Dataset) .where( Dataset.id == dataset_id, @@ -1683,7 +1501,106 @@ class RagPipelineService: ) if not dataset: raise ValueError("Dataset not found") - pipeline: Pipeline | None = session.scalar( + pipeline: Pipeline | None = self._session.scalar( + select(Pipeline) + .where( + Pipeline.id == dataset.pipeline_id, + Pipeline.tenant_id == tenant_id, + ) + .limit(1) + ) + if not pipeline: + raise ValueError("Pipeline not found") + + workflow: Workflow | None = None + if is_published: + workflow = self.get_published_workflow(pipeline=pipeline) + else: + workflow = self.get_draft_workflow(pipeline=pipeline) + if not pipeline or not workflow: + raise ValueError("Pipeline or workflow not found") + + datasource_nodes = workflow.graph_dict.get("nodes", []) + datasource_plugins = [] + for datasource_node in datasource_nodes: + if datasource_node.get("data", {}).get("type") == "datasource": + datasource_node_data = datasource_node["data"] + if not datasource_node_data: + continue + + variables = workflow.rag_pipeline_variables + if variables: + variables_map = {item["variable"]: item for item in variables} + else: + variables_map = {} + + datasource_parameters = datasource_node_data.get("datasource_parameters", {}) + user_input_variables_keys = [] + user_input_variables = [] + + for _, value in datasource_parameters.items(): + if value.get("value") and isinstance(value.get("value"), str): + pattern = r"\{\{#([a-zA-Z0-9_]{1,50}(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" + match = re.match(pattern, value["value"]) + if match: + full_path = match.group(1) + last_part = full_path.split(".")[-1] + user_input_variables_keys.append(last_part) + elif value.get("value") and isinstance(value.get("value"), list): + last_part = value.get("value")[-1] + user_input_variables_keys.append(last_part) + for key, value in variables_map.items(): + if key in user_input_variables_keys: + user_input_variables.append(value) + + # get credentials + datasource_provider_service: DatasourceProviderService = DatasourceProviderService() + credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( + tenant_id=tenant_id, + provider=datasource_node_data.get("provider_name"), + plugin_id=datasource_node_data.get("plugin_id"), + session=self._session, + ) + credential_info_list: list[Any] = [] + for credential in credentials: + credential_info_list.append( + { + "id": credential.get("id"), + "name": credential.get("name"), + "type": credential.get("type"), + "is_default": credential.get("is_default"), + } + ) + + datasource_plugins.append( + { + "node_id": datasource_node.get("id"), + "plugin_id": datasource_node_data.get("plugin_id"), + "provider_name": datasource_node_data.get("provider_name"), + "datasource_type": datasource_node_data.get("provider_type"), + "title": datasource_node_data.get("title"), + "user_input_variables": user_input_variables, + "credentials": credential_info_list, + } + ) + + return datasource_plugins + + def get_pipeline(self, tenant_id: str, dataset_id: str) -> Pipeline: + """ + Get pipeline + """ + dataset: Dataset | None = self._session.scalar( + select(Dataset) + .where( + Dataset.id == dataset_id, + Dataset.tenant_id == tenant_id, + ) + .limit(1) + ) + if not dataset: + raise ValueError("Dataset not found") + pipeline: Pipeline | None = self._session.scalar( select(Pipeline) .where( Pipeline.id == dataset.pipeline_id, diff --git a/api/services/rag_pipeline/rag_pipeline_dsl_service.py b/api/services/rag_pipeline/rag_pipeline_dsl_service.py index 5459c3e5f1f..d562a4b9adf 100644 --- a/api/services/rag_pipeline/rag_pipeline_dsl_service.py +++ b/api/services/rag_pipeline/rag_pipeline_dsl_service.py @@ -15,7 +15,7 @@ from Crypto.Util.Padding import pad, unpad from flask_login import current_user from pydantic import BaseModel from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.file import remote_fetcher from core.helper.name_generator import generate_incremental_name @@ -83,7 +83,7 @@ class RagPipelineDslService: when generated IDs are needed mid-operation; they never commit or rollback. """ - def __init__(self, session: Session | scoped_session): + def __init__(self, session: Session): self._session = session def import_rag_pipeline( diff --git a/api/services/rag_pipeline/rag_pipeline_transform_service.py b/api/services/rag_pipeline/rag_pipeline_transform_service.py index 1b922b3f7b9..6a7902c1908 100644 --- a/api/services/rag_pipeline/rag_pipeline_transform_service.py +++ b/api/services/rag_pipeline/rag_pipeline_transform_service.py @@ -96,7 +96,7 @@ class RagPipelineTransformService: # deal document data self._deal_document_data(dataset, session) - session.flush() + session.commit() return { "pipeline_id": pipeline.id, "dataset_id": dataset_id, @@ -194,6 +194,7 @@ class RagPipelineTransformService: def _create_pipeline( self, data: dict[str, Any], + *, session: Session, ) -> Pipeline: """Create a new app or update an existing one.""" @@ -291,7 +292,7 @@ class RagPipelineTransformService: logger.debug("Installing missing pipeline plugins %s", package_identifiers_to_install) PluginService.install_from_marketplace_pkg(tenant_id, package_identifiers_to_install) - def _transform_to_empty_pipeline(self, dataset: Dataset, session: Session): + def _transform_to_empty_pipeline(self, dataset: Dataset, *, session: Session): pipeline = Pipeline( tenant_id=dataset.tenant_id, name=dataset.name, @@ -306,7 +307,7 @@ class RagPipelineTransformService: dataset.updated_by = current_user.id dataset.updated_at = datetime.now(UTC).replace(tzinfo=None) session.add(dataset) - session.flush() + session.commit() return { "pipeline_id": pipeline.id, "dataset_id": dataset.id, diff --git a/api/services/recommend_app/buildin/buildin_retrieval.py b/api/services/recommend_app/buildin/buildin_retrieval.py index 03b72a4f57c..d29d754b67e 100644 --- a/api/services/recommend_app/buildin/buildin_retrieval.py +++ b/api/services/recommend_app/buildin/buildin_retrieval.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import Any, override from flask import current_app +from sqlalchemy.orm import Session from services.recommend_app.database.database_retrieval import DatabaseRecommendAppRetrieval from services.recommend_app.recommend_app_base import RecommendAppRetrievalBase @@ -22,17 +23,19 @@ class BuildInRecommendAppRetrieval(RecommendAppRetrievalBase): return RecommendAppType.BUILDIN @override - def get_recommended_apps_and_categories(self, language: str): + def get_recommended_apps_and_categories(self, language: str, *, session: Session): + del session result = self.fetch_recommended_apps_from_builtin(language) return result @override - def get_learn_dify_apps(self, language: str): - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language) + def get_learn_dify_apps(self, language: str, *, session: Session): + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language, session=session) return result @override - def get_recommend_app_detail(self, app_id: str): + def get_recommend_app_detail(self, app_id: str, *, session: Session): + del session result = self.fetch_recommended_app_detail_from_builtin(app_id) return result diff --git a/api/services/recommend_app/database/database_retrieval.py b/api/services/recommend_app/database/database_retrieval.py index f6786175896..08d902fdeb5 100644 --- a/api/services/recommend_app/database/database_retrieval.py +++ b/api/services/recommend_app/database/database_retrieval.py @@ -1,9 +1,9 @@ from typing import Any, NotRequired, TypedDict, override from sqlalchemy import select +from sqlalchemy.orm import Session from constants.languages import languages -from extensions.ext_database import db from models.model import App, RecommendedApp from services.app_dsl_service import AppDslService from services.recommend_app.category_order import order_categories @@ -45,18 +45,18 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): """ @override - def get_recommended_apps_and_categories(self, language: str) -> RecommendedAppsResultDict: - result = self.fetch_recommended_apps_from_db(language) + def get_recommended_apps_and_categories(self, language: str, *, session: Session) -> RecommendedAppsResultDict: + result = self.fetch_recommended_apps_from_db(language, session=session) return result @override - def get_learn_dify_apps(self, language: str) -> RecommendedAppsResultDict: - result = self.fetch_learn_dify_apps_from_db(language) + def get_learn_dify_apps(self, language: str, *, session: Session) -> RecommendedAppsResultDict: + result = self.fetch_learn_dify_apps_from_db(language, session=session) return result @override - def get_recommend_app_detail(self, app_id: str) -> RecommendedAppDetailDict | None: - result = self.fetch_recommended_app_detail_from_db(app_id) + def get_recommend_app_detail(self, app_id: str, *, session: Session) -> RecommendedAppDetailDict | None: + result = self.fetch_recommended_app_detail_from_db(app_id, session=session) return result @override @@ -64,42 +64,42 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): return RecommendAppType.DATABASE @classmethod - def fetch_recommended_apps_from_db(cls, language: str) -> RecommendedAppsResultDict: + def fetch_recommended_apps_from_db(cls, language: str, *, session: Session) -> RecommendedAppsResultDict: """ Fetch recommended apps from db. :param language: language :return: """ - recommended_apps = cls._fetch_listed_recommended_apps(language) + recommended_apps = cls._fetch_listed_recommended_apps(language, session=session) if len(recommended_apps) == 0: - recommended_apps = cls._fetch_listed_recommended_apps(languages[0]) + recommended_apps = cls._fetch_listed_recommended_apps(languages[0], session=session) return cls._format_recommended_apps(recommended_apps, language) @classmethod - def fetch_learn_dify_apps_from_db(cls, language: str) -> RecommendedAppsResultDict: + def fetch_learn_dify_apps_from_db(cls, language: str, *, session: Session) -> RecommendedAppsResultDict: """ Fetch listed recommended apps explicitly marked for the Learn Dify section. :param language: language :return: """ - recommended_apps = cls._fetch_listed_recommended_apps(language, is_learn_dify=True) + recommended_apps = cls._fetch_listed_recommended_apps(language, session=session, is_learn_dify=True) if len(recommended_apps) == 0 and language != languages[0]: - recommended_apps = cls._fetch_listed_recommended_apps(languages[0], is_learn_dify=True) + recommended_apps = cls._fetch_listed_recommended_apps(languages[0], session=session, is_learn_dify=True) return cls._format_recommended_apps(recommended_apps, language) @classmethod def _fetch_listed_recommended_apps( - cls, language: str, *, is_learn_dify: bool | None = None + cls, language: str, *, session: Session, is_learn_dify: bool | None = None ) -> list[RecommendedApp]: filters = [RecommendedApp.is_listed.is_(True), RecommendedApp.language == language] if is_learn_dify is not None: filters.append(RecommendedApp.is_learn_dify.is_(is_learn_dify)) - return list(db.session.scalars(select(RecommendedApp).where(*filters)).all()) + return list(session.scalars(select(RecommendedApp).where(*filters)).all()) @classmethod def _format_recommended_apps( @@ -146,14 +146,14 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): ) @classmethod - def fetch_recommended_app_detail_from_db(cls, app_id: str) -> RecommendedAppDetailDict | None: + def fetch_recommended_app_detail_from_db(cls, app_id: str, *, session: Session) -> RecommendedAppDetailDict | None: """ Fetch recommended app detail from db. :param app_id: App ID :return: """ # is in public recommended list - recommended_app = db.session.scalar( + recommended_app = session.scalar( select(RecommendedApp).where(RecommendedApp.is_listed == True, RecommendedApp.app_id == app_id).limit(1) ) @@ -161,7 +161,7 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): return None # get app detail - app_model = db.session.get(App, app_id) + app_model = session.get(App, app_id) if not app_model or not app_model.is_public: return None @@ -171,5 +171,5 @@ class DatabaseRecommendAppRetrieval(RecommendAppRetrievalBase): icon=app_model.icon, icon_background=app_model.icon_background, mode=app_model.mode, - export_data=AppDslService.export_dsl(app_model=app_model), + export_data=AppDslService.export_dsl(app_model=app_model, session=session), ) diff --git a/api/services/recommend_app/recommend_app_base.py b/api/services/recommend_app/recommend_app_base.py index f819cc3a937..821ad476c42 100644 --- a/api/services/recommend_app/recommend_app_base.py +++ b/api/services/recommend_app/recommend_app_base.py @@ -1,13 +1,15 @@ from typing import Any, Protocol +from sqlalchemy.orm import Session + class RecommendAppRetrievalBase(Protocol): """Interface for recommend app retrieval.""" - def get_recommended_apps_and_categories(self, language: str) -> Any: ... + def get_recommended_apps_and_categories(self, language: str, *, session: Session) -> Any: ... - def get_learn_dify_apps(self, language: str) -> Any: ... + def get_learn_dify_apps(self, language: str, *, session: Session) -> Any: ... - def get_recommend_app_detail(self, app_id: str) -> Any: ... + def get_recommend_app_detail(self, app_id: str, *, session: Session) -> Any: ... def get_type(self) -> str: ... diff --git a/api/services/recommend_app/remote/remote_retrieval.py b/api/services/recommend_app/remote/remote_retrieval.py index 2e3222bb978..c676ec907e0 100644 --- a/api/services/recommend_app/remote/remote_retrieval.py +++ b/api/services/recommend_app/remote/remote_retrieval.py @@ -3,6 +3,7 @@ from typing import Any, override import httpx from flask import has_request_context, request +from sqlalchemy.orm import Session from configs import dify_config from services.recommend_app.buildin.buildin_retrieval import BuildInRecommendAppRetrieval @@ -33,7 +34,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ @override - def get_recommend_app_detail(self, app_id: str): + def get_recommend_app_detail(self, app_id: str, *, session: Session): + del session try: result = self.fetch_recommended_app_detail_from_dify_official(app_id) except Exception as e: @@ -42,7 +44,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): return result @override - def get_recommended_apps_and_categories(self, language: str): + def get_recommended_apps_and_categories(self, language: str, *, session: Session): + del session try: result = self.fetch_recommended_apps_from_dify_official(language) except Exception as e: @@ -51,12 +54,12 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): return result @override - def get_learn_dify_apps(self, language: str): + def get_learn_dify_apps(self, language: str, *, session: Session): try: result = self.fetch_learn_dify_apps_from_dify_official(language) except Exception as e: logger.warning("fetch learn dify apps from dify official failed: %s, switch to database.", e) - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language) + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db(language, session=session) return result @override diff --git a/api/services/recommended_app_service.py b/api/services/recommended_app_service.py index 2d247ba5b71..813aa74754c 100644 --- a/api/services/recommended_app_service.py +++ b/api/services/recommended_app_service.py @@ -1,7 +1,7 @@ from typing import Any from sqlalchemy import select -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from configs import dify_config from models.model import AccountTrialAppRecord, TrialApp @@ -11,7 +11,7 @@ from services.recommend_app.recommend_app_factory import RecommendAppRetrievalFa class RecommendedAppService: @classmethod - def get_recommended_apps_and_categories(cls, session: scoped_session, language: str): + def get_recommended_apps_and_categories(cls, language: str, *, session: Session): """ Get recommended apps and categories. :param language: language @@ -19,7 +19,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result = retrieval_instance.get_recommended_apps_and_categories(language) + result = retrieval_instance.get_recommended_apps_and_categories(language, session=session) if not result.get("recommended_apps"): result = ( RecommendAppRetrievalFactory.get_buildin_recommend_app_retrieval().fetch_recommended_apps_from_builtin( @@ -35,7 +35,7 @@ class RecommendedAppService: return result @classmethod - def get_learn_dify_apps(cls, session: scoped_session, language: str) -> dict[str, Any]: + def get_learn_dify_apps(cls, language: str, *, session: Session) -> dict[str, Any]: """ Get recommended apps marked for the Learn Dify section. :param language: language @@ -43,7 +43,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result = retrieval_instance.get_learn_dify_apps(language) + result = retrieval_instance.get_learn_dify_apps(language, session=session) if FeatureService.get_system_features().enable_trial_app: for app in result["recommended_apps"]: @@ -52,7 +52,7 @@ class RecommendedAppService: return {"recommended_apps": result["recommended_apps"]} @classmethod - def get_recommend_app_detail(cls, session: scoped_session, app_id: str) -> dict[str, Any] | None: + def get_recommend_app_detail(cls, app_id: str, *, session: Session) -> dict[str, Any] | None: """ Get recommend app detail. :param app_id: app id @@ -60,7 +60,7 @@ class RecommendedAppService: """ mode = dify_config.HOSTED_FETCH_APP_TEMPLATES_MODE retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() - result: dict[str, Any] | None = retrieval_instance.get_recommend_app_detail(app_id) + result: dict[str, Any] | None = retrieval_instance.get_recommend_app_detail(app_id, session=session) if result is None: return None if FeatureService.get_system_features().enable_trial_app: @@ -69,7 +69,7 @@ class RecommendedAppService: return result @classmethod - def add_trial_app_record(cls, session: scoped_session, app_id: str, account_id: str): + def add_trial_app_record(cls, app_id: str, account_id: str, *, session: Session): """ Add trial app record. :param app_id: app id @@ -88,6 +88,6 @@ class RecommendedAppService: session.commit() @staticmethod - def _can_trial_app(session: scoped_session, app_id: str) -> bool: + def _can_trial_app(session: Session, app_id: str) -> bool: trial_app_model = session.scalar(select(TrialApp).where(TrialApp.app_id == app_id).limit(1)) return trial_app_model is not None diff --git a/api/services/saved_message_service.py b/api/services/saved_message_service.py index 9a65429748e..6165d74333f 100644 --- a/api/services/saved_message_service.py +++ b/api/services/saved_message_service.py @@ -12,7 +12,7 @@ from services.message_service import MessageService class SavedMessageService: @classmethod def pagination_by_last_id( - cls, session: Session, app_model: App, user: Account | EndUser | None, last_id: str | None, limit: int + cls, app_model: App, user: Account | EndUser | None, last_id: str | None, limit: int, *, session: Session ) -> InfiniteScrollPagination: if not user: raise ValueError("User is required") @@ -28,11 +28,16 @@ class SavedMessageService: message_ids = [sm.message_id for sm in saved_messages] return MessageService.pagination_by_last_id( - app_model=app_model, user=user, last_id=last_id, limit=limit, include_ids=message_ids + app_model=app_model, + user=user, + last_id=last_id, + limit=limit, + include_ids=message_ids, + session=session, ) @classmethod - def save(cls, session: Session, app_model: App, user: Account | EndUser | None, message_id: str): + def save(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): if not user: return saved_message = session.scalar( @@ -49,7 +54,7 @@ class SavedMessageService: if saved_message: return - message = MessageService.get_message(app_model=app_model, user=user, message_id=message_id) + message = MessageService.get_message(app_model=app_model, user=user, message_id=message_id, session=session) saved_message = SavedMessage( app_id=app_model.id, @@ -62,7 +67,7 @@ class SavedMessageService: session.commit() @classmethod - def delete(cls, session: Session, app_model: App, user: Account | EndUser | None, message_id: str): + def delete(cls, app_model: App, user: Account | EndUser | None, message_id: str, *, session: Session): if not user: return saved_message = session.scalar( diff --git a/api/services/snippet_service.py b/api/services/snippet_service.py index a54c9f6a069..64c1ec12370 100644 --- a/api/services/snippet_service.py +++ b/api/services/snippet_service.py @@ -6,9 +6,8 @@ from datetime import UTC, datetime from typing import Any from sqlalchemy import delete, func, select -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker -from core.db import session_factory from core.workflow.node_factory import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING from graphon.enums import BuiltinNodeTypes, NodeType from libs.infinite_scroll_pagination import InfiniteScrollPagination @@ -59,9 +58,8 @@ class SnippetService: session_maker = None if session is not None: session_maker = sessionmaker(bind=session.get_bind(), expire_on_commit=False) - elif session_maker is None: - session_maker = session_factory.get_session_maker() - assert session_maker is not None + if session_maker is None: + raise ValueError("SnippetService requires a session or session_maker.") self._session = session self._session_maker = session_maker self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( @@ -192,7 +190,7 @@ class SnippetService: self, *, tenant_id: str, - session: scoped_session, + session: Session, page: int = 1, limit: int = 20, keyword: str | None = None, diff --git a/api/services/summary_index_service.py b/api/services/summary_index_service.py index 3e065653bdf..3adc18dd2d5 100644 --- a/api/services/summary_index_service.py +++ b/api/services/summary_index_service.py @@ -7,7 +7,7 @@ from datetime import UTC, datetime from typing import TypedDict, cast from sqlalchemy import select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from core.db.session_factory import session_factory from core.model_manager import ModelManager @@ -94,6 +94,8 @@ class SummaryIndexService: dataset: Dataset, summary_content: str, status: SummaryStatus = SummaryStatus.GENERATING, + *, + session: Session, ) -> DocumentSegmentSummary: """ Create or update a DocumentSegmentSummary record. @@ -105,46 +107,48 @@ class SummaryIndexService: summary_content: Generated summary content status: Summary status (default: SummaryStatus.GENERATING) + Keyword Args: + session: SQLAlchemy session used for the summary record. + Returns: Created or updated DocumentSegmentSummary instance """ - with session_factory.create_session() as session: - # Check if summary record already exists - existing_summary = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + # Check if summary record already exists + existing_summary = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) - if existing_summary: - # Update existing record - existing_summary.summary_content = summary_content - existing_summary.status = status - existing_summary.error = None # Clear any previous errors - # Re-enable if it was disabled - if not existing_summary.enabled: - existing_summary.enabled = True - existing_summary.disabled_at = None - existing_summary.disabled_by = None - session.add(existing_summary) - session.flush() - return existing_summary - else: - # Create new record (enabled by default) - summary_record = DocumentSegmentSummary( - dataset_id=dataset.id, - document_id=segment.document_id, - chunk_id=segment.id, - summary_content=summary_content, - status=status, - enabled=True, # Explicitly set enabled to True - ) - session.add(summary_record) - session.flush() - return summary_record + if existing_summary: + # Update existing record + existing_summary.summary_content = summary_content + existing_summary.status = status + existing_summary.error = None # Clear any previous errors + # Re-enable if it was disabled + if not existing_summary.enabled: + existing_summary.enabled = True + existing_summary.disabled_at = None + existing_summary.disabled_by = None + session.add(existing_summary) + session.flush() + return existing_summary + else: + # Create new record (enabled by default) + summary_record = DocumentSegmentSummary( + dataset_id=dataset.id, + document_id=segment.document_id, + chunk_id=segment.id, + summary_content=summary_content, + status=status, + enabled=True, # Explicitly set enabled to True + ) + session.add(summary_record) + session.flush() + return summary_record @staticmethod def vectorize_summary( @@ -641,6 +645,8 @@ class SummaryIndexService: segment: DocumentSegment, dataset: Dataset, summary_index_setting: SummaryIndexSettingDict, + *, + session: Session, ) -> DocumentSegmentSummary: """ Generate summary for a segment and vectorize it. @@ -651,106 +657,101 @@ class SummaryIndexService: dataset: Dataset containing the segment summary_index_setting: Summary index configuration + Keyword Args: + session: SQLAlchemy session used for summary record updates. + Returns: Created DocumentSegmentSummary instance Raises: ValueError: If summary generation fails """ - with session_factory.create_session() as session: - try: - # Get or refresh summary record in this session - summary_record_in_session = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + try: + # Get or refresh summary record in this session + summary_record_in_session = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) - if not summary_record_in_session: - # If not found, create one - logger.warning("Summary record not found for segment %s, creating one", segment.id) - summary_record_in_session = DocumentSegmentSummary( - dataset_id=dataset.id, - document_id=segment.document_id, - chunk_id=segment.id, - summary_content="", - status=SummaryStatus.GENERATING, - enabled=True, - ) - session.add(summary_record_in_session) - session.flush() - - # Update status to "generating" - summary_record_in_session.status = SummaryStatus.GENERATING - summary_record_in_session.error = None - session.add(summary_record_in_session) - # Don't flush here - wait until after vectorization succeeds - - # Generate summary (returns summary_content and llm_usage) - summary_content, llm_usage = SummaryIndexService.generate_summary_for_segment( - segment, dataset, summary_index_setting + if not summary_record_in_session: + # If not found, create one + logger.warning("Summary record not found for segment %s, creating one", segment.id) + summary_record_in_session = DocumentSegmentSummary( + dataset_id=dataset.id, + document_id=segment.document_id, + chunk_id=segment.id, + summary_content="", + status=SummaryStatus.GENERATING, + enabled=True, ) - - # Update summary content - summary_record_in_session.summary_content = summary_content session.add(summary_record_in_session) - # Flush to ensure summary_content is saved before vectorize_summary queries it session.flush() - # Log LLM usage for summary generation - if llm_usage and llm_usage.total_tokens > 0: - logger.info( - "Summary generation for segment %s used %s tokens (prompt: %s, completion: %s)", - segment.id, - llm_usage.total_tokens, - llm_usage.prompt_tokens, - llm_usage.completion_tokens, - ) + # Update status to "generating" + summary_record_in_session.status = SummaryStatus.GENERATING + summary_record_in_session.error = None + session.add(summary_record_in_session) + # Don't flush here - wait until after vectorization succeeds - # Vectorize summary (will delete old vector if exists before creating new one) - # Pass the session-managed record to vectorize_summary - # vectorize_summary will update status to "completed" and tokens in its own session - # vectorize_summary will also ensure summary_content is preserved - try: - # Pass the session to vectorize_summary to avoid session isolation issues - SummaryIndexService.vectorize_summary(summary_record_in_session, segment, dataset, session=session) - # Refresh the object from database to get the updated status and tokens from vectorize_summary - session.refresh(summary_record_in_session) - # Commit the session - # (summary_record_in_session should have status="completed" and tokens from refresh) - session.commit() - logger.info("Successfully generated and vectorized summary for segment %s", segment.id) - return summary_record_in_session - except Exception as vectorize_error: - # If vectorization fails, update status to error in current session - logger.exception("Failed to vectorize summary for segment %s", segment.id) - summary_record_in_session.status = SummaryStatus.ERROR - summary_record_in_session.error = f"Vectorization failed: {str(vectorize_error)}" - session.add(summary_record_in_session) - session.commit() - raise + # Generate summary (returns summary_content and llm_usage) + summary_content, llm_usage = SummaryIndexService.generate_summary_for_segment( + segment, dataset, summary_index_setting + ) - except Exception as e: - logger.exception("Failed to generate summary for segment %s", segment.id) - # Update summary record with error status - summary_record_in_session = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + # Update summary content + summary_record_in_session.summary_content = summary_content + session.add(summary_record_in_session) + # Flush to ensure summary_content is saved before vectorize_summary queries it + session.flush() + + # Log LLM usage for summary generation + if llm_usage and llm_usage.total_tokens > 0: + logger.info( + "Summary generation for segment %s used %s tokens (prompt: %s, completion: %s)", + segment.id, + llm_usage.total_tokens, + llm_usage.prompt_tokens, + llm_usage.completion_tokens, ) - if summary_record_in_session: - summary_record_in_session.status = SummaryStatus.ERROR - summary_record_in_session.error = str(e) - session.add(summary_record_in_session) - session.commit() + + try: + SummaryIndexService.vectorize_summary(summary_record_in_session, segment, dataset, session=session) + # vectorize_summary mutates status and token fields; refresh before returning the ORM object. + session.refresh(summary_record_in_session) + session.commit() + logger.info("Successfully generated and vectorized summary for segment %s", segment.id) + return summary_record_in_session + except Exception as vectorize_error: + # If vectorization fails, update status to error in current session + logger.exception("Failed to vectorize summary for segment %s", segment.id) + summary_record_in_session.status = SummaryStatus.ERROR + summary_record_in_session.error = f"Vectorization failed: {str(vectorize_error)}" + session.add(summary_record_in_session) + session.commit() raise + except Exception as e: + logger.exception("Failed to generate summary for segment %s", segment.id) + # Update summary record with error status + summary_record_in_session = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, + ) + .limit(1) + ) + if summary_record_in_session: + summary_record_in_session.status = SummaryStatus.ERROR + summary_record_in_session.error = str(e) + session.add(summary_record_in_session) + session.commit() + raise + @staticmethod def generate_summaries_for_document( dataset: Dataset, @@ -840,7 +841,7 @@ class SummaryIndexService: try: summary_record = SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, summary_index_setting + segment, dataset, summary_index_setting, session=session ) summary_records.append(summary_record) except Exception as e: @@ -1048,6 +1049,8 @@ class SummaryIndexService: segment: DocumentSegment, dataset: Dataset, summary_content: str, + *, + session: Session, ) -> DocumentSegmentSummary | None: """ Update summary for a segment and re-vectorize it. @@ -1057,6 +1060,9 @@ class SummaryIndexService: dataset: Dataset containing the segment summary_content: New summary content + Keyword Args: + session: SQLAlchemy session used for summary record updates. + Returns: Updated DocumentSegmentSummary instance, or None if indexing technique is not high_quality """ @@ -1072,67 +1078,22 @@ class SummaryIndexService: if segment.document and segment.document.doc_form == "qa_model": return None - with session_factory.create_session() as session: - try: - # Check if summary_content is empty (whitespace-only strings are considered empty) - if not summary_content or not summary_content.strip(): - # If summary is empty, only delete existing summary vector and record - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) - ) - - if summary_record: - # Delete old vector if exists - old_summary_node_id = summary_record.summary_index_node_id - if old_summary_node_id: - try: - vector = Vector(dataset) - vector.delete_by_ids([old_summary_node_id]) - except Exception as e: - logger.warning( - "Failed to delete old summary vector for segment %s: %s", - segment.id, - str(e), - ) - - # Delete summary record since summary is empty - session.delete(summary_record) - session.commit() - logger.info("Deleted summary for segment %s (empty content provided)", segment.id) - return None - else: - # No existing summary record, nothing to do - logger.info("No summary record found for segment %s, nothing to delete", segment.id) - return None - - # Find existing summary record - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) + try: + summary_record = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, ) + .limit(1) + ) + # Check if summary_content is empty (whitespace-only strings are considered empty) + if not summary_content or not summary_content.strip(): + # If summary is empty, only delete existing summary vector and record if summary_record: - # Update existing summary + # Delete old vector if exists old_summary_node_id = summary_record.summary_index_node_id - - # Update summary content - summary_record.summary_content = summary_content - summary_record.status = SummaryStatus.GENERATING - summary_record.error = None # Clear any previous errors - session.add(summary_record) - # Flush to ensure summary_content is saved before vectorize_summary queries it - session.flush() - - # Delete old vector if exists (before vectorization) if old_summary_node_id: try: vector = Vector(dataset) @@ -1144,80 +1105,90 @@ class SummaryIndexService: str(e), ) - # Re-vectorize summary (this will update status to "completed" and tokens in its own session) - # vectorize_summary will also ensure summary_content is preserved - # Note: vectorize_summary may take time due to embedding API calls, but we need to complete it - # to ensure the summary is properly indexed - try: - # Pass the session to vectorize_summary to avoid session isolation issues - SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) - # Refresh the object from database to get the updated status and tokens from vectorize_summary - session.refresh(summary_record) - # Now commit the session (summary_record should have status="completed" and tokens from refresh) - session.commit() - logger.info("Successfully updated and re-vectorized summary for segment %s", segment.id) - return summary_record - except Exception as e: - # If vectorization fails, update status to error in current session - # Don't raise the exception - just log it and return the record with error status - # This allows the segment update to complete even if vectorization fails - summary_record.status = SummaryStatus.ERROR - summary_record.error = f"Vectorization failed: {str(e)}" - session.commit() - logger.exception("Failed to vectorize summary for segment %s", segment.id) - # Return the record with error status instead of raising - # The caller can check the status if needed - return summary_record - else: - # Create new summary record if doesn't exist - summary_record = SummaryIndexService.create_summary_record( - segment, dataset, summary_content, status=SummaryStatus.GENERATING - ) - # Re-vectorize summary (this will update status to "completed" and tokens in its own session) - # Note: summary_record was created in a different session, - # so we need to merge it into current session - try: - # Merge the record into current session first (since it was created in a different session) - summary_record = session.merge(summary_record) - # Pass the session to vectorize_summary - it will update the merged record - SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) - # Refresh to get updated status and tokens from database - session.refresh(summary_record) - # Commit the session to persist the changes - session.commit() - logger.info("Successfully created and vectorized summary for segment %s", segment.id) - return summary_record - except Exception as e: - # If vectorization fails, update status to error in current session - # Merge the record into current session first - error_record = session.merge(summary_record) - error_record.status = SummaryStatus.ERROR - error_record.error = f"Vectorization failed: {str(e)}" - session.commit() - logger.exception("Failed to vectorize summary for segment %s", segment.id) - # Return the record with error status instead of raising - return error_record - - except Exception as e: - logger.exception("Failed to update summary for segment %s", segment.id) - # Update summary record with error status if it exists - summary_record = session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment.id, - DocumentSegmentSummary.dataset_id == dataset.id, - ) - .limit(1) - ) - if summary_record: - summary_record.status = SummaryStatus.ERROR - summary_record.error = str(e) - session.add(summary_record) + # Delete summary record since summary is empty + session.delete(summary_record) session.commit() - raise + logger.info("Deleted summary for segment %s (empty content provided)", segment.id) + return None + else: + # No existing summary record, nothing to do + logger.info("No summary record found for segment %s, nothing to delete", segment.id) + return None + + if summary_record: + # Update existing summary + old_summary_node_id = summary_record.summary_index_node_id + + # Update summary content + summary_record.summary_content = summary_content + summary_record.status = SummaryStatus.GENERATING + summary_record.error = None # Clear any previous errors + session.add(summary_record) + # Flush to ensure summary_content is saved before vectorize_summary queries it + session.flush() + + # Delete old vector if exists (before vectorization) + if old_summary_node_id: + try: + vector = Vector(dataset) + vector.delete_by_ids([old_summary_node_id]) + except Exception as e: + logger.warning( + "Failed to delete old summary vector for segment %s: %s", + segment.id, + str(e), + ) + else: + # Create new summary record if doesn't exist + summary_record = SummaryIndexService.create_summary_record( + segment, + dataset, + summary_content, + status=SummaryStatus.GENERATING, + session=session, + ) + + try: + # Vectorization must finish here so the manual summary is searchable immediately. + SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) + session.refresh(summary_record) + session.commit() + logger.info("Successfully updated and re-vectorized summary for segment %s", segment.id) + return summary_record + except Exception as e: + # If vectorization fails, update status to error in current session. + # Return the record with error status so callers can still finish segment updates. + summary_record.status = SummaryStatus.ERROR + summary_record.error = f"Vectorization failed: {str(e)}" + session.commit() + logger.exception("Failed to vectorize summary for segment %s", segment.id) + return summary_record + + except Exception as e: + logger.exception("Failed to update summary for segment %s", segment.id) + # Update summary record with error status if it exists + summary_record = session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment.id, + DocumentSegmentSummary.dataset_id == dataset.id, + ) + .limit(1) + ) + if summary_record: + summary_record.status = SummaryStatus.ERROR + summary_record.error = str(e) + session.add(summary_record) + session.commit() + raise @staticmethod - def get_segment_summary(segment_id: str, dataset_id: str) -> DocumentSegmentSummary | None: + def get_segment_summary( + segment_id: str, + dataset_id: str, + *, + session: Session, + ) -> DocumentSegmentSummary | None: """ Get summary for a single segment. @@ -1225,22 +1196,29 @@ class SummaryIndexService: segment_id: Segment ID (chunk_id) dataset_id: Dataset ID + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: DocumentSegmentSummary instance if found, None otherwise """ - with session_factory.create_session() as session: - return session.scalar( - select(DocumentSegmentSummary) - .where( - DocumentSegmentSummary.chunk_id == segment_id, - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) - .limit(1) + return session.scalar( + select(DocumentSegmentSummary) + .where( + DocumentSegmentSummary.chunk_id == segment_id, + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), ) + .limit(1) + ) @staticmethod - def get_segments_summaries(segment_ids: list[str], dataset_id: str) -> dict[str, DocumentSegmentSummary]: + def get_segments_summaries( + segment_ids: list[str], + dataset_id: str, + *, + session: Session, + ) -> dict[str, DocumentSegmentSummary]: """ Get summaries for multiple segments. @@ -1248,26 +1226,31 @@ class SummaryIndexService: segment_ids: List of segment IDs (chunk_ids) dataset_id: Dataset ID + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: Dictionary mapping segment_id to DocumentSegmentSummary (only enabled summaries) """ if not segment_ids: return {} - with session_factory.create_session() as session: - summary_records = session.scalars( - select(DocumentSegmentSummary).where( - DocumentSegmentSummary.chunk_id.in_(segment_ids), - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) - ).all() - - return {summary.chunk_id: summary for summary in summary_records} + summaries = session.scalars( + select(DocumentSegmentSummary).where( + DocumentSegmentSummary.chunk_id.in_(segment_ids), + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), + ) + ).all() + return {summary.chunk_id: summary for summary in summaries} @staticmethod def get_document_summaries( - document_id: str, dataset_id: str, segment_ids: list[str] | None = None + document_id: str, + dataset_id: str, + segment_ids: list[str] | None = None, + *, + session: Session, ) -> list[DocumentSegmentSummary]: """ Get all summary records for a document. @@ -1277,23 +1260,31 @@ class SummaryIndexService: dataset_id: Dataset ID segment_ids: Optional list of segment IDs to filter by + Keyword Args: + session: SQLAlchemy session used to read summary records. + Returns: List of DocumentSegmentSummary instances (only enabled summaries) """ - with session_factory.create_session() as session: - stmt = select(DocumentSegmentSummary).where( - DocumentSegmentSummary.document_id == document_id, - DocumentSegmentSummary.dataset_id == dataset_id, - DocumentSegmentSummary.enabled.is_(True), # Only return enabled summaries - ) + stmt = select(DocumentSegmentSummary).where( + DocumentSegmentSummary.document_id == document_id, + DocumentSegmentSummary.dataset_id == dataset_id, + DocumentSegmentSummary.enabled.is_(True), + ) - if segment_ids: - stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) + if segment_ids: + stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) - return list(session.scalars(stmt).all()) + return list(session.scalars(stmt).all()) @staticmethod - def get_document_summary_index_status(document_id: str, dataset_id: str, tenant_id: str) -> str | None: + def get_document_summary_index_status( + document_id: str, + dataset_id: str, + tenant_id: str, + *, + session: Session, + ) -> str | None: """ Get summary_index_status for a single document. @@ -1302,26 +1293,28 @@ class SummaryIndexService: dataset_id: Dataset ID tenant_id: Tenant ID + Keyword Args: + session: SQLAlchemy session used to read summary status. + Returns: "SUMMARIZING" if there are pending summaries, None otherwise """ # Get all segments for this document (excluding qa_model and re_segment) - with session_factory.create_session() as session: - segment_ids = list( - session.scalars( - select(DocumentSegment.id).where( - DocumentSegment.document_id == document_id, - DocumentSegment.status != "re_segment", - DocumentSegment.tenant_id == tenant_id, - ) - ).all() - ) + segment_ids = list( + session.scalars( + select(DocumentSegment.id).where( + DocumentSegment.document_id == document_id, + DocumentSegment.status != "re_segment", + DocumentSegment.tenant_id == tenant_id, + ) + ).all() + ) if not segment_ids: return None # Get all summary records for these segments - summaries = SummaryIndexService.get_segments_summaries(segment_ids, dataset_id) + summaries = SummaryIndexService.get_segments_summaries(segment_ids, dataset_id, session=session) summary_status_map = {chunk_id: summary.status for chunk_id, summary in summaries.items()} # Check if there are any "not_started" or "generating" status summaries @@ -1335,7 +1328,11 @@ class SummaryIndexService: @staticmethod def get_documents_summary_index_status( - document_ids: list[str], dataset_id: str, tenant_id: str + document_ids: list[str], + dataset_id: str, + tenant_id: str, + *, + session: Session, ) -> dict[str, str | None]: """ Get summary_index_status for multiple documents. @@ -1345,6 +1342,9 @@ class SummaryIndexService: dataset_id: Dataset ID tenant_id: Tenant ID + Keyword Args: + session: SQLAlchemy session used to read summary status. + Returns: Dictionary mapping document_id to summary_index_status ("SUMMARIZING" or None) """ @@ -1352,14 +1352,13 @@ class SummaryIndexService: return {} # Get all segments for these documents (excluding qa_model and re_segment) - with session_factory.create_session() as session: - segments = session.execute( - select(DocumentSegment.id, DocumentSegment.document_id).where( - DocumentSegment.document_id.in_(document_ids), - DocumentSegment.status != "re_segment", - DocumentSegment.tenant_id == tenant_id, - ) - ).all() + segments = session.execute( + select(DocumentSegment.id, DocumentSegment.document_id).where( + DocumentSegment.document_id.in_(document_ids), + DocumentSegment.status != "re_segment", + DocumentSegment.tenant_id == tenant_id, + ) + ).all() # Group segments by document_id document_segments_map: dict[str, list[str]] = {} @@ -1371,7 +1370,7 @@ class SummaryIndexService: # Get all summary records for these segments all_segment_ids = [seg.id for seg in segments] - summaries = SummaryIndexService.get_segments_summaries(all_segment_ids, dataset_id) + summaries = SummaryIndexService.get_segments_summaries(all_segment_ids, dataset_id, session=session) summary_status_map = {chunk_id: summary.status for chunk_id, summary in summaries.items()} # Calculate summary_index_status for each document @@ -1407,7 +1406,7 @@ class SummaryIndexService: def get_document_summary_status_detail( document_id: str, dataset_id: str, - session: Session | scoped_session, + session: Session, ) -> DocumentSummaryStatusDetailDict: """ Get detailed summary status for a document. @@ -1448,6 +1447,7 @@ class SummaryIndexService: document_id=document_id, dataset_id=dataset_id, segment_ids=segment_ids, + session=session, ) # Create a mapping of chunk_id to summary diff --git a/api/services/tag_service.py b/api/services/tag_service.py index 2d89bafa920..f404ec0eb37 100644 --- a/api/services/tag_service.py +++ b/api/services/tag_service.py @@ -6,7 +6,7 @@ from flask_login import current_user from pydantic import BaseModel, Field from sqlalchemy import delete, func, select from sqlalchemy.engine import CursorResult -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from models.dataset import Dataset @@ -14,7 +14,6 @@ from models.enums import TagType from models.model import App, Tag, TagBinding from models.snippet import CustomizedSnippet -type _SessionLike = Session | scoped_session type _TagTypeLike = TagType | str @@ -41,7 +40,7 @@ class TagBindingDeletePayload(BaseModel): class TagService: @staticmethod - def get_tags(session: Session, tag_type: _TagTypeLike, current_tenant_id: str, keyword: str | None = None): + def get_tags(tag_type: _TagTypeLike, current_tenant_id: str, keyword: str | None = None, *, session: Session): stmt = ( select(Tag.id, Tag.type, Tag.name, func.count(TagBinding.id).label("binding_count")) .outerjoin(TagBinding, Tag.id == TagBinding.tag_id) @@ -61,7 +60,7 @@ class TagService: tag_type: _TagTypeLike, current_tenant_id: str, tag_ids: list[str], - session: _SessionLike, + session: Session, *, match_all: bool = False, ): @@ -107,7 +106,7 @@ class TagService: return tag_bindings @staticmethod - def get_tag_by_tag_name(tag_type: _TagTypeLike, current_tenant_id: str, tag_name: str, session: _SessionLike): + def get_tag_by_tag_name(tag_type: _TagTypeLike, current_tenant_id: str, tag_name: str, session: Session): if not tag_type or not tag_name: return [] tags = list( @@ -120,7 +119,7 @@ class TagService: return tags @staticmethod - def get_tags_by_target_id(tag_type: _TagTypeLike, current_tenant_id: str, target_id: str, session: _SessionLike): + def get_tags_by_target_id(tag_type: _TagTypeLike, current_tenant_id: str, target_id: str, session: Session): tags = session.scalars( select(Tag) .join(TagBinding, Tag.id == TagBinding.tag_id) @@ -135,7 +134,7 @@ class TagService: return tags or [] @staticmethod - def save_tags(payload: SaveTagPayload, session: _SessionLike) -> Tag: + def save_tags(payload: SaveTagPayload, session: Session) -> Tag: if TagService.get_tag_by_tag_name(payload.type, current_user.current_tenant_id, payload.name, session): raise ValueError("Tag name already exists") tag = Tag( @@ -151,7 +150,7 @@ class TagService: @staticmethod def update_tags( - payload: UpdateTagPayload, tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None + payload: UpdateTagPayload, tag_id: str, session: Session, *, tag_type: TagType | None = None ) -> Tag: current_tenant_id = current_user.current_tenant_id stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) @@ -178,7 +177,7 @@ class TagService: return tag @staticmethod - def get_tag_binding_count(tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None) -> int: + def get_tag_binding_count(tag_id: str, session: Session, *, tag_type: TagType | None = None) -> int: current_tenant_id = current_user.current_tenant_id stmt = ( select(func.count(TagBinding.id)) @@ -191,7 +190,7 @@ class TagService: return count @staticmethod - def delete_tag(tag_id: str, session: _SessionLike, *, tag_type: TagType | None = None): + def delete_tag(tag_id: str, session: Session, *, tag_type: TagType | None = None): current_tenant_id = current_user.current_tenant_id stmt = select(Tag).where(Tag.id == tag_id, Tag.tenant_id == current_tenant_id) if tag_type is not None: @@ -210,7 +209,7 @@ class TagService: session.commit() @staticmethod - def save_tag_binding(payload: TagBindingCreatePayload, session: _SessionLike): + def save_tag_binding(payload: TagBindingCreatePayload, session: Session): TagService.check_target_exists(payload.type, payload.target_id, session) valid_tag_ids = session.scalars( select(Tag.id).where( @@ -237,7 +236,7 @@ class TagService: session.commit() @staticmethod - def delete_tag_binding(payload: TagBindingDeletePayload, session: _SessionLike): + def delete_tag_binding(payload: TagBindingDeletePayload, session: Session): TagService.check_target_exists(payload.type, payload.target_id, session) result = cast( CursorResult, @@ -260,7 +259,7 @@ class TagService: session.commit() @staticmethod - def check_target_exists(type: _TagTypeLike, target_id: str, session: _SessionLike): + def check_target_exists(type: _TagTypeLike, target_id: str, session: Session): if type == "knowledge": dataset = session.scalar( select(Dataset) diff --git a/api/services/tools/builtin_tools_manage_service.py b/api/services/tools/builtin_tools_manage_service.py index e49ab8398f1..45480f71d1a 100644 --- a/api/services/tools/builtin_tools_manage_service.py +++ b/api/services/tools/builtin_tools_manage_service.py @@ -327,7 +327,7 @@ class BuiltinToolManageService: @staticmethod def generate_builtin_tool_provider_name( - session: Session, tenant_id: str, provider: str, credential_type: CredentialType + tenant_id: str, provider: str, credential_type: CredentialType, *, session: Session ) -> str: db_providers = session.scalars( select(BuiltinToolProvider) @@ -347,6 +347,7 @@ class BuiltinToolManageService: def get_builtin_tool_provider_credentials( tenant_id: str, provider_name: str, + session: Session, user: Account | None = None, include_credential_ids: list[str] | None = None, ) -> list[ToolProviderCredentialApiEntity]: @@ -367,7 +368,7 @@ class BuiltinToolManageService: from models.credential_permission import CredentialType as CredPermType from services.credential_permission_service import CredentialPermissionService - with db.session.no_autoflush: + with session.no_autoflush: base_filter = ( BuiltinToolProvider.tenant_id == tenant_id, BuiltinToolProvider.provider == provider_name, @@ -383,7 +384,7 @@ class BuiltinToolManageService: credential_type=CredPermType.BUILTIN_TOOL_PROVIDER, user=user, ) - visible_providers = list(db.session.scalars(visible_query).all()) + visible_providers = list(session.scalars(visible_query).all()) # Fetch any explicitly-included IDs that the visibility filter excluded. borrowed_ids: set[str] = set() @@ -397,7 +398,7 @@ class BuiltinToolManageService: .where(*base_filter, BuiltinToolProvider.id.in_(wanted_ids)) .order_by(*order) ) - borrowed_providers = list(db.session.scalars(borrowed_query).all()) + borrowed_providers = list(session.scalars(borrowed_query).all()) borrowed_ids = {p.id for p in borrowed_providers} providers = visible_providers + borrowed_providers @@ -427,7 +428,7 @@ class BuiltinToolManageService: if vis_str == "partial_members": credential_entity.partial_member_list = list( CredentialPermissionService.get_partial_member_list( - db.session, provider.id, CredPermType.BUILTIN_TOOL_PROVIDER + provider.id, CredPermType.BUILTIN_TOOL_PROVIDER, session=session ) ) if provider.id in borrowed_ids: @@ -439,6 +440,7 @@ class BuiltinToolManageService: def get_builtin_tool_provider_credential_info( tenant_id: str, provider: str, + session: Session, user: Account | None = None, include_credential_ids: list[str] | None = None, ) -> ToolProviderCredentialInfoApiEntity: @@ -450,6 +452,7 @@ class BuiltinToolManageService: credentials = BuiltinToolManageService.get_builtin_tool_provider_credentials( tenant_id, provider, + session=session, user=user, include_credential_ids=include_credential_ids, ) diff --git a/api/services/trigger/schedule_service.py b/api/services/trigger/schedule_service.py index a827222c1dc..495674248b1 100644 --- a/api/services/trigger/schedule_service.py +++ b/api/services/trigger/schedule_service.py @@ -26,10 +26,7 @@ logger = logging.getLogger(__name__) class ScheduleService: @staticmethod def create_schedule( - session: Session, - tenant_id: str, - app_id: str, - config: ScheduleConfig, + tenant_id: str, app_id: str, config: ScheduleConfig, *, session: Session ) -> WorkflowSchedulePlan: """ Create a new schedule with validated configuration. @@ -63,11 +60,7 @@ class ScheduleService: return schedule @staticmethod - def update_schedule( - session: Session, - schedule_id: str, - updates: SchedulePlanUpdate, - ) -> WorkflowSchedulePlan: + def update_schedule(schedule_id: str, updates: SchedulePlanUpdate, *, session: Session) -> WorkflowSchedulePlan: """ Update an existing schedule with validated configuration. @@ -110,10 +103,7 @@ class ScheduleService: return schedule @staticmethod - def delete_schedule( - session: Session, - schedule_id: str, - ) -> None: + def delete_schedule(schedule_id: str, *, session: Session) -> None: """ Delete a schedule plan. @@ -129,7 +119,7 @@ class ScheduleService: session.flush() @staticmethod - def get_tenant_owner(session: Session, tenant_id: str) -> Account: + def get_tenant_owner(tenant_id: str, *, session: Session) -> Account: """ Returns an account to execute scheduled workflows on behalf of the tenant. Prioritizes owner over admin to ensure proper authorization hierarchy. @@ -157,10 +147,7 @@ class ScheduleService: raise AccountNotFoundError(f"Account not found for tenant: {tenant_id}") @staticmethod - def update_next_run_at( - session: Session, - schedule_id: str, - ) -> datetime: + def update_next_run_at(schedule_id: str, *, session: Session) -> datetime: """ Advances the schedule to its next execution time after a successful trigger. Uses current time as base to prevent missing executions during delays. diff --git a/api/services/trigger/trigger_provider_service.py b/api/services/trigger/trigger_provider_service.py index b0a3de1cee8..8506c523a61 100644 --- a/api/services/trigger/trigger_provider_service.py +++ b/api/services/trigger/trigger_provider_service.py @@ -388,7 +388,7 @@ class TriggerProviderService: return subscription @classmethod - def delete_trigger_provider(cls, session: Session, tenant_id: str, subscription_id: str): + def delete_trigger_provider(cls, tenant_id: str, subscription_id: str, *, session: Session): """ Delete a trigger provider subscription within an existing session. diff --git a/api/services/trigger/trigger_subscription_operator_service.py b/api/services/trigger/trigger_subscription_operator_service.py index 5d7785549e6..491723c6ec2 100644 --- a/api/services/trigger/trigger_subscription_operator_service.py +++ b/api/services/trigger/trigger_subscription_operator_service.py @@ -40,12 +40,7 @@ class TriggerSubscriptionOperatorService: return list(subscribers) @classmethod - def delete_plugin_trigger_by_subscription( - cls, - session: Session, - tenant_id: str, - subscription_id: str, - ) -> None: + def delete_plugin_trigger_by_subscription(cls, tenant_id: str, subscription_id: str, *, session: Session) -> None: """Delete a plugin trigger by tenant_id and subscription_id within an existing session Args: diff --git a/api/services/trigger/webhook_service.py b/api/services/trigger/webhook_service.py index 23b3ac55b93..587048e2ccd 100644 --- a/api/services/trigger/webhook_service.py +++ b/api/services/trigger/webhook_service.py @@ -835,11 +835,7 @@ class WebhookService: # NOTE: don not use `with sessionmaker(bind=db.engine, expire_on_commit=False).begin()` # trigger_workflow_async need to handle multipe session commits internally with Session(db.engine, expire_on_commit=False) as session: - AsyncWorkflowService.trigger_workflow_async( - session, - end_user, - trigger_data, - ) + AsyncWorkflowService.trigger_workflow_async(end_user, trigger_data, session=session) quota_charge.commit() except Exception: quota_charge.refund() diff --git a/api/services/vector_service.py b/api/services/vector_service.py index 5b5088ec5a1..faf4fb085d6 100644 --- a/api/services/vector_service.py +++ b/api/services/vector_service.py @@ -1,6 +1,7 @@ import logging from sqlalchemy import delete, select +from sqlalchemy.orm import Session from core.model_manager import ModelInstance, ModelManager from core.rag.datasource.keyword.keyword_factory import Keyword @@ -11,7 +12,6 @@ from core.rag.index_processor.constant.index_type import IndexStructureType, Ind from core.rag.index_processor.index_processor_base import BaseIndexProcessor from core.rag.index_processor.index_processor_factory import IndexProcessorFactory from core.rag.models.document import AttachmentDocument, Document -from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from models import UploadFile from models.dataset import ChildChunk, Dataset, DatasetProcessRule, DocumentSegment, SegmentAttachmentBinding @@ -24,14 +24,20 @@ logger = logging.getLogger(__name__) class VectorService: @classmethod def create_segments_vector( - cls, keywords_list: list[list[str]] | None, segments: list[DocumentSegment], dataset: Dataset, doc_form: str + cls, + keywords_list: list[list[str]] | None, + segments: list[DocumentSegment], + dataset: Dataset, + doc_form: str, + session: Session, ): + """Create vector records for document segments using the caller's active DB session.""" documents: list[Document] = [] multimodal_documents: list[AttachmentDocument] = [] for segment in segments: if doc_form == IndexStructureType.PARENT_CHILD_INDEX: - dataset_document = db.session.get(DatasetDocument, segment.document_id) + dataset_document = session.get(DatasetDocument, segment.document_id) if not dataset_document: logger.warning( "Expected DatasetDocument record to exist, but none was found, document_id=%s, segment_id=%s", @@ -40,7 +46,7 @@ class VectorService: ) continue # get the process rule - processing_rule = db.session.get(DatasetProcessRule, dataset_document.dataset_process_rule_id) + processing_rule = session.get(DatasetProcessRule, dataset_document.dataset_process_rule_id) if not processing_rule: raise ValueError("No processing rule found.") # get embedding model instance @@ -63,7 +69,13 @@ class VectorService: else: raise ValueError("The knowledge base index technique is not high quality!") cls.generate_child_chunks( - segment, dataset_document, dataset, embedding_model_instance, processing_rule, False + segment, + dataset_document, + dataset, + embedding_model_instance, + processing_rule, + session, + False, ) else: rag_document = Document( @@ -136,8 +148,10 @@ class VectorService: dataset: Dataset, embedding_model_instance: ModelInstance, processing_rule: DatasetProcessRule, + session: Session, regenerate: bool = False, ): + """Generate child chunks and persist them with the caller's active DB session.""" index_processor = IndexProcessorFactory(dataset.doc_form).init_index_processor() assert segment.index_node_id if regenerate: @@ -184,8 +198,8 @@ class VectorService: type=SegmentType.AUTOMATIC, created_by=dataset_document.created_by, ) - db.session.add(child_segment) - db.session.commit() + session.add(child_segment) + session.commit() @classmethod def create_child_chunk_vector(cls, child_segment: ChildChunk, dataset: Dataset): @@ -255,7 +269,10 @@ class VectorService: vector.delete_by_ids([child_chunk.index_node_id]) @classmethod - def update_multimodel_vector(cls, segment: DocumentSegment, attachment_ids: list[str], dataset: Dataset): + def update_multimodel_vector( + cls, segment: DocumentSegment, attachment_ids: list[str], dataset: Dataset, session: Session + ): + """Update multimodal vectors and attachment bindings with the caller's active DB session.""" if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY: return @@ -274,19 +291,17 @@ class VectorService: vector.delete_by_ids(old_attachment_ids) # Delete existing segment attachment bindings in one operation - db.session.execute( - delete(SegmentAttachmentBinding).where(SegmentAttachmentBinding.segment_id == segment.id) - ) + session.execute(delete(SegmentAttachmentBinding).where(SegmentAttachmentBinding.segment_id == segment.id)) if not attachment_ids: - db.session.commit() + session.commit() return # Bulk fetch upload files - only fetch needed fields - upload_file_list = db.session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all() + upload_file_list = session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all() if not upload_file_list: - db.session.commit() + session.commit() return # Create a mapping for quick lookup @@ -329,16 +344,16 @@ class VectorService: # Bulk insert all bindings at once if bindings: - db.session.add_all(bindings) + session.add_all(bindings) # Add documents to vector store if any if documents and dataset.is_multimodal: vector.create_multimodal(documents) # Single commit for all operations - db.session.commit() + session.commit() except Exception: logger.exception("Failed to update multimodal vector for segment %s", segment.id) - db.session.rollback() + session.rollback() raise diff --git a/api/services/web_conversation_service.py b/api/services/web_conversation_service.py index 2c8a3be8631..96d95d5f5ac 100644 --- a/api/services/web_conversation_service.py +++ b/api/services/web_conversation_service.py @@ -2,7 +2,6 @@ from sqlalchemy import select from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account from models.enums import CreatorUserRole @@ -59,10 +58,10 @@ class WebConversationService: ) @classmethod - def pin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def pin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, session: Session): if not user: return - pinned_conversation = db.session.scalar( + pinned_conversation = session.scalar( select(PinnedConversation) .where( PinnedConversation.app_id == app_model.id, @@ -77,7 +76,7 @@ class WebConversationService: return conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) pinned_conversation = PinnedConversation( @@ -87,14 +86,14 @@ class WebConversationService: created_by=user.id, ) - db.session.add(pinned_conversation) - db.session.commit() + session.add(pinned_conversation) + session.commit() @classmethod - def unpin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None): + def unpin(cls, app_model: App, conversation_id: str, user: Account | EndUser | None, session: Session): if not user: return - pinned_conversation = db.session.scalar( + pinned_conversation = session.scalar( select(PinnedConversation) .where( PinnedConversation.app_id == app_model.id, @@ -108,5 +107,5 @@ class WebConversationService: if not pinned_conversation: return - db.session.delete(pinned_conversation) - db.session.commit() + session.delete(pinned_conversation) + session.commit() diff --git a/api/services/webapp_auth_service.py b/api/services/webapp_auth_service.py index 6ecc8eb8bc9..33267c53d5c 100644 --- a/api/services/webapp_auth_service.py +++ b/api/services/webapp_auth_service.py @@ -4,10 +4,10 @@ from datetime import UTC, datetime, timedelta from typing import Any from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound, Unauthorized from configs import dify_config -from extensions.ext_database import db from libs.helper import TokenManager from libs.passport import PassportService from libs.password import compare_password @@ -33,9 +33,9 @@ class WebAppAuthService: """Service for web app authentication.""" @staticmethod - def authenticate(email: str, password: str) -> Account: + def authenticate(email: str, password: str, session: Session) -> Account: """authenticate account with email and password""" - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) if not account: raise AccountNotFoundError() @@ -54,8 +54,8 @@ class WebAppAuthService: return access_token @classmethod - def get_user_through_email(cls, email: str): - account = AccountService.get_account_by_email_with_case_fallback(db.session, email) + def get_user_through_email(cls, email: str, session: Session): + account = AccountService.get_account_by_email_with_case_fallback(email, session=session) if not account: return None @@ -93,11 +93,11 @@ class WebAppAuthService: TokenManager.revoke_token(token, "email_code_login") @classmethod - def create_end_user(cls, app_code, email) -> EndUser: - site = db.session.scalar(select(Site).where(Site.code == app_code).limit(1)) + def create_end_user(cls, app_code, email, session: Session) -> EndUser: + site = session.scalar(select(Site).where(Site.code == app_code).limit(1)) if not site: raise NotFound("Site not found.") - app_model = db.session.get(App, site.app_id) + app_model = session.get(App, site.app_id) if not app_model: raise NotFound("App not found.") end_user = EndUser( @@ -109,8 +109,8 @@ class WebAppAuthService: name="enterpriseuser", external_user_id="enterpriseuser", ) - db.session.add(end_user) - db.session.commit() + session.add(end_user) + session.commit() return end_user @@ -133,7 +133,7 @@ class WebAppAuthService: @classmethod def is_app_require_permission_check( - cls, app_code: str | None = None, app_id: str | None = None, access_mode: str | None = None + cls, app_code: str | None = None, app_id: str | None = None, access_mode: str | None = None, *, session: Session ) -> bool: """ Check if the app requires permission check based on its access mode. @@ -145,7 +145,7 @@ class WebAppAuthService: raise ValueError("Either app_code or app_id must be provided.") if app_code: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=session) if not app_id: raise ValueError("App ID could not be determined from the provided app_code.") @@ -155,7 +155,9 @@ class WebAppAuthService: return False @classmethod - def get_app_auth_type(cls, app_code: str | None = None, access_mode: str | None = None) -> WebAppAuthType: + def get_app_auth_type( + cls, app_code: str | None = None, access_mode: str | None = None, *, session: Session + ) -> WebAppAuthType: """ Get the authentication type for the app based on its access mode. """ @@ -171,8 +173,8 @@ class WebAppAuthService: return WebAppAuthType.EXTERNAL if app_code: - app_id = AppService.get_app_id_by_code(app_code) + app_id = AppService.get_app_id_by_code(app_code, session=session) webapp_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=app_id) - return cls.get_app_auth_type(access_mode=webapp_settings.access_mode) + return cls.get_app_auth_type(access_mode=webapp_settings.access_mode, session=session) raise ValueError("Could not determine app authentication type.") diff --git a/api/services/workflow/node_output_inspector_service.py b/api/services/workflow/node_output_inspector_service.py index 66dcfec591f..5d6a8f1c675 100644 --- a/api/services/workflow/node_output_inspector_service.py +++ b/api/services/workflow/node_output_inspector_service.py @@ -52,9 +52,9 @@ from typing import Any from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.file_access import DatabaseFileAccessController -from core.db.session_factory import session_factory from core.workflow.nodes.agent_v2.binding_resolver import ( WorkflowAgentBindingError, WorkflowAgentBindingResolver, @@ -410,8 +410,8 @@ class NodeOutputInspectorService: The service is dependency-light: it holds a single :class:`WorkflowAgentBindingResolver` so agent v2 nodes can map to their declared outputs without re-implementing binding lookup. All other I/O - uses the global session factory so workflow runs / executions stay on the - repo-default code path. + receives an explicit SQLAlchemy session from its caller so transaction + ownership stays at the controller/task boundary. Tenancy is enforced via ``app_model.tenant_id`` + ``app_model.id`` on every load — the same scope guard regardless of trigger source. @@ -422,9 +422,13 @@ class NodeOutputInspectorService: # ── public API ──────────────────────────────────────────────────────── - def snapshot_workflow_run(self, *, app_model: App, workflow_run_id: str) -> WorkflowRunSnapshotView: + def snapshot_workflow_run( + self, *, app_model: App, workflow_run_id: str, session: Session + ) -> WorkflowRunSnapshotView: """Build the per-node snapshot for one debug workflow run.""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) executions_by_node = self._index_executions_by_node(executions) graph_nodes = _graph_nodes(workflow_run) @@ -447,9 +451,11 @@ class NodeOutputInspectorService: node_outputs=node_views, ) - def node_detail(self, *, app_model: App, workflow_run_id: str, node_id: str) -> NodeOutputsView: + def node_detail(self, *, app_model: App, workflow_run_id: str, node_id: str, session: Session) -> NodeOutputsView: """Per-node Inspector entry — returns one ``NodeOutputsView``.""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) graph_nodes = _graph_nodes(workflow_run) raw_node = next((n for n in graph_nodes if str(n.get("id")) == node_id), None) if raw_node is None: @@ -474,9 +480,12 @@ class NodeOutputInspectorService: workflow_run_id: str, node_id: str, output_name: str, + session: Session, ) -> OutputPreviewView: """Full payload for one declared output (with signed file URL).""" - workflow_run, executions = self._load_run_and_executions(app_model=app_model, workflow_run_id=workflow_run_id) + workflow_run, executions = self._load_run_and_executions( + app_model=app_model, workflow_run_id=workflow_run_id, session=session + ) graph_nodes = _graph_nodes(workflow_run) raw_node = next((n for n in graph_nodes if str(n.get("id")) == node_id), None) if raw_node is None: @@ -536,7 +545,7 @@ class NodeOutputInspectorService: # ── DB loading ──────────────────────────────────────────────────────── def _load_run_and_executions( - self, *, app_model: App, workflow_run_id: str + self, *, app_model: App, workflow_run_id: str, session: Session ) -> tuple[WorkflowRun, Sequence[WorkflowNodeExecutionModel]]: """Fetch the ``WorkflowRun`` row + every execution that belongs to it. @@ -548,24 +557,23 @@ class NodeOutputInspectorService: deliberately not checked here — D-1 was lifted 2026-05-26 and the Inspector now serves both draft and published runs. """ - with session_factory.create_session() as session: - workflow_run = session.scalar( - select(WorkflowRun).where( - WorkflowRun.id == workflow_run_id, - WorkflowRun.app_id == app_model.id, - WorkflowRun.tenant_id == app_model.tenant_id, - ) + workflow_run = session.scalar( + select(WorkflowRun).where( + WorkflowRun.id == workflow_run_id, + WorkflowRun.app_id == app_model.id, + WorkflowRun.tenant_id == app_model.tenant_id, ) - if workflow_run is None: - raise NodeOutputInspectorError("workflow_run_not_found", "Workflow run not found.") + ) + if workflow_run is None: + raise NodeOutputInspectorError("workflow_run_not_found", "Workflow run not found.") - executions = session.scalars( - select(WorkflowNodeExecutionModel).where( - WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, - WorkflowNodeExecutionModel.tenant_id == app_model.tenant_id, - WorkflowNodeExecutionModel.app_id == app_model.id, - ) - ).all() + executions = session.scalars( + select(WorkflowNodeExecutionModel).where( + WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, + WorkflowNodeExecutionModel.tenant_id == app_model.tenant_id, + WorkflowNodeExecutionModel.app_id == app_model.id, + ) + ).all() return workflow_run, executions diff --git a/api/services/workflow/workflow_converter.py b/api/services/workflow/workflow_converter.py index e279f1daaa3..5f787bb51cd 100644 --- a/api/services/workflow/workflow_converter.py +++ b/api/services/workflow/workflow_converter.py @@ -2,6 +2,7 @@ import json from typing import Any, TypedDict from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import ( DatasetEntity, @@ -18,7 +19,6 @@ from core.helper import encrypter from core.prompt.simple_prompt_transform import SimplePromptTransform from core.prompt.utils.prompt_template_parser import PromptTemplateParser from events.app_event import app_was_created -from extensions.ext_database import db from graphon.file import FileUploadConfig from graphon.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.utils.encoders import jsonable_encoder @@ -53,7 +53,14 @@ class WorkflowConverter: """ def convert_to_workflow( - self, app_model: App, account: Account, name: str, icon_type: str, icon: str, icon_background: str + self, + app_model: App, + account: Account, + name: str, + icon_type: str, + icon: str, + icon_background: str, + session: Session, ): """ Convert app to workflow @@ -77,7 +84,7 @@ class WorkflowConverter: raise ValueError("App model config is required") workflow = self.convert_app_model_config_to_workflow( - app_model=app_model, app_model_config=app_model.app_model_config, account_id=account.id + app_model=app_model, app_model_config=app_model.app_model_config, account_id=account.id, session=session ) # create new app @@ -97,17 +104,19 @@ class WorkflowConverter: new_app.created_by = account.id new_app.maintainer = account.id new_app.updated_by = account.id - db.session.add(new_app) - db.session.flush() + session.add(new_app) + session.flush() workflow.app_id = new_app.id - db.session.commit() + session.commit() app_was_created.send(new_app, account=account) return new_app - def convert_app_model_config_to_workflow(self, app_model: App, app_model_config: AppModelConfig, account_id: str): + def convert_app_model_config_to_workflow( + self, app_model: App, app_model_config: AppModelConfig, account_id: str, session: Session + ): """ Convert app model config to workflow mode :param app_model: App instance @@ -144,6 +153,7 @@ class WorkflowConverter: app_model=app_model, variables=app_config.variables, external_data_variables=app_config.external_data_variables, + session=session, ) for http_request_node in http_request_nodes: @@ -217,8 +227,8 @@ class WorkflowConverter: conversation_variables=[], ) - db.session.add(workflow) - db.session.commit() + session.add(workflow) + session.commit() return workflow @@ -262,7 +272,11 @@ class WorkflowConverter: } def _convert_to_http_request_node( - self, app_model: App, variables: list[VariableEntity], external_data_variables: list[ExternalDataVariableEntity] + self, + app_model: App, + variables: list[VariableEntity], + external_data_variables: list[ExternalDataVariableEntity], + session: Session, ) -> tuple[list[_NodeType], dict[str, str]]: """ Convert API Based Extension to HTTP Request Node @@ -290,7 +304,7 @@ class WorkflowConverter: # get api_based_extension api_based_extension = self._get_api_based_extension( - tenant_id=tenant_id, api_based_extension_id=api_based_extension_id + tenant_id=tenant_id, api_based_extension_id=api_based_extension_id, session=session ) # decrypt api_key @@ -650,14 +664,14 @@ class WorkflowConverter: else: return AppMode.ADVANCED_CHAT - def _get_api_based_extension(self, tenant_id: str, api_based_extension_id: str): + def _get_api_based_extension(self, tenant_id: str, api_based_extension_id: str, session: Session): """ Get API Based Extension :param tenant_id: tenant id :param api_based_extension_id: api based extension id :return: """ - api_based_extension = db.session.scalar( + api_based_extension = session.scalar( select(APIBasedExtension) .where(APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id) .limit(1) diff --git a/api/services/workflow_collaboration_service.py b/api/services/workflow_collaboration_service.py index bec61ce666d..5c635d7d66a 100644 --- a/api/services/workflow_collaboration_service.py +++ b/api/services/workflow_collaboration_service.py @@ -9,8 +9,8 @@ from collections.abc import Mapping from typing import Any, override from sqlalchemy import select +from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from models.account import Account from models.model import App from repositories.workflow_collaboration_repository import WorkflowCollaborationRepository, WorkflowSessionInfo @@ -94,20 +94,22 @@ class WorkflowCollaborationService: }, ) - def authorize_and_join_workflow_room(self, workflow_id: str, sid: str) -> tuple[str, bool] | None: + def authorize_and_join_workflow_room( + self, workflow_id: str, sid: str, *, session: Session + ) -> tuple[str, bool] | None: """ Join a collaboration room only after validating the socket session and tenant-scoped app access. The Socket.IO payload still calls the room key `workflow_id`, but the identifier is the workflow app's `App.id`. Returning `None` lets the controller reject the join before any Redis or room state is created. """ - session = self._socketio.get_session(sid) - user_id = session.get("user_id") - tenant_id = session.get("tenant_id") + socket_session = self._socketio.get_session(sid) + user_id = socket_session.get("user_id") + tenant_id = socket_session.get("tenant_id") if not user_id or not tenant_id: return None - if not self._can_access_workflow(workflow_id, str(tenant_id)): + if not self._can_access_workflow(workflow_id, str(tenant_id), session=session): logger.warning( "Workflow collaboration join rejected: workflow_id=%s tenant_id=%s user_id=%s sid=%s", workflow_id, @@ -121,8 +123,8 @@ class WorkflowCollaborationService: session_info: WorkflowSessionInfo = { "user_id": str(user_id), - "username": str(session.get("username", "Unknown")), - "avatar": session.get("avatar"), + "username": str(socket_session.get("username", "Unknown")), + "avatar": socket_session.get("avatar"), "sid": sid, "connected_at": int(time.time()), "server_id": self.server_id, @@ -140,10 +142,9 @@ class WorkflowCollaborationService: return str(user_id), is_leader - def _can_access_workflow(self, workflow_id: str, tenant_id: str) -> bool: + def _can_access_workflow(self, workflow_id: str, tenant_id: str, *, session: Session) -> bool: """Check room access without relying on Flask's app-context-bound scoped session.""" - with session_factory.create_session() as session: - app_id = session.scalar(select(App.id).where(App.id == workflow_id, App.tenant_id == tenant_id).limit(1)) + app_id = session.scalar(select(App.id).where(App.id == workflow_id, App.tenant_id == tenant_id).limit(1)) return app_id is not None def disconnect_session(self, sid: str) -> None: diff --git a/api/services/workflow_service.py b/api/services/workflow_service.py index 048b25c6bf9..95be0f7017a 100644 --- a/api/services/workflow_service.py +++ b/api/services/workflow_service.py @@ -7,7 +7,7 @@ from dataclasses import dataclass from typing import Any, cast from sqlalchemy import exists, select -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker from configs import dify_config from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager @@ -178,7 +178,7 @@ class WorkflowService: node_id=node_id, ) - def is_workflow_exist(self, app_model: App) -> bool: + def is_workflow_exist(self, app_model: App, *, session: Session) -> bool: stmt = select( exists().where( Workflow.tenant_id == app_model.tenant_id, @@ -186,23 +186,21 @@ class WorkflowService: Workflow.version == Workflow.VERSION_DRAFT, ) ) - return db.session.execute(stmt).scalar_one() + return session.execute(stmt).scalar_one() def get_draft_workflow( - self, app_model: App, workflow_id: str | None = None, session: Session | scoped_session | None = None + self, app_model: App, workflow_id: str | None = None, *, session: Session ) -> Workflow | None: """ Get draft workflow - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ if workflow_id: return self.get_published_workflow_by_id(app_model, workflow_id, session=session) # fetch draft workflow by app_model - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -215,18 +213,14 @@ class WorkflowService: # return draft workflow return workflow - def get_published_workflow_by_id( - self, app_model: App, workflow_id: str, session: Session | scoped_session | None = None - ) -> Workflow | None: + def get_published_workflow_by_id(self, app_model: App, workflow_id: str, *, session: Session) -> Workflow | None: """ fetch published workflow by workflow_id - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -244,20 +238,18 @@ class WorkflowService: ) return workflow - def get_published_workflow(self, app_model: App, session: Session | None = None) -> Workflow | None: + def get_published_workflow(self, app_model: App, *, session: Session) -> Workflow | None: """ Get published workflow - When ``session`` is provided, reuse it so callers that already hold a - Session avoid checking out an extra request-scoped ``db.session`` - connection. Falls back to ``db.session`` for backward compatibility. + Reuses the caller's active session so workflow reads stay in the same + transaction as the surrounding request or task. """ if not app_model.workflow_id: return None - bind = session if session is not None else db.session - workflow = bind.scalar( + workflow = session.scalar( select(Workflow) .where( Workflow.tenant_id == app_model.tenant_id, @@ -269,7 +261,7 @@ class WorkflowService: return workflow - def get_accessible_app_ids(self, app_ids: Sequence[str], tenant_id: str) -> set[str]: + def get_accessible_app_ids(self, app_ids: Sequence[str], tenant_id: str, *, session: Session) -> set[str]: """ Return app IDs that belong to the given tenant. """ @@ -277,7 +269,7 @@ class WorkflowService: return set() stmt = select(App.id).where(App.id.in_(app_ids), App.tenant_id == tenant_id) - return {str(app_id) for app_id in db.session.scalars(stmt).all()} + return {str(app_id) for app_id in session.scalars(stmt).all()} def get_all_published_workflow( self, @@ -327,13 +319,14 @@ class WorkflowService: account: Account, environment_variables: Sequence[VariableBase], conversation_variables: Sequence[VariableBase], + session: Session, ) -> Workflow: """ Sync draft workflow :raises WorkflowHashNotEqualError """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if workflow and workflow.unique_hash != unique_hash: raise WorkflowHashNotEqualError() @@ -357,7 +350,7 @@ class WorkflowService: environment_variables=environment_variables, conversation_variables=conversation_variables, ) - db.session.add(workflow) + session.add(workflow) # update draft workflow if found else: workflow.graph = json.dumps(graph) @@ -369,19 +362,19 @@ class WorkflowService: from services.agent.workflow_publish_service import WorkflowAgentPublishService - db.session.flush() + session.flush() WorkflowAgentPublishService.sync_agent_bindings_for_draft( - session=cast(Session, db.session), + session=session, draft_workflow=workflow, account_id=account.id, ) WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync( - session=cast(Session, db.session), + session=session, draft_workflow=workflow, ) # commit db session changes - db.session.commit() + session.commit() # trigger app workflow events app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=workflow) @@ -395,12 +388,13 @@ class WorkflowService: app_model: App, environment_variables: Sequence[VariableBase], account: Account, + session: Session, ): """ Update draft workflow environment variables """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -410,7 +404,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def update_draft_workflow_conversation_variables( self, @@ -418,12 +412,13 @@ class WorkflowService: app_model: App, conversation_variables: Sequence[VariableBase], account: Account, + session: Session, ): """ Update draft workflow conversation variables """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -433,7 +428,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def update_draft_workflow_features( self, @@ -441,12 +436,13 @@ class WorkflowService: app_model: App, features: dict, account: Account, + session: Session, ): """ Update draft workflow features """ # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model) + workflow = self.get_draft_workflow(app_model=app_model, session=session) if not workflow: raise ValueError("No draft workflow found.") @@ -459,7 +455,7 @@ class WorkflowService: workflow.updated_at = naive_utc_now() # commit db session changes - db.session.commit() + session.commit() def restore_published_workflow_to_draft( self, @@ -467,20 +463,23 @@ class WorkflowService: app_model: App, workflow_id: str, account: Account, + session: Session, ) -> Workflow: """Restore a published workflow snapshot into the draft workflow. Secret environment variables are copied server-side from the selected published workflow so the normal draft sync flow stays stateless. """ - source_workflow = self.get_published_workflow_by_id(app_model=app_model, workflow_id=workflow_id) + source_workflow = self.get_published_workflow_by_id( + app_model=app_model, workflow_id=workflow_id, session=session + ) if not source_workflow: raise WorkflowNotFoundError("Workflow not found.") self.validate_features_structure(app_model=app_model, features=source_workflow.normalized_features_dict) self.validate_graph_structure(graph=source_workflow.graph_dict) - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( tenant_id=app_model.tenant_id, app_id=app_model.id, @@ -491,9 +490,9 @@ class WorkflowService: ) if is_new_draft: - db.session.add(draft_workflow) + session.add(draft_workflow) - db.session.commit() + session.commit() app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=draft_workflow) return draft_workflow @@ -520,7 +519,7 @@ class WorkflowService: from services.feature_service import FeatureService if FeatureService.get_system_features().plugin_manager.enabled: - self._validate_workflow_credentials(draft_workflow) + self._validate_workflow_credentials(draft_workflow, session=session) # validate graph structure self.validate_graph_structure(graph=draft_workflow.graph_dict) @@ -577,7 +576,7 @@ class WorkflowService: # return new workflow return workflow - def _validate_workflow_credentials(self, workflow: Workflow) -> None: + def _validate_workflow_credentials(self, workflow: Workflow, *, session: Session) -> None: """ Validate all credentials in workflow nodes before publishing. @@ -609,7 +608,7 @@ class WorkflowService: ) else: # Check default workspace credential for this provider - self._check_default_tool_credential(workflow.tenant_id, provider) + self._check_default_tool_credential(workflow.tenant_id, provider, session=session) elif node_type == "agent": agent_params = node_data.get("agent_parameters", {}) @@ -622,7 +621,9 @@ class WorkflowService: # Validate load balancing credentials for agent model if load balancing is enabled agent_model_node_data = {"model": model_config} - self._validate_load_balancing_credentials(workflow, agent_model_node_data, node_id) + self._validate_load_balancing_credentials( + workflow, agent_model_node_data, node_id, session=session + ) # Validate agent tools tools = agent_params.get("tools", {}).get("value", []) @@ -636,7 +637,7 @@ class WorkflowService: check_credential_policy_compliance(credential_id, provider, PluginCredentialType.TOOL) else: - self._check_default_tool_credential(workflow.tenant_id, provider) + self._check_default_tool_credential(workflow.tenant_id, provider, session=session) elif node_type in ["llm", "knowledge_retrieval", "parameter_extractor", "question_classifier"]: model_config = node_data.get("model", {}) @@ -647,7 +648,7 @@ class WorkflowService: # Validate that the provider+model combination can fetch valid credentials self._validate_llm_model_config(workflow.tenant_id, provider, model_name) # Validate load balancing credentials if load balancing is enabled - self._validate_load_balancing_credentials(workflow, node_data, node_id) + self._validate_load_balancing_credentials(workflow, node_data, node_id, session=session) else: raise ValueError(f"Node {node_id} ({node_type}): Missing provider or model configuration") @@ -710,7 +711,7 @@ class WorkflowService: f"Failed to validate LLM model configuration (provider: {provider}, model: {model_name}): {str(e)}" ) - def _check_default_tool_credential(self, tenant_id: str, provider: str) -> None: + def _check_default_tool_credential(self, tenant_id: str, provider: str, *, session: Session) -> None: """ Check credential policy compliance for the default workspace credential of a tool provider. @@ -726,7 +727,7 @@ class WorkflowService: # Use the same fallback logic as runtime: get the first available credential # ordered by is_default DESC, created_at ASC (same as tool_manager.py) - default_provider = db.session.scalar( + default_provider = session.scalar( select(BuiltinToolProvider) .where( BuiltinToolProvider.tenant_id == tenant_id, @@ -753,7 +754,9 @@ class WorkflowService: except Exception as e: raise ValueError(f"Failed to validate default credential for tool provider {provider}: {str(e)}") - def _validate_load_balancing_credentials(self, workflow: Workflow, node_data: dict[str, Any], node_id: str) -> None: + def _validate_load_balancing_credentials( + self, workflow: Workflow, node_data: dict[str, Any], node_id: str, *, session: Session + ) -> None: """ Validate load balancing credentials for a workflow node. @@ -773,7 +776,9 @@ class WorkflowService: # Check if this model has load balancing enabled if self._is_load_balancing_enabled(workflow.tenant_id, provider, model_name): # Get all load balancing configurations for this model - load_balancing_configs = self._get_load_balancing_configs(workflow.tenant_id, provider, model_name) + load_balancing_configs = self._get_load_balancing_configs( + workflow.tenant_id, provider, model_name, session=session + ) # Validate each load balancing configuration try: for config in load_balancing_configs: @@ -817,7 +822,9 @@ class WorkflowService: # If we can't determine the status, assume load balancing is not enabled return False - def _get_load_balancing_configs(self, tenant_id: str, provider: str, model_name: str) -> list[dict[str, Any]]: + def _get_load_balancing_configs( + self, tenant_id: str, provider: str, model_name: str, *, session: Session + ) -> list[dict[str, Any]]: """ Get all load balancing configurations for a model. @@ -835,11 +842,17 @@ class WorkflowService: provider=provider, model=model_name, model_type="llm", # Load balancing is primarily used for LLM models + session=session, config_from="predefined-model", # Check both predefined and custom models ) _, custom_configs = model_load_balancing_service.get_load_balancing_configs( - tenant_id=tenant_id, provider=provider, model=model_name, model_type="llm", config_from="custom-model" + tenant_id=tenant_id, + provider=provider, + model=model_name, + model_type="llm", + session=session, + config_from="custom-model", ) all_configs = cast(list[dict[str, Any]], configs) + cast(list[dict[str, Any]], custom_configs) @@ -1047,6 +1060,7 @@ class WorkflowService: account: Account, node_id: str, inputs: Mapping[str, Any] | None = None, + session: Session, ) -> Mapping[str, Any]: """ Build a human input form preview for a draft workflow. @@ -1057,7 +1071,7 @@ class WorkflowService: node_id: Human input node ID. inputs: Values used to fill missing upstream variables referenced in form_content. """ - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1104,6 +1118,7 @@ class WorkflowService: form_inputs: Mapping[str, Any], inputs: Mapping[str, Any] | None = None, action: str, + session: Session, ) -> Mapping[str, Any]: """ Submit a human input form preview for a draft workflow. @@ -1116,7 +1131,7 @@ class WorkflowService: inputs: Values used to fill missing upstream variables referenced in form_content. action: Selected action ID. """ - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1189,8 +1204,9 @@ class WorkflowService: node_id: str, delivery_method_id: str, inputs: Mapping[str, Any] | None = None, + session: Session, ) -> None: - draft_workflow = self.get_draft_workflow(app_model=app_model) + draft_workflow = self.get_draft_workflow(app_model=app_model, session=session) if not draft_workflow: raise ValueError("Workflow not initialized") @@ -1529,7 +1545,7 @@ class WorkflowService: node_execution.status = WorkflowNodeExecutionStatus.FAILED node_execution.error = error - def convert_to_workflow(self, app_model: App, account: Account, args: dict[str, Any]) -> App: + def convert_to_workflow(self, app_model: App, account: Account, args: dict[str, Any], *, session: Session) -> App: """ Basic mode of chatbot app(expert mode) to workflow Completion App to Workflow App @@ -1553,6 +1569,7 @@ class WorkflowService: icon_type=args.get("icon_type", "emoji"), icon=args.get("icon", "🤖"), icon_background=args.get("icon_background", "#FFEAD5"), + session=session, ) return new_app diff --git a/api/services/workspace_service.py b/api/services/workspace_service.py index 180c077b88a..30853b2cc99 100644 --- a/api/services/workspace_service.py +++ b/api/services/workspace_service.py @@ -1,9 +1,9 @@ from flask_login import current_user from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from enums.cloud_plan import CloudPlan -from extensions.ext_database import db from models.account import Tenant, TenantAccountJoin, TenantAccountRole from services.account_service import TenantService from services.feature_service import FeatureService @@ -11,7 +11,7 @@ from services.feature_service import FeatureService class WorkspaceService: @classmethod - def get_tenant_info(cls, tenant: Tenant): + def get_tenant_info(cls, tenant: Tenant, session: Session): if not tenant: return None tenant_info: dict[str, object] = { @@ -25,7 +25,7 @@ class WorkspaceService: } # Get role of user - tenant_account_join = db.session.scalar( + tenant_account_join = session.scalar( select(TenantAccountJoin) .where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == current_user.id) .limit(1) @@ -37,7 +37,7 @@ class WorkspaceService: can_replace_logo = feature.can_replace_logo if can_replace_logo and TenantService.has_roles( - tenant, [TenantAccountRole.OWNER, TenantAccountRole.ADMIN], session=db.session + tenant, [TenantAccountRole.OWNER, TenantAccountRole.ADMIN], session=session ): base_url = dify_config.FILES_URL replace_webapp_logo = ( @@ -56,7 +56,7 @@ class WorkspaceService: from services.credit_pool_service import CreditPoolService - paid_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="paid") + paid_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="paid", session=session) # if the tenant is not on the sandbox plan and the paid pool is not full, use the paid pool if ( feature.billing.subscription.plan != CloudPlan.SANDBOX @@ -66,7 +66,7 @@ class WorkspaceService: tenant_info["trial_credits"] = paid_pool.quota_limit tenant_info["trial_credits_used"] = paid_pool.quota_used else: - trial_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="trial") + trial_pool = CreditPoolService.get_pool(tenant_id=tenant.id, pool_type="trial", session=session) if trial_pool: tenant_info["trial_credits"] = trial_pool.quota_limit tenant_info["trial_credits_used"] = trial_pool.quota_used diff --git a/api/tasks/batch_create_segment_to_index_task.py b/api/tasks/batch_create_segment_to_index_task.py index 9f19b03544d..0f92af21dc8 100644 --- a/api/tasks/batch_create_segment_to_index_task.py +++ b/api/tasks/batch_create_segment_to_index_task.py @@ -177,7 +177,7 @@ def batch_create_segment_to_index_task( with session_factory.create_session() as session: dataset = session.get(Dataset, dataset_id) if dataset: - VectorService.create_segments_vector(None, document_segments, dataset, document_config["doc_form"]) + VectorService.create_segments_vector(None, document_segments, dataset, document_config["doc_form"], session) redis_client.setex(indexing_cache_key, 600, "completed") end_at = time.perf_counter() diff --git a/api/tasks/regenerate_summary_index_task.py b/api/tasks/regenerate_summary_index_task.py index 16b59fdbba8..5cb8d4281f0 100644 --- a/api/tasks/regenerate_summary_index_task.py +++ b/api/tasks/regenerate_summary_index_task.py @@ -259,9 +259,8 @@ def regenerate_summary_index_task( # Regenerate both summary content and vectors (for summary_model change) SummaryIndexService.generate_and_vectorize_summary( - segment, dataset, summary_index_setting + segment, dataset, summary_index_setting, session=session ) - session.commit() total_segments_processed += 1 except Exception as e: diff --git a/api/tasks/retry_document_indexing_task.py b/api/tasks/retry_document_indexing_task.py index fa02afda15f..dddb7715d22 100644 --- a/api/tasks/retry_document_indexing_task.py +++ b/api/tasks/retry_document_indexing_task.py @@ -101,8 +101,9 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ session.commit() if dataset.runtime_mode == "rag_pipeline": - rag_pipeline_service = RagPipelineService() - rag_pipeline_service.retry_error_document(dataset, document, user) + with session_factory.create_session() as rag_session: + rag_pipeline_service = RagPipelineService(rag_session) + rag_pipeline_service.retry_error_document(dataset, document, user) else: indexing_runner = IndexingRunner() indexing_runner.run([document]) diff --git a/api/tasks/workflow_schedule_tasks.py b/api/tasks/workflow_schedule_tasks.py index 76386520000..38737f96e78 100644 --- a/api/tasks/workflow_schedule_tasks.py +++ b/api/tasks/workflow_schedule_tasks.py @@ -39,7 +39,7 @@ def run_schedule_trigger(schedule_id: str) -> None: if not schedule: raise ScheduleNotFoundError(f"Schedule {schedule_id} not found") - tenant_owner = ScheduleService.get_tenant_owner(session, schedule.tenant_id) + tenant_owner = ScheduleService.get_tenant_owner(schedule.tenant_id, session=session) if not tenant_owner: raise TenantOwnerNotFoundError(f"No owner or admin found for tenant {schedule.tenant_id}") diff --git a/api/tests/integration_tests/conftest.py b/api/tests/integration_tests/conftest.py index ea875e63fe8..25ee1974e92 100644 --- a/api/tests/integration_tests/conftest.py +++ b/api/tests/integration_tests/conftest.py @@ -84,7 +84,7 @@ def setup_account(request) -> Generator[Account, None, None]: password=secrets.token_hex(16), ip_address="localhost", language="en-US", - session=db.session, + session=db.session(), ) with _CACHED_APP.test_request_context(): diff --git a/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py b/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py index 23cbdd24b91..e6e681b426f 100644 --- a/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py +++ b/api/tests/integration_tests/services/plugin/test_plugin_lifecycle.py @@ -2,6 +2,7 @@ import pytest from sqlalchemy import delete, func, select from core.db.session_factory import session_factory +from extensions.ext_database import db from models import Tenant from models.account import ( TenantPluginAutoUpgradeCategory, @@ -39,17 +40,18 @@ def tenant(flask_req_ctx): class TestPluginPermissionLifecycle: def test_get_returns_none_for_new_tenant(self, tenant): - assert PluginPermissionService.get_permission(tenant) is None + assert PluginPermissionService.get_permission(tenant, session=db.session()) is None def test_change_creates_row(self, tenant): result = PluginPermissionService.change_permission( tenant, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.EVERYONE, + session=db.session, ) assert result is True - perm = PluginPermissionService.get_permission(tenant) + perm = PluginPermissionService.get_permission(tenant, session=db.session()) assert perm is not None assert perm.install_permission == TenantPluginInstallPermission.ADMINS assert perm.debug_permission == TenantPluginDebugPermission.EVERYONE @@ -59,13 +61,15 @@ class TestPluginPermissionLifecycle: tenant, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.NOBODY, + session=db.session, ) PluginPermissionService.change_permission( tenant, TenantPluginInstallPermission.EVERYONE, TenantPluginDebugPermission.ADMINS, + session=db.session, ) - perm = PluginPermissionService.get_permission(tenant) + perm = PluginPermissionService.get_permission(tenant, session=db.session()) assert perm is not None assert perm.install_permission == TenantPluginInstallPermission.EVERYONE assert perm.debug_permission == TenantPluginDebugPermission.ADMINS @@ -81,7 +85,7 @@ class TestPluginPermissionLifecycle: class TestPluginAutoUpgradeLifecycle: def test_get_returns_none_for_new_tenant(self, tenant): - assert PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) is None + assert PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) is None def test_change_creates_row(self, tenant): result = PluginAutoUpgradeService.change_strategy( @@ -92,10 +96,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) assert result is True - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert strategy.upgrade_time_of_day == 3 @@ -109,6 +114,7 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) PluginAutoUpgradeService.change_strategy( tenant, @@ -118,9 +124,10 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=["plugin-a"], category=PLUGIN_CATEGORY, + session=db.session(), ) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert strategy.upgrade_time_of_day == 12 @@ -128,9 +135,9 @@ class TestPluginAutoUpgradeLifecycle: assert strategy.include_plugins == ["plugin-a"] def test_exclude_plugin_creates_strategy_when_none_exists(self, tenant): - PluginAutoUpgradeService.exclude_plugin(tenant, "my-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "my-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert "my-plugin" in strategy.exclude_plugins @@ -144,10 +151,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=["existing"], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "new-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "new-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert "existing" in strategy.exclude_plugins assert "new-plugin" in strategy.exclude_plugins @@ -161,10 +169,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=["same-plugin"], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "same-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "same-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.exclude_plugins.count("same-plugin") == 1 @@ -177,10 +186,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=["p1", "p2"], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "p1", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "p1", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert "p1" not in strategy.include_plugins assert "p2" in strategy.include_plugins @@ -194,10 +204,11 @@ class TestPluginAutoUpgradeLifecycle: exclude_plugins=[], include_plugins=[], category=PLUGIN_CATEGORY, + session=db.session(), ) - PluginAutoUpgradeService.exclude_plugin(tenant, "excluded-plugin", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin(tenant, "excluded-plugin", PLUGIN_CATEGORY, session=db.session()) - strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY) + strategy = PluginAutoUpgradeService.get_strategy(tenant, PLUGIN_CATEGORY, session=db.session()) assert strategy is not None assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert "excluded-plugin" in strategy.exclude_plugins diff --git a/api/tests/integration_tests/services/test_node_output_inspector_service.py b/api/tests/integration_tests/services/test_node_output_inspector_service.py index 5a8c07e0434..c2253a20c15 100644 --- a/api/tests/integration_tests/services/test_node_output_inspector_service.py +++ b/api/tests/integration_tests/services/test_node_output_inspector_service.py @@ -219,6 +219,36 @@ def _stub_resolver(declared_outputs_payload: list[dict[str, Any]]): return _Resolver() +def _snapshot_workflow_run(service: NodeOutputInspectorService, *, app_model: Any, workflow_run_id: str): + with session_factory.create_session() as session: + return service.snapshot_workflow_run(app_model=app_model, workflow_run_id=workflow_run_id, session=session) + + +def _node_detail(service: NodeOutputInspectorService, *, app_model: Any, workflow_run_id: str, node_id: str): + with session_factory.create_session() as session: + return service.node_detail( + app_model=app_model, workflow_run_id=workflow_run_id, node_id=node_id, session=session + ) + + +def _output_preview( + service: NodeOutputInspectorService, + *, + app_model: Any, + workflow_run_id: str, + node_id: str, + output_name: str, +): + with session_factory.create_session() as session: + return service.output_preview( + app_model=app_model, + workflow_run_id=workflow_run_id, + node_id=node_id, + output_name=output_name, + session=session, + ) + + # ────────────────────────────────────────────────────────────────────────────── # Tests # ────────────────────────────────────────────────────────────────────────────── @@ -229,7 +259,8 @@ def test_snapshot_returns_agent_v2_declared_outputs_with_status_ready(seeded_run real ``WorkflowRun`` + ``WorkflowNodeExecutionModel`` rows.""" app_model, workflow_run, _ = seeded_run service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - snapshot = service.snapshot_workflow_run( + snapshot = _snapshot_workflow_run( + service, app_model=app_model, workflow_run_id=workflow_run.id, ) @@ -256,7 +287,7 @@ def test_snapshot_404s_for_missing_run(fake_app_model): """Service raises ``workflow_run_not_found`` when the row doesn't exist.""" service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=str(uuid.uuid4())) + _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=str(uuid.uuid4())) assert exc.value.code == "workflow_run_not_found" @@ -266,7 +297,7 @@ def test_snapshot_404s_for_cross_tenant_access(seeded_run): intruder = SimpleNamespace(id=str(uuid.uuid4()), tenant_id=str(uuid.uuid4())) service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=intruder, workflow_run_id=workflow_run.id) + _snapshot_workflow_run(service, app_model=intruder, workflow_run_id=workflow_run.id) assert exc.value.code == "workflow_run_not_found" @@ -286,7 +317,7 @@ def test_snapshot_404s_for_published_run_per_decision_d1(flask_req_ctx, fake_app try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([])) with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) assert exc.value.code == "published_run_inspector_not_implemented" finally: with session_factory.create_session() as session: @@ -328,7 +359,7 @@ def test_snapshot_surfaces_type_check_failure_from_metadata(flask_req_ctx, fake_ try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "summary", "type": "string"}])) - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.TYPE_CHECK_FAILED assert output.type_check is not None @@ -375,7 +406,7 @@ def test_snapshot_surfaces_output_check_failure_from_metadata(flask_req_ctx, fak "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", return_value="https://signed.example/report", ): - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.OUTPUT_CHECK_FAILED assert output.output_check is not None @@ -391,7 +422,8 @@ def test_snapshot_surfaces_output_check_failure_from_metadata(flask_req_ctx, fak def test_node_detail_serves_one_node(seeded_run): app_model, workflow_run, _ = seeded_run service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - view = service.node_detail( + view = _node_detail( + service, app_model=app_model, workflow_run_id=workflow_run.id, node_id="agent-node-1", @@ -421,7 +453,8 @@ def test_output_preview_for_file_renders_signed_url(seeded_run, fake_app_model): "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", return_value="https://signed.example/x.pdf", ): - preview = service.output_preview( + preview = _output_preview( + service, app_model=fake_app_model, workflow_run_id=workflow_run.id, node_id="agent-node-1", @@ -466,7 +499,7 @@ def test_keeps_latest_execution_per_node_by_index(flask_req_ctx, fake_app_model) try: service = NodeOutputInspectorService(binding_resolver=_stub_resolver([{"name": "text", "type": "string"}])) - snapshot = service.snapshot_workflow_run(app_model=fake_app_model, workflow_run_id=run_id) + snapshot = _snapshot_workflow_run(service, app_model=fake_app_model, workflow_run_id=run_id) assert snapshot.node_outputs[0].outputs[0].value_preview == "second attempt" finally: with session_factory.create_session() as session: diff --git a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py index ae37d305670..df9d655fbbc 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py +++ b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py @@ -500,7 +500,11 @@ class TestWorkflowDraftVariableEndpoints: api = workflow_draft_variable_module.WorkflowVariableCollectionApi() method = unwrap(api.get) - monkeypatch.setattr(workflow_draft_variable_module, "db", SimpleNamespace(engine=MagicMock())) + monkeypatch.setattr( + workflow_draft_variable_module, + "db", + SimpleNamespace(engine=MagicMock(), session=MagicMock()), + ) class DummySessionCtx: def __enter__(self): diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py index e55b46d38bf..ef8c0add709 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_data_source_bearer_auth.py @@ -85,7 +85,7 @@ def test_create_binding_successful( assert response.status_code == 200 assert response.get_json() == {"result": "success"} - create_auth.assert_called_once_with(ANY, tenant_id, payload) + create_auth.assert_called_once_with(tenant_id, payload, session=ANY) def test_create_binding_failure( diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py index 109332e16c9..d893e9e6efb 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_email_register.py @@ -270,7 +270,7 @@ def test_get_account_by_email_with_case_fallback_falls_back_to_lowercase(): second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Case@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Case@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py index 812aa299c1b..a7eba9d723c 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_forgot_password.py @@ -165,7 +165,7 @@ def test_get_account_by_email_with_case_fallback_falls_back_to_lowercase(): second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Mixed@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py index 464e0134a2f..484ca71ca59 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py +++ b/api/tests/test_containers_integration_tests/controllers/console/auth/test_oauth.py @@ -494,7 +494,7 @@ class TestAccountGeneration: second_result.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first_result, second_result] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Case@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Case@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index c34810c97d0..6c73b0010ed 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -5,7 +5,7 @@ from __future__ import annotations from collections.abc import Callable from inspect import unwrap from typing import cast -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch from uuid import uuid4 import pytest @@ -90,55 +90,51 @@ class TestPipelineTemplateDetailApi: "graph": {"nodes": nodes, "edges": edges, "viewport": viewport}, } - service = MagicMock() - service.get_pipeline_template_detail.return_value = template - with ( app.test_request_context("/?type=built-in"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=template, + ) as get_detail_mock, ): response, status = method(api, MagicMock(), "tpl-1") assert status == 200 assert response == {**template, "created_by": None} + get_detail_mock.assert_called_once_with("tpl-1", type="built-in", session=ANY) def test_get_returns_404_when_template_not_found(self, app: Flask) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - service = MagicMock() - service.get_pipeline_template_detail.return_value = None - with ( app.test_request_context("/?type=built-in"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=None, + ) as get_detail_mock, ): with pytest.raises(NotFound): method(api, MagicMock(), "non-existent-id") + get_detail_mock.assert_called_once_with("non-existent-id", type="built-in", session=ANY) + def test_get_returns_404_for_customized_type_not_found(self, app: Flask) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - service = MagicMock() - service.get_pipeline_template_detail.return_value = None - with ( app.test_request_context("/?type=customized"), patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService", - return_value=service, - ), + "controllers.console.datasets.rag_pipeline.rag_pipeline.RagPipelineService.get_pipeline_template_detail", + return_value=None, + ) as get_detail_mock, ): with pytest.raises(NotFound): method(api, MagicMock(), "non-existent-id") + get_detail_mock.assert_called_once_with("non-existent-id", type="customized", session=ANY) + class TestCustomizedPipelineTemplateApi: @pytest.fixture @@ -186,7 +182,7 @@ class TestCustomizedPipelineTemplateApi: ): response, status = method(api, tenant_id, "tpl-1") - delete_mock.assert_called_once_with("tpl-1", tenant_id) + delete_mock.assert_called_once_with("tpl-1", tenant_id, session=ANY) assert status == 204 assert response == "" diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py b/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py index e60558040a5..4cca4c2170f 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py +++ b/api/tests/test_containers_integration_tests/controllers/console/test_api_based_extension.py @@ -97,13 +97,13 @@ def test_list_scopes_api_based_extensions_to_authenticated_tenant( assert account_create_response.status_code == 201 APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=foreign_tenant_id, name="Foreign API", api_endpoint="https://foreign.example.com/hook", api_key="foreign-secret-12345", ), + session=db_session_with_containers, ) response = test_client_with_containers.get( diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py index 4cdbec3e30e..4222f49a28a 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_account_sessions.py @@ -30,7 +30,6 @@ def _mint_account_token( ) -> MintResult: """Mint a real, persisted ``dfoa_`` access token for ``account``.""" return mint_oauth_token( - db_session, redis_client, subject_email=account.email, subject_issuer=None, @@ -39,6 +38,7 @@ def _mint_account_token( device_label=device_label, prefix=PREFIX_OAUTH_ACCOUNT, ttl_days=14, + session=db_session, ) diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py index 2b9feeede14..4d9bfb5ea17 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_dsl.py @@ -96,7 +96,7 @@ def _app_and_account(db_session: Session, *, mode: str = "chat") -> tuple[App, A api_rph=100, api_rpm=10, ) - app_model = AppService().create_app(tenant.id, app_args, account) + app_model = AppService().create_app(tenant.id, app_args, account, session=db_session) return app_model, account diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py index 8e9278ad244..df4f3873b1e 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_app_run.py @@ -24,7 +24,7 @@ def _create_app(db_session: Session, account: Account, *, name: str = "Runner") icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) db_session.commit() return app_model diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py index 24580ae0a0e..ce1425d9e61 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_apps.py @@ -39,7 +39,7 @@ def _create_app( icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) # The openapi surface gate keys off ``enable_api``; flip it explicitly so # the test states the visibility precondition rather than relying on the # template default. diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py index 86cf70613c9..31a5485d2b3 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_files.py @@ -25,7 +25,7 @@ def _create_app(db_session: Session, account: Account, *, name: str = "Uploader" icon="🤖", icon_background="#FF6B6B", ) - app_model = AppService().create_app(tenant.id, params, account) + app_model = AppService().create_app(tenant.id, params, account, session=db_session) db_session.commit() return app_model diff --git a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py index 372157813cc..d670425be0c 100644 --- a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py +++ b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py @@ -734,7 +734,8 @@ class TestDatasetApiPatch: assert response["name"] == "Updated Dataset" assert response["partial_member_list"] == ["user-1"] mock_dataset_svc.update_dataset.assert_called_once() - session, _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + session = mock_dataset_svc.update_dataset.call_args.kwargs["session"] assert isinstance(session, (Session, scoped_session)) assert update_data["name"] == "Updated Dataset" assert update_data["permission"] == "partial_members" @@ -1013,7 +1014,7 @@ class TestDatasetTagsApiGet: assert status == 200 assert response == [{"id": "tag-1", "name": "Test Tag", "type": "knowledge", "binding_count": "0"}] - mock_tag_svc.get_tags.assert_called_once_with(SessionMatcher(), "knowledge", "tenant-1") + mock_tag_svc.get_tags.assert_called_once_with("knowledge", "tenant-1", session=SessionMatcher()) @patch("controllers.service_api.dataset.dataset.current_user") def test_list_tags_from_db( diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py b/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py index d568a1c0b04..cd754782df4 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_web_forgot_password.py @@ -57,7 +57,7 @@ class TestForgotPasswordSendEmailApi: response = ForgotPasswordSendEmailApi().post() assert response == {"result": "success", "data": "token-123"} - mock_get_account.assert_called_once_with(ANY, "User@Example.com") + mock_get_account.assert_called_once_with("User@Example.com", session=ANY) mock_send_mail.assert_called_once_with(account=mock_account, email="user@example.com", language="zh-Hans") mock_extract_ip.assert_called_once() mock_rate_limit.assert_called_once_with("127.0.0.1") @@ -177,7 +177,7 @@ class TestForgotPasswordResetApi: response = ForgotPasswordResetApi().post() assert response == {"result": "success"} - mock_get_account.assert_called_once_with(ANY, "User@Example.com") + mock_get_account.assert_called_once_with("User@Example.com", session=ANY) mock_update_account.assert_called_once() mock_revoke_token.assert_called_once_with("token-123") diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py index 3eab8ccbee5..aa85ac2ca7b 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py @@ -19,6 +19,8 @@ from controllers.web.wraps import ( decode_jwt_token, ) +pytestmark = pytest.mark.usefixtures("db_session_with_containers") + class TestValidateWebappToken: def test_enterprise_enabled_and_app_auth_requires_webapp_source(self) -> None: diff --git a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py b/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py index e2f8c8fc703..e22aa102328 100644 --- a/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py +++ b/api/tests/test_containers_integration_tests/services/auth/test_api_key_auth_service.py @@ -51,7 +51,7 @@ class TestApiKeyAuthService: self._create_binding(db_session_with_containers, tenant_id=tenant_id, category=category, provider=provider) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) assert len(result) >= 1 tenant_results = [r for r in result if r.tenant_id == tenant_id] @@ -61,7 +61,7 @@ class TestApiKeyAuthService: def test_get_provider_auth_list_empty( self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id ): - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] @@ -74,7 +74,7 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id) + result = ApiKeyAuthService.get_provider_auth_list(tenant_id, session=db_session_with_containers) tenant_results = [r for r in result if r.tenant_id == tenant_id] assert tenant_results == [] @@ -95,7 +95,7 @@ class TestApiKeyAuthService: mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.return_value = "encrypted_test_key_123" - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) mock_factory.assert_called_once() mock_auth_instance.validate_credentials.assert_called_once() @@ -118,7 +118,7 @@ class TestApiKeyAuthService: mock_auth_instance.validate_credentials.return_value = False mock_factory.return_value = mock_auth_instance - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) db_session_with_containers.expire_all() bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id).all() @@ -142,7 +142,7 @@ class TestApiKeyAuthService: original_key = mock_args["credentials"]["config"]["api_key"] - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=db_session_with_containers) assert mock_args["credentials"]["config"]["api_key"] == "encrypted_test_key_123" assert mock_args["credentials"]["config"]["api_key"] != original_key @@ -166,14 +166,18 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result == mock_credentials def test_get_auth_credentials_not_found( self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id, category, provider ): - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result is None @@ -190,7 +194,9 @@ class TestApiKeyAuthService: ) db_session_with_containers.expire_all() - result = ApiKeyAuthService.get_auth_credentials(db_session_with_containers, tenant_id, category, provider) + result = ApiKeyAuthService.get_auth_credentials( + tenant_id, category, provider, session=db_session_with_containers + ) assert result == special_credentials assert result["config"]["api_key"] == "key_with_中文_and_special_chars_!@#$%" @@ -204,7 +210,7 @@ class TestApiKeyAuthService: binding_id = binding.id db_session_with_containers.expire_all() - ApiKeyAuthService.delete_provider_auth(db_session_with_containers, tenant_id, binding_id) + ApiKeyAuthService.delete_provider_auth(tenant_id, binding_id, session=db_session_with_containers) db_session_with_containers.expire_all() remaining = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(id=binding_id).first() @@ -214,7 +220,7 @@ class TestApiKeyAuthService: self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id ): # Should not raise when binding not found - ApiKeyAuthService.delete_provider_auth(db_session_with_containers, tenant_id, str(uuid4())) + ApiKeyAuthService.delete_provider_auth(tenant_id, str(uuid4()), session=db_session_with_containers) def test_validate_api_key_auth_args_success(self, mock_args): ApiKeyAuthService.validate_api_key_auth_args(mock_args) @@ -291,13 +297,13 @@ class TestApiKeyAuthService: mock_session = MagicMock() mock_session.commit.side_effect = Exception("Database error") with pytest.raises(Exception, match="Database error"): - ApiKeyAuthService.create_provider_auth(mock_session, tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=mock_session) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") def test_create_provider_auth_factory_exception(self, mock_factory: MagicMock, tenant_id, mock_args): mock_factory.side_effect = Exception("Factory error") with pytest.raises(Exception, match="Factory error"): - ApiKeyAuthService.create_provider_auth(MagicMock(), tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) @patch("services.auth.api_key_auth_service.ApiKeyAuthFactory") @patch("services.auth.api_key_auth_service.encrypter") @@ -307,7 +313,7 @@ class TestApiKeyAuthService: mock_factory.return_value = mock_auth_instance mock_encrypter.encrypt_token.side_effect = Exception("Encryption error") with pytest.raises(Exception, match="Encryption error"): - ApiKeyAuthService.create_provider_auth(MagicMock(), tenant_id, mock_args) + ApiKeyAuthService.create_provider_auth(tenant_id, mock_args, session=MagicMock()) def test_validate_api_key_auth_args_none_input(self): with pytest.raises(TypeError): diff --git a/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py b/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py index 9b86ab41f2b..cd3ed01cbfa 100644 --- a/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py +++ b/api/tests/test_containers_integration_tests/services/auth/test_auth_integration.py @@ -57,7 +57,7 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_fc_test_key_123" args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) mock_http.assert_called_once() call_args = mock_http.call_args @@ -101,15 +101,15 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_key" args1 = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args1) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args1, session=db_session_with_containers) args2 = {"category": category, "provider": AuthType.JINA, "credentials": jina_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_2, args2) + ApiKeyAuthService.create_provider_auth(tenant_id_2, args2, session=db_session_with_containers) db_session_with_containers.expire_all() - result1 = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id_1) - result2 = ApiKeyAuthService.get_provider_auth_list(db_session_with_containers, tenant_id_2) + result1 = ApiKeyAuthService.get_provider_auth_list(tenant_id_1, session=db_session_with_containers) + result2 = ApiKeyAuthService.get_provider_auth_list(tenant_id_2, session=db_session_with_containers) assert len(result1) == 1 assert result1[0].tenant_id == tenant_id_1 @@ -120,7 +120,7 @@ class TestAuthIntegration: self, flask_app_with_containers: Flask, db_session_with_containers: Session, tenant_id_2, category ): result = ApiKeyAuthService.get_auth_credentials( - db_session_with_containers, tenant_id_2, category, AuthType.FIRECRAWL + tenant_id_2, category, AuthType.FIRECRAWL, session=db_session_with_containers ) assert result is None @@ -163,7 +163,7 @@ class TestAuthIntegration: "provider": AuthType.FIRECRAWL, "credentials": {"auth_type": "bearer", "config": {"api_key": "fc_test_key_123"}}, } - ApiKeyAuthService.create_provider_auth(db.session(), tenant_id_1, thread_args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, thread_args, session=db.session()) results.append("success") except Exception as e: exceptions.append(e) @@ -216,7 +216,7 @@ class TestAuthIntegration: args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} with pytest.raises(httpx.RequestError): - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) db_session_with_containers.expire_all() bindings = db_session_with_containers.query(DataSourceApiKeyAuthBinding).filter_by(tenant_id=tenant_id_1).all() @@ -253,12 +253,12 @@ class TestAuthIntegration: mock_encrypt.return_value = "encrypted_key" args = {"category": category, "provider": AuthType.FIRECRAWL, "credentials": firecrawl_credentials} - ApiKeyAuthService.create_provider_auth(db_session_with_containers, tenant_id_1, args) + ApiKeyAuthService.create_provider_auth(tenant_id_1, args, session=db_session_with_containers) db_session_with_containers.expire_all() result = ApiKeyAuthService.get_auth_credentials( - db_session_with_containers, tenant_id_1, category, AuthType.FIRECRAWL + tenant_id_1, category, AuthType.FIRECRAWL, session=db_session_with_containers ) assert result is not None assert result["config"]["api_key"] == "encrypted_key" diff --git a/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py b/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py index 646a0592630..0a34733adeb 100644 --- a/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py +++ b/api/tests/test_containers_integration_tests/services/enterprise/test_account_deletion_sync.py @@ -6,7 +6,7 @@ Redis queuing, error handling, and community vs enterprise behavior. from __future__ import annotations -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest @@ -118,7 +118,7 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = False - result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted") + result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted", session=MagicMock()) assert result is True mock_queue_task.assert_not_called() @@ -137,7 +137,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is True assert mock_queue_task.call_count == 3 @@ -151,7 +153,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted") + result = sync_account_deletion( + account_id=str(uuid4()), source="account_deleted", session=db_session_with_containers + ) assert result is True mock_queue_task.assert_not_called() @@ -176,7 +180,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is False assert mock_queue_task.call_count == 3 @@ -196,7 +202,9 @@ class TestSyncAccountDeletion: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: mock_config.ENTERPRISE_ENABLED = True - result = sync_account_deletion(account_id=account_id, source="account_deleted") + result = sync_account_deletion( + account_id=account_id, source="account_deleted", session=db_session_with_containers + ) assert result is False mock_queue_task.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py b/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py index 0a8f49bc7a7..a52458ac972 100644 --- a/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py +++ b/api/tests/test_containers_integration_tests/services/plugin/test_plugin_permission_service.py @@ -2,7 +2,6 @@ from __future__ import annotations from uuid import uuid4 -import pytest from sqlalchemy import func, select from sqlalchemy.orm import Session @@ -38,7 +37,7 @@ class TestGetPermission: db_session_with_containers.add(permission) db_session_with_containers.commit() - result = PluginPermissionService.get_permission(tenant_id) + result = PluginPermissionService.get_permission(tenant_id, session=db_session_with_containers) assert result is not None assert result.id == permission.id @@ -46,9 +45,8 @@ class TestGetPermission: assert result.install_permission == TenantPluginInstallPermission.ADMINS assert result.debug_permission == TenantPluginDebugPermission.EVERYONE - @pytest.mark.usefixtures("flask_app_with_containers") - def test_returns_none_when_not_found(self) -> None: - result = PluginPermissionService.get_permission(_tenant_id()) + def test_returns_none_when_not_found(self, db_session_with_containers: Session) -> None: + result = PluginPermissionService.get_permission(_tenant_id(), session=db_session_with_containers) assert result is None @@ -63,6 +61,7 @@ class TestChangePermission: tenant_id, TenantPluginInstallPermission.EVERYONE, TenantPluginDebugPermission.EVERYONE, + session=db_session_with_containers, ) permission = _get_permission(db_session_with_containers, tenant_id) @@ -85,6 +84,7 @@ class TestChangePermission: tenant_id, TenantPluginInstallPermission.ADMINS, TenantPluginDebugPermission.ADMINS, + session=db_session_with_containers, ) permission = _get_permission(db_session_with_containers, tenant_id) diff --git a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py index 2e7df67d266..75d127ce6b2 100644 --- a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py +++ b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py @@ -42,7 +42,9 @@ class TestRagPipelineServiceGetPipeline: yield db_session_with_containers.rollback() - def _make_service(self, flask_app_with_containers: Flask) -> RagPipelineService: + def _make_service( + self, flask_app_with_containers: Flask, db_session_with_containers: Session + ) -> RagPipelineService: with ( patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", @@ -54,7 +56,7 @@ class TestRagPipelineServiceGetPipeline: ), ): session_factory = sessionmaker(bind=flask_app_with_containers.extensions["sqlalchemy"].engine) - return RagPipelineService(session_maker=session_factory) + return RagPipelineService(db_session_with_containers, session_maker=session_factory) def _create_pipeline(self, db_session: Session, tenant_id: str, created_by: str) -> Pipeline: pipeline = Pipeline( @@ -85,7 +87,7 @@ class TestRagPipelineServiceGetPipeline: self, db_session_with_containers: Session, flask_app_with_containers: Flask ) -> None: """get_pipeline raises ValueError when dataset does not exist.""" - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) with pytest.raises(ValueError, match="Dataset not found"): service.get_pipeline(tenant_id=str(uuid4()), dataset_id=str(uuid4())) @@ -99,10 +101,10 @@ class TestRagPipelineServiceGetPipeline: dataset = self._create_dataset(db_session_with_containers, tenant_id, created_by, pipeline_id=None) db_session_with_containers.flush() - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) with pytest.raises(ValueError, match="Pipeline not found"): - service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) + service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) def test_get_pipeline_returns_pipeline_when_found( self, db_session_with_containers: Session, flask_app_with_containers: Flask @@ -115,9 +117,9 @@ class TestRagPipelineServiceGetPipeline: dataset = self._create_dataset(db_session_with_containers, tenant_id, created_by, pipeline_id=pipeline.id) db_session_with_containers.flush() - service = self._make_service(flask_app_with_containers) + service = self._make_service(flask_app_with_containers, db_session_with_containers) - result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) + result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) assert result.id == pipeline.id @@ -185,7 +187,9 @@ class TestUpdateCustomizedPipelineTemplate: icon_info=IconInfo(icon="📄"), ) with pytest.raises(ValueError, match="Customized pipeline template not found"): - RagPipelineService.update_customized_pipeline_template(str(uuid4()), info, account, tenant_id) + RagPipelineService.update_customized_pipeline_template( + str(uuid4()), info, account, tenant_id, session=db_session_with_containers + ) def test_update_template_raises_on_duplicate_name( self, db_session_with_containers: Session, flask_app_with_containers: Flask @@ -264,4 +268,6 @@ class TestDeleteCustomizedPipelineTemplate: tenant_id = str(uuid4()) with pytest.raises(ValueError, match="Customized pipeline template not found"): - RagPipelineService.delete_customized_pipeline_template(str(uuid4()), tenant_id) + RagPipelineService.delete_customized_pipeline_template( + str(uuid4()), tenant_id, session=db_session_with_containers + ) diff --git a/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py b/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py index 0f7c790ba14..1c366d3ee32 100644 --- a/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py +++ b/api/tests/test_containers_integration_tests/services/recommend_app/test_database_retrieval.py @@ -1,6 +1,6 @@ from __future__ import annotations -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 from flask import Flask @@ -82,8 +82,8 @@ class TestDatabaseRecommendAppRetrieval: "fetch_recommended_apps_from_db", return_value={"recommended_apps": [], "categories": []}, ) as mock_fetch: - result = DatabaseRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") - mock_fetch.assert_called_once_with("en-US") + result = DatabaseRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) + mock_fetch.assert_called_once() assert result == {"recommended_apps": [], "categories": []} def test_get_recommend_app_detail_delegates(self): @@ -92,8 +92,8 @@ class TestDatabaseRecommendAppRetrieval: "fetch_recommended_app_detail_from_db", return_value={"id": "app-1"}, ) as mock_fetch: - result = DatabaseRecommendAppRetrieval().get_recommend_app_detail("app-1") - mock_fetch.assert_called_once_with("app-1") + result = DatabaseRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) + mock_fetch.assert_called_once() assert result == {"id": "app-1"} @@ -112,7 +112,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id in app_ids @@ -135,7 +137,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) recommended_app = next(item for item in result["recommended_apps"] if item["app_id"] == created_app.id) assert recommended_app["categories"] == ["writing", "assistant"] @@ -160,7 +164,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) recommended_app = next(item for item in result["recommended_apps"] if item["app_id"] == created_app.id) assert "category" not in recommended_app @@ -177,7 +183,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("fr-FR") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "fr-FR", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id in app_ids @@ -190,7 +198,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id not in app_ids @@ -202,7 +212,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_recommended_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert app1.id not in app_ids @@ -235,7 +247,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db("en-US") + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db( + "en-US", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert learn_dify_app.id in app_ids @@ -261,7 +275,9 @@ class TestFetchRecommendedAppsFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db("fr-FR") + result = DatabaseRecommendAppRetrieval.fetch_learn_dify_apps_from_db( + "fr-FR", session=db_session_with_containers + ) app_ids = {r["app_id"] for r in result["recommended_apps"]} assert learn_dify_app.id in app_ids @@ -269,7 +285,9 @@ class TestFetchRecommendedAppsFromDb: class TestFetchRecommendedAppDetailFromDb: def test_returns_none_when_not_listed(self, flask_app_with_containers: Flask, db_session_with_containers: Session): - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(str(uuid4())) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + str(uuid4()), session=db_session_with_containers + ) assert result is None @@ -282,7 +300,9 @@ class TestFetchRecommendedAppDetailFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(app1.id) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + app1.id, session=db_session_with_containers + ) assert result is None @@ -298,7 +318,9 @@ class TestFetchRecommendedAppDetailFromDb: db_session_with_containers.expire_all() - result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db(app1.id) + result = DatabaseRecommendAppRetrieval.fetch_recommended_app_detail_from_db( + app1.id, session=db_session_with_containers + ) assert result is not None assert result["id"] == app1.id diff --git a/api/tests/test_containers_integration_tests/services/test_account_service.py b/api/tests/test_containers_integration_tests/services/test_account_service.py index 65a5b0a96bf..ac8ed39316b 100644 --- a/api/tests/test_containers_integration_tests/services/test_account_service.py +++ b/api/tests/test_containers_integration_tests/services/test_account_service.py @@ -1120,10 +1120,12 @@ class TestAccountService: mock_sync.return_value = True # Delete account - AccountService.delete_account(account) + AccountService.delete_account(account, session=db_session_with_containers) # Verify sync was called - mock_sync.assert_called_once_with(account_id=account.id, source="account_deleted") + mock_sync.assert_called_once_with( + account_id=account.id, source="account_deleted", session=db_session_with_containers + ) # Verify task was added to queue mock_delete_task.delay.assert_called_once_with(account.id) diff --git a/api/tests/test_containers_integration_tests/services/test_agent_service.py b/api/tests/test_containers_integration_tests/services/test_agent_service.py index 0ee0cb84e75..00b4a1563ff 100644 --- a/api/tests/test_containers_integration_tests/services/test_agent_service.py +++ b/api/tests/test_containers_integration_tests/services/test_agent_service.py @@ -132,7 +132,7 @@ class TestAgentService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Update the app model config to set agent_mode for agent-chat mode if app.mode == AppMode.AGENT_CHAT and app.app_model_config: @@ -295,7 +295,7 @@ class TestAgentService: agent_thoughts = self._create_test_agent_thoughts(db_session_with_containers, message) # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result structure assert result is not None @@ -355,7 +355,7 @@ class TestAgentService: # Execute the method under test with non-existent conversation with pytest.raises(ValueError, match="Conversation not found"): - AgentService.get_agent_logs(app, fake.uuid4(), fake.uuid4()) + AgentService.get_agent_logs(app, fake.uuid4(), fake.uuid4(), db_session_with_containers) def test_get_agent_logs_message_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -371,7 +371,7 @@ class TestAgentService: # Execute the method under test with non-existent message with pytest.raises(ValueError, match="Message not found"): - AgentService.get_agent_logs(app, conversation.id, fake.uuid4()) + AgentService.get_agent_logs(app, conversation.id, fake.uuid4(), db_session_with_containers) def test_get_agent_logs_with_end_user( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -452,7 +452,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -524,7 +524,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -569,7 +569,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -593,7 +593,7 @@ class TestAgentService: conversation, message = self._create_test_conversation_and_message(db_session_with_containers, app, account) # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -655,7 +655,7 @@ class TestAgentService: # Execute the method under test with pytest.raises(ValueError, match="App model config not found"): - AgentService.get_agent_logs(app, conversation.id, message.id) + AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) def test_get_agent_logs_agent_config_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -674,7 +674,7 @@ class TestAgentService: # Execute the method under test with pytest.raises(ValueError, match="Agent config not found"): - AgentService.get_agent_logs(app, conversation.id, message.id) + AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) def test_list_agent_providers_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -804,7 +804,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -899,7 +899,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -927,7 +927,7 @@ class TestAgentService: mock_external_service_dependencies["current_user"].timezone = "Asia/Shanghai" # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -968,7 +968,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result assert result is not None @@ -1009,7 +1009,7 @@ class TestAgentService: db_session_with_containers.commit() # Execute the method under test - result = AgentService.get_agent_logs(app, conversation.id, message.id) + result = AgentService.get_agent_logs(app, conversation.id, message.id, db_session_with_containers) # Verify the result - should handle malformed JSON gracefully assert result is not None diff --git a/api/tests/test_containers_integration_tests/services/test_annotation_service.py b/api/tests/test_containers_integration_tests/services/test_annotation_service.py index 94d72b19be8..2710df5e56c 100644 --- a/api/tests/test_containers_integration_tests/services/test_annotation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_annotation_service.py @@ -101,7 +101,7 @@ class TestAnnotationService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Setup current_user mock self._mock_current_user(mock_external_service_dependencies, account.id, tenant.id) @@ -207,7 +207,9 @@ class TestAnnotationService: } # Insert annotation directly - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -241,7 +243,9 @@ class TestAnnotationService: } with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) def test_insert_app_annotation_directly_app_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -263,7 +267,9 @@ class TestAnnotationService: # Try to insert annotation with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.insert_app_annotation_directly(annotation_args, non_existent_app_id) + AppAnnotationService.insert_app_annotation_directly( + annotation_args, non_existent_app_id, session=db_session_with_containers + ) def test_update_app_annotation_directly_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -279,7 +285,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(original_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + original_args, app.id, session=db_session_with_containers + ) # Update the annotation updated_args = { @@ -328,7 +336,9 @@ class TestAnnotationService: } # Insert annotation from message - annotation = AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, app.id) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -361,7 +371,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - initial_annotation = AppAnnotationService.up_insert_app_annotation_from_message(initial_args, app.id) + initial_annotation = AppAnnotationService.up_insert_app_annotation_from_message( + initial_args, app.id, session=db_session_with_containers + ) # Update the annotation updated_args = { @@ -369,7 +381,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - updated_annotation = AppAnnotationService.up_insert_app_annotation_from_message(updated_args, app.id) + updated_annotation = AppAnnotationService.up_insert_app_annotation_from_message( + updated_args, app.id, session=db_session_with_containers + ) # Verify annotation was updated correctly (same ID) assert updated_annotation.id == initial_annotation.id @@ -402,7 +416,9 @@ class TestAnnotationService: # Try to insert annotation with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, non_existent_app_id) + AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, non_existent_app_id, session=db_session_with_containers + ) def test_get_annotation_list_by_app_id_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -420,12 +436,18 @@ class TestAnnotationService: "question": f"Question {i}: {fake.sentence()}", "answer": f"Answer {i}: {fake.text(max_nb_chars=200)}", } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotations.append(annotation) # Get annotation list annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="" + app.id, + page=1, + limit=10, + keyword="", + session=db_session_with_containers, ) # Verify results @@ -452,18 +474,22 @@ class TestAnnotationService: "question": f"Question with {unique_keyword} keyword", "answer": f"Answer with {unique_keyword} keyword", } - AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id, session=db_session_with_containers) # Create another annotation without the keyword other_args = { "question": "Different question without special term", "answer": "Different answer without special content", } - AppAnnotationService.insert_app_annotation_directly(other_args, app.id) + AppAnnotationService.insert_app_annotation_directly(other_args, app.id, session=db_session_with_containers) # Search with keyword annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword=unique_keyword + app.id, + page=1, + limit=10, + keyword=unique_keyword, + session=db_session_with_containers, ) # Verify only matching annotations are returned @@ -490,30 +516,42 @@ class TestAnnotationService: "question": "Question with 50% discount", "answer": "Answer about 50% discount offer", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_percent, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_percent, app.id, session=db_session_with_containers + ) annotation_with_underscore = { "question": "Question with test_data", "answer": "Answer about test_data value", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_underscore, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_underscore, app.id, session=db_session_with_containers + ) annotation_with_backslash = { "question": "Question with path\\to\\file", "answer": "Answer about path\\to\\file location", } - AppAnnotationService.insert_app_annotation_directly(annotation_with_backslash, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_with_backslash, app.id, session=db_session_with_containers + ) # Create annotation that should NOT match (contains % but as part of different text) annotation_no_match = { "question": "Question with 100% different", "answer": "Answer about 100% different content", } - AppAnnotationService.insert_app_annotation_directly(annotation_no_match, app.id) + AppAnnotationService.insert_app_annotation_directly( + annotation_no_match, app.id, session=db_session_with_containers + ) # Test 1: Search with % character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="50%" + app.id, + page=1, + limit=10, + keyword="50%", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -521,7 +559,11 @@ class TestAnnotationService: # Test 2: Search with _ character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="test_data" + app.id, + page=1, + limit=10, + keyword="test_data", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -529,7 +571,11 @@ class TestAnnotationService: # Test 3: Search with \ character - should find exact match only annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="path\\to\\file" + app.id, + page=1, + limit=10, + keyword="path\\to\\file", + session=db_session_with_containers, ) assert total == 1 assert len(annotation_list) == 1 @@ -537,7 +583,11 @@ class TestAnnotationService: # Test 4: Search with % should NOT match 100% (verifies escaping works) annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, page=1, limit=10, keyword="50%" + app.id, + page=1, + limit=10, + keyword="50%", + session=db_session_with_containers, ) # Should only find the 50% annotation, not the 100% one assert total == 1 @@ -557,7 +607,9 @@ class TestAnnotationService: # Try to get annotation list with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.get_annotation_list_by_app_id(non_existent_app_id, page=1, limit=10, keyword="") + AppAnnotationService.get_annotation_list_by_app_id( + non_existent_app_id, page=1, limit=10, keyword="", session=db_session_with_containers + ) def test_delete_app_annotation_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -573,7 +625,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotation_id = annotation.id # Delete the annotation @@ -728,7 +782,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Add some hit histories for i in range(3): @@ -742,6 +798,7 @@ class TestAnnotationService: message_id=fake.uuid4(), from_source=ConversationFromSource.CONSOLE, score=0.8 + (i * 0.1), + session=db_session_with_containers, ) # Get hit histories @@ -749,6 +806,7 @@ class TestAnnotationService: self._annotation_ref(app, annotation.id), page=1, limit=10, + session=db_session_with_containers, ) # Verify results @@ -775,7 +833,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Get initial hit count initial_hit_count = annotation.hit_count @@ -795,6 +855,7 @@ class TestAnnotationService: message_id=message_id, from_source=ConversationFromSource.CONSOLE, score=score, + session=db_session_with_containers, ) # Verify hit count was incremented @@ -834,10 +895,14 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - created_annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + created_annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Get annotation by ID - retrieved_annotation = AppAnnotationService.get_annotation_by_id(created_annotation.id) + retrieved_annotation = AppAnnotationService.get_annotation_by_id( + created_annotation.id, session=db_session_with_containers + ) # Verify annotation was retrieved correctly assert retrieved_annotation is not None @@ -880,7 +945,9 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify result structure assert "job_id" in result @@ -920,7 +987,9 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify error result assert "error_msg" in result @@ -966,7 +1035,9 @@ class TestAnnotationService: ].get_features.return_value.annotation_quota_limit.size = 0 # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage) + result = AppAnnotationService.batch_import_app_annotations( + app.id, file_storage, session=db_session_with_containers + ) # Verify error result assert "error_msg" in result @@ -1008,7 +1079,7 @@ class TestAnnotationService: db_session_with_containers.commit() # Get annotation setting - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=db_session_with_containers) # Verify result structure assert result["enabled"] is True @@ -1027,7 +1098,7 @@ class TestAnnotationService: app, account = self._create_test_app_and_account(db_session_with_containers, mock_external_service_dependencies) # Get annotation setting (no setting exists) - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=db_session_with_containers) # Verify result structure assert result["enabled"] is False @@ -1072,7 +1143,9 @@ class TestAnnotationService: "score_threshold": 0.9, } - result = AppAnnotationService.update_app_annotation_setting(app.id, annotation_setting.id, update_args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, annotation_setting.id, update_args, session=db_session_with_containers + ) # Verify result structure assert result["enabled"] is True @@ -1101,11 +1174,15 @@ class TestAnnotationService: "question": f"Question {i}: {fake.sentence()}", "answer": f"Answer {i}: {fake.text(max_nb_chars=200)}", } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotations.append(annotation) # Export annotation list - exported_annotations = AppAnnotationService.export_annotation_list_by_app_id(app.id) + exported_annotations = AppAnnotationService.export_annotation_list_by_app_id( + app.id, session=db_session_with_containers + ) # Verify results assert len(exported_annotations) == 3 @@ -1132,7 +1209,9 @@ class TestAnnotationService: # Try to export annotation list with non-existent app with pytest.raises(NotFound, match="App not found"): - AppAnnotationService.export_annotation_list_by_app_id(non_existent_app_id) + AppAnnotationService.export_annotation_list_by_app_id( + non_existent_app_id, session=db_session_with_containers + ) def test_insert_app_annotation_directly_with_setting_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1176,7 +1255,9 @@ class TestAnnotationService: } # Insert annotation directly - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id @@ -1235,7 +1316,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(original_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + original_args, app.id, session=db_session_with_containers + ) # Reset mock to clear previous calls mock_external_service_dependencies["update_task"].delay.reset_mock() @@ -1312,7 +1395,9 @@ class TestAnnotationService: "question": fake.sentence(), "answer": fake.text(max_nb_chars=200), } - annotation = AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id) + annotation = AppAnnotationService.insert_app_annotation_directly( + annotation_args, app.id, session=db_session_with_containers + ) annotation_id = annotation.id # Reset mock to clear previous calls @@ -1382,7 +1467,9 @@ class TestAnnotationService: } # Insert annotation from message - annotation = AppAnnotationService.up_insert_app_annotation_from_message(annotation_args, app.id) + annotation = AppAnnotationService.up_insert_app_annotation_from_message( + annotation_args, app.id, session=db_session_with_containers + ) # Verify annotation was created correctly assert annotation.app_id == app.id diff --git a/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py b/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py index 1f88ce90621..de51f5077e6 100644 --- a/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py +++ b/api/tests/test_containers_integration_tests/services/test_api_based_extension_service.py @@ -82,7 +82,7 @@ class TestAPIBasedExtensionService: ) # Save extension - saved_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Verify extension was saved correctly assert saved_extension.id is not None @@ -120,21 +120,21 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test empty api_endpoint extension_data.name = fake.company() extension_data.api_endpoint = "" with pytest.raises(ValueError, match="api_endpoint must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test empty api_key extension_data.api_endpoint = f"https://{fake.domain_name()}/api" extension_data.api_key = "" with pytest.raises(ValueError, match="api_key must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_all_by_tenant_id_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -158,11 +158,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - saved_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) extensions.append(saved_extension) # Get all extensions for tenant - extension_list = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant.id) + extension_list = APIBasedExtensionService.get_all_by_tenant_id(tenant.id, session=db_session_with_containers) # Verify results assert len(extension_list) == 3 @@ -192,11 +192,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Get extension by ID retrieved_extension = APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, created_extension.id + tenant.id, created_extension.id, session=db_session_with_containers ) # Verify extension was retrieved correctly @@ -223,7 +223,7 @@ class TestAPIBasedExtensionService: # Try to get non-existent extension with pytest.raises(ValueError, match="API based extension is not found"): APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, non_existent_extension_id + tenant.id, non_existent_extension_id, session=db_session_with_containers ) def test_delete_extension_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -243,11 +243,11 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) extension_id = created_extension.id # Delete the extension - APIBasedExtensionService.delete(db_session_with_containers, created_extension) + APIBasedExtensionService.delete(created_extension, session=db_session_with_containers) # Verify extension was deleted @@ -275,7 +275,7 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - APIBasedExtensionService.save(db_session_with_containers, extension_data1) + APIBasedExtensionService.save(extension_data1, session=db_session_with_containers) # Try to create second extension with same name extension_data2 = APIBasedExtension( tenant_id=tenant.id, @@ -285,7 +285,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must be unique, it is already existed"): - APIBasedExtensionService.save(db_session_with_containers, extension_data2) + APIBasedExtensionService.save(extension_data2, session=db_session_with_containers) def test_save_extension_update_existing( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -306,7 +306,7 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Save original values for later comparison original_name = created_extension.name @@ -325,7 +325,7 @@ class TestAPIBasedExtensionService: created_extension.api_endpoint = new_endpoint created_extension.api_key = new_api_key - updated_extension = APIBasedExtensionService.save(db_session_with_containers, created_extension) + updated_extension = APIBasedExtensionService.save(created_extension, session=db_session_with_containers) # Verify extension was updated correctly assert updated_extension.id == created_extension.id @@ -342,7 +342,7 @@ class TestAPIBasedExtensionService: # Verify the update by retrieving the extension again retrieved_extension = APIBasedExtensionService.get_with_tenant_id( - db_session_with_containers, tenant.id, created_extension.id + tenant.id, created_extension.id, session=db_session_with_containers ) assert retrieved_extension.name == new_name assert retrieved_extension.api_endpoint == new_endpoint @@ -374,7 +374,7 @@ class TestAPIBasedExtensionService: # Try to save extension with connection error with pytest.raises(ValueError, match="connection error: request timeout"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_invalid_api_key_length( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -397,7 +397,7 @@ class TestAPIBasedExtensionService: # Try to save extension with short API key with pytest.raises(ValueError, match="api_key must be at least 5 characters"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_empty_fields(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -417,21 +417,21 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="name must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test with None api_endpoint extension_data.name = fake.company() extension_data.api_endpoint = None with pytest.raises(ValueError, match="api_endpoint must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Test with None api_key extension_data.api_endpoint = f"https://{fake.domain_name()}/api" extension_data.api_key = None with pytest.raises(ValueError, match="api_key must not be empty"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_all_by_tenant_id_empty_list( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -445,7 +445,7 @@ class TestAPIBasedExtensionService: ) # Get all extensions for tenant (none exist) - extension_list = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant.id) + extension_list = APIBasedExtensionService.get_all_by_tenant_id(tenant.id, session=db_session_with_containers) # Verify empty list is returned assert len(extension_list) == 0 @@ -475,7 +475,7 @@ class TestAPIBasedExtensionService: # Try to save extension with invalid ping response with pytest.raises(ValueError, match="{'result': 'invalid'}"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_missing_ping_result( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -501,7 +501,7 @@ class TestAPIBasedExtensionService: # Try to save extension with missing ping result with pytest.raises(ValueError, match="{'status': 'ok'}"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_get_with_tenant_id_wrong_tenant( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -527,11 +527,13 @@ class TestAPIBasedExtensionService: api_key=fake.password(length=20), ) - created_extension = APIBasedExtensionService.save(db_session_with_containers, extension_data) + created_extension = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) # Try to get extension with wrong tenant ID with pytest.raises(ValueError, match="API based extension is not found"): - APIBasedExtensionService.get_with_tenant_id(db_session_with_containers, tenant2.id, created_extension.id) + APIBasedExtensionService.get_with_tenant_id( + tenant2.id, created_extension.id, session=db_session_with_containers + ) def test_save_extension_api_key_exactly_four_chars_rejected( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -551,7 +553,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="api_key must be at least 5 characters"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_api_key_exactly_five_chars_accepted( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -570,7 +572,7 @@ class TestAPIBasedExtensionService: api_key="12345", ) - saved = APIBasedExtensionService.save(db_session_with_containers, extension_data) + saved = APIBasedExtensionService.save(extension_data, session=db_session_with_containers) assert saved.id is not None def test_save_extension_requestor_constructor_error( @@ -593,7 +595,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="connection error: bad config"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_network_exception( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -617,7 +619,7 @@ class TestAPIBasedExtensionService: ) with pytest.raises(ValueError, match="connection error: network failure"): - APIBasedExtensionService.save(db_session_with_containers, extension_data) + APIBasedExtensionService.save(extension_data, session=db_session_with_containers) def test_save_extension_update_duplicate_name_rejected( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -630,28 +632,28 @@ class TestAPIBasedExtensionService: assert tenant is not None ext1 = APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant.id, name="Extension Alpha", api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) ext2 = APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant.id, name="Extension Beta", api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) # Try to rename ext2 to ext1's name ext2.name = "Extension Alpha" with pytest.raises(ValueError, match="name must be unique, it is already existed"): - APIBasedExtensionService.save(db_session_with_containers, ext2) + APIBasedExtensionService.save(ext2, session=db_session_with_containers) def test_get_all_returns_empty_for_different_tenant( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -667,15 +669,15 @@ class TestAPIBasedExtensionService: assert tenant1 is not None APIBasedExtensionService.save( - db_session_with_containers, APIBasedExtension( tenant_id=tenant1.id, name=fake.company(), api_endpoint=f"https://{fake.domain_name()}/api", api_key=fake.password(length=20), ), + session=db_session_with_containers, ) assert tenant2 is not None - result = APIBasedExtensionService.get_all_by_tenant_id(db_session_with_containers, tenant2.id) + result = APIBasedExtensionService.get_all_by_tenant_id(tenant2.id, session=db_session_with_containers) assert result == [] diff --git a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py index cee08c4c33e..24c14637296 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py @@ -162,7 +162,7 @@ class TestAppDslService: api_rpm=10, ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account def _create_simple_yaml_content(self, app_name: str = "Test App", app_mode: str = "chat") -> str: @@ -841,7 +841,7 @@ class TestAppDslService: # ── Export ───────────────────────────────────────────────────────── - def test_export_dsl_delegates_by_mode(self, monkeypatch: pytest.MonkeyPatch): + def test_export_dsl_delegates_by_mode(self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session): workflow_calls: list[bool] = [] model_calls: list[bool] = [] monkeypatch.setattr( @@ -859,7 +859,7 @@ class TestAppDslService: mode=AppMode.WORKFLOW, icon_type="emoji", ) - AppDslService.export_dsl(workflow_app) + AppDslService.export_dsl(workflow_app, session=db_session_with_containers) assert workflow_calls == [True] chat_app = _app_stub( @@ -867,10 +867,12 @@ class TestAppDslService: icon_type="emoji", app_model_config=SimpleNamespace(to_dict=lambda: {"agent_mode": {"tools": []}}), ) - AppDslService.export_dsl(chat_app) + AppDslService.export_dsl(chat_app, session=db_session_with_containers) assert model_calls == [True] - def test_export_dsl_preserves_icon_and_icon_type(self, monkeypatch: pytest.MonkeyPatch): + def test_export_dsl_preserves_icon_and_icon_type( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): monkeypatch.setattr( AppDslService, "_append_workflow_export_data", @@ -886,7 +888,7 @@ class TestAppDslService: description="App with emoji icon", use_icon_as_answer_icon=True, ) - yaml_output = AppDslService.export_dsl(emoji_app) + yaml_output = AppDslService.export_dsl(emoji_app, session=db_session_with_containers) data = yaml.safe_load(yaml_output) assert data["app"]["icon"] == "🎨" assert data["app"]["icon_type"] == "emoji" @@ -901,7 +903,7 @@ class TestAppDslService: description="App with image icon", use_icon_as_answer_icon=False, ) - yaml_output = AppDslService.export_dsl(image_app) + yaml_output = AppDslService.export_dsl(image_app, session=db_session_with_containers) data = yaml.safe_load(yaml_output) assert data["app"]["icon"] == "https://example.com/icon.png" assert data["app"]["icon_type"] == "image" @@ -936,7 +938,7 @@ class TestAppDslService: db_session_with_containers.add(model_config) db_session_with_containers.commit() - exported_dsl = AppDslService.export_dsl(app, include_secret=False) + exported_dsl = AppDslService.export_dsl(app, include_secret=False, session=db_session_with_containers) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -972,7 +974,7 @@ class TestAppDslService: "workflow_service" ].return_value.get_draft_workflow.return_value = mock_workflow - exported_dsl = AppDslService.export_dsl(app, include_secret=False) + exported_dsl = AppDslService.export_dsl(app, include_secret=False, session=db_session_with_containers) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -1006,7 +1008,7 @@ class TestAppDslService: workflow_id = str(uuid4()) - def mock_get_draft_workflow(app_model, wf_id=None): + def mock_get_draft_workflow(app_model, wf_id=None, **_kwargs): if wf_id == workflow_id: return mock_workflow return None @@ -1015,7 +1017,9 @@ class TestAppDslService: "workflow_service" ].return_value.get_draft_workflow.side_effect = mock_get_draft_workflow - exported_dsl = AppDslService.export_dsl(app, include_secret=False, workflow_id=workflow_id) + exported_dsl = AppDslService.export_dsl( + app, include_secret=False, workflow_id=workflow_id, session=db_session_with_containers + ) exported_data = yaml.safe_load(exported_dsl) assert exported_data["kind"] == "app" @@ -1034,11 +1038,15 @@ class TestAppDslService: WorkflowNotFoundError, match="Missing draft workflow configuration, please check.", ): - AppDslService.export_dsl(app, include_secret=False, workflow_id=str(uuid4())) + AppDslService.export_dsl( + app, include_secret=False, workflow_id=str(uuid4()), session=db_session_with_containers + ) # ── Workflow Export Data ─────────────────────────────────────────── - def test_append_workflow_export_data_filters_and_overrides(self, monkeypatch: pytest.MonkeyPatch): + def test_append_workflow_export_data_filters_and_overrides( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): workflow_dict = { "graph": { "nodes": [ @@ -1123,6 +1131,7 @@ class TestAppDslService: app_model=_app_stub(), include_secret=False, workflow_id=None, + session=db_session_with_containers, ) nodes = export_data["workflow"]["graph"]["nodes"] @@ -1138,7 +1147,9 @@ class TestAppDslService: assert nodes[5]["data"]["subscription_id"] == "" assert export_data["dependencies"] == [{"tenant": _DEFAULT_TENANT_ID, "dep": "dep-1"}] - def test_append_workflow_export_data_missing_workflow_raises(self, monkeypatch: pytest.MonkeyPatch): + def test_append_workflow_export_data_missing_workflow_raises( + self, monkeypatch: pytest.MonkeyPatch, db_session_with_containers: Session + ): workflow_service = MagicMock() workflow_service.get_draft_workflow.return_value = None monkeypatch.setattr(app_dsl_service, "WorkflowService", lambda: workflow_service) @@ -1149,6 +1160,7 @@ class TestAppDslService: app_model=_app_stub(), include_secret=False, workflow_id=None, + session=db_session_with_containers, ) # ── Model Config Export Data ────────────────────────────────────── diff --git a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py index 473111f364a..89cc7715d1c 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py @@ -187,7 +187,7 @@ class TestAppGenerateService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -234,12 +234,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -267,12 +267,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -298,12 +298,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -329,12 +329,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -362,12 +362,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -399,12 +399,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -431,12 +431,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -461,12 +461,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=db_session_with_containers, ) # Verify the result @@ -503,12 +503,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -535,12 +535,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -574,12 +574,12 @@ class TestAppGenerateService: # StatementError (from EnumText validation during autoflush) with pytest.raises((ValueError, sa.exc.StatementError)): AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) def test_generate_with_workflow_id_format_error( @@ -603,12 +603,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowIdFormatError with pytest.raises(WorkflowIdFormatError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -642,12 +642,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowNotFoundError with pytest.raises(WorkflowNotFoundError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -673,12 +673,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -704,12 +704,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -731,7 +731,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -758,7 +763,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -786,7 +796,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate_single_iteration( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -808,7 +823,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -835,7 +855,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -861,7 +886,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate_single_loop( - app_model=app, user=account, node_id=node_id, args=args, streaming=True + app_model=app, + user=account, + node_id=node_id, + args=args, + streaming=True, + session=db_session_with_containers, ) # Verify error message @@ -1021,12 +1051,12 @@ class TestAppGenerateService: # Execute the method under test and expect exception with pytest.raises(Exception) as exc_info: AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify exception message @@ -1054,12 +1084,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -1094,12 +1124,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=invoke_from, streaming=True, + session=db_session_with_containers, ) # Verify the result @@ -1137,12 +1167,12 @@ class TestAppGenerateService: mock_exec_params.new.return_value = mock_payload result = AppGenerateService.generate( - session=db_session_with_containers, app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=db_session_with_containers, ) # Verify the result diff --git a/api/tests/test_containers_integration_tests/services/test_app_service.py b/api/tests/test_containers_integration_tests/services/test_app_service.py index f9df99c5594..8deaf6d462d 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_service.py @@ -84,7 +84,7 @@ class TestAppService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Verify app was created correctly assert app.name == app_params.name @@ -144,7 +144,7 @@ class TestAppService: icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Verify app mode was set correctly assert app.mode == mode @@ -183,7 +183,7 @@ class TestAppService: ) app_service = AppService() - created_app = app_service.create_app(tenant.id, app_params, account) + created_app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Get app using the service - needs current_user mock mock_current_user = create_autospec(Account, instance=True) @@ -234,7 +234,7 @@ class TestAppService: icon="📱", icon_background="#96CEB4", ) - app_service.create_app(tenant.id, app_params, account) + app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Get paginated apps params = AppListParams(page=1, limit=10, mode="chat") @@ -277,16 +277,19 @@ class TestAppService: tenant.id, CreateAppParams(name="Oldest Created", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) newest_modified = app_service.create_app( tenant.id, CreateAppParams(name="Newest Modified", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) newest_created = app_service.create_app( tenant.id, CreateAppParams(name="Newest Created", mode="chat", icon_type="emoji", icon="3"), account, + session=db_session_with_containers, ) timestamp_by_app_id = { @@ -362,15 +365,17 @@ class TestAppService: tenant.id, CreateAppParams(name="Starred App", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) unstarred_app = app_service.create_app( tenant.id, CreateAppParams(name="Unstarred App", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) - app_service.star_app(db_session_with_containers, app=starred_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=starred_app, account_id=account.id) + app_service.star_app(app=starred_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=starred_app, account_id=account.id, session=db_session_with_containers) db_session_with_containers.commit() star_count = db_session_with_containers.scalar( @@ -386,7 +391,7 @@ class TestAppService: assert starred_by_app_id[starred_app.id] is True assert starred_by_app_id[unstarred_app.id] is False - app_service.unstar_app(db_session_with_containers, app=starred_app, account_id=account.id) + app_service.unstar_app(app=starred_app, account_id=account.id, session=db_session_with_containers) db_session_with_containers.commit() paginated_apps = app_service.get_paginate_apps( @@ -422,26 +427,30 @@ class TestAppService: tenant.id, CreateAppParams(name="Oldest Created Starred App", mode="chat", icon_type="emoji", icon="1"), account, + session=db_session_with_containers, ) newest_modified_app = app_service.create_app( tenant.id, CreateAppParams(name="Newest Modified Starred App", mode="chat", icon_type="emoji", icon="2"), account, + session=db_session_with_containers, ) newest_created_app = app_service.create_app( tenant.id, CreateAppParams(name="Newest Created Starred App", mode="chat", icon_type="emoji", icon="3"), account, + session=db_session_with_containers, ) unstarred_app = app_service.create_app( tenant.id, CreateAppParams(name="Unstarred App", mode="chat", icon_type="emoji", icon="4"), account, + session=db_session_with_containers, ) - app_service.star_app(db_session_with_containers, app=oldest_created_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=newest_modified_app, account_id=account.id) - app_service.star_app(db_session_with_containers, app=newest_created_app, account_id=account.id) + app_service.star_app(app=oldest_created_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=newest_modified_app, account_id=account.id, session=db_session_with_containers) + app_service.star_app(app=newest_created_app, account_id=account.id, session=db_session_with_containers) timestamp_by_app_id = { oldest_created_app.id: (datetime(2026, 1, 1, 10, 0, 0), datetime(2026, 1, 1, 10, 0, 0)), @@ -535,8 +544,10 @@ class TestAppService: icon_background="#4ECDC4", ) - chat_app = app_service.create_app(tenant.id, chat_app_params, account) - completion_app = app_service.create_app(tenant.id, completion_app_params, account) + chat_app = app_service.create_app(tenant.id, chat_app_params, account, session=db_session_with_containers) + completion_app = app_service.create_app( + tenant.id, completion_app_params, account, session=db_session_with_containers + ) # Test filter by mode chat_apps = app_service.get_paginate_apps( @@ -599,7 +610,7 @@ class TestAppService: icon="💬", icon_background="#FF6B6B", ) - app_service.create_app(tenant.id, app_params, first_account) + app_service.create_app(tenant.id, app_params, first_account, session=db_session_with_containers) other_app_params = CreateAppParams( name="Second Creator App", description="Created by the second account", @@ -608,7 +619,7 @@ class TestAppService: icon="✍️", icon_background="#4ECDC4", ) - app_service.create_app(tenant.id, other_app_params, second_account) + app_service.create_app(tenant.id, other_app_params, second_account, session=db_session_with_containers) filtered_apps = app_service.get_paginate_apps( first_account.id, @@ -654,7 +665,7 @@ class TestAppService: icon="🏷️", icon_background="#FFEAA7", ) - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Mock TagService to return the app ID for tag filtering with patch("services.app_service.TagService.get_target_ids_by_tag_ids") as mock_tag_service: @@ -717,7 +728,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original values original_name = app.name @@ -741,7 +752,7 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app(app, update_args) + updated_app = app_service.update_app(app, update_args, session=db_session_with_containers) # Verify updated fields assert updated_app.name == update_args["name"] @@ -788,6 +799,7 @@ class TestAppService: icon_background="#45B7D1", ), account, + session=db_session_with_containers, ) mock_current_user = create_autospec(Account, instance=True) @@ -805,6 +817,7 @@ class TestAppService: "icon_background": "#FF8C42", "use_icon_as_answer_icon": True, }, + session=db_session_with_containers, ) assert updated_app.icon_type == IconType.EMOJI @@ -841,6 +854,7 @@ class TestAppService: icon_background="#45B7D1", ), account, + session=db_session_with_containers, ) mock_current_user = create_autospec(Account, instance=True) @@ -859,6 +873,7 @@ class TestAppService: "icon_background": "#FF8C42", "use_icon_as_answer_icon": True, }, + session=db_session_with_containers, ) def test_update_app_name_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -892,7 +907,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original name original_name = app.name @@ -904,7 +919,7 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_name(app, new_name) + updated_app = app_service.update_app_name(app, new_name, session=db_session_with_containers) assert updated_app.name == new_name assert updated_app.updated_by == account.id @@ -946,7 +961,7 @@ class TestAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_params, account) + app = app_service.create_app(tenant.id, app_params, account, session=db_session_with_containers) # Store original values original_icon = app.icon @@ -961,7 +976,9 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_icon(app, new_icon, new_icon_background, new_icon_type) + updated_app = app_service.update_app_icon( + app, new_icon, new_icon_background, new_icon_type, session=db_session_with_containers + ) assert updated_app.icon == new_icon assert updated_app.icon_background == new_icon_background @@ -1007,7 +1024,7 @@ class TestAppService: icon_background="#74B9FF", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original site status original_site_status = app.enable_site @@ -1018,13 +1035,13 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_site_status(app, False) + updated_app = app_service.update_app_site_status(app, False, session=db_session_with_containers) assert updated_app.enable_site is False assert updated_app.updated_by == account.id # Update site status back to enabled with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_site_status(updated_app, True) + updated_app = app_service.update_app_site_status(updated_app, True, session=db_session_with_containers) assert updated_app.enable_site is True assert updated_app.updated_by == account.id @@ -1067,7 +1084,7 @@ class TestAppService: icon_background="#A29BFE", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original API status original_api_status = app.enable_api @@ -1078,13 +1095,13 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_api_status(app, False) + updated_app = app_service.update_app_api_status(app, False, session=db_session_with_containers) assert updated_app.enable_api is False assert updated_app.updated_by == account.id # Update API status back to enabled with patch("services.app_service.current_user", mock_current_user): - updated_app = app_service.update_app_api_status(updated_app, True) + updated_app = app_service.update_app_api_status(updated_app, True, session=db_session_with_containers) assert updated_app.enable_api is True assert updated_app.updated_by == account.id @@ -1127,14 +1144,14 @@ class TestAppService: icon_background="#FD79A8", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store original values original_site_status = app.enable_site original_updated_at = app.updated_at # Update site status to the same value (no change) - updated_app = app_service.update_app_site_status(app, original_site_status) + updated_app = app_service.update_app_site_status(app, original_site_status, session=db_session_with_containers) # Verify app is returned unchanged assert updated_app.id == app.id @@ -1178,7 +1195,7 @@ class TestAppService: icon_background="#E17055", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store app ID for verification app_id = app.id @@ -1188,7 +1205,7 @@ class TestAppService: mock_delete_task.delay.return_value = None # Delete app - app_service.delete_app(app) + app_service.delete_app(app, session=db_session_with_containers) # Verify async deletion task was called mock_delete_task.delay.assert_called_once_with(tenant_id=tenant.id, app_id=app_id) @@ -1230,7 +1247,7 @@ class TestAppService: icon_background="#00B894", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Store app ID for verification app_id = app.id @@ -1245,7 +1262,7 @@ class TestAppService: mock_delete_task.delay.return_value = None # Delete app - app_service.delete_app(app) + app_service.delete_app(app, session=db_session_with_containers) # Verify webapp auth cleanup was called mock_external_service_dependencies["enterprise_service"].WebAppAuth.cleanup_webapp.assert_called_once_with( @@ -1290,10 +1307,10 @@ class TestAppService: icon_background="#6C5CE7", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Get app metadata - app_meta = app_service.get_app_meta(app) + app_meta = app_service.get_app_meta(app, session=db_session_with_containers) # Verify metadata contains expected fields assert "tool_icons" in app_meta @@ -1329,10 +1346,10 @@ class TestAppService: icon_background="#FDCB6E", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Get app code by ID - app_code = AppService.get_app_code_by_id(app.id) + app_code = AppService.get_app_code_by_id(app.id, session=db_session_with_containers) # Verify app code was retrieved correctly # Note: Site would be created when App is created, site.code is auto-generated @@ -1369,7 +1386,7 @@ class TestAppService: icon_background="#E84393", ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create a site for the app site = Site() @@ -1384,7 +1401,7 @@ class TestAppService: db_session_with_containers.commit() # Get app ID by code - app_id = AppService.get_app_id_by_code(site.code) + app_id = AppService.get_app_id_by_code(site.code, session=db_session_with_containers) # Verify app ID was retrieved correctly assert app_id == app.id @@ -1462,6 +1479,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) app_with_underscore = app_service.create_app( @@ -1477,6 +1495,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) app_with_backslash = app_service.create_app( @@ -1492,6 +1511,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) # Create app that should NOT match @@ -1508,6 +1528,7 @@ class TestAppService: api_rpm=10, ), account, + session=db_session_with_containers, ) # Test 1: Search with % character @@ -1560,7 +1581,7 @@ class TestAppService: from services.app_service import AppService with pytest.raises(ValueError, match="not found"): - AppService.get_app_code_by_id(str(uuid4())) + AppService.get_app_code_by_id(str(uuid4()), session=db_session_with_containers) def test_get_app_id_by_code_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1569,7 +1590,7 @@ class TestAppService: from services.app_service import AppService with pytest.raises(ValueError, match="not found"): - AppService.get_app_id_by_code("nonexistent-code") + AppService.get_app_id_by_code("nonexistent-code", session=db_session_with_containers) def test_get_app_meta_returns_empty_when_workflow_missing( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1582,7 +1603,7 @@ class TestAppService: app_service = AppService() workflow_app = SimpleNamespace(mode="workflow", workflow=None) - meta = app_service.get_app_meta(workflow_app) + meta = app_service.get_app_meta(workflow_app, session=db_session_with_containers) assert meta == {"tool_icons": {}} def test_get_app_meta_returns_empty_when_model_config_missing( @@ -1596,5 +1617,5 @@ class TestAppService: app_service = AppService() chat_app = SimpleNamespace(mode="chat", app_model_config=None) - meta = app_service.get_app_meta(chat_app) + meta = app_service.get_app_meta(chat_app, session=db_session_with_containers) assert meta == {"tool_icons": {}} diff --git a/api/tests/test_containers_integration_tests/services/test_billing_service.py b/api/tests/test_containers_integration_tests/services/test_billing_service.py index a3a4a0e6edd..777fb7721b4 100644 --- a/api/tests/test_containers_integration_tests/services/test_billing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_billing_service.py @@ -417,7 +417,7 @@ class TestBillingServiceIsTenantOwnerOrAdmin: account, _ = self._create_account_with_tenant_role(db_session_with_containers, TenantAccountRole.EDITOR) with pytest.raises(ValueError, match="Only team owner or team admin can perform this action"): - BillingService.is_tenant_owner_or_admin(db_session_with_containers, account) + BillingService.is_tenant_owner_or_admin(account, session=db_session_with_containers) def test_is_tenant_owner_or_admin_dataset_operator_raises_error(self, db_session_with_containers: Session) -> None: """is_tenant_owner_or_admin raises ValueError for DATASET_OPERATOR role.""" @@ -426,4 +426,4 @@ class TestBillingServiceIsTenantOwnerOrAdmin: ) with pytest.raises(ValueError, match="Only team owner or team admin can perform this action"): - BillingService.is_tenant_owner_or_admin(db_session_with_containers, account) + BillingService.is_tenant_owner_or_admin(account, session=db_session_with_containers) diff --git a/api/tests/test_containers_integration_tests/services/test_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_conversation_service.py index b19b6b9c984..19dd4d6cf70 100644 --- a/api/tests/test_containers_integration_tests/services/test_conversation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_conversation_service.py @@ -350,6 +350,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, # No starting point specified limit=10, + session=db_session_with_containers, ) # Assert - Verify the results @@ -395,6 +396,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=first_message.id, limit=10, + session=db_session_with_containers, ) # Assert - Verify the results @@ -426,6 +428,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=str(uuid4()), limit=10, + session=db_session_with_containers, ) def test_pagination_with_has_more_flag(self, db_session_with_containers: Session): @@ -461,6 +464,7 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, limit=limit, + session=db_session_with_containers, ) # Assert @@ -498,7 +502,8 @@ class TestConversationServiceMessageCreation: conversation_id=conversation.id, first_id=None, limit=10, - order="asc", # Ascending order + order="asc", # Ascending order, + session=db_session_with_containers, ) # Assert @@ -547,7 +552,7 @@ class TestConversationServiceSummarization: mock_llm_generator.return_value = generated_name # Act - result = ConversationService.auto_generate_name(app_model, conversation) + result = ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) # Assert assert conversation.name == generated_name # Name updated on conversation object @@ -572,7 +577,7 @@ class TestConversationServiceSummarization: # Act & Assert with pytest.raises(MessageNotExistsError): - ConversationService.auto_generate_name(app_model, conversation) + ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) @patch("services.conversation_service.LLMGenerator.generate_conversation_name") def test_auto_generate_name_handles_llm_failure_gracefully( @@ -604,7 +609,7 @@ class TestConversationServiceSummarization: mock_llm_generator.side_effect = Exception("LLM service unavailable") # Act - result = ConversationService.auto_generate_name(app_model, conversation) + result = ConversationService.auto_generate_name(app_model, conversation, session=db_session_with_containers) # Assert assert conversation.name == original_name # Name remains unchanged @@ -637,6 +642,7 @@ class TestConversationServiceSummarization: user=user, name=new_name, auto_generate=False, + session=db_session_with_containers, ) # Assert @@ -671,6 +677,7 @@ class TestConversationServiceSummarization: user=user, name=None, auto_generate=True, + session=db_session_with_containers, ) # Assert @@ -719,7 +726,9 @@ class TestConversationServiceMessageAnnotation: args = {"message_id": message.id, "answer": "AI is artificial intelligence"} # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.message_id == message.id @@ -753,7 +762,9 @@ class TestConversationServiceMessageAnnotation: } # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.message_id is None @@ -802,7 +813,9 @@ class TestConversationServiceMessageAnnotation: args = {"message_id": message.id, "answer": "Updated annotation content"} # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app_model.id) + result = AppAnnotationService.up_insert_app_annotation_from_message( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.id == existing_annotation.id @@ -838,7 +851,11 @@ class TestConversationServiceMessageAnnotation: # Act result_items, result_total = AppAnnotationService.get_annotation_list_by_app_id( - app_id=app_model.id, page=1, limit=10, keyword="" + app_id=app_model.id, + page=1, + limit=10, + keyword="", + session=db_session_with_containers, ) # Assert @@ -886,7 +903,8 @@ class TestConversationServiceMessageAnnotation: app_id=app_model.id, page=1, limit=10, - keyword="machine", # Search keyword + keyword="machine", # Search keyword, + session=db_session_with_containers, ) # Assert @@ -914,7 +932,9 @@ class TestConversationServiceMessageAnnotation: } # Act - result = AppAnnotationService.insert_app_annotation_directly(args, app_model.id) + result = AppAnnotationService.insert_app_annotation_directly( + args, app_model.id, session=db_session_with_containers + ) # Assert assert result.question == args["question"] @@ -942,7 +962,9 @@ class TestConversationServiceExport: ) # Act - result = ConversationService.get_conversation(app_model=app_model, conversation_id=conversation.id, user=user) + result = ConversationService.get_conversation( + app_model=app_model, conversation_id=conversation.id, user=user, session=db_session_with_containers + ) # Assert assert result == conversation @@ -956,7 +978,12 @@ class TestConversationServiceExport: # Act & Assert with pytest.raises(ConversationNotExistsError): - ConversationService.get_conversation(app_model=app_model, conversation_id=str(uuid4()), user=user) + ConversationService.get_conversation( + app_model=app_model, + conversation_id=str(uuid4()), + user=user, + session=db_session_with_containers, + ) @patch("services.annotation_service.current_account_with_tenant") def test_export_annotation_list(self, mock_current_account, db_session_with_containers: Session): @@ -982,7 +1009,7 @@ class TestConversationServiceExport: mock_current_account.return_value = (account, app_model.tenant_id) # Act - result = AppAnnotationService.export_annotation_list_by_app_id(app_model.id) + result = AppAnnotationService.export_annotation_list_by_app_id(app_model.id, session=db_session_with_containers) # Assert assert len(result) == 10 @@ -1006,7 +1033,9 @@ class TestConversationServiceExport: ) # Act - result = MessageService.get_message(app_model=app_model, user=user, message_id=message.id) + result = MessageService.get_message( + app_model=app_model, user=user, message_id=message.id, session=db_session_with_containers + ) # Assert assert result == message @@ -1020,7 +1049,9 @@ class TestConversationServiceExport: # Act & Assert with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app_model, user=user, message_id=str(uuid4())) + MessageService.get_message( + app_model=app_model, user=user, message_id=str(uuid4()), session=db_session_with_containers + ) def test_get_conversation_for_end_user(self, db_session_with_containers: Session): """ @@ -1041,7 +1072,10 @@ class TestConversationServiceExport: # Act result = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation.id, user=end_user + app_model=app_model, + conversation_id=conversation.id, + user=end_user, + session=db_session_with_containers, ) # Assert @@ -1069,7 +1103,9 @@ class TestConversationServiceExport: conversation_id = conversation.id # Act - Delete the conversation - ConversationService.delete(app_model=app_model, conversation_id=conversation_id, user=user) + ConversationService.delete( + app_model=app_model, conversation_id=conversation_id, user=user, session=db_session_with_containers + ) # Assert - Verify two-step deletion process # Step 1: Immediate database deletion @@ -1104,6 +1140,7 @@ class TestConversationServiceExport: app_model=app_model, conversation_id=conversation.id, user=other_account, + session=db_session_with_containers, ) # Verify no deletion and no async cleanup trigger @@ -1129,9 +1166,14 @@ class TestConversationServiceExport: conversation_id = conversation.id # Act — force an error during the delete to exercise the rollback path - with patch("services.conversation_service.db.session.delete", side_effect=Exception("DB error")): + with patch.object(db_session_with_containers, "delete", side_effect=Exception("DB error")): with pytest.raises(Exception, match="DB error"): - ConversationService.delete(app_model=app_model, conversation_id=conversation_id, user=user) + ConversationService.delete( + app_model=app_model, + conversation_id=conversation_id, + user=user, + session=db_session_with_containers, + ) # Assert — async cleanup must NOT have been scheduled mock_delete_task.delay.assert_not_called() diff --git a/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py b/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py index 33d4563904e..9a725b06b64 100644 --- a/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py +++ b/api/tests/test_containers_integration_tests/services/test_conversation_service_variables.py @@ -6,10 +6,9 @@ from uuid import uuid4 import pytest from flask import Flask -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from graphon.variables import FloatVariable, IntegerVariable, StringVariable from models.account import Account, Tenant, TenantAccountJoin from models.enums import ConversationFromSource, EndUserType @@ -152,13 +151,6 @@ class ConversationServiceVariableIntegrationFactory: @pytest.fixture def real_conversation_service_session_factory(flask_app_with_containers: Flask): del flask_app_with_containers - real_session_maker = sessionmaker(bind=db.engine, expire_on_commit=False) - - with ( - patch("services.conversation_service.session_factory.create_session", side_effect=lambda: real_session_maker()), - patch("services.conversation_service.session_factory.get_session_maker", return_value=real_session_maker), - ): - yield class TestConversationServiceVariables: @@ -193,6 +185,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=None, + session=db_session_with_containers, ) assert [item["id"] for item in result.data] == [first_variable.id, second_variable.id] @@ -237,6 +230,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=first_variable.id, + session=db_session_with_containers, ) assert [item["id"] for item in result.data] == [second_variable.id, third_variable.id] @@ -257,6 +251,7 @@ class TestConversationServiceVariables: user=account, limit=10, last_id=str(uuid4()), + session=db_session_with_containers, ) def test_get_conversational_variable_sets_has_more( @@ -282,6 +277,7 @@ class TestConversationServiceVariables: user=account, limit=2, last_id=None, + session=db_session_with_containers, ) assert len(result.data) == 2 @@ -309,6 +305,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value="support", + session=db_session_with_containers, ) db_session_with_containers.expire_all() @@ -335,6 +332,7 @@ class TestConversationServiceVariables: variable_id=str(uuid4()), user=account, new_value="support", + session=db_session_with_containers, ) def test_update_conversation_variable_type_mismatch_raises_error( @@ -358,6 +356,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value="wrong-type", + session=db_session_with_containers, ) def test_update_conversation_variable_integer_number_compatibility( @@ -380,6 +379,7 @@ class TestConversationServiceVariables: variable_id=existing.id, user=account, new_value=42, + session=db_session_with_containers, ) db_session_with_containers.expire_all() diff --git a/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py b/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py index de8e6ba612c..9cbe5252bbf 100644 --- a/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py +++ b/api/tests/test_containers_integration_tests/services/test_credit_pool_service.py @@ -4,10 +4,8 @@ from unittest.mock import patch from uuid import uuid4 import pytest -from flask import has_app_context from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from core.errors.error import QuotaExceededError from models import TenantCreditPool from models.enums import ProviderQuotaType @@ -35,11 +33,10 @@ class TestCreditPoolService: db_session.add(pool) db_session.commit() - @pytest.mark.usefixtures("db_session_with_containers") - def test_create_default_pool(self) -> None: + def test_create_default_pool(self, db_session_with_containers: Session) -> None: tenant_id = self._create_tenant_id() - pool = CreditPoolService.create_default_pool(tenant_id) + pool = CreditPoolService.create_default_pool(tenant_id, session=db_session_with_containers) assert isinstance(pool, TenantCreditPool) assert pool.tenant_id == tenant_id @@ -51,43 +48,46 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=0) - result = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + result = CreditPoolService.get_pool( + tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is not None assert result.tenant_id == tenant_id assert result.pool_type == ProviderQuotaType.TRIAL - @pytest.mark.usefixtures("flask_app_with_containers") - def test_get_pool_uses_configured_session_factory_without_flask_app_context(self) -> None: + def test_get_pool_uses_provided_session(self, db_session_with_containers: Session) -> None: tenant_id = self._create_tenant_id() - session_maker = session_factory.get_session_maker() - with session_maker.begin() as session: - session.add( - TenantCreditPool( - tenant_id=tenant_id, - pool_type=ProviderQuotaType.TRIAL, - quota_limit=10, - quota_used=2, - ) + db_session_with_containers.add( + TenantCreditPool( + tenant_id=tenant_id, + pool_type=ProviderQuotaType.TRIAL, + quota_limit=10, + quota_used=2, ) + ) + db_session_with_containers.commit() - assert not has_app_context() - result = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + result = CreditPoolService.get_pool( + tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is not None assert result.tenant_id == tenant_id assert result.pool_type == ProviderQuotaType.TRIAL assert result.quota_used == 2 - @pytest.mark.usefixtures("flask_app_with_containers") - def test_get_pool_returns_none_when_not_exists(self) -> None: - result = CreditPoolService.get_pool(tenant_id=self._create_tenant_id(), pool_type=ProviderQuotaType.TRIAL) + def test_get_pool_returns_none_when_not_exists(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.get_pool( + tenant_id=self._create_tenant_id(), pool_type=ProviderQuotaType.TRIAL, session=db_session_with_containers + ) assert result is None - @pytest.mark.usefixtures("flask_app_with_containers") - def test_check_credits_available_returns_false_when_no_pool(self) -> None: - result = CreditPoolService.check_credits_available(tenant_id=self._create_tenant_id(), credits_required=10) + def test_check_credits_available_returns_false_when_no_pool(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.check_credits_available( + tenant_id=self._create_tenant_id(), credits_required=10, session=db_session_with_containers + ) assert result is False @@ -95,7 +95,9 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=0) - result = CreditPoolService.check_credits_available(tenant_id=tenant_id, credits_required=10) + result = CreditPoolService.check_credits_available( + tenant_id=tenant_id, credits_required=10, session=db_session_with_containers + ) assert result is True @@ -103,14 +105,17 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) - result = CreditPoolService.check_credits_available(tenant_id=tenant_id, credits_required=1) + result = CreditPoolService.check_credits_available( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) assert result is False - @pytest.mark.usefixtures("flask_app_with_containers") - def test_check_and_deduct_credits_raises_when_no_pool(self) -> None: + def test_check_and_deduct_credits_raises_when_no_pool(self, db_session_with_containers: Session) -> None: with pytest.raises(QuotaExceededError, match="Credit pool not found"): - CreditPoolService.check_and_deduct_credits(tenant_id=self._create_tenant_id(), credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=self._create_tenant_id(), credits_required=1, session=db_session_with_containers + ) def test_check_and_deduct_credits_returns_zero_for_non_positive_request( self, db_session_with_containers: Session @@ -118,10 +123,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) - result = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=0) + result = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=0, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -130,9 +137,11 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) with pytest.raises(QuotaExceededError, match="No credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -141,10 +150,12 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) credits_required = 3 - result = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=credits_required) + result = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=credits_required, session=db_session_with_containers + ) assert result == credits_required - pool = CreditPoolService.get_pool(tenant_id=tenant_id) + pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert pool is not None assert pool.quota_used == 5 @@ -155,9 +166,11 @@ class TestCreditPoolService: self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=9) with pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=3, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 9 @@ -171,9 +184,11 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -181,10 +196,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=9) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=3) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=3, session=db_session_with_containers + ) assert result == 1 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -194,16 +211,19 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=2) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=0) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=0, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 - @pytest.mark.usefixtures("flask_app_with_containers") - def test_deduct_credits_capped_returns_zero_when_no_pool(self) -> None: - result = CreditPoolService.deduct_credits_capped(tenant_id=self._create_tenant_id(), credits_required=1) + def test_deduct_credits_capped_returns_zero_when_no_pool(self, db_session_with_containers: Session) -> None: + result = CreditPoolService.deduct_credits_capped( + tenant_id=self._create_tenant_id(), credits_required=1, session=db_session_with_containers + ) assert result == 0 @@ -211,10 +231,12 @@ class TestCreditPoolService: tenant_id = self._create_tenant_id() self._create_pool(db_session_with_containers, tenant_id=tenant_id, quota_limit=10, quota_used=10) - result = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + result = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) assert result == 0 - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 10 @@ -226,9 +248,11 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 @@ -240,8 +264,10 @@ class TestCreditPoolService: patch.object(CreditPoolService, "_get_locked_pool", side_effect=QuotaExceededError("quota unavailable")), pytest.raises(QuotaExceededError, match="quota unavailable"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=db_session_with_containers + ) - updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id) + updated_pool = CreditPoolService.get_pool(tenant_id=tenant_id, session=db_session_with_containers) assert updated_pool is not None assert updated_pool.quota_used == 2 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_service.py index 40c00267043..912e00b0b7d 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service.py @@ -602,10 +602,7 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act / Assert with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( - db_session_with_containers, - source_dataset.id, - {"name": "Existing Dataset"}, - account, + source_dataset.id, {"name": "Existing Dataset"}, account, session=db_session_with_containers ) def test_delete_dataset_with_documents_success(self, db_session_with_containers: Session): @@ -728,7 +725,7 @@ class TestDatasetServiceRetrievalConfiguration: } # Act - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, account) + result = DatasetService.update_dataset(dataset.id, update_data, account, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index ced144e8d6e..6b32273624b 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -574,11 +574,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="does not have permission"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.ALL_TEAM, - [], + user, dataset, DatasetPermissionEnum.ALL_TEAM, [], session=db_session_with_containers ) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode( @@ -589,11 +585,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.ONLY_ME, - [], + user, dataset, DatasetPermissionEnum.ONLY_ME, [], session=db_session_with_containers ) def test_check_permission_requires_partial_member_list_for_partial_members_mode( @@ -604,11 +596,7 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(ValueError, match="Partial member list is required"): DatasetPermissionService.check_permission( - db_session_with_containers, - user, - dataset, - DatasetPermissionEnum.PARTIAL_TEAM, - [], + user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [], session=db_session_with_containers ) def test_check_permission_rejects_dataset_operator_member_list_changes(self, db_session_with_containers: Session): @@ -618,11 +606,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-2"}], + session=db_session_with_containers, ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged( @@ -633,11 +621,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( - db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-1"}], + session=db_session_with_containers, ) def test_clear_partial_member_list_deletes_permissions_and_commits(self, db_session_with_containers: Session): diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py index f719a465dbd..d9fb23e8e33 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py @@ -189,7 +189,7 @@ class TestDatasetServiceUpdateDataset: "external_knowledge_api_id": external_api.id, } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) updated_binding = db_session_with_containers.query(ExternalKnowledgeBindings).filter_by(id=binding_id).first() @@ -221,7 +221,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_api_id": str(uuid4())} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge id is required" in str(context.value) db_session_with_containers.rollback() @@ -245,7 +245,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_id": "knowledge_id"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge api id is required" in str(context.value) db_session_with_containers.rollback() @@ -272,7 +272,7 @@ class TestDatasetServiceUpdateDataset: } with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "External knowledge binding not found" in str(context.value) db_session_with_containers.rollback() @@ -303,7 +303,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": "text-embedding-ada-002", } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -338,7 +338,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": None, } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -371,7 +371,7 @@ class TestDatasetServiceUpdateDataset: } with patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task: - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_task.delay.assert_called_once_with(dataset.id, "remove") db_session_with_containers.refresh(dataset) @@ -418,7 +418,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -462,7 +462,7 @@ class TestDatasetServiceUpdateDataset: "retrieval_model": "new_model", } - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -514,7 +514,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -545,7 +545,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(db_session_with_containers, str(uuid4()), update_data, user) + DatasetService.update_dataset(str(uuid4()), update_data, user, session=db_session_with_containers) assert "Dataset not found" in str(context.value) @@ -568,7 +568,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(NoPermissionError): - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, outsider) + DatasetService.update_dataset(dataset.id, update_data, outsider, session=db_session_with_containers) def test_update_internal_dataset_embedding_model_error(self, db_session_with_containers: Session): """Test error when embedding model is not available.""" @@ -595,6 +595,6 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.side_effect = Exception("No Embedding Model available") with pytest.raises(Exception) as context: - DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) + DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) assert "No Embedding Model available".lower() in str(context.value).lower() diff --git a/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py b/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py index 5eb84f805aa..a5445a17297 100644 --- a/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py +++ b/api/tests/test_containers_integration_tests/services/test_file_service_zip_and_lookup.py @@ -69,7 +69,7 @@ def test_build_upload_files_zip_tempfile_sanitizes_and_dedupes_names(monkeypatch def test_get_upload_files_by_ids_returns_empty_when_no_ids(db_session_with_containers: Session) -> None: """Ensure empty input returns an empty mapping without hitting the database.""" - assert FileService.get_upload_files_by_ids(db_session_with_containers, str(uuid4()), []) == {} + assert FileService.get_upload_files_by_ids(str(uuid4()), [], session=db_session_with_containers) == {} def test_get_upload_files_by_ids_returns_id_keyed_mapping(db_session_with_containers: Session) -> None: @@ -78,7 +78,9 @@ def test_get_upload_files_by_ids_returns_id_keyed_mapping(db_session_with_contai file1 = _create_upload_file(db_session_with_containers, tenant_id=tenant_id, key="k1", name="file1.txt") file2 = _create_upload_file(db_session_with_containers, tenant_id=tenant_id, key="k2", name="file2.txt") - result = FileService.get_upload_files_by_ids(db_session_with_containers, tenant_id, [file1.id, file1.id, file2.id]) + result = FileService.get_upload_files_by_ids( + tenant_id, [file1.id, file1.id, file2.id], session=db_session_with_containers + ) assert set(result.keys()) == {file1.id, file2.id} assert result[file1.id].id == file1.id @@ -92,6 +94,6 @@ def test_get_upload_files_by_ids_filters_by_tenant(db_session_with_containers: S file_a = _create_upload_file(db_session_with_containers, tenant_id=tenant_a, key="ka", name="a.txt") _create_upload_file(db_session_with_containers, tenant_id=tenant_b, key="kb", name="b.txt") - result = FileService.get_upload_files_by_ids(db_session_with_containers, tenant_a, [file_a.id]) + result = FileService.get_upload_files_by_ids(tenant_a, [file_a.id], session=db_session_with_containers) assert set(result.keys()) == {file_a.id} diff --git a/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py b/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py index 4a73f98f50e..67b3a2d3e57 100644 --- a/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_hit_testing_service.py @@ -192,7 +192,7 @@ class TestHitTestingService: mock_format.return_value = [mock_record] response = _RetrieveResponse.model_validate( - HitTestingService.compact_retrieve_response(db_session_with_containers, query, [mock_doc]) + HitTestingService.compact_retrieve_response(query, [mock_doc], session=db_session_with_containers) ) assert response.query.content == query @@ -246,12 +246,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.external_retrieve( - db_session_with_containers, dataset=dataset, query='test "query"', account=account, external_retrieval_model={"model": "test"}, metadata_filtering_conditions={"key": "val"}, + session=db_session_with_containers, ) ) @@ -276,7 +276,7 @@ class TestHitTestingService: account = MagicMock() response = _RetrieveResponse.model_validate( - HitTestingService.external_retrieve(db_session_with_containers, dataset, "test query", account) + HitTestingService.external_retrieve(dataset, "test query", account, session=db_session_with_containers) ) assert response.query.content == "test query" @@ -300,12 +300,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=None, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) ) @@ -343,12 +343,12 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) mock_get_meta.assert_called_once() @@ -380,12 +380,12 @@ class TestHitTestingService: response = _RetrieveResponse.model_validate( HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) ) @@ -412,13 +412,13 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, attachment_ids=attachment_ids, + session=db_session_with_containers, ) mock_retrieve.assert_called_once_with( @@ -472,12 +472,12 @@ class TestHitTestingService: mock_retrieve.return_value = retrieved_documents HitTestingService.retrieve( - db_session_with_containers, dataset=dataset, query="test query", account=account, retrieval_model=retrieval_model, external_retrieval_model=external_retrieval_model, + session=db_session_with_containers, ) mock_retrieve.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py b/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py index c1188d3d0f9..84a0226ba17 100644 --- a/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py +++ b/api/tests/test_containers_integration_tests/services/test_human_input_delivery_test.py @@ -119,6 +119,7 @@ def test_human_input_delivery_test_sends_email( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) assert send_mock.call_count == 1 @@ -145,6 +146,7 @@ def test_human_input_delivery_test_form_accepts_file_upload( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) form = db_session_with_containers.scalar( @@ -213,6 +215,7 @@ def test_human_input_delivery_test_form_accepts_remote_file_upload( account=account, node_id="human-node", delivery_method_id=str(delivery_method_id), + session=db_session_with_containers, ) form = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_message_service.py b/api/tests/test_containers_integration_tests/services/test_message_service.py index f2d682be3bf..702812b96de 100644 --- a/api/tests/test_containers_integration_tests/services/test_message_service.py +++ b/api/tests/test_containers_integration_tests/services/test_message_service.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -117,7 +117,7 @@ class TestMessageService: # Create app app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Setup current_user mock self._mock_current_user(mock_external_service_dependencies, account.id, tenant.id) @@ -222,6 +222,7 @@ class TestMessageService: first_id=messages[2].id, # Use middle message as first_id limit=2, order="asc", + session=db_session_with_containers, ) # Verify results @@ -243,7 +244,12 @@ class TestMessageService: # Test pagination with no user result = MessageService.pagination_by_first_id( - app_model=app, user=None, conversation_id=fake.uuid4(), first_id=None, limit=10 + app_model=app, + user=None, + conversation_id=fake.uuid4(), + first_id=None, + limit=10, + session=db_session_with_containers, ) # Verify empty result @@ -262,7 +268,12 @@ class TestMessageService: # Test pagination with no conversation ID result = MessageService.pagination_by_first_id( - app_model=app, user=account, conversation_id="", first_id=None, limit=10 + app_model=app, + user=account, + conversation_id="", + first_id=None, + limit=10, + session=db_session_with_containers, ) # Verify empty result @@ -291,6 +302,7 @@ class TestMessageService: conversation_id=conversation.id, first_id=fake.uuid4(), # Non-existent message ID limit=10, + session=db_session_with_containers, ) def test_pagination_by_last_id_success( @@ -316,6 +328,7 @@ class TestMessageService: last_id=messages[2].id, # Use middle message as last_id limit=2, conversation_id=conversation.id, + session=db_session_with_containers, ) # Verify results @@ -345,7 +358,12 @@ class TestMessageService: # Test pagination with include_ids include_ids = [messages[0].id, messages[1].id, messages[2].id] result = MessageService.pagination_by_last_id( - app_model=app, user=account, last_id=messages[1].id, limit=2, include_ids=include_ids + app_model=app, + user=account, + last_id=messages[1].id, + limit=2, + include_ids=include_ids, + session=db_session_with_containers, ) # Verify results @@ -364,8 +382,10 @@ class TestMessageService: fake = Faker() app, account = self._create_test_app_and_account(db_session_with_containers, mock_external_service_dependencies) - # Test pagination with no user - result = MessageService.pagination_by_last_id(app_model=app, user=None, last_id=None, limit=10) + # Test pagination with no user, + result = MessageService.pagination_by_last_id( + app_model=app, user=None, last_id=None, limit=10, session=db_session_with_containers + ) # Verify empty result assert result.limit == 10 @@ -393,6 +413,7 @@ class TestMessageService: last_id=fake.uuid4(), # Non-existent message ID limit=10, conversation_id=conversation.id, + session=db_session_with_containers, ) def test_create_feedback_success(self, db_session_with_containers: Session, mock_external_service_dependencies): @@ -410,7 +431,12 @@ class TestMessageService: rating = FeedbackRating.LIKE content = fake.text(max_nb_chars=100) feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=rating, content=content + app_model=app, + message_id=message.id, + user=account, + rating=rating, + content=content, + session=db_session_with_containers, ) # Verify feedback was created correctly @@ -442,6 +468,7 @@ class TestMessageService: user=None, rating=FeedbackRating.LIKE, content=fake.text(max_nb_chars=100), + session=db_session_with_containers, ) def test_create_feedback_update_existing( @@ -461,14 +488,24 @@ class TestMessageService: initial_rating = FeedbackRating.LIKE initial_content = fake.text(max_nb_chars=100) feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=initial_rating, content=initial_content + app_model=app, + message_id=message.id, + user=account, + rating=initial_rating, + content=initial_content, + session=db_session_with_containers, ) # Update feedback updated_rating = FeedbackRating.DISLIKE updated_content = fake.text(max_nb_chars=100) updated_feedback = MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=updated_rating, content=updated_content + app_model=app, + message_id=message.id, + user=account, + rating=updated_rating, + content=updated_content, + session=db_session_with_containers, ) # Verify feedback was updated correctly @@ -498,10 +535,18 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE, content=fake.text(max_nb_chars=100), + session=db_session_with_containers, ) - # Delete feedback by setting rating to None - MessageService.create_feedback(app_model=app, message_id=message.id, user=account, rating=None, content=None) + # Delete feedback by setting rating to None, + MessageService.create_feedback( + app_model=app, + message_id=message.id, + user=account, + rating=None, + content=None, + session=db_session_with_containers, + ) # Verify feedback was deleted @@ -526,7 +571,12 @@ class TestMessageService: # Test creating feedback with no rating when no feedback exists with pytest.raises(ValueError, match="rating cannot be None when feedback not exists"): MessageService.create_feedback( - app_model=app, message_id=message.id, user=account, rating=None, content=None + app_model=app, + message_id=message.id, + user=account, + rating=None, + content=None, + session=db_session_with_containers, ) def test_get_all_messages_feedbacks_success( @@ -550,11 +600,12 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE if i % 2 == 0 else FeedbackRating.DISLIKE, content=f"Feedback {i}: {fake.text(max_nb_chars=50)}", + session=db_session_with_containers, ) feedbacks.append(feedback) - # Get all feedbacks - result = MessageService.get_all_messages_feedbacks(app, page=1, limit=10) + # Get all feedbacks, + result = MessageService.get_all_messages_feedbacks(app, page=1, limit=10, session=db_session_with_containers) # Verify results assert len(result) == 3 @@ -583,11 +634,16 @@ class TestMessageService: user=account, rating=FeedbackRating.LIKE, content=f"Feedback {i}", + session=db_session_with_containers, ) # Get feedbacks with pagination - result_page_1 = MessageService.get_all_messages_feedbacks(app, page=1, limit=3) - result_page_2 = MessageService.get_all_messages_feedbacks(app, page=2, limit=3) + result_page_1 = MessageService.get_all_messages_feedbacks( + app, page=1, limit=3, session=db_session_with_containers + ) + result_page_2 = MessageService.get_all_messages_feedbacks( + app, page=2, limit=3, session=db_session_with_containers + ) # Verify pagination results assert len(result_page_1) == 3 @@ -609,8 +665,10 @@ class TestMessageService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) message = self._create_test_message(db_session_with_containers, app, conversation, account, fake) - # Get message - retrieved_message = MessageService.get_message(app_model=app, user=account, message_id=message.id) + # Get message, + retrieved_message = MessageService.get_message( + app_model=app, user=account, message_id=message.id, session=db_session_with_containers + ) # Verify message was retrieved correctly assert retrieved_message.id == message.id @@ -628,7 +686,9 @@ class TestMessageService: # Test getting non-existent message with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=account, message_id=fake.uuid4()) + MessageService.get_message( + app_model=app, user=account, message_id=fake.uuid4(), session=db_session_with_containers + ) def test_get_message_wrong_user(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -657,7 +717,9 @@ class TestMessageService: # Test getting message with different user with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=other_account, message_id=message.id) + MessageService.get_message( + app_model=app, user=other_account, message_id=message.id, session=db_session_with_containers + ) def test_get_suggested_questions_after_answer_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -682,7 +744,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) # Verify results @@ -714,7 +780,11 @@ class TestMessageService: with pytest.raises(ValueError, match="user cannot be None"): MessageService.get_suggested_questions_after_answer( - app_model=app, user=None, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=None, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) def test_get_suggested_questions_after_answer_disabled( @@ -740,7 +810,11 @@ class TestMessageService: with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) def test_get_suggested_questions_after_answer_no_workflow( @@ -763,7 +837,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.SERVICE_API + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.SERVICE_API, + session=db_session_with_containers, ) # Verify empty result @@ -792,7 +870,11 @@ class TestMessageService: from core.app.entities.app_invoke_entities import InvokeFrom result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=account, message_id=message.id, invoke_from=InvokeFrom.DEBUGGER + app_model=app, + user=account, + message_id=message.id, + invoke_from=InvokeFrom.DEBUGGER, + session=db_session_with_containers, ) # Verify results @@ -800,7 +882,7 @@ class TestMessageService: # Verify draft workflow was used instead of published workflow mock_external_service_dependencies["workflow_service"].return_value.get_draft_workflow.assert_called_once_with( - app_model=app + app_model=app, session=ANY ) # Verify TraceQueueManager was called diff --git a/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py b/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py index 6a9046acd4a..9da20d96f06 100644 --- a/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py +++ b/api/tests/test_containers_integration_tests/services/test_message_service_execution_extra_content.py @@ -20,6 +20,7 @@ def test_pagination_returns_extra_contents(db_session_with_containers: Session): conversation_id=fixture.conversation.id, first_id=None, limit=10, + session=db_session_with_containers, ) assert pagination.data @@ -59,6 +60,7 @@ def test_pagination_returns_waiting_human_input_extra_contents(db_session_with_c conversation_id=fixture.conversation.id, first_id=None, limit=10, + session=db_session_with_containers, ) assert pagination.data diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py index fbdc265265d..a9399985307 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py @@ -95,7 +95,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() updated_doc = db_session_with_containers.get(Document, document.id) @@ -126,7 +128,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() updated_doc = db_session_with_containers.get(Document, document.id) @@ -168,7 +172,9 @@ class TestMetadataPartialUpdate: ) metadata_args = MetadataOperationData(operation_data=[operation]) - MetadataService.update_documents_metadata(db_session_with_containers, dataset, metadata_args, current_account) + MetadataService.update_documents_metadata( + dataset, metadata_args, current_account, session=db_session_with_containers + ) db_session_with_containers.expire_all() bindings = db_session_with_containers.scalars( @@ -205,5 +211,5 @@ class TestMetadataPartialUpdate: with patch.object(db_session_with_containers, "commit", side_effect=RuntimeError("database connection lost")): with pytest.raises(RuntimeError, match="database connection lost"): MetadataService.update_documents_metadata( - db_session_with_containers, dataset, metadata_args, current_account + dataset, metadata_args, current_account, session=db_session_with_containers ) diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_service.py b/api/tests/test_containers_integration_tests/services/test_metadata_service.py index 7cc9fc7e696..00afe7f8467 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_service.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_service.py @@ -184,7 +184,7 @@ class TestMetadataService: # Act: Execute the method under test result = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -220,7 +220,9 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."): - MetadataService.create_metadata(db_session_with_containers, dataset.id, metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers + ) def test_create_metadata_name_already_exists( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -238,7 +240,9 @@ class TestMetadataService: # Create first metadata first_metadata_args = MetadataArgs(type="string", name="duplicate_name") - MetadataService.create_metadata(db_session_with_containers, dataset.id, first_metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, first_metadata_args, account, tenant.id, session=db_session_with_containers + ) # Try to create second metadata with same name second_metadata_args = MetadataArgs(type="number", name="duplicate_name") @@ -246,7 +250,7 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name already exists."): MetadataService.create_metadata( - db_session_with_containers, dataset.id, second_metadata_args, account, tenant.id + dataset.id, second_metadata_args, account, tenant.id, session=db_session_with_containers ) def test_create_metadata_name_conflicts_with_built_in_field( @@ -269,7 +273,9 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."): - MetadataService.create_metadata(db_session_with_containers, dataset.id, metadata_args, account, tenant.id) + MetadataService.create_metadata( + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers + ) def test_update_metadata_name_success( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -288,13 +294,13 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test new_name = "new_name" result = MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, new_name, account, tenant.id + dataset.id, metadata.id, new_name, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -325,7 +331,7 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update with too long name @@ -334,7 +340,7 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, long_name, account, tenant.id + dataset.id, metadata.id, long_name, account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_already_exists( @@ -354,18 +360,18 @@ class TestMetadataService: # Create two metadata entries first_metadata_args = MetadataArgs(type="string", name="first_metadata") first_metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, first_metadata_args, account, tenant.id + dataset.id, first_metadata_args, account, tenant.id, session=db_session_with_containers ) second_metadata_args = MetadataArgs(type="number", name="second_metadata") second_metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, second_metadata_args, account, tenant.id + dataset.id, second_metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update first metadata with second metadata's name with pytest.raises(ValueError, match="Metadata name already exists."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, first_metadata.id, "second_metadata", account, tenant.id + dataset.id, first_metadata.id, "second_metadata", account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_conflicts_with_built_in_field( @@ -385,7 +391,7 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="old_name") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Try to update with built-in field name @@ -393,7 +399,7 @@ class TestMetadataService: with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."): MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, metadata.id, built_in_field_name, account, tenant.id + dataset.id, metadata.id, built_in_field_name, account, tenant.id, session=db_session_with_containers ) def test_update_metadata_name_not_found( @@ -418,7 +424,7 @@ class TestMetadataService: # Act: Execute the method under test result = MetadataService.update_metadata_name( - db_session_with_containers, dataset.id, fake_metadata_id, new_name, account, tenant.id + dataset.id, fake_metadata_id, new_name, account, tenant.id, session=db_session_with_containers ) # Assert: Verify the method returns None when metadata is not found @@ -441,11 +447,11 @@ class TestMetadataService: # Create metadata first metadata_args = MetadataArgs(type="string", name="to_be_deleted") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, metadata.id) + result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -476,7 +482,7 @@ class TestMetadataService: fake_metadata_id = str(uuid.uuid4()) # Use valid UUID format # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, fake_metadata_id) + result = MetadataService.delete_metadata(dataset.id, fake_metadata_id, session=db_session_with_containers) # Assert: Verify the method returns None when metadata is not found assert result is None @@ -501,7 +507,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create metadata binding @@ -522,7 +528,7 @@ class TestMetadataService: db_session_with_containers.commit() # Act: Execute the method under test - result = MetadataService.delete_metadata(db_session_with_containers, dataset.id, metadata.id) + result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -587,7 +593,7 @@ class TestMetadataService: assert dataset.built_in_field_enabled is False # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -623,7 +629,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the method returns early without changes db_session_with_containers.refresh(dataset) @@ -649,7 +655,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.enable_built_in_field(db_session_with_containers, dataset) + MetadataService.enable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -696,7 +702,7 @@ class TestMetadataService: ] # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes db_session_with_containers.refresh(dataset) @@ -728,7 +734,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the method returns early without changes @@ -761,7 +767,7 @@ class TestMetadataService: ]() # Act: Execute the method under test - MetadataService.disable_built_in_field(db_session_with_containers, dataset) + MetadataService.disable_built_in_field(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes db_session_with_containers.refresh(dataset) @@ -787,7 +793,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Mock DocumentService.get_document @@ -807,7 +813,7 @@ class TestMetadataService: operation_data = MetadataOperationData(operation_data=[operation]) # Act: Execute the method under test - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata(dataset, operation_data, account, session=db_session_with_containers) # Assert: Verify the expected outcomes @@ -853,7 +859,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Mock DocumentService.get_document @@ -873,7 +879,7 @@ class TestMetadataService: operation_data = MetadataOperationData(operation_data=[operation]) # Act: Execute the method under test - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata(dataset, operation_data, account, session=db_session_with_containers) # Assert: Verify the expected outcomes # Verify document metadata was updated with both custom and built-in fields @@ -902,7 +908,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create metadata operation data @@ -924,7 +930,9 @@ class TestMetadataService: # Act & Assert: The method should raise ValueError("Document not found.") # because the exception is now re-raised after rollback with pytest.raises(ValueError, match="Document not found"): - MetadataService.update_documents_metadata(db_session_with_containers, dataset, operation_data, account) + MetadataService.update_documents_metadata( + dataset, operation_data, account, session=db_session_with_containers + ) def test_knowledge_base_metadata_lock_check_dataset_id( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps @@ -1021,7 +1029,7 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Create document and metadata binding @@ -1041,7 +1049,7 @@ class TestMetadataService: db_session_with_containers.commit() # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -1082,11 +1090,11 @@ class TestMetadataService: # Create metadata metadata_args = MetadataArgs(type="string", name="test_metadata") metadata = MetadataService.create_metadata( - db_session_with_containers, dataset.id, metadata_args, account, tenant.id + dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -1115,7 +1123,7 @@ class TestMetadataService: ) # Act: Execute the method under test - result = MetadataService.get_dataset_metadatas(db_session_with_containers, dataset) + result = MetadataService.get_dataset_metadatas(dataset, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None diff --git a/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py b/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py index aca38391353..71d2c1c6800 100644 --- a/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py +++ b/api/tests/test_containers_integration_tests/services/test_model_load_balancing_service.py @@ -339,7 +339,11 @@ class TestModelLoadBalancingService: # Act: Execute the method under test service = ModelLoadBalancingService() is_enabled, configs = service.get_load_balancing_configs( - tenant_id=tenant.id, provider="openai", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="openai", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -381,7 +385,11 @@ class TestModelLoadBalancingService: service = ModelLoadBalancingService() with pytest.raises(ValueError) as exc_info: service.get_load_balancing_configs( - tenant_id=tenant.id, provider="nonexistent_provider", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="nonexistent_provider", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Verify correct error message @@ -443,7 +451,11 @@ class TestModelLoadBalancingService: # Act: Execute the method under test service = ModelLoadBalancingService() is_enabled, configs = service.get_load_balancing_configs( - tenant_id=tenant.id, provider="openai", model="gpt-3.5-turbo", model_type="llm" + tenant_id=tenant.id, + provider="openai", + model="gpt-3.5-turbo", + model_type="llm", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes diff --git a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py b/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py index 0969198ecf3..d397c62b6a8 100644 --- a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py +++ b/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py @@ -4,7 +4,7 @@ from __future__ import annotations import uuid from typing import cast -from unittest.mock import ANY, MagicMock, patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest @@ -159,17 +159,19 @@ class TestOAuthServerServiceTokenOperations: def test_validate_access_token_returns_none_when_not_found(self, mock_redis): mock_redis.get.return_value = None + session = MagicMock() - result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token") + result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session) assert result is None def test_validate_access_token_loads_user_when_exists(self, mock_redis): mock_redis.get.return_value = b"user-88" expected_user = MagicMock() + session = MagicMock() with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load: - result = OAuthServerService.validate_oauth_access_token("client-1", "access-token") + result = OAuthServerService.validate_oauth_access_token("client-1", "access-token", session) assert result is expected_user - mock_load.assert_called_once_with("user-88", ANY) + mock_load.assert_called_once_with("user-88", session) diff --git a/api/tests/test_containers_integration_tests/services/test_ops_service.py b/api/tests/test_containers_integration_tests/services/test_ops_service.py index 9643fb61d44..b4b8521fb2e 100644 --- a/api/tests/test_containers_integration_tests/services/test_ops_service.py +++ b/api/tests/test_containers_integration_tests/services/test_ops_service.py @@ -67,6 +67,7 @@ class TestOpsService: icon_background="#FF6B6B", ), account, + session=db_session_with_containers, ) return app, account @@ -91,13 +92,13 @@ class TestOpsService: # ── get_tracing_app_config ───────────────────────────────────────── def test_get_tracing_app_config_no_config(self, db_session_with_containers: Session, mock_ops_trace_manager): - result = OpsService.get_tracing_app_config(str(uuid.uuid4()), "arize") + result = OpsService.get_tracing_app_config(str(uuid.uuid4()), "arize", db_session_with_containers) assert result is None def test_get_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): fake_app_id = str(uuid.uuid4()) self._insert_trace_config(db_session_with_containers, fake_app_id, "arize") - result = OpsService.get_tracing_app_config(fake_app_id, "arize") + result = OpsService.get_tracing_app_config(fake_app_id, "arize", db_session_with_containers) assert result is None def test_get_tracing_app_config_none_config( @@ -107,7 +108,7 @@ class TestOpsService: self._insert_trace_config(db_session_with_containers, app.id, "arize", tracing_config=None) with pytest.raises(ValueError, match="Tracing config cannot be None."): - OpsService.get_tracing_app_config(app.id, "arize") + OpsService.get_tracing_app_config(app.id, "arize", db_session_with_containers) @pytest.mark.parametrize( ("provider", "default_url"), @@ -135,7 +136,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, provider) - result = OpsService.get_tracing_app_config(app.id, provider) + result = OpsService.get_tracing_app_config(app.id, provider, db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == default_url @@ -155,7 +156,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, provider) - result = OpsService.get_tracing_app_config(app.id, provider) + result = OpsService.get_tracing_app_config(app.id, provider, db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "success_url" @@ -171,7 +172,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "langfuse") - result = OpsService.get_tracing_app_config(app.id, "langfuse") + result = OpsService.get_tracing_app_config(app.id, "langfuse", db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "https://api.langfuse.com/project/key" @@ -187,7 +188,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "langfuse") - result = OpsService.get_tracing_app_config(app.id, "langfuse") + result = OpsService.get_tracing_app_config(app.id, "langfuse", db_session_with_containers) assert result is not None assert result["tracing_config"]["project_url"] == "https://api.langfuse.com/" @@ -195,7 +196,9 @@ class TestOpsService: # ── create_tracing_app_config ────────────────────────────────────── def test_create_tracing_app_config_invalid_provider(self, db_session_with_containers: Session): - result = OpsService.create_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}) + result = OpsService.create_tracing_app_config( + str(uuid.uuid4()), "invalid_provider", {}, db_session_with_containers + ) assert result == {"error": "Invalid tracing provider: invalid_provider"} def test_create_tracing_app_config_invalid_credentials( @@ -203,7 +206,10 @@ class TestOpsService: ): mock_ops_trace_manager.check_trace_config_is_effective.return_value = False result = OpsService.create_tracing_app_config( - str(uuid.uuid4()), TracingProviderEnum.LANGFUSE, {"public_key": "p", "secret_key": "s"} + str(uuid.uuid4()), + TracingProviderEnum.LANGFUSE, + {"public_key": "p", "secret_key": "s"}, + db_session_with_containers, ) assert result == {"error": "Invalid Credentials"} @@ -228,7 +234,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(provider)) - result = OpsService.create_tracing_app_config(app.id, provider, config) + result = OpsService.create_tracing_app_config(app.id, provider, config, db_session_with_containers) assert result is None @@ -245,6 +251,7 @@ class TestOpsService: app.id, TracingProviderEnum.LANGFUSE, {"public_key": "p", "secret_key": "s", "host": "https://api.langfuse.com"}, + db_session_with_containers, ) assert result == {"result": "success"} @@ -258,13 +265,17 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_create_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): mock_ops_trace_manager.check_trace_config_is_effective.return_value = True - result = OpsService.create_tracing_app_config(str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_create_tracing_app_config_with_empty_other_keys( @@ -277,7 +288,9 @@ class TestOpsService: mock_otm.encrypt_tracing_config.return_value = {} app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {"project": ""}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {"project": ""}, db_session_with_containers + ) assert result == {"result": "success"} @@ -290,7 +303,9 @@ class TestOpsService: mock_otm.encrypt_tracing_config.return_value = {"encrypted": "config"} app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) - result = OpsService.create_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.create_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result == {"result": "success"} @@ -298,17 +313,21 @@ class TestOpsService: def test_update_tracing_app_config_invalid_provider(self, db_session_with_containers: Session): with pytest.raises(ValueError, match="Invalid tracing provider: invalid_provider"): - OpsService.update_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}) + OpsService.update_tracing_app_config(str(uuid.uuid4()), "invalid_provider", {}, db_session_with_containers) def test_update_tracing_app_config_no_config(self, db_session_with_containers: Session, mock_ops_trace_manager): - result = OpsService.update_tracing_app_config(str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + str(uuid.uuid4()), TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_update_tracing_app_config_no_app(self, db_session_with_containers: Session, mock_ops_trace_manager): fake_app_id = str(uuid.uuid4()) self._insert_trace_config(db_session_with_containers, fake_app_id, str(TracingProviderEnum.ARIZE)) mock_ops_trace_manager.encrypt_tracing_config.return_value = {} - result = OpsService.update_tracing_app_config(fake_app_id, TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + fake_app_id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is None def test_update_tracing_app_config_invalid_credentials( @@ -323,7 +342,7 @@ class TestOpsService: self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) with pytest.raises(ValueError, match="Invalid Credentials"): - OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers) def test_update_tracing_app_config_success( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -336,7 +355,9 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, str(TracingProviderEnum.ARIZE)) - result = OpsService.update_tracing_app_config(app.id, TracingProviderEnum.ARIZE, {}) + result = OpsService.update_tracing_app_config( + app.id, TracingProviderEnum.ARIZE, {}, db_session_with_containers + ) assert result is not None assert result["app_id"] == app.id @@ -344,7 +365,7 @@ class TestOpsService: # ── delete_tracing_app_config ────────────────────────────────────── def test_delete_tracing_app_config_no_config(self, db_session_with_containers: Session): - result = OpsService.delete_tracing_app_config(str(uuid.uuid4()), "arize") + result = OpsService.delete_tracing_app_config(str(uuid.uuid4()), "arize", db_session_with_containers) assert result is None def test_delete_tracing_app_config_success( @@ -353,7 +374,7 @@ class TestOpsService: app, _ = self._create_app(db_session_with_containers, mock_external_service_dependencies) self._insert_trace_config(db_session_with_containers, app.id, "arize") - result = OpsService.delete_tracing_app_config(app.id, "arize") + result = OpsService.delete_tracing_app_config(app.id, "arize", db_session_with_containers) assert result is True remaining = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py b/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py index 9b8eec08ef4..f27132b0fe9 100644 --- a/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_recommended_app_service.py @@ -14,6 +14,8 @@ from models.model import AccountTrialAppRecord, TrialApp from services import recommended_app_service as service_module from services.recommended_app_service import RecommendedAppService +pytestmark = pytest.mark.usefixtures("db_session_with_containers") + class RecommendedAppPayload(TypedDict, total=False): id: str @@ -118,13 +120,13 @@ class TestRecommendedAppServiceGetApps: mock_factory = MagicMock(return_value=mock_instance) mock_factory_class.get_recommend_app_factory.return_value = mock_factory - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == expected assert len(result["recommended_apps"]) == 2 assert len(result["categories"]) == 3 mock_factory_class.get_recommend_app_factory.assert_called_once_with("remote") - mock_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US") + mock_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US", session=db.session()) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") @@ -143,7 +145,7 @@ class TestRecommendedAppServiceGetApps: mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "zh-CN") + result = RecommendedAppService.get_recommended_apps_and_categories("zh-CN", session=db.session()) assert result == builtin_response assert result["recommended_apps"][0]["id"] == "builtin-1" @@ -164,7 +166,7 @@ class TestRecommendedAppServiceGetApps: mock_builtin_instance.fetch_recommended_apps_from_builtin.return_value = builtin_response mock_factory_class.get_buildin_recommend_app_retrieval.return_value = mock_builtin_instance - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == builtin_response mock_builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once() @@ -182,10 +184,10 @@ class TestRecommendedAppServiceGetApps: mock_instance.get_recommended_apps_and_categories.return_value = lang_response mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, language) + result = RecommendedAppService.get_recommended_apps_and_categories(language, session=db.session()) assert result["recommended_apps"][0]["id"] == f"app-{language}" - mock_instance.get_recommended_apps_and_categories.assert_called_with(language) + mock_instance.get_recommended_apps_and_categories.assert_called_with(language, session=db.session()) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") @@ -197,7 +199,7 @@ class TestRecommendedAppServiceGetApps: mock_instance.get_recommended_apps_and_categories.return_value = response mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) mock_factory_class.get_recommend_app_factory.assert_called_with(mode) @@ -237,10 +239,10 @@ class TestRecommendedAppServiceGetDetail: mock_instance.get_recommend_app_detail.return_value = expected mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, app_id) + result = RecommendedAppService.get_recommend_app_detail(app_id, session=db.session()) assert result == expected - mock_instance.get_recommend_app_detail.assert_called_once_with(app_id) + mock_instance.get_recommend_app_detail.assert_called_once_with(app_id, session=db.session()) @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @@ -256,10 +258,10 @@ class TestRecommendedAppServiceGetDetail: mock_instance.get_recommend_app_detail.return_value = detail mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, "test-app") + result = RecommendedAppService.get_recommend_app_detail("test-app", session=db.session()) assert result is not None - mock_instance.get_recommend_app_detail.assert_called_with("test-app") + mock_instance.get_recommend_app_detail.assert_called_with("test-app", session=db.session()) mock_factory_class.get_recommend_app_factory.assert_called_with(mode) @@ -283,11 +285,11 @@ class TestRecommendedAppServiceGetLearnDifyApps: } mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_learn_dify_apps(db.session, "en-US") + result = RecommendedAppService.get_learn_dify_apps("en-US", session=db.session()) assert result == {"recommended_apps": [expected_app]} mock_factory_class.get_recommend_app_factory.assert_called_once_with("remote") - mock_instance.get_learn_dify_apps.assert_called_once_with("en-US") + mock_instance.get_learn_dify_apps.assert_called_once_with("en-US", session=db.session()) @patch("services.recommended_app_service.dify_config") def test_sets_can_trial_when_trial_feature_enabled( @@ -314,10 +316,10 @@ class TestRecommendedAppServiceGetLearnDifyApps: can_trial_mock = MagicMock(return_value=True) monkeypatch.setattr(RecommendedAppService, "_can_trial_app", can_trial_mock) - result = RecommendedAppService.get_learn_dify_apps(db.session, "en-US") + result = RecommendedAppService.get_learn_dify_apps("en-US", session=db.session()) assert result["recommended_apps"][0]["can_trial"] is True - can_trial_mock.assert_called_once_with(db.session, "app-1") + can_trial_mock.assert_called_once_with(db.session(), "app-1") # ── Integration tests: trial app features (real DB) ──────────────────── @@ -333,10 +335,10 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=False)), ) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "en-US") + result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=db.session()) assert result == expected - retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US") + retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US", session=db.session()) builtin_instance.fetch_recommended_apps_from_builtin.assert_not_called() def test_get_apps_should_enrich_can_trial_when_enabled( @@ -364,7 +366,7 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=True)), ) - result = RecommendedAppService.get_recommended_apps_and_categories(db.session, "ja-JP") + result = RecommendedAppService.get_recommended_apps_and_categories("ja-JP", session=db.session()) builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once_with("en-US") assert result["recommended_apps"][0]["can_trial"] is True @@ -400,7 +402,7 @@ class TestRecommendedAppServiceTrialFeatures: MagicMock(return_value=SimpleNamespace(enable_trial_app=True)), ) - result = RecommendedAppService.get_recommend_app_detail(db.session, app_id) + result = RecommendedAppService.get_recommend_app_detail(app_id, session=db.session()) assert result is not None detail_result = cast(RecommendedAppPayload, result) @@ -421,10 +423,10 @@ class TestRecommendedAppServiceTrialFeatures: mock_instance.get_recommend_app_detail.return_value = None mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - result = RecommendedAppService.get_recommend_app_detail(db.session, "nonexistent") + result = RecommendedAppService.get_recommend_app_detail("nonexistent", session=db.session()) assert result is None - mock_instance.get_recommend_app_detail.assert_called_once_with("nonexistent") + mock_instance.get_recommend_app_detail.assert_called_once_with("nonexistent", session=db.session()) mock_feature_service.get_system_features.assert_not_called() def test_add_trial_app_record_increments_count_for_existing(self, db_session_with_containers: Session) -> None: @@ -434,7 +436,7 @@ class TestRecommendedAppServiceTrialFeatures: db_session_with_containers.add(AccountTrialAppRecord(app_id=app_id, account_id=account_id, count=3)) db_session_with_containers.commit() - RecommendedAppService.add_trial_app_record(db.session, app_id, account_id) + RecommendedAppService.add_trial_app_record(app_id, account_id, session=db.session()) db_session_with_containers.expire_all() record = db_session_with_containers.scalar( @@ -449,7 +451,7 @@ class TestRecommendedAppServiceTrialFeatures: app_id = str(uuid.uuid4()) account_id = str(uuid.uuid4()) - RecommendedAppService.add_trial_app_record(db.session, app_id, account_id) + RecommendedAppService.add_trial_app_record(app_id, account_id, session=db.session()) db_session_with_containers.expire_all() record = db_session_with_containers.scalar( diff --git a/api/tests/test_containers_integration_tests/services/test_saved_message_service.py b/api/tests/test_containers_integration_tests/services/test_saved_message_service.py index cfd1d4e86b4..92741ac56cb 100644 --- a/api/tests/test_containers_integration_tests/services/test_saved_message_service.py +++ b/api/tests/test_containers_integration_tests/services/test_saved_message_service.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -86,7 +86,7 @@ class TestSavedMessageService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -222,7 +222,7 @@ class TestSavedMessageService: # Act: Execute the method under test result = SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=account, last_id=None, limit=10 + app_model=app, user=account, last_id=None, limit=10, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -297,7 +297,7 @@ class TestSavedMessageService: # Act: Execute the method under test result = SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=end_user, last_id="test_last_id", limit=5 + app_model=app, user=end_user, last_id="test_last_id", limit=5, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -347,7 +347,7 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message # Act: Execute the method under test - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) # Assert: Verify the expected outcomes # Check if saved message was created in database @@ -372,7 +372,7 @@ class TestSavedMessageService: # Verify MessageService.get_message was called mock_external_service_dependencies["message_service"].get_message.assert_called_once_with( - app_model=app, user=account, message_id=message.id + app_model=app, user=account, message_id=message.id, session=ANY ) # Verify database state @@ -397,7 +397,7 @@ class TestSavedMessageService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: SavedMessageService.pagination_by_last_id( - db_session_with_containers, app_model=app, user=None, last_id=None, limit=10 + app_model=app, user=None, last_id=None, limit=10, session=db_session_with_containers ) assert "User is required" in str(exc_info.value) @@ -417,7 +417,9 @@ class TestSavedMessageService: message = self._create_test_message(db_session_with_containers, app, account) # Act: Execute the method under test with None user - result = SavedMessageService.save(db_session_with_containers, app_model=app, user=None, message_id=message.id) + result = SavedMessageService.save( + app_model=app, user=None, message_id=message.id, session=db_session_with_containers + ) # Assert: Verify the expected outcomes assert result is None @@ -476,7 +478,9 @@ class TestSavedMessageService: ) # Act: Execute the method under test - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=account, message_id=message.id, session=db_session_with_containers + ) # Assert: Verify the expected outcomes # Check if saved message was deleted from database @@ -506,7 +510,9 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message - SavedMessageService.save(db_session_with_containers, app_model=app, user=end_user, message_id=message.id) + SavedMessageService.save( + app_model=app, user=end_user, message_id=message.id, session=db_session_with_containers + ) saved = ( db_session_with_containers.query(SavedMessage) @@ -527,9 +533,9 @@ class TestSavedMessageService: mock_external_service_dependencies["message_service"].get_message.return_value = message # Save once - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) # Save again - SavedMessageService.save(db_session_with_containers, app_model=app, user=account, message_id=message.id) + SavedMessageService.save(app_model=app, user=account, message_id=message.id, session=db_session_with_containers) count = ( db_session_with_containers.query(SavedMessage) @@ -552,7 +558,7 @@ class TestSavedMessageService: db_session_with_containers.add(saved) db_session_with_containers.commit() - SavedMessageService.delete(db_session_with_containers, app_model=app, user=None, message_id=message.id) + SavedMessageService.delete(app_model=app, user=None, message_id=message.id, session=db_session_with_containers) # Should still exist assert ( @@ -571,7 +577,9 @@ class TestSavedMessageService: # Should not raise — use a valid UUID that doesn't exist in DB from uuid import uuid4 - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account, message_id=str(uuid4())) + SavedMessageService.delete( + app_model=app, user=account, message_id=str(uuid4()), session=db_session_with_containers + ) def test_delete_for_end_user(self, db_session_with_containers: Session, mock_external_service_dependencies): """Test deleting a saved message for an EndUser.""" @@ -585,7 +593,9 @@ class TestSavedMessageService: db_session_with_containers.add(saved) db_session_with_containers.commit() - SavedMessageService.delete(db_session_with_containers, app_model=app, user=end_user, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=end_user, message_id=message.id, session=db_session_with_containers + ) assert ( db_session_with_containers.query(SavedMessage) @@ -615,7 +625,9 @@ class TestSavedMessageService: db_session_with_containers.commit() # Delete only account1's saved message - SavedMessageService.delete(db_session_with_containers, app_model=app, user=account1, message_id=message.id) + SavedMessageService.delete( + app_model=app, user=account1, message_id=message.id, session=db_session_with_containers + ) # Account's saved message should be gone assert ( diff --git a/api/tests/test_containers_integration_tests/services/test_tag_service.py b/api/tests/test_containers_integration_tests/services/test_tag_service.py index 748cca6c845..86b635ac23d 100644 --- a/api/tests/test_containers_integration_tests/services/test_tag_service.py +++ b/api/tests/test_containers_integration_tests/services/test_tag_service.py @@ -205,7 +205,7 @@ def test_get_tags_success(db_session_with_containers: Session, current_user_stub db_session_with_containers, tags=tags[:2], target_id=dataset.id, tenant_id=tenant.id, user_id=account.id ) - result = TagService.get_tags(db_session_with_containers, TagType.KNOWLEDGE, tenant.id) + result = TagService.get_tags(TagType.KNOWLEDGE, tenant.id, session=db_session_with_containers) assert result is not None assert len(result) == 3 @@ -235,7 +235,7 @@ def test_get_tags_with_keyword_filter(db_session_with_containers: Session, curre tags[2].name = "web_development" db_session_with_containers.flush() - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="development") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="development", session=db_session_with_containers) assert result is not None assert len(result) == 2 @@ -243,7 +243,9 @@ def test_get_tags_with_keyword_filter(db_session_with_containers: Session, curre for tag_result in result: assert "development" in tag_result.name.lower() - result_no_match = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="nonexistent") + result_no_match = TagService.get_tags( + TagType.APP, tenant.id, keyword="nonexistent", session=db_session_with_containers + ) assert result_no_match == [] @@ -291,19 +293,19 @@ def test_get_tags_with_special_characters_in_keyword( db_session_with_containers.flush() - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="50%") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="50%", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "50% discount" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="test_data") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="test_data", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "test_data_tag" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="path\\to\\tag") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="path\\to\\tag", session=db_session_with_containers) assert len(result) == 1 assert result[0].name == "path\\to\\tag" - result = TagService.get_tags(db_session_with_containers, TagType.APP, tenant.id, keyword="50%") + result = TagService.get_tags(TagType.APP, tenant.id, keyword="50%", session=db_session_with_containers) assert len(result) == 1 assert all("50%" in item.name for item in result) @@ -312,7 +314,7 @@ def test_get_tags_empty_result(db_session_with_containers: Session, current_user account, tenant = _create_account_with_tenant(db_session_with_containers) _set_current_user(current_user_stub, account, tenant) - result = TagService.get_tags(db_session_with_containers, TagType.KNOWLEDGE, tenant.id) + result = TagService.get_tags(TagType.KNOWLEDGE, tenant.id, session=db_session_with_containers) assert result == [] diff --git a/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py index 664c1167994..ed063ceaccc 100644 --- a/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_web_conversation_service.py @@ -90,7 +90,7 @@ class TestWebConversationService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -312,7 +312,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify the conversation was pinned @@ -346,10 +346,10 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first time - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Pin the conversation again - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify only one pinned conversation record exists @@ -380,7 +380,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, end_user, fake) # Pin the conversation - WebConversationService.pin(app, conversation.id, end_user) + WebConversationService.pin(app, conversation.id, end_user, db_session_with_containers) # Verify the conversation was pinned @@ -412,7 +412,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify it was pinned @@ -430,7 +430,7 @@ class TestWebConversationService: assert pinned_conversation is not None # Unpin the conversation - WebConversationService.unpin(app, conversation.id, account) + WebConversationService.unpin(app, conversation.id, account, db_session_with_containers) # Verify it was unpinned pinned_conversation = ( @@ -459,7 +459,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Try to unpin a conversation that was never pinned - WebConversationService.unpin(app, conversation.id, account) + WebConversationService.unpin(app, conversation.id, account, db_session_with_containers) # Verify no pinned conversation record exists @@ -509,7 +509,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Try to pin with None user - WebConversationService.pin(app, conversation.id, None) + WebConversationService.pin(app, conversation.id, None, db_session_with_containers) # Verify no pinned conversation was created @@ -537,7 +537,7 @@ class TestWebConversationService: conversation = self._create_test_conversation(db_session_with_containers, app, account, fake) # Pin the conversation first - WebConversationService.pin(app, conversation.id, account) + WebConversationService.pin(app, conversation.id, account, db_session_with_containers) # Verify it was pinned @@ -555,7 +555,7 @@ class TestWebConversationService: assert pinned_conversation is not None # Try to unpin with None user - WebConversationService.unpin(app, conversation.id, None) + WebConversationService.unpin(app, conversation.id, None, db_session_with_containers) # Verify the conversation is still pinned pinned_conversation = ( diff --git a/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py b/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py index 7825f502f77..52d1fde7927 100644 --- a/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py +++ b/api/tests/test_containers_integration_tests/services/test_webapp_auth_service.py @@ -1,6 +1,6 @@ import time import uuid -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from faker import Faker @@ -223,7 +223,7 @@ class TestWebAppAuthService: ) # Act: Execute authentication - result = WebAppAuthService.authenticate(account.email, password) + result = WebAppAuthService.authenticate(account.email, password, db_session_with_containers) # Assert: Verify successful authentication assert result is not None @@ -260,7 +260,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountNotFoundError): - WebAppAuthService.authenticate(non_existent_email, "any_password") + WebAppAuthService.authenticate(non_existent_email, "any_password", db_session_with_containers) def test_authenticate_account_banned(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -297,7 +297,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountLoginError) as exc_info: - WebAppAuthService.authenticate(account.email, password) + WebAppAuthService.authenticate(account.email, password, db_session_with_containers) assert "Account is banned." in str(exc_info.value) @@ -318,7 +318,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with wrong password with pytest.raises(AccountPasswordError) as exc_info: - WebAppAuthService.authenticate(account.email, "wrong_password") + WebAppAuthService.authenticate(account.email, "wrong_password", db_session_with_containers) assert "Invalid email or password." in str(exc_info.value) @@ -350,7 +350,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(AccountPasswordError) as exc_info: - WebAppAuthService.authenticate(account.email, "any_password") + WebAppAuthService.authenticate(account.email, "any_password", db_session_with_containers) assert "Invalid email or password." in str(exc_info.value) @@ -403,7 +403,7 @@ class TestWebAppAuthService: ) # Act: Execute user retrieval - result = WebAppAuthService.get_user_through_email(account.email) + result = WebAppAuthService.get_user_through_email(account.email, db_session_with_containers) # Assert: Verify successful retrieval assert result is not None @@ -430,7 +430,7 @@ class TestWebAppAuthService: non_existent_email = f"nonexistent_{uuid.uuid4().hex}@example.com" # Act: Execute user retrieval - result = WebAppAuthService.get_user_through_email(non_existent_email) + result = WebAppAuthService.get_user_through_email(non_existent_email, db_session_with_containers) # Assert: Verify proper handling assert result is None @@ -463,7 +463,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(Unauthorized) as exc_info: - WebAppAuthService.get_user_through_email(account.email) + WebAppAuthService.get_user_through_email(account.email, db_session_with_containers) assert "Account is banned." in str(exc_info.value) @@ -659,7 +659,7 @@ class TestWebAppAuthService: ) # Act: Execute end user creation - result = WebAppAuthService.create_end_user(site.code, "test@example.com") + result = WebAppAuthService.create_end_user(site.code, "test@example.com", db_session_with_containers) # Assert: Verify successful creation assert result is not None @@ -694,7 +694,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(NotFound) as exc_info: - WebAppAuthService.create_end_user(non_existent_code, "test@example.com") + WebAppAuthService.create_end_user(non_existent_code, "test@example.com", db_session_with_containers) assert "Site not found." in str(exc_info.value) @@ -732,7 +732,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(NotFound) as exc_info: - WebAppAuthService.create_end_user(site.code, "test@example.com") + WebAppAuthService.create_end_user(site.code, "test@example.com", db_session_with_containers) assert "App not found." in str(exc_info.value) @@ -750,7 +750,9 @@ class TestWebAppAuthService: # Arrange: Setup test with private access mode # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(access_mode="private") + result = WebAppAuthService.is_app_require_permission_check( + access_mode="private", session=db_session_with_containers + ) # Assert: Verify correct result assert result is True @@ -769,7 +771,9 @@ class TestWebAppAuthService: # Arrange: Setup test with public access mode # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(access_mode="public") + result = WebAppAuthService.is_app_require_permission_check( + access_mode="public", session=db_session_with_containers + ) # Assert: Verify correct result assert result is False @@ -789,13 +793,17 @@ class TestWebAppAuthService: mock_external_service_dependencies["app_service"].get_app_id_by_code.return_value = "mock_app_id" # Act: Execute permission check requirement test - result = WebAppAuthService.is_app_require_permission_check(app_code="mock_app_code") + result = WebAppAuthService.is_app_require_permission_check( + app_code="mock_app_code", session=db_session_with_containers + ) # Assert: Verify correct result assert result is True # Verify mock service was called correctly - mock_external_service_dependencies["app_service"].get_app_id_by_code.assert_called_once_with("mock_app_code") + mock_external_service_dependencies["app_service"].get_app_id_by_code.assert_called_once_with( + "mock_app_code", session=ANY + ) mock_external_service_dependencies[ "enterprise_service" ].WebAppAuth.get_app_access_mode_by_id.assert_called_once_with("mock_app_id") @@ -814,7 +822,7 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: - WebAppAuthService.is_app_require_permission_check() + WebAppAuthService.is_app_require_permission_check(session=db_session_with_containers) assert "Either app_code or app_id must be provided." in str(exc_info.value) @@ -832,7 +840,7 @@ class TestWebAppAuthService: # Arrange: Setup test with public access mode # Act: Execute authentication type determination - result = WebAppAuthService.get_app_auth_type(access_mode="public") + result = WebAppAuthService.get_app_auth_type(access_mode="public", session=db_session_with_containers) # Assert: Verify correct result assert result == WebAppAuthType.PUBLIC @@ -851,7 +859,7 @@ class TestWebAppAuthService: # Arrange: Setup test with private access mode # Act: Execute authentication type determination - result = WebAppAuthService.get_app_auth_type(access_mode="private") + result = WebAppAuthService.get_app_auth_type(access_mode="private", session=db_session_with_containers) # Assert: Verify correct result assert result == WebAppAuthType.INTERNAL @@ -875,7 +883,9 @@ class TestWebAppAuthService: ].WebAppAuth.get_app_access_mode_by_id.return_value = setting # Act: Execute authentication type determination - result: WebAppAuthType = WebAppAuthService.get_app_auth_type(app_code="mock_app_code") + result: WebAppAuthType = WebAppAuthService.get_app_auth_type( + app_code="mock_app_code", session=db_session_with_containers + ) # Assert: Verify correct result assert result == WebAppAuthType.EXTERNAL @@ -899,6 +909,6 @@ class TestWebAppAuthService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError) as exc_info: - WebAppAuthService.get_app_auth_type() + WebAppAuthService.get_app_auth_type(session=db_session_with_containers) assert "Either app_code or access_mode must be provided." in str(exc_info.value) diff --git a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py index c699d39dde1..902134e053d 100644 --- a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py +++ b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py @@ -350,9 +350,10 @@ class TestWebhookServiceTriggerExecutionWithContainers: quota_charge.commit.assert_called_once() mock_trigger.assert_called_once() trigger_args = mock_trigger.call_args.args - assert trigger_args[1] is end_user - assert trigger_args[2].workflow_id == workflow.id - assert trigger_args[2].root_node_id == webhook_trigger.node_id + assert trigger_args[0] is end_user + assert trigger_args[1].workflow_id == workflow.id + assert trigger_args[1].root_node_id == webhook_trigger.node_id + assert mock_trigger.call_args.kwargs["session"] is not None def test_trigger_workflow_execution_marks_tenant_rate_limited_when_quota_exceeded( self, db_session_with_containers: Session, flask_app_with_containers: Flask diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py index cf76afb303c..f553b0f72a0 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py @@ -99,7 +99,7 @@ class TestWorkflowAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -164,7 +164,7 @@ class TestWorkflowAppService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py index 726c360d77e..7c528f06b10 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_run_service.py @@ -92,7 +92,7 @@ class TestWorkflowRunService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) return app, account @@ -544,7 +544,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow run without node executions workflow_run = self._create_test_workflow_run(db_session_with_containers, app, account, "debugging") @@ -596,7 +596,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Use invalid workflow run ID invalid_workflow_run_id = str(uuid.uuid4()) @@ -648,7 +648,7 @@ class TestWorkflowRunService: icon="🚀", icon_background="#4ECDC4", ) - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow run workflow_run = self._create_test_workflow_run(db_session_with_containers, app, account, "debugging") diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_service.py index 349aac1be36..6531ed4fbb0 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_service.py @@ -227,7 +227,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=db_session_with_containers) # Assert assert result is True @@ -247,7 +247,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=db_session_with_containers) # Assert assert result is False @@ -269,7 +269,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=db_session_with_containers) # Assert assert result is not None @@ -293,7 +293,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=db_session_with_containers) # Assert assert result is None @@ -320,7 +320,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow_by_id(app, workflow.id) + result = workflow_service.get_published_workflow_by_id(app, workflow.id, session=db_session_with_containers) # Assert assert result is not None @@ -349,7 +349,7 @@ class TestWorkflowService: from services.errors.app import IsDraftWorkflowError with pytest.raises(IsDraftWorkflowError): - workflow_service.get_published_workflow_by_id(app, workflow.id) + workflow_service.get_published_workflow_by_id(app, workflow.id, session=db_session_with_containers) def test_get_published_workflow_by_id_not_found(self, db_session_with_containers: Session): """ @@ -366,7 +366,9 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow_by_id(app, non_existent_workflow_id) + result = workflow_service.get_published_workflow_by_id( + app, non_existent_workflow_id, session=db_session_with_containers + ) # Assert assert result is None @@ -393,7 +395,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=db_session_with_containers) # Assert assert result is not None @@ -416,7 +418,7 @@ class TestWorkflowService: workflow_service = WorkflowService() # Act - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=db_session_with_containers) # Assert assert result is None @@ -714,6 +716,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) # Assert @@ -778,6 +781,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) # Assert @@ -838,6 +842,7 @@ class TestWorkflowService: account=account, environment_variables=environment_variables, conversation_variables=conversation_variables, + session=db_session_with_containers, ) def test_publish_workflow_success(self, db_session_with_containers: Session): @@ -979,9 +984,7 @@ class TestWorkflowService: workflow_service = WorkflowService() restored_workflow = workflow_service.restore_published_workflow_to_draft( - app_model=app, - workflow_id=published_workflow.id, - account=account, + app_model=app, workflow_id=published_workflow.id, account=account, session=db_session_with_containers ) db_session_with_containers.expire_all() @@ -1130,7 +1133,9 @@ class TestWorkflowService: } # Act - result = workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + result = workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) # Assert assert result is not None @@ -1190,7 +1195,9 @@ class TestWorkflowService: } # Act - result = workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + result = workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) # Assert assert result is not None @@ -1222,7 +1229,9 @@ class TestWorkflowService: # Act & Assert with pytest.raises(ValueError, match="Current App mode: workflow is not supported convert to workflow"): - workflow_service.convert_to_workflow(app_model=app, account=account, args=conversion_args) + workflow_service.convert_to_workflow( + app_model=app, account=account, args=conversion_args, session=db_session_with_containers + ) def test_validate_features_structure_advanced_chat(self, db_session_with_containers: Session): """ diff --git a/api/tests/test_containers_integration_tests/services/test_workspace_service.py b/api/tests/test_containers_integration_tests/services/test_workspace_service.py index 4e89d906f16..d7cbfa91ab7 100644 --- a/api/tests/test_containers_integration_tests/services/test_workspace_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workspace_service.py @@ -104,7 +104,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -151,7 +151,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -206,7 +206,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -261,7 +261,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -291,7 +291,7 @@ class TestWorkspaceService: # Arrange: No test data needed for this test # Act: Execute the method under test with None tenant - result = WorkspaceService.get_tenant_info(None) + result = WorkspaceService.get_tenant_info(None, db_session_with_containers) # Assert: Verify the expected outcomes assert result is None @@ -341,7 +341,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -398,7 +398,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -448,7 +448,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -513,7 +513,7 @@ class TestWorkspaceService: # Mock current_user for flask_login with patch("services.workspace_service.current_user", account): # Act: Execute the method under test - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -553,7 +553,7 @@ class TestWorkspaceService: # No TenantAccountJoin created with patch("services.workspace_service.current_user", account): with pytest.raises(AssertionError, match="TenantAccountJoin not found"): - WorkspaceService.get_tenant_info(tenant) + WorkspaceService.get_tenant_info(tenant, db_session_with_containers) def test_get_tenant_info_should_set_replace_webapp_logo_to_none_when_flag_absent( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -572,7 +572,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = True with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["custom_config"]["replace_webapp_logo"] is None @@ -596,7 +596,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = True with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["custom_config"]["replace_webapp_logo"].startswith(custom_base) @@ -615,7 +615,7 @@ class TestWorkspaceService: mock_external_service_dependencies["tenant_service"].has_roles.return_value = False with patch("services.workspace_service.current_user", account): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert "next_credit_reset_date" not in result @@ -642,7 +642,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=None), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["next_credit_reset_date"] == "2025-02-01" @@ -669,7 +669,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=paid_pool), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 1000 @@ -697,7 +697,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, None]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == -1 @@ -726,7 +726,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 100 @@ -754,7 +754,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[None, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 50 @@ -785,7 +785,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[paid_pool, trial_pool]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert result["trial_credits"] == 200 @@ -811,7 +811,7 @@ class TestWorkspaceService: patch("services.workspace_service.current_user", account), patch("services.credit_pool_service.CreditPoolService.get_pool", side_effect=[None, None]), ): - result = WorkspaceService.get_tenant_info(tenant) + result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None assert "trial_credits" not in result diff --git a/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py b/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py index 6f342e63dc8..b12472c586c 100644 --- a/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py +++ b/api/tests/test_containers_integration_tests/services/tools/test_workflow_tools_manage_service.py @@ -107,7 +107,7 @@ class TestWorkflowToolManageService: ) app_service = AppService() - app = app_service.create_app(tenant.id, app_args, account) + app = app_service.create_app(tenant.id, app_args, account, session=db_session_with_containers) # Create workflow for the app workflow = WorkflowModel( diff --git a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py index ce5c2bd162f..8cd9526f6a0 100644 --- a/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py +++ b/api/tests/test_containers_integration_tests/services/workflow/test_workflow_converter.py @@ -217,6 +217,7 @@ class TestWorkflowConverter: icon_type="emoji", icon="🚀", icon_background="#4CAF50", + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -291,6 +292,7 @@ class TestWorkflowConverter: icon_type="emoji", icon="🚀", icon_background="#4CAF50", + session=db_session_with_containers, ) # Verify database state remains unchanged @@ -325,6 +327,7 @@ class TestWorkflowConverter: app_model=app, app_model_config=app.app_model_config, account_id=account.id, + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -467,6 +470,7 @@ class TestWorkflowConverter: app_model=app, variables=variables, external_data_variables=external_data_variables, + session=db_session_with_containers, ) # Assert: Verify the expected outcomes @@ -569,7 +573,7 @@ class TestConvertToHttpRequestNodeVariants: """Tests for chatbot vs workflow differences in HTTP request node conversion.""" @staticmethod - def _setup(app_mode, default_variables): + def _setup(app_mode, default_variables, db_session_with_containers: Session): app_model = App( tenant_id="tenant_id", mode=app_mode, @@ -598,19 +602,20 @@ class TestConvertToHttpRequestNodeVariants: app_model=app_model, variables=default_variables, external_data_variables=ext_vars, + session=db_session_with_containers, ) return nodes - def test_chatbot_query_uses_sys_query(self, default_variables): - nodes = self._setup(AppMode.CHAT, default_variables) + def test_chatbot_query_uses_sys_query(self, default_variables, db_session_with_containers: Session): + nodes = self._setup(AppMode.CHAT, default_variables, db_session_with_containers) body = json.loads(nodes[0]["data"]["body"]["data"]) assert body["params"]["query"] == "{{#sys.query#}}" assert body["point"] == APIBasedExtensionPoint.APP_EXTERNAL_DATA_TOOL_QUERY assert nodes[1]["data"]["type"] == "code" - def test_workflow_query_is_empty(self, default_variables): - nodes = self._setup(AppMode.WORKFLOW, default_variables) + def test_workflow_query_is_empty(self, default_variables, db_session_with_containers: Session): + nodes = self._setup(AppMode.WORKFLOW, default_variables, db_session_with_containers) body = json.loads(nodes[0]["data"]["body"]["data"]) assert body["params"]["query"] == "" diff --git a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py index 9c20118e278..b6865510adf 100644 --- a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py +++ b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py @@ -194,7 +194,7 @@ def test_webhook_trigger_creates_trigger_log( db_session_with_containers.add_all([webhook_trigger, app_trigger]) db_session_with_containers.commit() - def _fake_trigger_workflow_async(session: Session, user: Any, trigger_data: Any) -> SimpleNamespace: + def _fake_trigger_workflow_async(user: Any, trigger_data: Any, *, session: Session) -> SimpleNamespace: log = WorkflowTriggerLog( tenant_id=trigger_data.tenant_id, app_id=trigger_data.app_id, @@ -575,7 +575,7 @@ def test_schedule_trigger_creates_trigger_log( db_session_with_containers.commit() # Mock AsyncWorkflowService to create WorkflowTriggerLog - def _fake_trigger_workflow_async(session: Session, user: Any, trigger_data: Any) -> SimpleNamespace: + def _fake_trigger_workflow_async(user: Any, trigger_data: Any, *, session: Session) -> SimpleNamespace: log = WorkflowTriggerLog( tenant_id=trigger_data.tenant_id, app_id=trigger_data.app_id, diff --git a/api/tests/unit_tests/commands/test_data_migration_commands.py b/api/tests/unit_tests/commands/test_data_migration_commands.py index b7f92f3291a..84e39d19a2d 100644 --- a/api/tests/unit_tests/commands/test_data_migration_commands.py +++ b/api/tests/unit_tests/commands/test_data_migration_commands.py @@ -109,8 +109,8 @@ def test_export_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): package = MigrationPackage.from_mapping({"metadata": {"version": "1", "source_scope": "single"}}) class FakeMigrationExportService: - def export(self, export_session, selection): - captured["session"] = export_session + def export(self, selection, *, session): + captured["session"] = session captured["selection"] = selection return ExportResult(package=package, report_items=[], report_context=ReportContext()) @@ -156,8 +156,8 @@ def test_import_command_uses_cli_owned_session(monkeypatch, tmp_path: Path): ) class FakeMigrationImportService: - def import_package(self, import_session, request): - captured["session"] = import_session + def import_package(self, request, *, session): + captured["session"] = session captured["request"] = request return ImportResult(report_items=[], report_context=ReportContext(target_tenant="target")) diff --git a/api/tests/unit_tests/controllers/common/test_app_access.py b/api/tests/unit_tests/controllers/common/test_app_access.py index d070cc6e0fc..60a576346a0 100644 --- a/api/tests/unit_tests/controllers/common/test_app_access.py +++ b/api/tests/unit_tests/controllers/common/test_app_access.py @@ -152,7 +152,7 @@ class TestResolveAppAccessFilter: self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[])) monkeypatch.setattr( f"{_RBAC_MODULE}.RBACService.MyPermissions.get", - lambda tenant_id, account_id: _permissions(workspace_keys=["app.create_and_management"]), + lambda tenant_id, account_id, session: _permissions(workspace_keys=["app.create_and_management"]), ) flt = resolve_app_access_filter("tenant-1", "acc-1") diff --git a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py index 2e5851349d4..8f294293c07 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py @@ -236,7 +236,7 @@ def test_agent_app_list_and_create_use_agent_route( items=[_app_detail_obj(id="app-list", bound_agent_id="agent-list")], ) - def create_app(self, tenant_id: str, params, current_user: object) -> object: + def create_app(self, tenant_id: str, params, current_user: object, *, session: object) -> object: captured["create"] = {"tenant_id": tenant_id, "params": params, "current_user": current_user} return _app_detail_obj(id="app-created", bound_agent_id="agent-created") @@ -392,7 +392,8 @@ def test_agent_app_create_omits_optional_role_as_empty_string( captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params: object, account: object) -> object: + def create_app(self, tenant_id: str, params: object, account: object, *, session: object) -> object: + del session captured["create"] = {"tenant_id": tenant_id, "params": params, "account": account} return _app_detail_obj(id="app-created", bound_agent_id="agent-created") @@ -472,11 +473,11 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( captured["get_app"] = app_obj return app_obj - def update_app(self, app_obj: object, args: dict[str, object]) -> object: + def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: captured["update"] = {"app": app_obj, "args": args} return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) - def delete_app(self, app_obj: object) -> None: + def delete_app(self, app_obj: object, *, session: object) -> None: captured["delete"] = app_obj monkeypatch.setattr(roster_controller, "AppService", FakeAppService) @@ -661,18 +662,26 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( discard_agent_app_build_draft, ) + def assert_call_without_session(key: str, expected: dict[str, object]) -> None: + call = dict(captured[key]) # type: ignore[arg-type] + assert call.pop("session", None) is not None + assert call == expected + with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/publish", json={"version_note": "publish v1"}, ): published = unwrap(AgentPublishApi.post)(AgentPublishApi(), "tenant-1", current_user, agent_id) assert published["active_config_snapshot_id"] == "version-1" - assert captured["publish"] == { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "version_note": "publish v1", - } + assert_call_without_session( + "publish", + { + "tenant_id": "tenant-1", + "agent_id": agent_id, + "account_id": account_id, + "version_note": "publish v1", + }, + ) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/checkout", @@ -682,17 +691,20 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( AgentBuildDraftCheckoutApi(), "tenant-1", current_user, agent_id ) assert checked_out["draft"]["id"] == "build-draft-1" - assert captured["checkout"] == { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "force": True, - } + assert_call_without_session( + "checkout", + { + "tenant_id": "tenant-1", + "agent_id": agent_id, + "account_id": account_id, + "force": True, + }, + ) with app.test_request_context("/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft"): loaded = unwrap(AgentBuildDraftApi.get)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) assert loaded["draft"]["id"] == "build-draft-1" - assert captured["load"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("load", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", @@ -711,7 +723,7 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( ): applied = unwrap(AgentBuildDraftApplyApi.post)(AgentBuildDraftApplyApi(), "tenant-1", current_user, agent_id) assert applied == {"result": "success", "draft": {"id": "draft-1"}} - assert captured["apply"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("apply", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", @@ -719,7 +731,7 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( ): discarded = unwrap(AgentBuildDraftApi.delete)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) assert discarded == {"result": "success"} - assert captured["discard"] == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} + assert_call_without_session("discard", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata( @@ -775,7 +787,7 @@ def test_agent_api_status_and_key_routes_resolve_backing_app( monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda app_id: 1) class FakeAppService: - def update_app_api_status(self, app_obj: object, enable_api: bool) -> object: + def update_app_api_status(self, app_obj: object, enable_api: bool, *, session: object) -> object: captured["enable"] = {"app": app_obj, "enable_api": enable_api} app_model.enable_api = enable_api return app_model @@ -890,7 +902,7 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk def get_app(self, app_obj: object) -> object: return app_obj - def update_app(self, app_obj: object, args: dict[str, object]) -> object: + def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: captured["update"] = {"app": app_obj, "args": args} return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) @@ -1292,6 +1304,7 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk ) assert result["binding"]["binding_type"] == "inline_agent" + assert captured.pop("session") is not None assert captured == { "tenant_id": "tenant-1", "app_id": "app-1", @@ -1896,8 +1909,18 @@ def test_list_agent_chat_messages_uses_current_user_conversation( captured.update(kwargs) return conversation + class SessionProxy: + def __call__(self): + return session + + def scalar(self, stmt: object): + return session.scalar(stmt) + + def scalars(self, stmt: object): + return session.scalars(stmt) + monkeypatch.setattr(message_controller.ConversationService, "get_conversation", get_conversation) - monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=session)) + monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=SessionProxy())) monkeypatch.setattr(message_controller, "attach_message_extra_contents", lambda messages: None) monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) @@ -1905,6 +1928,7 @@ def test_list_agent_chat_messages_uses_current_user_conversation( result = message_controller._list_chat_messages(app_model=app_model, current_user=current_user) assert result == {"data": [message_id], "limit": 20, "has_more": False} + assert captured.pop("session") is session assert captured == {"app_model": app_model, "conversation_id": conversation_id, "user": current_user} diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py index 8086f578956..0ab8814f368 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py @@ -48,6 +48,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> SandboxListResponse: self.calls.append(("list", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return SandboxListResponse(path=path, entries=[], truncated=False) @@ -61,6 +62,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> SandboxReadResponse: self.calls.append(("read", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return SandboxReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") @@ -74,6 +76,7 @@ class _WorkflowService: node_id: str, node_execution_id: str | None, path: str, + session, ) -> AgentSandboxUploadDownload: self.calls.append(("upload", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) return AgentSandboxUploadDownload(url="https://files.example/upload.txt") diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py index 8a6094b94b8..cc95f7f8a94 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py @@ -2,7 +2,7 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask @@ -142,7 +142,7 @@ class TestConsoleAnnotationRefBoundaries: assert response == "" assert status == 204 - delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"]) + delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], session=ANY) def test_update_uses_annotation_ref(self, app: Flask): api = annotation_module.AnnotationUpdateDeleteApi() @@ -216,4 +216,4 @@ class TestConsoleAnnotationRefBoundaries: response = handler(api, "app-1", "ann-1") assert response["total"] == 1 - hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5) + hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5, session=ANY) diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_security.py b/api/tests/unit_tests/controllers/console/app/test_annotation_security.py index bfa4048191f..6a22d8769bc 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_security.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_security.py @@ -193,9 +193,7 @@ class TestAnnotationImportServiceValidation: @pytest.fixture def mock_db_session(self): - """Mock database session.""" - with patch("services.annotation_service.db.session") as mock: - yield mock + return MagicMock() def test_max_records_limit_enforced(self, mock_app, mock_db_session): """Test that files with too many records are rejected.""" @@ -214,7 +212,7 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.FeatureService") as mock_features: mock_features.get_features.return_value.billing.enabled = False - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) # Should return error about too many records assert "error_msg" in result @@ -231,7 +229,7 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.current_account_with_tenant") as mock_auth: mock_auth.return_value = (MagicMock(id="user_id"), "tenant_id") - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) # Should return error about insufficient records assert "error_msg" in result @@ -250,7 +248,7 @@ class TestAnnotationImportServiceValidation: ): mock_auth.return_value = (MagicMock(id="user_id"), "tenant_id") - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) assert "error_msg" in result assert "malformed" in result["error_msg"].lower() @@ -271,7 +269,9 @@ class TestAnnotationImportServiceValidation: with patch("services.annotation_service.batch_import_annotations_task") as mock_task: with patch("services.annotation_service.redis_client"): - result = AppAnnotationService.batch_import_app_annotations("app_id", file) + result = AppAnnotationService.batch_import_app_annotations( + "app_id", file, session=mock_db_session + ) # Should return success response assert "job_id" in result diff --git a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py index e7784f8fd94..79f109fd383 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py @@ -6,7 +6,7 @@ from datetime import datetime from importlib import util from pathlib import Path from types import ModuleType, SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -500,7 +500,8 @@ def test_app_list_uses_injected_session_for_draft_workflows( ) session = MagicMock() session.execute.return_value.scalars.return_value.all.return_value = [workflow] - scoped_session = SimpleNamespace(execute=MagicMock(side_effect=AssertionError("db.session should not be used"))) + scoped_session = MagicMock() + scoped_session.execute.side_effect = AssertionError("db.session should not be used") monkeypatch.setattr( app_module, @@ -515,7 +516,7 @@ def test_app_list_uses_injected_session_for_draft_workflows( monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( @@ -563,12 +564,12 @@ def test_app_create_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module, "AppService", - lambda: SimpleNamespace(create_app=lambda tenant_id, params, user: app_obj), + lambda: SimpleNamespace(create_app=lambda tenant_id, params, user, session: app_obj), ) monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppPermissions, "batch_get", - lambda tenant_id, account_id, app_ids: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, + lambda tenant_id, account_id, app_ids, session: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, ) resp, status = method(app_module.AppListApi(), "tenant-1", SimpleNamespace(id="acct-1")) @@ -611,7 +612,7 @@ def test_app_list_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( default_permission_keys=["app.preview", "app.acl.view_layout"], overrides=[ @@ -655,7 +656,7 @@ def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permis monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( workspace=app_module.enterprise_rbac_service.WorkspacePermissionSnapshot( permission_keys=["app.create_and_management"] ) @@ -698,7 +699,7 @@ def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission( monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse( + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( @@ -754,7 +755,7 @@ def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permiss monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.MyPermissions, "get", - lambda tenant_id, account_id: app_module.enterprise_rbac_service.MyPermissionsResponse(), + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse(), ) monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppAccess, @@ -820,7 +821,7 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): resp = method(app_module.AppApi(), "tenant-1", SimpleNamespace(id="acct-1"), app_model=app_obj) - get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1") + get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=ANY) assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit", "app.acl.monitor"] @@ -861,7 +862,7 @@ def test_app_copy_api_attaches_permission_keys(app, app_module): "get_system_features", lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=object(), session=lambda: MagicMock())) monkeypatch.setattr( app_module, "Session", @@ -870,7 +871,7 @@ def test_app_copy_api_attaches_permission_keys(app, app_module): monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppPermissions, "batch_get", - lambda tenant_id, account_id, app_ids: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, + lambda tenant_id, account_id, app_ids, session: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, ) resp, status = method( diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow.py b/api/tests/unit_tests/controllers/console/app/test_workflow.py index 2f971eaf74f..2d811deb916 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -621,7 +621,7 @@ def test_workflow_online_users_filters_inaccessible_workflow(app: Flask, monkeyp monkeypatch.setattr( workflow_module, "WorkflowService", - lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id: {app_id_1}), + lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id, session: {app_id_1}), ) monkeypatch.setattr(workflow_module.file_helpers, "get_signed_file_url", sign_avatar) @@ -703,7 +703,7 @@ def test_workflow_online_users_batches_redis_reads(app: Flask, monkeypatch: pyte monkeypatch.setattr( workflow_module, "WorkflowService", - lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id: set(app_ids)), + lambda: SimpleNamespace(get_accessible_app_ids=lambda app_ids, tenant_id, session: set(app_ids)), ) first_pipeline = Mock() diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py index f04ab6d6e7c..956706eafb6 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py @@ -2,7 +2,7 @@ from __future__ import annotations from dataclasses import dataclass from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -94,6 +94,7 @@ def test_human_input_preview_delegates_to_service( account=account, node_id="node-42", inputs={"topic": "tech"}, + session=ANY, ) @@ -144,6 +145,7 @@ def test_human_input_submit_forwards_payload(app: Flask, monkeypatch: pytest.Mon form_inputs={"answer": "42"}, inputs={"#node-1.result#": "LLM output"}, action="approve", + session=ANY, ) @@ -193,6 +195,7 @@ def test_human_input_delivery_test_calls_service( node_id="node-7", delivery_method_id="delivery-123", inputs={}, + session=ANY, ) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py index e66ae5246bc..dfe35a89f57 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py @@ -25,7 +25,7 @@ from __future__ import annotations import json from collections.abc import Iterator from typing import Any -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock from uuid import UUID import pytest @@ -382,7 +382,9 @@ def test_serve_snapshot_happy_path(patch_service, app_model, run_id): result = ctrl._serve_snapshot(app_model, run_id) assert isinstance(result, dict) assert result["workflow_run_id"] == "00000000-0000-0000-0000-0000000000aa" - patch_service.snapshot_workflow_run.assert_called_once_with(app_model=app_model, workflow_run_id=str(run_id)) + patch_service.snapshot_workflow_run.assert_called_once_with( + app_model=app_model, workflow_run_id=str(run_id), session=ANY + ) def test_serve_snapshot_translates_inspector_error_to_404(patch_service, app_model, run_id): @@ -399,7 +401,7 @@ def test_serve_node_detail_happy_path(patch_service, app_model, run_id): result = ctrl._serve_node_detail(app_model, run_id, "agent-1") assert result["node_id"] == "agent-1" patch_service.node_detail.assert_called_once_with( - app_model=app_model, workflow_run_id=str(run_id), node_id="agent-1" + app_model=app_model, workflow_run_id=str(run_id), node_id="agent-1", session=ANY ) @@ -431,6 +433,7 @@ def test_serve_output_preview_happy_path(patch_service, app_model, run_id): workflow_run_id=str(run_id), node_id="agent-1", output_name="text", + session=ANY, ) diff --git a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py index ebae7de6c15..001ca0bf8fb 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py +++ b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py @@ -597,7 +597,7 @@ class TestActivateApi: assert response["result"] == "success" mock_create_tenant_member.assert_called_once_with( - mock_invitation["tenant"], mock_account, mock_db.session, role=TenantAccountRole.ADMIN + mock_invitation["tenant"], mock_account, mock_db.session(), role=TenantAccountRole.ADMIN ) mock_switch_tenant.assert_called_once_with(mock_account, mock_invitation["tenant"].id, session=ANY) mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") diff --git a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py index 21d1932f820..b231826aeac 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py +++ b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py @@ -43,7 +43,7 @@ def test_list_data_source_auth_uses_injected_tenant_id() -> None: ): result = method(api, "tenant-1") - get_provider_auth_list.assert_called_once_with(ANY, "tenant-1") + get_provider_auth_list.assert_called_once_with("tenant-1", session=ANY) assert result["sources"][0]["id"] == "binding-1" assert result["sources"][0]["provider"] == "custom" @@ -65,7 +65,7 @@ def test_create_data_source_auth_binding_uses_injected_tenant_id() -> None: ): result, status = method(api, "tenant-1") - create_auth.assert_called_once_with(ANY, "tenant-1", payload) + create_auth.assert_called_once_with("tenant-1", payload, session=ANY) assert result == {"result": "success"} assert status == 200 @@ -82,6 +82,6 @@ def test_delete_data_source_auth_binding_uses_injected_tenant_id() -> None: ): result, status = method(api, "tenant-1", "binding-1") - delete_provider_auth.assert_called_once_with(ANY, "tenant-1", "binding-1") + delete_provider_auth.assert_called_once_with("tenant-1", "binding-1", session=ANY) assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py index 33aaa19e640..8f66ca5c993 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py @@ -1,6 +1,6 @@ import inspect from datetime import UTC, datetime -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -508,6 +508,7 @@ class TestDatasourceAuthDeleteApi: auth_id="cred-1", provider="notion", plugin_id="langgenius/notion_datasource", + session=ANY, ) def test_delete_missing_credential_id(self, app: Flask): diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index 2a1970d3837..39a6fa65a06 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -64,10 +64,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates( - session: Mock, template_type: str, language: str, current_tenant_id: str - ) -> dict[str, object]: - service_calls.append((template_type, language, current_tenant_id)) + def get_pipeline_templates(*, type: str, language: str, current_tenant_id: str, session) -> dict[str, object]: + del session + service_calls.append((type, language, current_tenant_id)) return {"pipeline_templates": [_template_item()]} with ( @@ -94,10 +93,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates( - session: Mock, template_type: str, language: str, current_tenant_id: str - ) -> dict[str, object]: - service_calls.append((template_type, language, current_tenant_id)) + def get_pipeline_templates(*, type: str, language: str, current_tenant_id: str, session) -> dict[str, object]: + del session + service_calls.append((type, language, current_tenant_id)) return {"pipeline_templates": []} with ( @@ -117,16 +115,18 @@ class TestPipelineTemplateDetailApi: method = unwrap(api.get) service_calls: list[tuple[str, str]] = [] - class Service: - def get_pipeline_template_detail( - self, session: Mock, template_id: str, template_type: str - ) -> dict[str, object]: - service_calls.append((template_id, template_type)) - return _template_detail() + def get_pipeline_template_detail(template_id: str, type: str, *, session) -> dict[str, object]: + del session + service_calls.append((template_id, type)) + return _template_detail() with ( app.test_request_context("/rag/pipeline/templates/template-1?type=customized"), - patch.object(module, "RagPipelineService", Service), + patch.object( + module.RagPipelineService, + "get_pipeline_template_detail", + side_effect=get_pipeline_template_detail, + ), ): response, status = method(api, Mock(), "template-1") @@ -138,13 +138,16 @@ class TestPipelineTemplateDetailApi: api = PipelineTemplateDetailApi() method = unwrap(api.get) - class Service: - def get_pipeline_template_detail(self, session: Mock, template_id: str, template_type: str) -> None: - return None + def get_pipeline_template_detail(template_id: str, type: str, *, session) -> None: + del template_id, type, session with ( app.test_request_context("/rag/pipeline/templates/missing"), - patch.object(module, "RagPipelineService", Service), + patch.object( + module.RagPipelineService, + "get_pipeline_template_detail", + side_effect=get_pipeline_template_detail, + ), ): with pytest.raises(NotFound): method(api, Mock(), "missing") @@ -160,8 +163,14 @@ class TestCustomizedPipelineTemplateApi: service_calls: list[tuple[str, PipelineTemplateInfoEntity, Account, str]] = [] def update_template( - template_id: str, template_info: PipelineTemplateInfoEntity, current_user: Account, current_tenant_id: str + template_id: str, + template_info: PipelineTemplateInfoEntity, + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((template_id, template_info, current_user, current_tenant_id)) with ( @@ -198,8 +207,14 @@ class TestCustomizedPipelineTemplateApi: service_calls: list[tuple[str, PipelineTemplateInfoEntity, Account, str]] = [] def update_template( - template_id: str, template_info: PipelineTemplateInfoEntity, current_user: Account, current_tenant_id: str + template_id: str, + template_info: PipelineTemplateInfoEntity, + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((template_id, template_info, current_user, current_tenant_id)) with ( @@ -228,7 +243,8 @@ class TestCustomizedPipelineTemplateApi: tenant_id = "tenant-1" deleted_templates: list[tuple[str, str]] = [] - def delete_template(template_id: str, current_tenant_id: str) -> None: + def delete_template(template_id: str, current_tenant_id: str, *, session) -> None: + del session deleted_templates.append((template_id, current_tenant_id)) with ( @@ -325,9 +341,19 @@ class TestPublishCustomizedPipelineTemplateApi: service_calls: list[tuple[str, dict[str, object], Account, str]] = [] class Service: + def __init__(self, *args, **kwargs) -> None: + pass + def publish_customized_pipeline_template( - self, pipeline_id: str, data: dict[str, object], current_user: Account, current_tenant_id: str + self, + pipeline_id: str, + data: dict[str, object], + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((pipeline_id, data, current_user, current_tenant_id)) with ( @@ -352,9 +378,19 @@ class TestPublishCustomizedPipelineTemplateApi: service_calls: list[tuple[str, dict[str, object], Account, str]] = [] class Service: + def __init__(self, *args, **kwargs) -> None: + pass + def publish_customized_pipeline_template( - self, pipeline_id: str, data: dict[str, object], current_user: Account, current_tenant_id: str + self, + pipeline_id: str, + data: dict[str, object], + current_user: Account, + current_tenant_id: str, + *, + session, ) -> None: + del session service_calls.append((pipeline_id, data, current_user, current_tenant_id)) with ( diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py index 19dc90ed8a4..e344a4c8bab 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py @@ -52,7 +52,9 @@ def _pipeline() -> Pipeline: def test_draft_rag_pipeline_workflow_get_serializes_response_model(monkeypatch: pytest.MonkeyPatch) -> None: workflow = _make_workflow() monkeypatch.setattr( - module, "RagPipelineService", lambda: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow) + module, + "RagPipelineService", + lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow), ) api = module.DraftRagPipelineApi() @@ -97,12 +99,12 @@ def test_published_rag_pipeline_workflows_serialize_items_before_session_closes( assert session_state["open"] is True return getattr(base_workflow, name) - monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object())) monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker()) monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)), + lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False)), ) with app.test_request_context( @@ -132,12 +134,12 @@ def test_rag_pipeline_workflow_patch_serializes_response_model(app: Flask, monke def begin(self): return _SessionContext() - monkeypatch.setattr(module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(module, "db", SimpleNamespace(engine=object(), session=lambda: object())) monkeypatch.setattr(module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker()) monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(update_workflow=lambda **_kwargs: workflow), + lambda *_args, **_kwargs: SimpleNamespace(update_workflow=lambda **_kwargs: workflow), ) payload: dict[str, object] = {"marked_name": "Updated release"} @@ -165,7 +167,7 @@ def test_default_rag_pipeline_block_configs_serializes_root_response(monkeypatch monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_default_block_configs=lambda: block_configs), + lambda *_args, **_kwargs: SimpleNamespace(get_default_block_configs=lambda: block_configs), ) api = module.DefaultRagPipelineBlockConfigsApi() @@ -190,7 +192,7 @@ def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, mon monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables), + lambda *_args, **_kwargs: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables), ) api = module.DraftRagPipelineSecondStepApi() @@ -210,7 +212,7 @@ def test_rag_pipeline_recommended_plugins_serializes_known_envelope(app, monkeyp monkeypatch.setattr( module, "RagPipelineService", - lambda: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins), + lambda *_args, **_kwargs: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins), ) api = module.RagPipelineRecommendedPluginApi() diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py index 53f4f139937..6913825d599 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py @@ -3,7 +3,7 @@ import json from contextlib import ExitStack from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock, PropertyMock, patch +from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask @@ -63,6 +63,18 @@ def dataset_model_property_defaults(): for name, value in properties.items(): property_mock = stack.enter_context(patch.object(Dataset, name, new_callable=PropertyMock)) property_mock.return_value = value + stack.enter_context( + patch( + "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.MyPermissions.get", + return_value=enterprise_rbac_service.MyPermissionsResponse(), + ) + ) + stack.enter_context( + patch( + "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetPermissions.batch_get", + return_value={}, + ) + ) yield @@ -245,7 +257,7 @@ class TestDatasetList: ): resp, status = method(api, "tenant-1", current_user) - get_permissions.assert_called_once_with("tenant-1", current_user.id) + get_permissions.assert_called_once_with("tenant-1", current_user.id, session=ANY) assert status == 200 assert resp["data"][0]["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"] @@ -742,7 +754,7 @@ class TestDatasetApiGet: data, status = method(api, tenant_id, user, dataset_id) - get_permissions.assert_called_once_with(tenant_id, user.id, dataset_id=dataset_id) + get_permissions.assert_called_once_with(tenant_id, user.id, dataset_id=dataset_id, session=ANY) assert status == 200 assert data["permission_keys"] == ["dataset.acl.readonly", "dataset.acl.edit"] diff --git a/api/tests/unit_tests/controllers/console/datasets/test_external.py b/api/tests/unit_tests/controllers/console/datasets/test_external.py index 8ac40f03d3b..1cffc90ae23 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -1,5 +1,6 @@ import inspect -from unittest.mock import MagicMock, PropertyMock, patch +from types import SimpleNamespace +from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask @@ -7,6 +8,7 @@ from werkzeug.exceptions import Forbidden, NotFound import services from controllers.console import console_ns +from controllers.console.datasets import external as external_module from controllers.console.datasets.error import DatasetNameDuplicateError from controllers.console.datasets.external import ( BedrockRetrievalApi, @@ -142,7 +144,7 @@ class TestExternalApiUseCheckApi: assert status == 200 assert response == {"is_using": True, "count": 2} - mock_use_check.assert_called_once_with(session, "api-id", "tenant-1") + mock_use_check.assert_called_once_with("api-id", "tenant-1", session=ANY) class TestExternalDatasetCreateApi: @@ -186,6 +188,7 @@ class TestExternalDatasetCreateApi: "create_external_dataset", return_value=dataset, ), + patch.object(external_module, "db", SimpleNamespace(session=lambda: MagicMock())), ): _, status = method(api, MagicMock(), "tenant-1", current_user) diff --git a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py index 8a2e14cce9b..4adeaaa90dd 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py +++ b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py @@ -32,7 +32,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "en-US") + service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data def test_get_fallback_to_user_language(self, app: Flask): @@ -51,7 +51,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "fr-FR") + service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data def test_get_fallback_to_default_language(self, app: Flask): @@ -70,7 +70,7 @@ class TestRecommendedAppListApi: ): result = method(api, make_account(None)) - service_mock.assert_called_once_with(ANY, module.languages[0]) + service_mock.assert_called_once_with(module.languages[0], session=ANY) assert result == result_data @@ -91,7 +91,7 @@ class TestLearnDifyAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "en-US") + service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data def test_get_fallback_to_user_language(self, app: Flask): @@ -110,7 +110,7 @@ class TestLearnDifyAppListApi: ): result = method(api, make_account("fr-FR")) - service_mock.assert_called_once_with(ANY, "fr-FR") + service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data @@ -131,7 +131,7 @@ class TestRecommendedAppApi: ): result = method(api, "11111111-1111-1111-1111-111111111111") - service_mock.assert_called_once_with(ANY, "11111111-1111-1111-1111-111111111111") + service_mock.assert_called_once_with("11111111-1111-1111-1111-111111111111", session=ANY) assert result == result_data diff --git a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py index ae05b8f6a0e..f210d0d5d04 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py @@ -63,7 +63,7 @@ class TestSavedMessageListApi: result = method(api, current_user, installed_app) pagination_mock.assert_called_once() - assert pagination_mock.call_args.args[2] is current_user + assert pagination_mock.call_args.args[1] is current_user assert result["limit"] == 20 assert result["has_more"] is False assert len(result["data"]) == 2 @@ -96,7 +96,7 @@ class TestSavedMessageListApi: result = method(api, current_user, installed_app) save_mock.assert_called_once() - assert save_mock.call_args.args[2] is current_user + assert save_mock.call_args.args[1] is current_user assert result == {"result": "success"} def test_post_message_not_exists(self, app: Flask, payload_patch): @@ -136,7 +136,7 @@ class TestSavedMessageApi: result, status = method(api, current_user, installed_app, str(uuid4())) delete_mock.assert_called_once() - assert delete_mock.call_args.args[2] is current_user + assert delete_mock.call_args.args[1] is current_user assert status == 204 assert result == "" diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py index 8785ce85109..98b538800ac 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py @@ -38,7 +38,10 @@ def _snippet(**overrides) -> CustomizedSnippet: @pytest.fixture(autouse=True) def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch) -> None: def factory(): - return snippet_workflow_module.SnippetService() + try: + return snippet_workflow_module.SnippetService(snippet_workflow_module._snippet_session_maker()) + except TypeError: + return snippet_workflow_module.SnippetService() monkeypatch.setattr(snippet_workflow_module, "_snippet_service", factory) monkeypatch.setattr(snippet_workflow_module, "_snippet_session_maker", Mock(return_value=Mock())) diff --git a/api/tests/unit_tests/controllers/console/tag/test_tags.py b/api/tests/unit_tests/controllers/console/tag/test_tags.py index 2da11afa1f7..8aaebeb124a 100644 --- a/api/tests/unit_tests/controllers/console/tag/test_tags.py +++ b/api/tests/unit_tests/controllers/console/tag/test_tags.py @@ -3,7 +3,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden import controllers.console.tag.tags as module @@ -22,7 +22,7 @@ from services.tag_service import UpdateTagPayload class SessionMatcher: def __eq__(self, other): - return isinstance(other, Session | scoped_session) + return isinstance(other, Session) def unwrap(func): @@ -131,7 +131,7 @@ class TestTagListApi: ): result, status = method(api, "tenant-1") - get_tags_mock.assert_called_once_with(SessionMatcher(), "snippet", "tenant-1", None) + get_tags_mock.assert_called_once_with("snippet", "tenant-1", None, session=SessionMatcher()) assert status == 200 assert result == [{"id": "1", "name": "snippet-tag", "type": "snippet", "binding_count": "1"}] @@ -224,7 +224,7 @@ class TestTagUpdateDeleteApi: update_payload, tag_id, session = update_tags_mock.call_args.args assert update_payload == UpdateTagPayload(name="updated") assert tag_id == "tag-1" - assert session == module.db.session + assert session == SessionMatcher() assert result["binding_count"] == "3" def test_patch_forbidden(self, app: Flask, readonly_user, payload_patch): @@ -250,7 +250,7 @@ class TestTagUpdateDeleteApi: ): result, status = method(api, "tag-1") - delete_mock.assert_called_once_with("tag-1", module.db.session) + delete_mock.assert_called_once_with("tag-1", SessionMatcher()) assert status == 204 def test_delete_snippet_tag_checks_type_in_current_tenant(self, app: Flask, admin_user): @@ -278,7 +278,7 @@ class TestTagUpdateDeleteApi: scene=module.RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False, ) - delete_mock.assert_called_once_with("tag-1", module.db.session) + delete_mock.assert_called_once_with("tag-1", SessionMatcher()) assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/test_extension.py b/api/tests/unit_tests/controllers/console/test_extension.py index bab825ca6f0..8ea327dfdce 100644 --- a/api/tests/unit_tests/controllers/console/test_extension.py +++ b/api/tests/unit_tests/controllers/console/test_extension.py @@ -114,7 +114,7 @@ def test_api_based_extension_get_returns_tenant_extensions(app: Flask, monkeypat assert response[0]["name"] == "Weather API" assert response[0]["api_endpoint"] == extension.api_endpoint assert response[0]["api_key"].startswith(extension.api_key[:3]) - service_mock.assert_called_once_with(ANY, "tenant-123") + service_mock.assert_called_once_with("tenant-123", session=ANY) def test_api_based_extension_post_creates_extension(app: Flask, monkeypatch: pytest.MonkeyPatch): @@ -132,7 +132,7 @@ def test_api_based_extension_post_creates_extension(app: Flask, monkeypatch: pyt response, status = APIBasedExtensionAPI().post() args, _ = save_mock.call_args - created_extension: APIBasedExtension = args[1] + created_extension: APIBasedExtension = args[0] assert created_extension.tenant_id == "tenant-123" assert created_extension.name == payload["name"] assert created_extension.api_endpoint == payload["api_endpoint"] @@ -157,7 +157,7 @@ def test_api_based_extension_detail_get_fetches_extension(app: Flask, monkeypatc assert response["id"] == extension.id assert response["name"] == extension.name - service_mock.assert_called_once_with(ANY, "tenant-123", str(extension_id)) + service_mock.assert_called_once_with("tenant-123", str(extension_id), session=ANY) def test_api_based_extension_detail_post_keeps_hidden_api_key(app: Flask, monkeypatch: pytest.MonkeyPatch): @@ -187,7 +187,7 @@ def test_api_based_extension_detail_post_keeps_hidden_api_key(app: Flask, monkey assert existing_extension.name == payload["name"] assert existing_extension.api_endpoint == payload["api_endpoint"] assert existing_extension.api_key == "keep-me" - save_mock.assert_called_once_with(ANY, existing_extension) + save_mock.assert_called_once_with(existing_extension, session=ANY) assert response["name"] == payload["name"] assert response["api_key"] == _masked_api_key("keep-me") @@ -217,7 +217,7 @@ def test_api_based_extension_detail_post_updates_api_key_when_provided(app: Flas response = APIBasedExtensionDetailAPI().post(extension_id) assert existing_extension.api_key == "new-secret" - save_mock.assert_called_once_with(ANY, existing_extension) + save_mock.assert_called_once_with(existing_extension, session=ANY) assert response["name"] == payload["name"] assert response["api_key"] == _masked_api_key(payload["api_key"]) @@ -239,6 +239,6 @@ def test_api_based_extension_detail_delete_removes_extension(app: Flask, monkeyp ): response, status = APIBasedExtensionDetailAPI().delete(extension_id) - delete_mock.assert_called_once_with(ANY, existing_extension) + delete_mock.assert_called_once_with(existing_extension, session=ANY) assert status == 204 assert response == "" diff --git a/api/tests/unit_tests/controllers/console/test_workspace_account.py b/api/tests/unit_tests/controllers/console/test_workspace_account.py index 5f36e805baa..39a3b2485dd 100644 --- a/api/tests/unit_tests/controllers/console/test_workspace_account.py +++ b/api/tests/unit_tests/controllers/console/test_workspace_account.py @@ -692,7 +692,7 @@ def test_get_account_by_email_with_case_fallback_uses_lowercase_lookup(): second.scalar_one_or_none.return_value = expected_account mock_session.execute.side_effect = [first, second] - result = AccountService.get_account_by_email_with_case_fallback(mock_session, "Mixed@Test.com") + result = AccountService.get_account_by_email_with_case_fallback("Mixed@Test.com", session=mock_session) assert result is expected_account assert mock_session.execute.call_count == 2 diff --git a/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py b/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py index a1d08849ee3..7d034f90642 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_load_balancing_config.py @@ -6,7 +6,7 @@ import builtins import importlib import sys from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import ANY, MagicMock import pytest from flask import Flask @@ -92,6 +92,7 @@ def test_validate_credentials_success(app: Flask, load_balancing_module, monkeyp model="gpt-4o", model_type=ModelType.LLM, credentials={"api_key": "sk-***"}, + session=ANY, ) @@ -143,5 +144,6 @@ def test_validate_credentials_with_config_id(app: Flask, load_balancing_module, model="gpt-4o", model_type=ModelType.LLM, credentials={"api_key": "sk-***"}, + session=ANY, config_id="cfg-1", ) diff --git a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py index 2a576d1c920..5a578c42603 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py @@ -7,7 +7,7 @@ import importlib from contextlib import ExitStack, contextmanager from inspect import unwrap from types import ModuleType, SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -186,6 +186,7 @@ def test_builtin_provider_credentials_get(app: Flask, controller_module, monkeyp service_mock.assert_called_once_with( tenant_id="tenant-cred", provider_name="demo", + session=ANY, user=user, include_credential_ids=None, ) @@ -210,6 +211,7 @@ def test_builtin_provider_credentials_get_reads_repeated_include_ids( service_mock.assert_called_once_with( tenant_id="tenant-cred", provider_name="demo", + session=ANY, user=user, include_credential_ids=["cred-1", "cred-2"], ) diff --git a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py index 71381e6a2b4..ad84eed1f5e 100644 --- a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py +++ b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py @@ -6,7 +6,7 @@ in test_auth_wraps.py; handler tests use inspect.unwrap() to bypass them. """ import inspect -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -19,7 +19,7 @@ from controllers.inner_api.app.dsl import ( _get_active_account, ) from models.account import AccountStatus -from services.app_dsl_service import ImportStatus +from services.app_dsl_service import Import, ImportStatus class TestInnerAppDSLImportPayload: @@ -117,9 +117,7 @@ class TestEnterpriseAppDSLImport: mock_dsl_cls.return_value = self._mock_dsl yield - def _make_import_result(self, status: ImportStatus, **kwargs) -> "Import": - from services.app_dsl_service import Import - + def _make_import_result(self, status: ImportStatus, **kwargs) -> Import: result = Import( id="import-id", status=status, @@ -224,7 +222,7 @@ class TestEnterpriseAppDSLExport: body, status_code = result assert status_code == 200 assert body["data"] == "version: 0.6.0\nkind: app\n" - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, include_secret=False) + mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=False) @patch("controllers.inner_api.app.dsl.AppDslService") @patch("controllers.inner_api.app.dsl.db") @@ -239,7 +237,7 @@ class TestEnterpriseAppDSLExport: body, status_code = result assert status_code == 200 - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, include_secret=True) + mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=True) @patch("controllers.inner_api.app.dsl.db") def test_export_app_not_found_returns_404(self, mock_db, api_instance, app: Flask): diff --git a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py index 8c38564b3d0..8289a575050 100644 --- a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py +++ b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_drive.py @@ -9,7 +9,7 @@ from __future__ import annotations import inspect from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch import pytest from flask import Flask @@ -33,7 +33,7 @@ def test_manifest_parses_query_and_returns_items(): result = raw(AgentDriveManifestApi(), "agent-agent-1") assert result == {"items": [{"key": "docs/a.txt"}]} svc.return_value.manifest.assert_called_once_with( - tenant_id="tenant-1", agent_id="agent-1", prefix="docs/", include_download_url=True + tenant_id="tenant-1", agent_id="agent-1", prefix="docs/", include_download_url=True, session=ANY ) @@ -85,7 +85,11 @@ def test_skills_requires_tenant_id_and_returns_items(): } ] } - assert svc.return_value.list_skills.call_args.kwargs == {"tenant_id": "tenant-1", "agent_id": "agent-1"} + assert svc.return_value.list_skills.call_args.kwargs == { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "session": ANY, + } def test_commit_parses_body_and_returns_items(): diff --git a/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py b/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py index a6626adc420..bda25bb2fa8 100644 --- a/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/inner_api/workspace/test_workspace.py @@ -117,7 +117,7 @@ class TestEnterpriseWorkspace: assert result["tenant"]["name"] == "My Workspace" mock_tenant_svc.create_tenant.assert_called_once_with("My Workspace", is_from_dashboard=True, session=ANY) mock_tenant_svc.create_tenant_member.assert_called_once_with( - mock_tenant, mock_account, mock_db.session, role="owner" + mock_tenant, mock_account, mock_db.session(), role="owner" ) mock_event.send.assert_called_once_with(mock_tenant) diff --git a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py index b78473fadda..86d26420253 100644 --- a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py +++ b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py @@ -151,8 +151,8 @@ def _tenant_service(**overrides) -> SimpleNamespace: "get_tenant_members": Mock(return_value=[]), "remove_member_from_tenant": Mock(), "update_member_role": Mock(), - "get_tenant_by_id": lambda session, tenant_id: session.get(None, tenant_id), - "find_workspace_for_account": lambda session, account_id, workspace_id: session.execute(None).first(), + "get_tenant_by_id": lambda tenant_id, *, session: session.get(None, tenant_id), + "find_workspace_for_account": lambda account_id, workspace_id, *, session: session.execute(None).first(), } methods.update(overrides) return SimpleNamespace(**methods) @@ -162,12 +162,18 @@ def _account_service(**overrides) -> SimpleNamespace: """AccountService double; ``get_account_by_id`` delegates to the injected session (see :func:`_tenant_service`).""" methods: dict = { - "get_account_by_id": lambda session, account_id: session.get(None, account_id), + "get_account_by_id": lambda account_id, *, session: session.get(None, account_id), } methods.update(overrides) return SimpleNamespace(**methods) +def _db_mock() -> MagicMock: + mock_db = MagicMock() + mock_db.session.return_value = mock_db.session + return mock_db + + # --------------------------------------------------------------------------- # Route registration # --------------------------------------------------------------------------- @@ -272,7 +278,7 @@ def test_switch_returns_workspace_detail_with_current_true( acct_id = uuid.uuid4() api = WorkspaceSwitchApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _account(account_id=str(acct_id)) membership = SimpleNamespace(role=TenantAccountRole.OWNER, current=True) mock_db.session.execute.return_value.first.return_value = (_tenant(ws_id), membership) @@ -304,7 +310,7 @@ def test_switch_404s_when_service_raises_account_not_link_tenant( acct_id = uuid.uuid4() api = WorkspaceSwitchApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _account(account_id=str(acct_id)) monkeypatch.setattr( @@ -339,7 +345,7 @@ def test_members_list_returns_normalized_rows(app: Flask, bypass_pipeline, monke role=TenantAccountRole.ADMIN, ) - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr( @@ -381,7 +387,7 @@ def test_members_list_paginates_with_query_params(app: Flask, bypass_pipeline, m for i in range(5) ] - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr( @@ -409,7 +415,7 @@ def test_members_list_rejects_unknown_query_param(app: Flask, bypass_pipeline, m acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = _tenant(ws_id) monkeypatch.setattr(sys.modules["controllers.openapi.workspaces"], "db", mock_db) @@ -433,7 +439,7 @@ def test_invite_happy_path_returns_invite_url_and_member_id( invited = _account(account_id="new-1", email="new@example.com") - mock_db = MagicMock() + mock_db = _db_mock() # session.get is called twice: once for inviter Account, once for Tenant mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] @@ -514,7 +520,7 @@ def test_invite_blocked_by_saas_members_cap(app: Flask, bypass_pipeline, monkeyp acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] invite_mock = Mock() @@ -552,7 +558,7 @@ def test_invite_blocked_by_ee_workspace_members_license(app: Flask, bypass_pipel acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] invite_mock = Mock() @@ -592,7 +598,7 @@ def test_invite_ce_passes_when_both_caps_disabled(app: Flask, bypass_pipeline, m api = WorkspaceMembersApi() invited = _account(account_id="new-1", email="new@example.com") - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( @@ -625,7 +631,7 @@ def test_invite_400_when_already_in_tenant(app: Flask, bypass_pipeline, monkeypa acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( @@ -656,7 +662,7 @@ def test_delete_member_happy_path(app: Flask, bypass_pipeline, monkeypatch: pyte acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), # operator _tenant(ws_id), # tenant @@ -698,7 +704,7 @@ def test_delete_member_exception_mapping(app: Flask, bypass_pipeline, monkeypatc acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -731,7 +737,7 @@ def test_delete_member_404_when_member_missing(app: Flask, bypass_pipeline, monk acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -763,7 +769,7 @@ def test_update_role_happy_path(app: Flask, bypass_pipeline, monkeypatch: pytest acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -809,7 +815,7 @@ def test_update_role_exception_mapping(app: Flask, bypass_pipeline, monkeypatch, acct_id = uuid.uuid4() api = WorkspaceMemberApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [ _account(account_id=str(acct_id)), _tenant(ws_id), @@ -851,7 +857,7 @@ def test_load_tenant_rejects_archived_workspace(app: Flask, bypass_pipeline, mon api = WorkspaceMembersApi() archived = SimpleNamespace(id=ws_id, name="WS", status="archive", created_at=datetime(2026, 5, 18, tzinfo=UTC)) - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.return_value = archived monkeypatch.setattr( @@ -878,7 +884,7 @@ def test_invite_400_when_register_error(app: Flask, bypass_pipeline, monkeypatch acct_id = uuid.uuid4() api = WorkspaceMembersApi() - mock_db = MagicMock() + mock_db = _db_mock() mock_db.session.get.side_effect = [_account(account_id=str(acct_id)), _tenant(ws_id)] monkeypatch.setattr( diff --git a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py index 810101fb0a5..1ff925cba7e 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py @@ -15,7 +15,7 @@ Note: API endpoint tests for annotation controllers are complex due to: import uuid from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import ANY, Mock import pytest from flask import Flask @@ -264,7 +264,7 @@ class TestAnnotationListApi: assert response["page"] == 1 assert response["limit"] == 20 - get_mock.assert_called_once_with("app", 1, 20, "") + get_mock.assert_called_once_with("app", 1, 20, "", session=ANY) def test_get_accepts_valid_numeric_strings(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: annotation = SimpleNamespace(id="a1", question="q", content="a", created_at=0) @@ -281,7 +281,7 @@ class TestAnnotationListApi: assert response["total"] == 1 assert response["page"] == 2 assert response["limit"] == 5 - get_mock.assert_called_once_with("app", 2, 5, "refund") + get_mock.assert_called_once_with("app", 2, 5, "refund", session=ANY) @pytest.mark.parametrize("query_string", ["page=abc&limit=5", "page=1&limit=abc", "page=&limit=5", "limit=0"]) def test_get_rejects_invalid_explicit_pagination_value( diff --git a/api/tests/unit_tests/controllers/service_api/app/test_app.py b/api/tests/unit_tests/controllers/service_api/app/test_app.py index 30979a26980..04e9220ad55 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_app.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_app.py @@ -3,7 +3,7 @@ Unit tests for Service API App controllers """ import uuid -from unittest.mock import Mock, patch +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask @@ -368,7 +368,7 @@ class TestAppMetaApi: response = api.get() # Assert - mock_service_instance.get_app_meta.assert_called_once_with(mock_app_model) + mock_service_instance.get_app_meta.assert_called_once_with(mock_app_model, session=ANY) assert response == {"tool_icons": {}, "AgentIcons": {}} diff --git a/api/tests/unit_tests/controllers/service_api/app/test_completion.py b/api/tests/unit_tests/controllers/service_api/app/test_completion.py index 9f2a2edeff8..65652594294 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_completion.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_completion.py @@ -252,7 +252,12 @@ class TestAppGenerateService: mock_generate.return_value = expected result = AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={"query": "Hi"}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={"query": "Hi"}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) assert result == expected @@ -264,7 +269,12 @@ class TestAppGenerateService: with pytest.raises(services.errors.conversation.ConversationNotExistsError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) @patch.object(AppGenerateService, "generate") @@ -274,7 +284,12 @@ class TestAppGenerateService: with pytest.raises(QuotaExceededError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) @patch.object(AppGenerateService, "generate") @@ -284,7 +299,12 @@ class TestAppGenerateService: with pytest.raises(InvokeError): AppGenerateService.generate( - app_model=Mock(spec=App), user=Mock(spec=EndUser), args={}, invoke_from=Mock(), streaming=False + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + args={}, + invoke_from=Mock(), + session=Mock(), + streaming=False, ) diff --git a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py index 97873c631ae..3197812bc31 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py @@ -475,6 +475,7 @@ class TestConversationService: user=Mock(spec=EndUser), name="New Name", auto_generate=False, + session=Mock(), ) assert result.name == "New Name" diff --git a/api/tests/unit_tests/controllers/service_api/app/test_message.py b/api/tests/unit_tests/controllers/service_api/app/test_message.py index d8d5c61bcb3..0400abe0e65 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_message.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_message.py @@ -266,6 +266,7 @@ class TestMessageService: conversation_id=str(uuid.uuid4()), first_id=None, limit=20, + session=Mock(), ) assert hasattr(result, "data") @@ -281,7 +282,12 @@ class TestMessageService: with pytest.raises(services.errors.conversation.ConversationNotExistsError): MessageService.pagination_by_first_id( - app_model=Mock(spec=App), user=Mock(spec=EndUser), conversation_id="invalid_id", first_id=None, limit=20 + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + conversation_id="invalid_id", + first_id=None, + limit=20, + session=Mock(), ) @patch.object(MessageService, "pagination_by_first_id") @@ -296,6 +302,7 @@ class TestMessageService: conversation_id=str(uuid.uuid4()), first_id="invalid_first_id", limit=20, + session=Mock(), ) @patch.object(MessageService, "create_feedback") @@ -309,6 +316,7 @@ class TestMessageService: user=Mock(spec=EndUser), rating=FeedbackRating.LIKE, content="Great response!", + session=Mock(), ) mock_create_feedback.assert_called_once() @@ -325,6 +333,7 @@ class TestMessageService: user=Mock(spec=EndUser), rating=FeedbackRating.LIKE, content=None, + session=Mock(), ) @patch.object(MessageService, "get_all_messages_feedbacks") @@ -336,7 +345,7 @@ class TestMessageService: ] mock_get_feedbacks.return_value = mock_feedbacks - result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20) + result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20, session=Mock()) assert len(result) == 2 assert result[0]["rating"] == "like" @@ -348,7 +357,11 @@ class TestMessageService: mock_get_questions.return_value = mock_questions result = MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id=str(uuid.uuid4()), invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id=str(uuid.uuid4()), + invoke_from=Mock(), + session=Mock(), ) assert len(result) == 3 @@ -361,7 +374,11 @@ class TestMessageService: with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id=str(uuid.uuid4()), invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id=str(uuid.uuid4()), + invoke_from=Mock(), + session=Mock(), ) @patch.object(MessageService, "get_suggested_questions_after_answer") @@ -371,7 +388,11 @@ class TestMessageService: with pytest.raises(MessageNotExistsError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), user=Mock(spec=EndUser), message_id="invalid_message_id", invoke_from=Mock() + app_model=Mock(spec=App), + user=Mock(spec=EndUser), + message_id="invalid_message_id", + invoke_from=Mock(), + session=Mock(), ) diff --git a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py index 3cabfe43ddc..2115bb85526 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py @@ -18,7 +18,7 @@ import uuid from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -167,26 +167,6 @@ class TestWorkflowLogQuery: query_max_limit = WorkflowLogQuery(limit=100) assert query_max_limit.limit == 100 - def test_query_rejects_page_below_minimum(self): - """Test query rejects page < 1.""" - with pytest.raises(ValueError): - WorkflowLogQuery(page=0) - - def test_query_rejects_page_above_maximum(self): - """Test query rejects page > 99999.""" - with pytest.raises(ValueError): - WorkflowLogQuery(page=100000) - - def test_query_rejects_limit_below_minimum(self): - """Test query rejects limit < 1.""" - with pytest.raises(ValueError): - WorkflowLogQuery(limit=0) - - def test_query_rejects_limit_above_maximum(self): - """Test query rejects limit > 100.""" - with pytest.raises(ValueError): - WorkflowLogQuery(limit=101) - def test_query_with_keyword_search(self): """Test query with keyword filter.""" query = WorkflowLogQuery(keyword="workflow execution") @@ -263,7 +243,7 @@ class TestAppGenerateServiceWorkflow: """Test AppGenerateService workflow integration.""" @patch.object(AppGenerateService, "generate") - def test_generate_accepts_workflow_args(self, mock_generate): + def test_generate_accepts_workflow_args(self, mock_generate: MagicMock): """Test generate accepts workflow-specific args.""" mock_generate.return_value = {"result": "success"} @@ -272,6 +252,7 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"inputs": {"key": "value"}, "workflow_id": "workflow_123"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @@ -279,7 +260,7 @@ class TestAppGenerateServiceWorkflow: mock_generate.assert_called_once() @patch.object(AppGenerateService, "generate") - def test_generate_raises_workflow_not_found_error(self, mock_generate): + def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock): """Test generate raises WorkflowNotFoundError.""" mock_generate.side_effect = WorkflowNotFoundError("Workflow not found") @@ -289,11 +270,12 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"workflow_id": "invalid_id"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @patch.object(AppGenerateService, "generate") - def test_generate_raises_is_draft_workflow_error(self, mock_generate): + def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock): """Test generate raises IsDraftWorkflowError.""" mock_generate.side_effect = IsDraftWorkflowError("Workflow is draft") @@ -303,11 +285,12 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"workflow_id": "draft_workflow"}, invoke_from=Mock(), + session=MagicMock(), streaming=False, ) @patch.object(AppGenerateService, "generate") - def test_generate_supports_streaming_mode(self, mock_generate): + def test_generate_supports_streaming_mode(self, mock_generate: MagicMock): """Test generate supports streaming response mode.""" mock_stream = Mock() mock_generate.return_value = mock_stream @@ -317,6 +300,7 @@ class TestAppGenerateServiceWorkflow: user=Mock(), args={"inputs": {}, "response_mode": "streaming"}, invoke_from=Mock(), + session=MagicMock(), streaming=True, ) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py index 43cc2450db5..406037e268d 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py @@ -549,16 +549,14 @@ class TestPipelineRunApiPost: new_callable=lambda: Mock(spec=Account), ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService") - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_success_streaming( - self, mock_ns, mock_db, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app - ): + def test_post_success_streaming(self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app): """Test successful pipeline run with streaming response.""" tenant_id = str(uuid.uuid4()) dataset_id = str(uuid.uuid4()) - mock_db.session.scalar.return_value = Mock() + session = Mock() + session.scalar.return_value = Mock() mock_ns.payload = { "inputs": {"key": "val"}, @@ -579,27 +577,33 @@ class TestPipelineRunApiPost: with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() - response = api.post(tenant_id=tenant_id, dataset_id=dataset_id) + response = api.post.__wrapped__(api, session, tenant_id=tenant_id, dataset_id=dataset_id) assert response == {"result": "ok"} + mock_svc_cls.assert_called_once_with(session) mock_gen_svc.generate.assert_called_once() - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") - def test_post_not_found(self, mock_db, app: Flask): + def test_post_not_found(self, app: Flask): """Test NotFound when dataset check fails.""" - mock_db.session.scalar.return_value = None + session = Mock() + session.scalar.return_value = None with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() with pytest.raises(NotFound): - api.post(tenant_id=str(uuid.uuid4()), dataset_id=str(uuid.uuid4())) + api.post.__wrapped__( + api, + session, + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account") - @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_forbidden_non_account_user(self, mock_ns, mock_db, app: Flask): + def test_post_forbidden_non_account_user(self, mock_ns, app: Flask): """Test Forbidden when current_user is not an Account.""" - mock_db.session.scalar.return_value = Mock() + session = Mock() + session.scalar.return_value = Mock() mock_ns.payload = { "inputs": {}, "datasource_type": "online_document", @@ -612,7 +616,12 @@ class TestPipelineRunApiPost: with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() with pytest.raises(Forbidden): - api.post(tenant_id=str(uuid.uuid4()), dataset_id=str(uuid.uuid4())) + api.post.__wrapped__( + api, + session, + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + ) class TestFileUploadApiPost: diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py index a95baf1b482..0b1ca8741a9 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py @@ -1193,7 +1193,7 @@ class TestDatasetSegmentApiDelete: # Assert assert response == ("", 204) - mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset, mock_db.session) + mock_seg_svc.delete_segment.assert_called_once_with(mock_segment, mock_doc, mock_dataset, mock_db.session()) @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index dd2caf4f3fc..e83724f955f 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -783,7 +783,7 @@ class TestDocumentApiDelete: # Assert assert response == ("", 204) - mock_doc_svc.delete_document.assert_called_once_with(mock_document, mock_db.session) + mock_doc_svc.delete_document.assert_called_once_with(mock_document, mock_db.session()) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py index b77c783ae16..dd1322a6344 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py @@ -408,7 +408,7 @@ class TestDatasetMetadataBuiltInFieldAction: assert status == 200 assert response["result"] == "success" - mock_meta_svc.enable_built_in_field.assert_called_once_with(ANY, mock_dataset) + mock_meta_svc.enable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -439,7 +439,7 @@ class TestDatasetMetadataBuiltInFieldAction: ) assert status == 200 - mock_meta_svc.disable_built_in_field.assert_called_once_with(ANY, mock_dataset) + mock_meta_svc.disable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_action_dataset_not_found( diff --git a/api/tests/unit_tests/controllers/web/test_app.py b/api/tests/unit_tests/controllers/web/test_app.py index 542ee111e1b..73f308dc749 100644 --- a/api/tests/unit_tests/controllers/web/test_app.py +++ b/api/tests/unit_tests/controllers/web/test_app.py @@ -3,7 +3,7 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -148,7 +148,7 @@ class TestAppAccessMode: with app.test_request_context("/webapp/access-mode?appCode=code1"): result = AppAccessMode().get() - mock_resolve.assert_called_once_with("code1") + mock_resolve.assert_called_once_with("code1", session=ANY) mock_access.assert_called_once_with("resolved-id") assert result == {"accessMode": "external"} diff --git a/api/tests/unit_tests/controllers/web/test_message_list.py b/api/tests/unit_tests/controllers/web/test_message_list.py index 2bb425cdba2..b5d74df65ef 100644 --- a/api/tests/unit_tests/controllers/web/test_message_list.py +++ b/api/tests/unit_tests/controllers/web/test_message_list.py @@ -6,7 +6,7 @@ import builtins import uuid from datetime import datetime from types import ModuleType, SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch from uuid import uuid4 import pytest @@ -158,7 +158,7 @@ def test_message_list_mapping(app: Flask) -> None: ): response = MessageListApi().get(app_model, end_user) - mock_page.assert_called_once_with(app_model, end_user, conversation_id, None, 20) + mock_page.assert_called_once_with(app_model, end_user, conversation_id, None, 20, session=ANY) assert response["limit"] == 20 assert response["has_more"] is False assert len(response["data"]) == 1 diff --git a/api/tests/unit_tests/controllers/web/test_web_login.py b/api/tests/unit_tests/controllers/web/test_web_login.py index 984be6ddba9..a91d4253aa8 100644 --- a/api/tests/unit_tests/controllers/web/test_web_login.py +++ b/api/tests/unit_tests/controllers/web/test_web_login.py @@ -1,7 +1,7 @@ import base64 import logging from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -66,7 +66,7 @@ class TestEmailCodeLoginSendEmailApi: response = EmailCodeLoginSendEmailApi().post() assert response == {"result": "success", "data": "token-123"} - mock_get_user.assert_called_once_with("User@Example.com") + mock_get_user.assert_called_once_with("User@Example.com", ANY) mock_send_email.assert_called_once_with(account=mock_account, language="en-US") @@ -96,7 +96,7 @@ class TestEmailCodeLoginApi: response = EmailCodeLoginApi().post() assert response == {"result": "success", "data": {"access_token": "new-access-token"}} - mock_get_user.assert_called_once_with("User@Example.com") + mock_get_user.assert_called_once_with("User@Example.com", ANY) mock_revoke_token.assert_called_once_with("token-123") mock_login.assert_called_once() mock_reset_login_rate.assert_called_once_with("user@example.com") diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py index cda5178e30c..41e14af72de 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py @@ -147,7 +147,7 @@ class TestAdvancedChatAppGeneratorInternals: ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", - SimpleNamespace(engine=object(), session=SimpleNamespace(close=lambda: None)), + SimpleNamespace(engine=object(), session=lambda: SimpleNamespace(close=lambda: None)), ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace() diff --git a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py index 9b89b108207..ef12f0be965 100644 --- a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py @@ -138,7 +138,7 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", - SimpleNamespace(engine=object()), + SimpleNamespace(engine=object(), session=lambda: MagicMock()), ) trace_manager = object.__new__(TraceQueueManager) monkeypatch.setattr( diff --git a/api/tests/unit_tests/core/app/test_llm_quota.py b/api/tests/unit_tests/core/app/test_llm_quota.py index 13bdf765358..ec6ac134443 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -1,7 +1,7 @@ from collections.abc import Generator from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from sqlalchemy import create_engine, select @@ -28,8 +28,19 @@ from models.provider import Provider, ProviderType @contextmanager def _patched_credit_pool_session_factory(engine: Engine) -> Generator[None, None, None]: session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield + sessions = [] + + def _session(): + session = session_maker() + sessions.append(session) + return session + + with patch("core.app.llm.quota.db", SimpleNamespace(session=_session)): + try: + yield + finally: + for session in sessions: + session.close() def test_ensure_llm_quota_available_for_model_raises_when_system_model_is_exhausted() -> None: @@ -122,6 +133,7 @@ def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None: mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=42, + session=ANY, ) @@ -241,6 +253,7 @@ def test_deduct_llm_quota_for_model_uses_credit_configuration() -> None: mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=9, + session=ANY, ) @@ -276,6 +289,7 @@ def test_deduct_llm_quota_for_model_uses_single_charge_for_times_quota() -> None mock_deduct_credits.assert_called_once_with( tenant_id="tenant-id", credits_required=1, + session=ANY, ) @@ -313,6 +327,7 @@ def test_deduct_llm_quota_for_model_uses_paid_billing_pool() -> None: tenant_id="tenant-id", credits_required=5, pool_type="paid", + session=ANY, ) diff --git a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py index ddb33f0758f..1c4a6e2db7c 100644 --- a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py +++ b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py @@ -149,12 +149,12 @@ class TestWorkflowServiceInterface: from core.llm_generator.llm_generator import WorkflowServiceInterface class MockService(WorkflowServiceInterface): - def get_draft_workflow(self, app_model, workflow_id=None): - return super().get_draft_workflow(app_model, workflow_id) + def get_draft_workflow(self, app_model, workflow_id=None, *, session): + return super().get_draft_workflow(app_model, workflow_id, session=session) def get_node_last_run(self, app_model, workflow, node_id): return super().get_node_last_run(app_model, workflow, node_id) service = MockService() - service.get_draft_workflow(None) + service.get_draft_workflow(None, session=None) service.get_node_last_run(None, None, "node") diff --git a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py index 7c672570bfa..d8452d91e2c 100644 --- a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py +++ b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py @@ -227,13 +227,12 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.ExternalDatasetService.fetch_external_knowledge_retrieval") @patch("core.rag.datasource.retrieval_service.MetadataFilteringCondition.model_validate") - @patch("core.rag.datasource.retrieval_service.db.session.scalar") - def test_external_retrieve_with_metadata_conditions(self, mock_scalar, mock_validate, mock_fetch): - mock_scalar.return_value = SimpleNamespace(tenant_id="tenant-1") + def test_external_retrieve_with_metadata_conditions(self, mock_validate, mock_fetch): mock_validate.return_value = "validated-condition" expected_documents = [create_mock_document("external-doc", "external-1", 0.8, provider="external")] mock_fetch.return_value = expected_documents session = MagicMock() + session.scalar.return_value = SimpleNamespace(tenant_id="tenant-1") results = RetrievalService.external_retrieve( session=session, @@ -246,19 +245,19 @@ class TestRetrievalServiceInternals: assert results == expected_documents mock_validate.assert_called_once() mock_fetch.assert_called_once_with( - session, - "tenant-1", - "dataset-1", - "test query", - {"top_k": 3}, + tenant_id="tenant-1", + dataset_id="dataset-1", + query="test query", + external_retrieval_parameters={"top_k": 3}, metadata_condition="validated-condition", + session=session, ) - @patch("core.rag.datasource.retrieval_service.db.session.scalar") - def test_external_retrieve_returns_empty_when_dataset_not_found(self, mock_scalar): - mock_scalar.return_value = None + def test_external_retrieve_returns_empty_when_dataset_not_found(self): + session = MagicMock() + session.scalar.return_value = None - results = RetrievalService.external_retrieve(session=MagicMock(), dataset_id="missing", query="q") + results = RetrievalService.external_retrieve(session=session, dataset_id="missing", query="q") assert results == [] diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py index 302ababb48f..f5761b5ba3d 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py @@ -209,7 +209,7 @@ class TestParagraphIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, ["node-1"], delete_summaries=True) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_economy_deletes_summaries_and_keywords( @@ -225,7 +225,7 @@ class TestParagraphIndexProcessor: ): processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) mock_keyword_cls.return_value.delete.assert_called_once() def test_clean_deletes_keywords_by_ids(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py index 7d339a7701f..672764e5336 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py @@ -278,7 +278,7 @@ class TestParentChildIndexProcessor: ): processor.clean(dataset, ["node-1"], delete_summaries=True, precomputed_child_node_ids=[]) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) def test_clean_deletes_all_summaries_when_node_ids_missing( self, processor: ParentChildIndexProcessor, dataset: Mock @@ -291,7 +291,7 @@ class TestParentChildIndexProcessor: ): processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) def test_split_child_nodes_requires_subchunk_segmentation(self, processor: ParentChildIndexProcessor) -> None: rules = Rule(subchunk_segmentation=None) diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py index 6e5a4fabbb0..5dde1623d2d 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py @@ -243,7 +243,7 @@ class TestQAIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, ["node-1"], delete_summaries=True) - mock_summary.assert_called_once_with(dataset, ["seg-1"]) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_handles_dataset_wide_cleanup(self, processor: QAIndexProcessor, dataset: Mock) -> None: @@ -256,7 +256,7 @@ class TestQAIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, None, delete_summaries=True) - mock_summary.assert_called_once_with(dataset, None) + mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) vector.delete.assert_called_once() def test_index_adds_documents_and_vectors_for_high_quality( diff --git a/api/tests/unit_tests/events/test_app_event_signals.py b/api/tests/unit_tests/events/test_app_event_signals.py index 29582a50f6d..a6059fadbcf 100644 --- a/api/tests/unit_tests/events/test_app_event_signals.py +++ b/api/tests/unit_tests/events/test_app_event_signals.py @@ -44,7 +44,7 @@ def _make_collector(target: list): @pytest.mark.usefixtures("mock_db", "_mock_deps") class TestAppWasDeletedSignal: - def test_sends_signal(self, app_model): + def test_sends_signal(self, app_model, mock_db): from events.app_event import app_was_deleted from services.app_service import AppService @@ -52,7 +52,7 @@ class TestAppWasDeletedSignal: handler = _make_collector(received) app_was_deleted.connect(handler) try: - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=mock_db.session) finally: app_was_deleted.disconnect(handler) @@ -71,7 +71,7 @@ class TestAppWasDeletedSignal: mock_db.session.delete.side_effect = lambda _: call_order.append("db_delete") try: - AppService().delete_app(app_model) + AppService().delete_app(app_model, session=mock_db.session) finally: app_was_deleted.disconnect(handler) @@ -80,7 +80,7 @@ class TestAppWasDeletedSignal: @pytest.mark.usefixtures("mock_db") class TestAppWasUpdatedSignal: - def test_update_app(self, app_model): + def test_update_app(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -101,13 +101,14 @@ class TestAppWasUpdatedSignal: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_name(self, app_model): + def test_update_app_name(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -117,13 +118,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: - AppService().update_app_name(app_model, "New Name") + AppService().update_app_name(app_model, "New Name", session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_icon(self, app_model): + def test_update_app_icon(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -133,13 +134,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: - AppService().update_app_icon(app_model, "🎉", "#000") + AppService().update_app_icon(app_model, "🎉", "#000", session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_site_status_sends_when_changed(self, app_model): + def test_update_app_site_status_sends_when_changed(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -150,13 +151,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: app_model.enable_site = False - AppService().update_app_site_status(app_model, True) + AppService().update_app_site_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_site_status_skips_when_unchanged(self, app_model): + def test_update_app_site_status_skips_when_unchanged(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -166,13 +167,13 @@ class TestAppWasUpdatedSignal: try: app_model.enable_site = True - AppService().update_app_site_status(app_model, True) + AppService().update_app_site_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [] - def test_update_app_api_status_sends_when_changed(self, app_model): + def test_update_app_api_status_sends_when_changed(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -183,13 +184,13 @@ class TestAppWasUpdatedSignal: with patch("services.app_service.current_user", MagicMock(id="user-1")): try: app_model.enable_api = False - AppService().update_app_api_status(app_model, True) + AppService().update_app_api_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) assert received == [app_model] - def test_update_app_api_status_skips_when_unchanged(self, app_model): + def test_update_app_api_status_skips_when_unchanged(self, app_model, mock_db): from events.app_event import app_was_updated from services.app_service import AppService @@ -199,7 +200,7 @@ class TestAppWasUpdatedSignal: try: app_model.enable_api = True - AppService().update_app_api_status(app_model, True) + AppService().update_app_api_status(app_model, True, session=mock_db.session) finally: app_was_updated.disconnect(handler) diff --git a/api/tests/unit_tests/events/test_update_provider_when_message_created.py b/api/tests/unit_tests/events/test_update_provider_when_message_created.py index f9ac5d9678e..327c80323b4 100644 --- a/api/tests/unit_tests/events/test_update_provider_when_message_created.py +++ b/api/tests/unit_tests/events/test_update_provider_when_message_created.py @@ -1,7 +1,7 @@ from collections.abc import Generator from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import ANY, patch from uuid import uuid4 import pytest @@ -19,8 +19,19 @@ from models.provider import ProviderType @contextmanager def _patched_credit_pool_session_factory(engine: Engine) -> Generator[None, None, None]: session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield + sessions = [] + + def _session(): + session = session_maker() + sessions.append(session) + return session + + with patch("events.event_handlers.update_provider_when_message_created.db", SimpleNamespace(session=_session)): + try: + yield + finally: + for session in sessions: + session.close() def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_insufficient() -> None: @@ -140,5 +151,6 @@ def test_capped_credit_pool_accounting_skips_exhaustion_warning_when_full_amount tenant_id="tenant-id", credits_required=3, pool_type="trial", + session=ANY, ) assert "Credit pool exhausted during message-created accounting" not in caplog.text diff --git a/api/tests/unit_tests/services/agent/test_agent_services.py b/api/tests/unit_tests/services/agent/test_agent_services.py index 6fee58bb29d..f896a761904 100644 --- a/api/tests/unit_tests/services/agent/test_agent_services.py +++ b/api/tests/unit_tests/services/agent/test_agent_services.py @@ -117,7 +117,9 @@ def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPa monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) - result = AgentComposerService.load_workflow_composer(tenant_id="tenant-1", app_id="app-1", node_id="node-1") + result = AgentComposerService.load_workflow_composer( + tenant_id="tenant-1", app_id="app-1", node_id="node-1", session=composer_service.db.session + ) assert result["binding"] is None assert result["save_options"] == ["node_job_only", "save_to_roster"] @@ -155,7 +157,9 @@ def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest. lambda **kwargs: {"agent": kwargs["agent"].id, "version": kwargs["version"].id}, ) - result = AgentComposerService.load_workflow_composer(tenant_id="tenant-1", app_id="app-1", node_id="node-1") + result = AgentComposerService.load_workflow_composer( + tenant_id="tenant-1", app_id="app-1", node_id="node-1", session=composer_service.db.session + ) assert result == {"agent": "agent-1", "version": "version-1"} @@ -190,6 +194,7 @@ def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest app_id="app-1", node_id="node-1", snapshot_id="preview-version", + session=composer_service.db.session, ) assert result == {"binding_snapshot_id": "binding-version", "version": "preview-version"} @@ -232,6 +237,7 @@ def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest app_id="app-1", node_id="node-1", snapshot_id="inline-preview-version", + session=composer_service.db.session, ) assert result == {"agent": "inline-agent-1", "version": "inline-preview-version"} @@ -258,6 +264,7 @@ def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch binding=binding, agent=agent, account_id="account-1", + session="session-1", ) assert debug_conversation_id == "debug-conversation-1" @@ -279,6 +286,7 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p binding=SimpleNamespace(binding_type=WorkflowAgentBindingType.ROSTER_AGENT), agent=SimpleNamespace(id="agent-1", scope=AgentScope.ROSTER), account_id="account-1", + session="session-1", ) is None ) @@ -288,6 +296,7 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p binding=SimpleNamespace(binding_type=WorkflowAgentBindingType.INLINE_AGENT), agent=SimpleNamespace(id="inline-agent-1", scope=AgentScope.WORKFLOW_ONLY), account_id=None, + session="session-1", ) is None ) @@ -303,6 +312,7 @@ def test_load_workflow_composer_rejects_preview_without_binding(monkeypatch: pyt app_id="app-1", node_id="node-1", snapshot_id="preview-version", + session=composer_service.db.session, ) @@ -361,7 +371,12 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, ) result = AgentComposerService.save_workflow_composer( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -382,7 +397,12 @@ def test_save_workflow_composer_rejects_agent_app_variant(): with pytest.raises(ValueError): AgentComposerService.save_workflow_composer( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) @@ -444,7 +464,11 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. ) result = AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -475,7 +499,9 @@ def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.Mon monkeypatch.setattr(AgentComposerService, "_serialize_version", lambda _version: None) monkeypatch.setattr(AgentComposerService, "_serialize_draft", lambda _draft: {"id": "draft-1"}) - result = AgentComposerService.load_agent_app_composer(tenant_id="tenant-1", app_id="app-1") + result = AgentComposerService.load_agent_app_composer( + tenant_id="tenant-1", app_id="app-1", session=composer_service.db.session + ) assert result["save_options"] == [ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION.value] @@ -495,6 +521,7 @@ def test_save_agent_app_composer_rejects_version_save_strategy(): app_id="app-1", account_id="account-1", payload=payload, + session=composer_service.db.session, ) @@ -528,7 +555,11 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey ) result = AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} @@ -570,7 +601,7 @@ def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snaps ) AgentComposerService.save_agent_app_composer( - tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload + tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload, session=fake_session ) assert agent.active_config_is_published is True @@ -617,6 +648,7 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke agent_id="agent-1", account_id="account-1", version_note="ship it", + session=fake_session, ) assert exc_info.value.error_code == "agent_model_not_configured" @@ -665,6 +697,7 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. agent_id="agent-1", account_id="account-1", version_note="ship it", + session=composer_service.db.session, ) assert result["result"] == "success" @@ -708,6 +741,7 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=composer_service.db.session, ) build_draft = fake_session.added[0] @@ -729,6 +763,7 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=composer_service.db.session, ) assert applied["result"] == "success" @@ -787,6 +822,7 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", + session=fake_session, ) assert normal_draft.config_snapshot_dict == build_draft.config_snapshot_dict @@ -812,12 +848,21 @@ def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatc monkeypatch.setattr(AgentComposerService, "_workspace_dify_tools", lambda **kwargs: []) workflow_candidates = AgentComposerService.get_workflow_candidates( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", user_id="account-1" + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + user_id="account-1", + session=composer_service.db.session, ) agent_app_candidates = AgentComposerService.get_agent_app_candidates( - tenant_id="tenant-1", agent_id="agent-1", user_id="account-1" + tenant_id="tenant-1", + agent_id="agent-1", + user_id="account-1", + session=composer_service.db.session, + ) + impact = AgentComposerService.calculate_impact( + tenant_id="tenant-1", current_snapshot_id="version-1", session=composer_service.db.session ) - impact = AgentComposerService.calculate_impact(tenant_id="tenant-1", current_snapshot_id="version-1") assert workflow_candidates["variant"] == "workflow" assert workflow_candidates["allowed_node_job_candidates"]["previous_node_outputs"] == [] @@ -854,7 +899,9 @@ def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pyt version = AgentConfigSnapshot(id="version-1", version=1, config_snapshot='{"prompt":{"system_prompt":"x"}}') monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) - state = AgentComposerService._serialize_workflow_state(binding=binding, agent=agent, version=version) + state = AgentComposerService._serialize_workflow_state( + binding=binding, agent=agent, version=version, session=composer_service.db.session + ) assert state["soul_lock"]["locked"] is True assert state["agent"]["role"] == "Tender Analyst" @@ -893,7 +940,9 @@ def test_serialize_workflow_state_passes_user_declared_outputs_through_effective version = AgentConfigSnapshot(id="version-1", version=1, config_snapshot='{"prompt":{"system_prompt":"x"}}') monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) - state = AgentComposerService._serialize_workflow_state(binding=binding, agent=agent, version=version) + state = AgentComposerService._serialize_workflow_state( + binding=binding, agent=agent, version=version, session=composer_service.db.session + ) # When the user has declared outputs, effective_declared_outputs is the same # list (no defaults injected). @@ -943,6 +992,7 @@ def test_serialize_workflow_state_includes_inline_debug_conversation_message_sta agent=agent, version=version, account_id="account-1", + session=composer_service.db.session, ) assert state["debug_conversation_id"] == "debug-conversation-1" @@ -1024,6 +1074,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=existing_binding, payload=payload, + session=composer_service.db.session, ) inline_binding = AgentComposerService._save_node_job_only( tenant_id="tenant-1", @@ -1033,6 +1084,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=None, payload=payload, + session=composer_service.db.session, ) new_agent_binding = AgentComposerService._save_as_new_agent( tenant_id="tenant-1", @@ -1042,6 +1094,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk account_id="account-1", binding=None, payload=payload, + session=composer_service.db.session, ) save_to_roster_binding = AgentComposerService._save_to_roster( tenant_id="tenant-1", @@ -1055,12 +1108,14 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk current_snapshot_id="inline-version-1", ), payload=payload, + session=composer_service.db.session, ) new_version_binding = AgentComposerService._save_as_new_version( tenant_id="tenant-1", account_id="account-1", binding=WorkflowAgentNodeBinding(agent_id="roster-agent-1", current_snapshot_id="source-version-1"), payload=payload, + session=composer_service.db.session, ) assert updated_binding.updated_by == "account-1" @@ -1085,6 +1140,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk "account_id": "account-1", "agent_soul": payload.agent_soul, "node_job": payload.node_job, + "session": composer_service.db.session, } ] @@ -1151,6 +1207,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) assert updated_binding.current_snapshot_id == "inline-version-2" @@ -1203,6 +1260,7 @@ def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pyte account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) assert updated_binding is binding @@ -1252,6 +1310,7 @@ def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_i account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) @@ -1299,6 +1358,7 @@ def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypat account_id="account-1", binding=binding, payload=payload, + session=composer_service.db.session, ) @@ -1385,6 +1445,7 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n account_id="account-1", source_agent_id="roster-agent-1", source_snapshot_id="roster-version-2", + session=composer_service.db.session, ) assert state["binding"]["binding_type"] == WorkflowAgentBindingType.INLINE_AGENT.value @@ -1445,6 +1506,7 @@ def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkey account_id="account-1", source_agent_id="roster-agent-1", source_snapshot_id="roster-version-1", + session=composer_service.db.session, ) @@ -1495,6 +1557,7 @@ def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(mo account_id="account-1", source_agent_id="roster-agent-1", idempotency_key="same-click", + session=composer_service.db.session, ) assert state == {"binding_type": WorkflowAgentBindingType.INLINE_AGENT.value} @@ -1573,6 +1636,7 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( node_id="node-1", account_id="account-1", source_agent_id="roster-agent-1", + session=composer_service.db.session, ) @@ -1629,6 +1693,7 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest account_id="account-1", agent_soul=agent_soul, node_job=node_job, + session=composer_service.db.session, ) copied = [row for row in fake_session.added if isinstance(row, AgentDriveFile)] @@ -1654,6 +1719,7 @@ def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys(monkeypatch: target_agent_id="inline-agent-1", account_id="account-1", agent_soul=agent_soul, + session=composer_service.db.session, ) assert fake_session.added == [] @@ -1680,6 +1746,7 @@ def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.Mo target_agent_id="inline-agent-1", account_id="account-1", agent_soul=agent_soul, + session=composer_service.db.session, ) assert [row for row in fake_session.added if isinstance(row, AgentDriveFile)] == [] @@ -1743,7 +1810,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes ) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): created_apps.append((tenant_id, params, account)) return SimpleNamespace(id="app-agent-1") @@ -1781,6 +1848,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes node_id="node-1", account_id="account-1", agent_soul=_agent_soul_with_model(), + session=composer_service.db.session, ) roster_agent = AgentComposerService._create_roster_agent_for_composer( tenant_id="tenant-1", @@ -1789,6 +1857,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) assert workflow_agent.active_config_snapshot_id == "version-with-model" @@ -1810,14 +1879,14 @@ def test_composer_require_account(monkeypatch: pytest.MonkeyPatch): account = SimpleNamespace(id="account-1") monkeypatch.setattr(composer_service.db, "session", SimpleNamespace(get=lambda model, account_id: account)) - assert AgentComposerService._require_account(account_id="account-1") is account + assert AgentComposerService._require_account(account_id="account-1", session=composer_service.db.session) is account def test_composer_require_account_raises_when_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(composer_service.db, "session", SimpleNamespace(get=lambda model, account_id: None)) with pytest.raises(ValueError, match="Account not found"): - AgentComposerService._require_account(account_id="missing-account") + AgentComposerService._require_account(account_id="missing-account", session=composer_service.db.session) def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pytest.MonkeyPatch): @@ -1825,7 +1894,7 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte monkeypatch.setattr(composer_service.db, "session", fake_session) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): raise IntegrityError("insert apps", params, Exception("duplicate")) monkeypatch.setattr(composer_service, "AppService", FakeAppService) @@ -1839,6 +1908,7 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) assert fake_session.rollbacks == 1 @@ -1849,7 +1919,7 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa monkeypatch.setattr(composer_service.db, "session", fake_session) class FakeAppService: - def create_app(self, tenant_id, params, account): + def create_app(self, tenant_id, params, account, session): return SimpleNamespace(id="app-agent-1") class FakeAgentRosterService: @@ -1871,6 +1941,7 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + session=composer_service.db.session, ) @@ -1892,6 +1963,7 @@ def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeyp tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + session=fake_session, ) is False ) @@ -1915,6 +1987,7 @@ def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monke tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + session=fake_session, ) is True ) @@ -1945,6 +2018,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note="note", + session=composer_service.db.session, ) updated_snapshot = AgentComposerService._update_current_version( current_snapshot=AgentConfigSnapshot( @@ -1958,21 +2032,40 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_CURRENT_VERSION, version_note="updated", + session=composer_service.db.session, + ) + workflow = AgentComposerService._get_draft_workflow( + tenant_id="tenant-1", app_id="app-1", session=composer_service.db.session ) - workflow = AgentComposerService._get_draft_workflow(tenant_id="tenant-1", app_id="app-1") with pytest.raises(ValueError): - AgentComposerService._get_draft_workflow(tenant_id="tenant-1", app_id="missing") - assert AgentComposerService._require_agent(tenant_id="tenant-1", agent_id="agent-1").id == "agent-1" - with pytest.raises(composer_service.AgentNotFoundError): - AgentComposerService._require_agent(tenant_id="tenant-1", agent_id=None) - assert AgentComposerService._get_agent_if_present(tenant_id="tenant-1", agent_id="agent-1") is None + AgentComposerService._get_draft_workflow( + tenant_id="tenant-1", app_id="missing", session=composer_service.db.session + ) assert ( - AgentComposerService._require_version(tenant_id="tenant-1", agent_id="agent-1", version_id="version-1").id + AgentComposerService._require_agent( + tenant_id="tenant-1", agent_id="agent-1", session=composer_service.db.session + ).id + == "agent-1" + ) + with pytest.raises(composer_service.AgentNotFoundError): + AgentComposerService._require_agent(tenant_id="tenant-1", agent_id=None, session=composer_service.db.session) + assert ( + AgentComposerService._get_agent_if_present( + tenant_id="tenant-1", agent_id="agent-1", session=composer_service.db.session + ) + is None + ) + assert ( + AgentComposerService._require_version( + tenant_id="tenant-1", agent_id="agent-1", version_id="version-1", session=composer_service.db.session + ).id == "version-1" ) with pytest.raises(composer_service.AgentVersionNotFoundError): - AgentComposerService._require_version(tenant_id="tenant-1", agent_id="agent-1", version_id="missing") + AgentComposerService._require_version( + tenant_id="tenant-1", agent_id="agent-1", version_id="missing", session=composer_service.db.session + ) assert version.version == 2 assert updated_snapshot.version == 3 @@ -2006,7 +2099,11 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc ) result = AgentComposerService._save_to_current_version( - tenant_id="tenant-1", account_id="account-1", binding=binding, payload=payload + tenant_id="tenant-1", + account_id="account-1", + binding=binding, + payload=payload, + session=composer_service.db.session, ) assert result.updated_by == "account-1" @@ -2027,6 +2124,7 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc "save_strategy": ComposerSaveStrategy.SAVE_AS_NEW_AGENT.value, } ), + session=composer_service.db.session, ) @@ -3172,7 +3270,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: captured["tenant_id"] = tenant_id captured["params"] = params captured["account"] = account @@ -3241,7 +3339,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: captured["params"] = params return target_app @@ -3303,7 +3401,7 @@ class TestAgentAppBackingAgent: monkeypatch.setattr(service, "_next_duplicate_agent_name", lambda **_: "Iris copy") class FakeAppService: - def create_app(self, tenant_id: str, params, account: object) -> object: + def create_app(self, tenant_id: str, params, account: object, session: object) -> object: return target_app access_mode_updates = [] @@ -4307,7 +4405,33 @@ def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): assert captured == {} -def test_validate_knowledge_datasets_rejects_malformed_ids_without_dataset_lookup(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + ("variant", "save_call"), + [ + ( + ComposerVariant.AGENT_APP, + lambda payload: AgentComposerService.save_agent_app_composer( + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ( + ComposerVariant.WORKFLOW, + lambda payload: AgentComposerService.save_workflow_composer( + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ], +) +def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pytest.MonkeyPatch, variant, save_call): captured = {"calls": 0} def fake_get_datasets_by_ids(ids, tenant_id): @@ -4342,7 +4466,35 @@ def test_validate_knowledge_datasets_rejects_malformed_ids_without_dataset_looku assert captured == {"calls": 0} -def test_validate_knowledge_datasets_rejects_missing_or_out_of_scope_datasets(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + ("variant", "save_call"), + [ + ( + ComposerVariant.AGENT_APP, + lambda payload: AgentComposerService.save_agent_app_composer( + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ( + ComposerVariant.WORKFLOW, + lambda payload: AgentComposerService.save_workflow_composer( + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + session=composer_service.db.session, + ), + ), + ], +) +def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets( + monkeypatch: pytest.MonkeyPatch, variant, save_call +): captured = {} missing_dataset_id = "550e8400-e29b-41d4-a716-446655440000" @@ -4431,6 +4583,7 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte agent_id="agent-1", account_id="account-1", payload=payload, + session=fake_session, ) assert result["loaded"] is True @@ -4519,6 +4672,7 @@ def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyP tenant_id="tenant-1", agent_id="agent-1", prompt=_drive_soul().prompt.system_prompt, + session=composer_service.db.session, ) assert [(f["code"], f["id"]) for f in findings] == [("mention_target_missing", "files/sample.pdf")] @@ -4534,6 +4688,7 @@ def test_drive_mention_findings_clean_when_all_keys_exist(monkeypatch: pytest.Mo tenant_id="tenant-1", agent_id="agent-1", prompt=_drive_soul().prompt.system_prompt, + session=composer_service.db.session, ) == [] ) @@ -4546,6 +4701,7 @@ def test_drive_mention_findings_skips_prompt_without_drive_mentions(monkeypatch: tenant_id="tenant-1", agent_id="agent-1", prompt=soul.prompt.system_prompt, + session=composer_service.db.session, ) assert findings == [] @@ -4565,7 +4721,10 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c ) findings = AgentComposerService.collect_validation_findings( - tenant_id="tenant-1", payload=payload, agent_id="agent-1" + tenant_id="tenant-1", + payload=payload, + agent_id="agent-1", + session=composer_service.db.session, ) codes = {w["code"] for w in findings["warnings"]} @@ -4575,7 +4734,9 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c "files/sample.pdf", } # without agent context the drive check is skipped entirely - findings_no_agent = AgentComposerService.collect_validation_findings(tenant_id="tenant-1", payload=payload) + findings_no_agent = AgentComposerService.collect_validation_findings( + tenant_id="tenant-1", payload=payload, session=composer_service.db.session + ) assert all(w["code"] != "mention_target_missing" for w in findings_no_agent["warnings"]) @@ -4588,7 +4749,12 @@ def test_resolve_bound_agent_id_queries_active_roster_agent(monkeypatch: pytest. import services.agent.composer_service as module monkeypatch.setattr(module.db, "session", SimpleNamespace(scalar=lambda stmt: "agent-9")) - assert AgentComposerService.resolve_bound_agent_id(tenant_id="t-1", app_id="app-1") == "agent-9" + assert ( + AgentComposerService.resolve_bound_agent_id( + tenant_id="t-1", app_id="app-1", session=composer_service.db.session + ) + == "agent-9" + ) def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(monkeypatch: pytest.MonkeyPatch): @@ -4598,20 +4764,35 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon raise ValueError("no draft workflow") monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", classmethod(boom)) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") is None + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + is None + ) monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", classmethod(lambda cls, **kwargs: None)) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") is None + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + is None + ) monkeypatch.setattr( AgentComposerService, "_get_workflow_binding", classmethod(lambda cls, **kwargs: SimpleNamespace(agent_id="agent-7")), ) - assert AgentComposerService.resolve_workflow_node_agent_id(tenant_id="t", app_id="a", node_id="n") == "agent-7" + assert ( + AgentComposerService.resolve_workflow_node_agent_id( + tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session + ) + == "agent-7" + ) def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only(monkeypatch: pytest.MonkeyPatch): @@ -4654,7 +4835,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( ) guarded: dict[str, str] = {} - def fake_collect(cls, *, tenant_id, payload, agent_id=None): + def fake_collect(cls, *, tenant_id, payload, agent_id=None, session=None): guarded["tenant_id"] = tenant_id guarded["agent_id"] = agent_id return {"warnings": [{"code": "mention_target_missing", "id": "files/sample.pdf"}]} @@ -4662,7 +4843,12 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( monkeypatch.setattr(AgentComposerService, "collect_validation_findings", classmethod(fake_collect)) result = AgentComposerService.save_workflow_composer( - tenant_id="t-1", app_id="app-1", node_id="n-1", account_id="acc-1", payload=payload + tenant_id="t-1", + app_id="app-1", + node_id="n-1", + account_id="acc-1", + payload=payload, + session=composer_service.db.session, ) assert result == { @@ -4712,14 +4898,19 @@ def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only( ) captured: dict[str, str | None] = {} - def fake_collect(cls, *, tenant_id, payload, agent_id=None): + def fake_collect(cls, *, tenant_id, payload, agent_id=None, session=None): captured["agent_id"] = agent_id return {"warnings": []} monkeypatch.setattr(AgentComposerService, "collect_validation_findings", classmethod(fake_collect)) result = AgentComposerService.save_workflow_composer( - tenant_id="t-1", app_id="app-1", node_id="n-1", account_id="acc-1", payload=payload + tenant_id="t-1", + app_id="app-1", + node_id="n-1", + account_id="acc-1", + payload=payload, + session=composer_service.db.session, ) assert result == {"state": "ok", "validation": {"warnings": []}} diff --git a/api/tests/unit_tests/services/agent/test_skill_standardize_service.py b/api/tests/unit_tests/services/agent/test_skill_standardize_service.py index 074ac59bb1c..5b3ade55721 100644 --- a/api/tests/unit_tests/services/agent/test_skill_standardize_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_standardize_service.py @@ -50,6 +50,7 @@ def test_standardize_creates_drive_owned_toolfiles_and_commits_archive_manifest( tenant_id="tenant-1", user_id="user-1", agent_id="agent-1", + session=MagicMock(), ) # ToolFiles: SKILL.md and the full archive. Archive members stay lazy. diff --git a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py index 25678bb4f4d..cfb32d63a92 100644 --- a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py @@ -38,14 +38,17 @@ def test_infer_returns_suggestions_with_inferred_from(monkeypatch): ' "env_suggestions": [{"key": "OPENAI_API_KEY", "reason": "whisper call", "secret_likely": true}]}]}' ) with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + session = MagicMock() + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=session) assert result["inferable"] is True tool = result["cli_tools"][0] assert tool["name"] == "ffmpeg" assert tool["inferred_from"] == "audio-transcribe" assert tool["env_suggestions"] == [{"key": "OPENAI_API_KEY", "reason": "whisper call", "secret_likely": True}] - drive.preview.assert_called_once_with(tenant_id="t-1", agent_id="a-1", key="audio-transcribe/SKILL.md") + drive.preview.assert_called_once_with( + tenant_id="t-1", agent_id="a-1", key="audio-transcribe/SKILL.md", session=session + ) def test_infer_threads_skill_md_into_the_prompt(monkeypatch): @@ -57,7 +60,7 @@ def test_infer_threads_skill_md_into_the_prompt(monkeypatch): return '{"inferable": false, "cli_tools": [], "reason": "none"}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(fake_invoke)): - service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert "Files inside the skill package" not in captured["prompt"] assert "ffmpeg" in captured["prompt"] # SKILL.md body present @@ -67,7 +70,7 @@ def test_infer_not_inferable_passes_reason_through(monkeypatch): service, _ = _service() raw = '{"inferable": false, "cli_tools": [], "reason": "SKILL.md 未描述任何外部命令依赖"}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert result == {"inferable": False, "cli_tools": [], "reason": "SKILL.md 未描述任何外部命令依赖"} @@ -81,7 +84,7 @@ def test_infer_retries_once_then_422(monkeypatch): with patch.object(SkillToolInferenceService, "_invoke", staticmethod(bad_invoke)): with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert len(calls) == 2 # one retry assert exc_info.value.code == "inference_failed" @@ -92,7 +95,7 @@ def test_infer_repairs_slightly_malformed_json(monkeypatch): service, _ = _service() raw = 'Here you go: {"inferable": true, "cli_tools": [], "reason": null,}' with patch.object(SkillToolInferenceService, "_invoke", staticmethod(lambda **kwargs: raw)): - result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe") + result = service.infer(tenant_id="t-1", agent_id="a-1", slug="audio-transcribe", session=MagicMock()) assert result["inferable"] is True @@ -102,7 +105,7 @@ def test_missing_skill_maps_to_404(): service = SkillToolInferenceService(drive_service=drive) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="ghost") + service.infer(tenant_id="t-1", agent_id="a-1", slug="ghost", session=MagicMock()) assert exc_info.value.code == "skill_not_found" assert exc_info.value.status_code == 404 @@ -110,7 +113,7 @@ def test_missing_skill_maps_to_404(): def test_binary_skill_md_maps_to_404(): service, _ = _service(preview={"key": "x/SKILL.md", "size": 1, "truncated": False, "binary": True, "text": None}) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="x") + service.infer(tenant_id="t-1", agent_id="a-1", slug="x", session=MagicMock()) assert exc_info.value.code == "skill_not_found" @@ -160,5 +163,5 @@ def test_load_skill_md_passes_through_non_missing_drive_errors(): service = SkillToolInferenceService(drive_service=drive) with pytest.raises(SkillToolInferenceError) as exc_info: - service.infer(tenant_id="t-1", agent_id="a-1", slug="x") + service.infer(tenant_id="t-1", agent_id="a-1", slug="x", session=MagicMock()) assert exc_info.value.code == "agent_not_found" diff --git a/api/tests/unit_tests/services/data_migration/test_export_service.py b/api/tests/unit_tests/services/data_migration/test_export_service.py index f5480ff52af..5479de5ba22 100644 --- a/api/tests/unit_tests/services/data_migration/test_export_service.py +++ b/api/tests/unit_tests/services/data_migration/test_export_service.py @@ -1,3 +1,5 @@ +from unittest.mock import MagicMock + import pytest from services.data_migration.dependency_discovery_service import DiscoveredDependency @@ -126,13 +128,13 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): report_items = [] service._export_mcp_tools( - object(), tenant_id="tenant-1", provider_ids=["mcp-1"], include_secrets=False, exported_mcp_tools=mcp_tools, dependencies=dependencies, report_items=report_items, + session=MagicMock(), ) assert mcp_tools == [] @@ -151,12 +153,14 @@ def test_secret_free_mcp_dependencies_are_dependency_only(): def test_get_mcp_provider_does_not_compare_non_uuid_identifier_to_uuid_id(): statements = [] - class StubSession: - def scalar(self, statement): - statements.append(str(statement)) + def capture_scalar(statement): + statements.append(str(statement)) + + session = MagicMock() + session.scalar.side_effect = capture_scalar with pytest.raises(MigrationDataError, match="MCP provider not found"): - MigrationExportService()._get_mcp_provider(StubSession(), "tenant-1", "my-test-mcp") + MigrationExportService()._get_mcp_provider("tenant-1", "my-test-mcp", session=session) assert len(statements) == 1 assert "tool_mcp_providers.id =" not in statements[0] diff --git a/api/tests/unit_tests/services/data_migration/test_import_service.py b/api/tests/unit_tests/services/data_migration/test_import_service.py index 2b11d575ed6..10460aba470 100644 --- a/api/tests/unit_tests/services/data_migration/test_import_service.py +++ b/api/tests/unit_tests/services/data_migration/test_import_service.py @@ -3,6 +3,7 @@ import yaml from models.tools import MCPToolProvider, WorkflowToolProvider from services.app_dsl_service import Import +from services.data_migration import import_service from services.data_migration.entities import ( ConflictStrategy, IdStrategy, @@ -92,8 +93,12 @@ def test_package_target_tenant_id_ignores_invalid_uuid(monkeypatch): return EmptyResult() + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + with pytest.raises(MigrationDataError, match="Target tenant not found"): - ImportTargetResolver().resolve(StubSession(), ImportRequest(package=package)) + ImportTargetResolver().resolve(ImportRequest(package=package), session=import_service.db.session) def test_options_override_replaces_package_defaults(): @@ -113,7 +118,7 @@ def test_options_override_replaces_package_defaults(): captured_options: list[ImportOptions] = [] class StubResolver(ImportTargetResolver): - def resolve(self, session, request: ImportRequest) -> ImportTarget: + def resolve(self, request: ImportRequest, session) -> ImportTarget: return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -124,7 +129,6 @@ def test_options_override_replaces_package_defaults(): class CapturingImportService(MigrationImportService): def _import_workflows( self, - session, package: MigrationPackage, target: ImportTarget, options: ImportOptions, @@ -137,7 +141,8 @@ def test_options_override_replaces_package_defaults(): override = ImportOptions(create_app_api_token_on_import=False, conflict_strategy=ConflictStrategy.SKIP) CapturingImportService(target_resolver=StubResolver()).import_package( - object(), ImportRequest(package=package, options_override=override) + ImportRequest(package=package, options_override=override), + session=import_service.db.session, ) assert captured_options == [override] @@ -150,37 +155,53 @@ def test_only_preserve_id_strategy_reuses_source_app_id(): assert service._should_preserve_source_app_id(ImportOptions(id_strategy=IdStrategy.GENERATE_NEW_ID)) is False -def test_find_existing_app_ignores_invalid_uuid(): +def test_find_existing_app_ignores_invalid_uuid(monkeypatch): class StubSession: def scalar(self, statement): raise AssertionError("invalid UUID should not be queried against App.id") - assert MigrationImportService()._find_existing_app(StubSession(), "not-a-uuid", "tenant-1") is None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + + assert ( + MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1", session=import_service.db.session) is None + ) -def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(): +def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + MigrationImportService()._find_existing_workflow_tool( - StubSession(), "tenant-1", "not-a-uuid", "tool-name", "app-id" + "tenant-1", "not-a-uuid", "tool-name", "app-id", session=import_service.db.session ) where_clause = str(captured[0].whereclause) assert f"{WorkflowToolProvider.__tablename__}.id" not in where_clause -def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(): +def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) - MigrationImportService()._find_existing_mcp_tool(StubSession(), "tenant-1", "my-test-mcp", "my-test-mcp") + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + + MigrationImportService()._find_existing_mcp_tool( + "tenant-1", "my-test-mcp", "my-test-mcp", session=import_service.db.session + ) where_clause = str(captured[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause @@ -211,16 +232,17 @@ def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "AppDslService", StubAppDslService) imported_app_id = MigrationImportService()._import_workflow_app( - session=StubSession(), account=object(), workflow_data={"name": "main_chatflow"}, dsl_content="app:\n mode: workflow\n", app_id="source-app-id", existing_app=None, options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), + session=import_service.db.session, ) assert imported_app_id == "imported-app-id" @@ -317,19 +339,20 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch return account class PublishingImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): if ("created", app_id) in events: return type("WorkflowToolProvider", (), {"id": workflow_tool_id or "created-workflow-tool-id"})() return None - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): events.append(("published", app_id)) from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -337,7 +360,6 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch ) PublishingImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -360,6 +382,7 @@ def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch {}, [], [], + session=import_service.db.session, ) assert events == [("published", "workflow-app-1"), ("created", "workflow-app-1")] @@ -384,17 +407,18 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP return account class StrategyImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): return target_provider if created_kwargs else None - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.WorkflowToolManageService, "create_workflow_tool", @@ -402,7 +426,6 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP ) StrategyImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -425,6 +448,7 @@ def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyP id_mapping, id_mapping_details, [], + session=import_service.db.session, ) assert created_kwargs[0]["import_id"] == expected_import_id @@ -449,17 +473,20 @@ def test_workflow_tool_skip_records_id_mapping(monkeypatch): return account class SkipImportService(MigrationImportService): - def _find_existing_app(self, session, app_id, tenant_id): + def _find_existing_app(self, app_id, tenant_id, session): return object() - def _find_existing_workflow_tool(self, session, tenant_id, workflow_tool_id, tool_name, app_id): + def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): return existing_provider - def _ensure_workflow_app_is_published(self, session, target, account, app_id): + def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + SkipImportService()._import_workflow_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -482,6 +509,7 @@ def test_workflow_tool_skip_records_id_mapping(monkeypatch): id_mapping, [], [], + session=import_service.db.session, ) assert id_mapping["source-workflow-tool-id"] == "existing-workflow-tool-id" @@ -495,22 +523,18 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str report_items = [] class ExistingApiImportService(MigrationImportService): - def _find_api_tool_provider(self, session, tenant_id, provider_name): - return target_provider - - class StubSession: - def scalar(self, statement): + def _find_api_tool_provider(self, tenant_id, provider_name, session): return target_provider from services.data_migration import import_service + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: target_provider) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "update_api_tool_provider", lambda **kwargs: None) ExistingApiImportService()._import_api_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -528,6 +552,7 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str id_mapping, id_mapping_details, {"weather": {"source-api-provider-id-from-dsl"}}, + session=import_service.db.session, ) assert id_mapping == { @@ -549,18 +574,18 @@ def test_api_tool_create_records_id_mapping(monkeypatch): return None class CreatedApiImportService(MigrationImportService): - def _find_api_tool_provider(self, session, tenant_id, provider_name): + def _find_api_tool_provider(self, tenant_id, provider_name, session): return target_provider from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", lambda **kwargs: None) CreatedApiImportService()._import_api_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -578,6 +603,7 @@ def test_api_tool_create_records_id_mapping(monkeypatch): id_mapping, [], {}, + session=import_service.db.session, ) assert id_mapping["source-api-provider-id"] == "target-api-provider-id" @@ -605,10 +631,10 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) MigrationImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -634,6 +660,7 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): report_items, {}, [], + session=import_service.db.session, ) assert provider.tools == '[{"name": "echo"}]' @@ -653,7 +680,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str return None class ExistingMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): return provider class StubMCPToolManageService: @@ -665,10 +692,10 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) ExistingMCPImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -694,6 +721,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str [], id_mapping, id_mapping_details, + session=import_service.db.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -715,7 +743,7 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): return None class CreatedMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, session, tenant_id, provider_id, server_identifier): + def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): return provider if provider_created else None class StubMCPToolManageService: @@ -728,10 +756,10 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): from services.data_migration import import_service + monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) CreatedMCPImportService()._import_mcp_tools( - StubSession(), MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -756,6 +784,7 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): [], id_mapping, [], + session=import_service.db.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -800,12 +829,11 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work } ) - class StubSession: - def scalar(self, statement): - return None + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: None) MigrationImportService()._preflight_dependency_only_mcp( - StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -814,6 +842,7 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work operator_email="owner@example.com", ), report_items, + session=import_service.db.session, ) assert report_items == [ @@ -828,18 +857,22 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work ] -def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(): +def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(monkeypatch): captured = [] class StubSession: def scalar(self, statement): captured.append(statement) + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db, "session", StubSession()) + MigrationImportService()._find_dependency_only_mcp_provider( - StubSession(), "tenant-1", "my-test-mcp-server", "my-test-mcp", + session=import_service.db.session, ) where_clause = str(captured[0].whereclause) @@ -860,12 +893,11 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp {"id": "target-provider-id", "name": "my-test-mcp", "server_identifier": "my-test-mcp-server"}, )() - class StubSession: - def scalar(self, statement): - return provider + from services.data_migration import import_service + + monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: provider) MigrationImportService()._preflight_dependency_only_mcp( - StubSession(), package, ImportTarget( tenant_id="tenant-1", @@ -874,6 +906,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp operator_email="owner@example.com", ), report_items, + session=import_service.db.session, ) assert report_items == [ @@ -891,7 +924,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): events = [] class StubResolver(ImportTargetResolver): - def resolve(self, session, request): + def resolve(self, request, session): return ImportTarget( tenant_id="tenant-1", tenant_name="target", @@ -902,7 +935,6 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): class OrderedImportService(MigrationImportService): def _import_api_tools( self, - session, package, target, options, @@ -910,12 +942,13 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): id_mapping, id_mapping_details, source_provider_ids_by_name, + *, + session=None, ): events.append(("api_tools", "imported")) def _import_workflows( self, - session, package, target, options, @@ -926,6 +959,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): imported_workflow_ids=None, only_app_ids=None, skip_app_ids=None, + session=None, ): only_app_ids = set(only_app_ids or []) skip_app_ids = set(skip_app_ids or []) @@ -941,11 +975,13 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): imported_workflow_ids.add(app_id) def _import_workflow_tools( - self, session, package, target, options, id_mapping, id_mapping_details, report_items + self, package, target, options, id_mapping, id_mapping_details, report_items, *, session=None ): events.append(("workflow_tool", package.workflow_tools[0]["id"])) - def _import_mcp_tools(self, session, package, target, options, report_items, id_mapping, id_mapping_details): + def _import_mcp_tools( + self, package, target, options, report_items, id_mapping, id_mapping_details, *, session=None + ): events.append(("mcp_tools", "imported")) package = MigrationPackage.from_mapping( @@ -959,7 +995,9 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): } ) - OrderedImportService(target_resolver=StubResolver()).import_package(object(), ImportRequest(package=package)) + OrderedImportService(target_resolver=StubResolver()).import_package( + ImportRequest(package=package), session=import_service.db.session + ) assert events == [ ("api_tools", "imported"), diff --git a/api/tests/unit_tests/services/enterprise/test_rbac_service.py b/api/tests/unit_tests/services/enterprise/test_rbac_service.py index fdf921265b2..85638b11fff 100644 --- a/api/tests/unit_tests/services/enterprise/test_rbac_service.py +++ b/api/tests/unit_tests/services/enterprise/test_rbac_service.py @@ -558,7 +558,7 @@ class TestMyPermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" @@ -613,11 +613,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = role - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) mock_send.assert_not_called() assert out.workspace.permission_keys == workspace_keys @@ -655,11 +652,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = role - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) actual_snippet_keys = { permission_key for permission_key in out.workspace.permission_keys if permission_key.startswith("snippets.") @@ -672,11 +666,8 @@ class TestMyPermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = None - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", session=mock_session) mock_send.assert_not_called() assert out.workspace.permission_keys == [] @@ -694,7 +685,7 @@ class TestMyPermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1") + out = svc.RBACService.MyPermissions.get("tenant-1", "acct-1", app_id="app-1", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" @@ -716,7 +707,7 @@ class TestMemberRoles: ], } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2") + out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=MagicMock()) call = _call_args(mock_send) assert call.method == "GET" assert call.endpoint == "/rbac/members/rbac-roles" @@ -728,12 +719,8 @@ class TestMemberRoles: session = MagicMock() session.scalar.return_value = svc.TenantAccountRole.EDITOR - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session") as create_session, - ): - create_session.return_value.__enter__.return_value = session - out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2") + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.get("tenant-1", "acct-1", "acct-2", session=session) mock_send.assert_not_called() assert out.account_id == "acct-2" @@ -755,7 +742,11 @@ class TestMemberRoles: mock_send.return_value = {"account_id": "acct-2", "roles": []} with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): svc.RBACService.MemberRoles.replace( - "tenant-1", "acct-1", "acct-2", role_ids=["workspace.owner", "workspace.editor"] + "tenant-1", + "acct-1", + "acct-2", + role_ids=["workspace.owner", "workspace.editor"], + session=MagicMock(), ) call = _call_args(mock_send) assert call.method == "PUT" @@ -769,11 +760,10 @@ class TestMemberRoles: target_join = SimpleNamespace(role=svc.TenantAccountRole.NORMAL, account_id="acct-2") session.scalar.return_value = target_join - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=session), - ): - out = svc.RBACService.MemberRoles.replace("tenant-1", "acct-1", "acct-2", role_ids=["editor"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.replace( + "tenant-1", "acct-1", "acct-2", role_ids=["editor"], session=session + ) mock_send.assert_not_called() session.commit.assert_called_once() @@ -789,11 +779,10 @@ class TestMemberRoles: owner_join = SimpleNamespace(role=svc.TenantAccountRole.OWNER, account_id="acct-owner") session.scalar.side_effect = [target_join, owner_join] - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=session), - ): - out = svc.RBACService.MemberRoles.replace("tenant-1", "acct-1", "acct-2", role_ids=["owner"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.MemberRoles.replace( + "tenant-1", "acct-1", "acct-2", role_ids=["owner"], session=session + ) mock_send.assert_not_called() session.commit.assert_called_once() @@ -832,7 +821,9 @@ class TestResourcePermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.AppPermissions.batch_get("tenant-1", "acct-1", ["app-1", "app-2"]) + out = svc.RBACService.AppPermissions.batch_get( + "tenant-1", "acct-1", ["app-1", "app-2"], session=MagicMock() + ) call = _call_args(mock_send) assert call.method == "POST" @@ -847,11 +838,10 @@ class TestResourcePermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = "editor" - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.AppPermissions.batch_get("tenant-1", "acct-1", ["app-1", "app-2"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.AppPermissions.batch_get( + "tenant-1", "acct-1", ["app-1", "app-2"], session=mock_session + ) mock_send.assert_not_called() assert out == { @@ -868,7 +858,9 @@ class TestResourcePermissions: } with patch(f"{MODULE}.dify_config.RBAC_ENABLED", True): - out = svc.RBACService.DatasetPermissions.batch_get("tenant-1", "acct-1", ["ds-1", "ds-2"]) + out = svc.RBACService.DatasetPermissions.batch_get( + "tenant-1", "acct-1", ["ds-1", "ds-2"], session=MagicMock() + ) call = _call_args(mock_send) assert call.method == "POST" @@ -883,11 +875,10 @@ class TestResourcePermissions: mock_session = MagicMock() mock_session.__enter__.return_value = mock_session mock_session.scalar.return_value = "dataset_operator" - with ( - patch(f"{MODULE}.dify_config.RBAC_ENABLED", False), - patch(f"{MODULE}.session_factory.create_session", return_value=mock_session), - ): - out = svc.RBACService.DatasetPermissions.batch_get("tenant-1", "acct-1", ["ds-1", "ds-2"]) + with patch(f"{MODULE}.dify_config.RBAC_ENABLED", False): + out = svc.RBACService.DatasetPermissions.batch_get( + "tenant-1", "acct-1", ["ds-1", "ds-2"], session=mock_session + ) mock_send.assert_not_called() assert out == { diff --git a/api/tests/unit_tests/services/hit_service.py b/api/tests/unit_tests/services/hit_service.py index ae19daba898..ffeb158e37a 100644 --- a/api/tests/unit_tests/services/hit_service.py +++ b/api/tests/unit_tests/services/hit_service.py @@ -186,7 +186,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -234,7 +234,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -292,7 +292,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -337,7 +337,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -380,7 +380,7 @@ class TestHitTestingServiceRetrieve: # Act result = HitTestingService.retrieve( - mock_db_session, dataset, query, account, retrieval_model, external_retrieval_model + dataset, query, account, retrieval_model, external_retrieval_model, session=mock_db_session ) # Assert @@ -438,7 +438,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -469,7 +474,7 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, query, account, external_retrieval_model, metadata_filtering_conditions, session=mock_db_session ) # Assert @@ -504,7 +509,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -538,7 +548,12 @@ class TestHitTestingServiceExternalRetrieve: # Act result = HitTestingService.external_retrieve( - mock_db_session, dataset, query, account, external_retrieval_model, metadata_filtering_conditions + dataset, + query, + account, + external_retrieval_model, + metadata_filtering_conditions, + session=mock_db_session, ) # Assert @@ -579,7 +594,7 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = mock_records # Act - result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents) + result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) # Assert assert result["query"]["content"] == query @@ -605,7 +620,7 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = [] # Act - result = HitTestingService.compact_retrieve_response(MagicMock(), query, documents) + result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) # Assert assert result["query"]["content"] == query diff --git a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py index 0e793fae7ef..e66bb3fff04 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py @@ -15,47 +15,40 @@ PLUGIN_CATEGORY = TenantPluginAutoUpgradeCategory.TOOL def _patched_session(): - """Patch session_factory.create_session() to return a mock session as context manager.""" + """Return a mock SQLAlchemy session for service calls.""" session = MagicMock() - session.__enter__ = MagicMock(return_value=session) - session.__exit__ = MagicMock(return_value=False) - mock_factory = MagicMock() - mock_factory.create_session.return_value = session - patcher = patch(f"{MODULE}.session_factory", mock_factory) - return patcher, session + return session class TestGetStrategy: def test_returns_strategy_when_found(self): - p1, session = _patched_session() + session = _patched_session() strategy = MagicMock() session.scalar.return_value = strategy - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) assert result is strategy def test_returns_none_when_not_found(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) assert result is None class TestChangeStrategy: def test_creates_new_strategy(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None - with p1, patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: strat_cls.return_value = MagicMock() from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService @@ -67,28 +60,29 @@ class TestChangeStrategy: [], [], category=PLUGIN_CATEGORY, + session=session, ) assert result is True session.add.assert_called_once() def test_updates_existing_strategy(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() session.scalar.return_value = existing - with p1: - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.change_strategy( - "t1", - TenantPluginAutoUpgradeStrategySetting.LATEST, - 5, - TenantPluginAutoUpgradeMode.PARTIAL, - ["p1"], - ["p2"], - category=PLUGIN_CATEGORY, - ) + result = PluginAutoUpgradeService.change_strategy( + "t1", + TenantPluginAutoUpgradeStrategySetting.LATEST, + 5, + TenantPluginAutoUpgradeMode.PARTIAL, + ["p1"], + ["p2"], + category=PLUGIN_CATEGORY, + session=session, + ) assert result is True assert existing.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST @@ -100,11 +94,10 @@ class TestChangeStrategy: class TestExcludePlugin: def test_creates_default_strategy_when_none_exists(self): - p1, session = _patched_session() + session = _patched_session() session.scalar.return_value = None with ( - p1, patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy"), ): @@ -114,74 +107,87 @@ class TestExcludePlugin: "t1", "plugin-1", PLUGIN_CATEGORY, + session=session, ) assert result is True session.add.assert_called_once() def test_appends_to_exclude_list_in_exclude_mode(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE existing.exclude_plugins = ["p-existing"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY, session=session) assert result is True assert existing.exclude_plugins == ["p-existing", "p-new"] def test_removes_from_include_list_in_partial_mode(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.PARTIAL existing.include_plugins = ["p1", "p2"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert result is True assert existing.include_plugins == ["p2"] def test_switches_to_exclude_mode_from_all(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.ALL session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert result is True assert existing.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert existing.exclude_plugins == ["p1"] def test_no_duplicate_in_exclude_list(self): - p1, session = _patched_session() + session = _patched_session() existing = MagicMock() existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE existing.exclude_plugins = ["p1"] session.scalar.return_value = existing - with p1, patch(f"{MODULE}.select"): + with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: + strat_cls.UpgradeMode.EXCLUDE = "exclude" + strat_cls.UpgradeMode.PARTIAL = "partial" + strat_cls.UpgradeMode.ALL = "all" from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY) + PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) assert existing.exclude_plugins == ["p1"] class TestBackfillStrategyCategories: def test_creates_default_missing_categories_without_fetching_daemon(self): - p1, session = _patched_session() + session = _patched_session() tool_strategy = SimpleNamespace( category=TenantPluginAutoUpgradeCategory.TOOL, strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, @@ -193,10 +199,10 @@ class TestBackfillStrategyCategories: session.scalars.return_value.all.return_value = [tool_strategy] installer = MagicMock() - with p1, patch(f"{MODULE}.PluginInstaller", return_value=installer): + with patch(f"{MODULE}.PluginInstaller", return_value=installer): from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.backfill_strategy_categories("t1") + result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1") assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 1 @@ -219,7 +225,7 @@ class TestBackfillStrategyCategories: assert 0 <= default_time < 24 * 60 * 60 def test_creates_missing_categories_and_splits_known_plugins(self, caplog: pytest.LogCaptureFixture): - p1, session = _patched_session() + session = _patched_session() tool_strategy = SimpleNamespace( category=TenantPluginAutoUpgradeCategory.TOOL, strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, @@ -252,13 +258,12 @@ class TestBackfillStrategyCategories: installer.list_plugins.return_value = installed_plugins with ( - p1, patch(f"{MODULE}.PluginInstaller", return_value=installer), caplog.at_level(logging.WARNING, logger=MODULE), ): from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - result = PluginAutoUpgradeService.backfill_strategy_categories("t1") + result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 2 assert result.normalized is True diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py index 5bc41fdb5bd..a7bb8cfeed6 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py @@ -22,8 +22,9 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: }, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - templates = retrieval.get_pipeline_templates(mocker.Mock(), "en-US") + templates = retrieval.get_pipeline_templates("en-US", session=session) assert templates == {"pipeline_templates": [{"id": "tpl-1"}]} @@ -39,8 +40,9 @@ def test_get_pipeline_template_detail(mocker: MockerFixture) -> None: }, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - detail = retrieval.get_pipeline_template_detail(mocker.Mock(), "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session) assert detail == {"id": "tpl-1", "name": "Template 1"} @@ -52,8 +54,9 @@ def test_get_pipeline_templates_missing_language_returns_empty_dict(mocker: Mock return_value={"pipeline_templates": {}}, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_templates(mocker.Mock(), "fr-FR") + result = retrieval.get_pipeline_templates("fr-FR", session=session) assert result == {} @@ -65,8 +68,9 @@ def test_get_pipeline_template_detail_returns_none_for_unknown_id(mocker: Mocker return_value={"pipeline_templates": {"tpl-1": {"id": "tpl-1"}}}, ) retrieval = BuiltInPipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_template_detail(mocker.Mock(), "nonexistent-id") + result = retrieval.get_pipeline_template_detail("nonexistent-id", session=session) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py index b3ef79961d3..b3befeb41fd 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py @@ -21,7 +21,7 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: session_mock.scalars.return_value = scalars_mock retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates(session_mock, "en-US", "tenant-id") + result = retrieval.get_pipeline_templates("en-US", "tenant-id", session=session_mock) assert retrieval.get_type() == PipelineTemplateType.CUSTOMIZED assert result == { @@ -51,7 +51,7 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N ) retrieval = CustomizedPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session_mock) assert detail == { "id": "tpl-1", @@ -70,6 +70,6 @@ def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: Mocker session_mock.get.return_value = None retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(session_mock, "missing") + result = retrieval.get_pipeline_template_detail("missing", session=session_mock) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py index cae79175b1c..48ae26ce3aa 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py @@ -23,7 +23,7 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: session_mock.scalars.return_value = scalars_mock retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates(session_mock, "en-US") + result = retrieval.get_pipeline_templates("en-US", session=session_mock) assert retrieval.get_type() == PipelineTemplateType.DATABASE assert result == { @@ -54,7 +54,7 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N ) retrieval = DatabasePipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") + detail = retrieval.get_pipeline_template_detail("tpl-1", session=session_mock) assert detail == { "id": "tpl-1", @@ -72,6 +72,6 @@ def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: Mocker session_mock.get.return_value = None retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(session_mock, "missing") + result = retrieval.get_pipeline_template_detail("missing", session=session_mock) assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py index c8af1869732..17cd5db7ab3 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py @@ -4,11 +4,11 @@ from services.rag_pipeline.pipeline_template.pipeline_template_base import Pipel class DummyRetrieval(PipelineTemplateRetrievalBase): - def get_pipeline_templates(self, session: Mock, language: str, current_tenant_id: str | None = None) -> dict: - del session, current_tenant_id + def get_pipeline_templates(self, language: str, *, session) -> dict: + del session return {"language": language} - def get_pipeline_template_detail(self, session: Mock, template_id: str) -> dict | None: + def get_pipeline_template_detail(self, template_id: str, *, session) -> dict | None: del session return {"id": template_id} @@ -20,6 +20,6 @@ def test_pipeline_template_retrieval_base_concrete_implementation() -> None: retrieval = DummyRetrieval() session = Mock() - assert retrieval.get_pipeline_templates(session, "en-US") == {"language": "en-US"} - assert retrieval.get_pipeline_template_detail(session, "tpl-1") == {"id": "tpl-1"} + assert retrieval.get_pipeline_templates("en-US", session=session) == {"language": "en-US"} + assert retrieval.get_pipeline_template_detail("tpl-1", session=session) == {"id": "tpl-1"} assert retrieval.get_type() == "dummy" diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py index 8f55b4b1c2f..78e46d272c2 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py @@ -20,12 +20,12 @@ def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFix retrieval = RemotePipelineTemplateRetrieval() session = mocker.Mock() - result = retrieval.get_pipeline_templates(session, "en-US") + result = retrieval.get_pipeline_templates("en-US", session=session) assert retrieval.get_type() == PipelineTemplateType.REMOTE assert result == {"pipeline_templates": [{"id": "db-1"}]} fetch_mock.assert_called_once_with("en-US") - fallback_mock.assert_called_once_with(session, "en-US") + fallback_mock.assert_called_once_with("en-US", session=session) def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: MockerFixture) -> None: @@ -42,11 +42,11 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: Moc retrieval = RemotePipelineTemplateRetrieval() session = mocker.Mock() - result = retrieval.get_pipeline_template_detail(session, "tpl-1") + result = retrieval.get_pipeline_template_detail("tpl-1", session=session) assert result == {"id": "db-1"} fetch_mock.assert_called_once_with("tpl-1") - fallback_mock.assert_called_once_with(session, "tpl-1") + fallback_mock.assert_called_once_with("tpl-1", session=session) def test_fetch_pipeline_templates_from_dify_official(mocker: MockerFixture) -> None: diff --git a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py index 0ae2ba97f1a..c8992585653 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py @@ -1,30 +1,15 @@ -from collections.abc import Iterator from types import SimpleNamespace from typing import cast -from uuid import uuid4 import pytest from pytest_mock import MockerFixture -from sqlalchemy import create_engine, func, select -from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom -from models.dataset import Document, Pipeline -from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus +from models.dataset import Pipeline from models.model import Account, App, EndUser from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService -@pytest.fixture -def document_session() -> Iterator[Session]: - engine = create_engine("sqlite:///:memory:") - Document.__table__.create(engine) - session_factory = sessionmaker(bind=engine, expire_on_commit=False) - with session_factory() as session: - yield session - engine.dispose() - - def test_get_max_active_requests_uses_smallest_non_zero_limit(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_DEFAULT_ACTIVE_REQUESTS", 5) mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_MAX_ACTIVE_REQUESTS", 3) @@ -62,12 +47,13 @@ def test_get_workflow(mocker: MockerFixture, invoke_from, workflow, expected_err rag_pipeline_service.get_published_workflow.return_value = workflow pipeline = cast(Pipeline, SimpleNamespace(id="pipeline-1")) + session = mocker.Mock() if expected_error: with pytest.raises(ValueError, match=expected_error): - PipelineGenerateService._get_workflow(pipeline, invoke_from) + PipelineGenerateService._get_workflow(pipeline, invoke_from, session) else: - result = PipelineGenerateService._get_workflow(pipeline, invoke_from) + result = PipelineGenerateService._get_workflow(pipeline, invoke_from, session) assert result == workflow @@ -75,10 +61,10 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke pipeline = cast(Pipeline, SimpleNamespace(id="pipeline-1")) user = cast(Account | EndUser, SimpleNamespace(id="user-1")) args = {"original_document_id": "doc-1", "query": "hello"} + session_mock = mocker.Mock() mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) update_status_mock = mocker.patch.object(PipelineGenerateService, "update_document_status") - session = mocker.Mock() generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") generator_instance = generator_cls.return_value @@ -86,49 +72,39 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke generator_cls.convert_to_event_stream.return_value = "stream-events" result = PipelineGenerateService.generate( - session=session, pipeline=pipeline, user=user, args=args, invoke_from=InvokeFrom.WEB_APP, streaming=True, + session=session_mock, ) assert result == "stream-events" - update_status_mock.assert_called_once_with("doc-1", session) + update_status_mock.assert_called_once_with("doc-1", session=session_mock) -def test_update_document_status_updates_existing_document(document_session: Session) -> None: - session = document_session - document_id = str(uuid4()) - document = Document( - id=document_id, - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - position=1, - data_source_type=DataSourceType.UPLOAD_FILE, - batch="batch-1", - name="Doc", - created_from=DocumentCreatedFrom.WEB, - created_by=str(uuid4()), - indexing_status=IndexingStatus.COMPLETED, - ) - session.add(document) - session.commit() +def test_update_document_status_updates_existing_document(mocker: MockerFixture) -> None: + document = SimpleNamespace(indexing_status="completed") - PipelineGenerateService.update_document_status(document_id, session) + session_mock = mocker.Mock() + session_mock.get.return_value = document + add_mock = session_mock.add - updated_document = session.get(Document, document_id) - assert updated_document is not None - assert updated_document.indexing_status == IndexingStatus.WAITING + PipelineGenerateService.update_document_status("doc-1", session=session_mock) + + assert document.indexing_status == "waiting" + add_mock.assert_called_once_with(document) -def test_update_document_status_skips_when_document_missing(document_session: Session) -> None: - session = document_session +def test_update_document_status_skips_when_document_missing(mocker: MockerFixture) -> None: + session_mock = mocker.Mock() + session_mock.get.return_value = None + add_mock = session_mock.add - PipelineGenerateService.update_document_status(str(uuid4()), session) + PipelineGenerateService.update_document_status("missing", session=session_mock) - assert session.scalar(select(func.count()).select_from(Document)) == 0 + add_mock.assert_not_called() # --- generate_single_iteration --- @@ -144,8 +120,9 @@ def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) + session = mocker.Mock() - result = PipelineGenerateService.generate_single_iteration(pipeline, user, "node-1", {"key": "val"}) + result = PipelineGenerateService.generate_single_iteration(pipeline, user, "node-1", {"key": "val"}, session) assert result == "stream-iter" generator_instance.single_iteration_generate.assert_called_once() @@ -164,8 +141,9 @@ def test_generate_single_loop_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) + session = mocker.Mock() - result = PipelineGenerateService.generate_single_loop(pipeline, user, "node-1", {"key": "val"}) + result = PipelineGenerateService.generate_single_loop(pipeline, user, "node-1", {"key": "val"}, session) assert result == "stream-loop" generator_instance.single_loop_generate.assert_called_once() diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py index 37141c97c83..0d74b3abf9c 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py @@ -49,9 +49,8 @@ def rag_pipeline_service(mocker: MockerFixture) -> RagPipelineServiceTestContext ) session = mocker.Mock() session_maker = _make_mock_session_maker(mocker, session) - mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock())) - service = RagPipelineService(session_maker=session_maker) + service = RagPipelineService(session=session, session_maker=session_maker) return RagPipelineServiceTestContext(service=service, session=session, session_maker=session_maker) @@ -156,10 +155,10 @@ def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_resul builtin_retrieval.fetch_pipeline_templates_from_builtin.return_value = {"pipeline_templates": [{"id": "builtin-1"}]} factory_mock.get_built_in_pipeline_template_retrieval.return_value = builtin_retrieval - result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="ja-JP") + result = RagPipelineService.get_pipeline_templates(type="built-in", language="ja-JP", session=session) assert result == {"pipeline_templates": [{"id": "builtin-1"}]} - remote_retrieval.get_pipeline_templates.assert_called_once_with(session, "ja-JP", None) + remote_retrieval.get_pipeline_templates.assert_called_once_with("ja-JP", None, session=session) builtin_retrieval.fetch_pipeline_templates_from_builtin.assert_called_once_with("en-US") @@ -171,11 +170,11 @@ def test_get_pipeline_templates_customized_mode_uses_customized_factory(mocker: factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_templates(session, type="customized", language="en-US") + result = RagPipelineService.get_pipeline_templates(type="customized", language="en-US", session=session) assert result == {"pipeline_templates": [{"id": "custom-1"}]} factory_mock.get_pipeline_template_factory.assert_called_with("customized") - retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) + retrieval.get_pipeline_templates.assert_called_once_with("en-US", None, session=session) @pytest.mark.parametrize("template_type", ["built-in", "customized"]) @@ -188,12 +187,12 @@ def test_get_pipeline_template_detail_uses_expected_mode(mocker: MockerFixture, factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_template_detail(session, "tpl-1", type=template_type) + result = RagPipelineService.get_pipeline_template_detail("tpl-1", type=template_type, session=session) assert result == {"id": "tpl-1"} expected_mode = "remote" if template_type == "built-in" else "customized" factory_mock.get_pipeline_template_factory.assert_called_with(expected_mode) - retrieval.get_pipeline_template_detail.assert_called_once_with(session, "tpl-1") + retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", session=session) def test_get_published_workflow_returns_none_when_pipeline_has_no_workflow_id( @@ -845,13 +844,14 @@ def test_publish_customized_pipeline_template_success( # 2. Run test args = {"name": "New Template", "description": "Desc", "icon_info": {"icon": "star"}, "tags": ["tag1"]} - rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1", session=session) # 3. Assertions # Verify a new template was added to session or similar? # Since we can't easily check the session inside the context manager with Mock, # we just check that no error was raised and DSL was exported. - mock_dsl_service.export_rag_pipeline_dsl.assert_called_once() + pipeline.retrieve_dataset.assert_called_once_with(session=session) + mock_dsl_service.export_rag_pipeline_dsl.assert_called_once_with(pipeline=pipeline, include_secret=True) # --- get_datasource_plugins --- @@ -863,7 +863,7 @@ def test_get_datasource_plugins_success( # 1. Setup mocks dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = mocker.Mock() workflow.graph_dict = { @@ -996,7 +996,7 @@ def test_set_datasource_variables_success( # --- Utility Methods --- -def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext) -> None: +def test_get_draft_workflow_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: # 1. Setup mocks pipeline = _make_pipeline() @@ -1011,9 +1011,7 @@ def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: assert result == workflow -def test_get_published_workflow_success( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext -) -> None: +def test_get_published_workflow_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: # 1. Setup mocks pipeline = _make_pipeline(workflow_id="wf-pub") @@ -1406,10 +1404,7 @@ def test_get_node_last_run_delegates_to_repository( ) -> None: repo = mocker.Mock() repo.get_node_last_execution.return_value = "node-exec" - mocker.patch( - "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", - return_value=repo, - ) + rag_pipeline_service.service._node_execution_service_repo = repo pipeline = _make_pipeline() workflow = _make_workflow(workflow_id="wf1") @@ -1785,21 +1780,25 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( def test_publish_customized_pipeline_template_raises_for_missing_pipeline( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - rag_pipeline_service.session.get.return_value = None + session = mocker.Mock() + session.get.return_value = None with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_publish_customized_pipeline_template_raises_for_missing_workflow_id( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id=None) - rag_pipeline_service.session.get.return_value = pipeline + session = mocker.Mock() + session.get.return_value = pipeline with pytest.raises(ValueError, match="Pipeline workflow not found"): rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {"name": "template-name"}, _make_account(), "t1" + "p1", {"name": "template-name"}, _make_account(), "t1", session=session ) @@ -1824,10 +1823,8 @@ def test_get_pipeline_raises_when_pipeline_missing( def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None: default_session_maker = mocker.Mock() - mocker.patch( - "services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", - return_value=default_session_maker, - ) + mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker", return_value=default_session_maker) + mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock(), session=mocker.Mock())) create_exec_repo = mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository" ) @@ -1835,7 +1832,7 @@ def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_run_repository" ) - RagPipelineService(session_maker=None) + RagPipelineService(session=mocker.Mock(), session_maker=None) create_exec_repo.assert_called_once_with(default_session_maker) create_run_repo.assert_called_once_with(default_session_maker) @@ -1849,11 +1846,12 @@ def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture) factory = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory.get_pipeline_template_factory.return_value.return_value = retrieval builtin = factory.get_built_in_pipeline_template_retrieval.return_value + session = mocker.Mock() - result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="en-US") + result = RagPipelineService.get_pipeline_templates(type="built-in", language="en-US", session=session) assert result == {"pipeline_templates": []} - retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) + retrieval.get_pipeline_templates.assert_called_once_with("en-US", None, session=session) builtin.fetch_pipeline_templates_from_builtin.assert_not_called() @@ -1861,14 +1859,14 @@ def test_update_customized_pipeline_template_commits_when_name_empty(mocker: Moc template = _make_customized_template() session = mocker.Mock() session.scalar.return_value = template - session_maker = _make_mock_session_maker(mocker, session) - mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) info = PipelineTemplateInfoEntity(name="", description="updated", icon_info=IconInfo(icon="i")) - result = RagPipelineService.update_customized_pipeline_template("tpl-1", info, _make_account(), "t1") + result = RagPipelineService.update_customized_pipeline_template( + "tpl-1", info, _make_account(), "t1", session=session + ) assert result.description == "updated" - session_maker.begin.assert_called_once() + session.commit.assert_called_once() def test_get_all_published_workflow_without_filters_has_no_more( @@ -2102,10 +2100,13 @@ def test_publish_customized_pipeline_template_raises_when_workflow_missing( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") - rag_pipeline_service.session.get.side_effect = [pipeline, None] + session = mocker.Mock() + session.get.side_effect = [pipeline, None] with pytest.raises(ValueError, match="Workflow not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_publish_customized_pipeline_template_raises_when_dataset_missing( @@ -2113,11 +2114,14 @@ def test_publish_customized_pipeline_template_raises_when_dataset_missing( ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") workflow = _make_workflow(workflow_id="wf-1") + session = rag_pipeline_service.session + session.get.side_effect = [pipeline, workflow] pipeline.retrieve_dataset = mocker.Mock(return_value=None) - rag_pipeline_service.session.get.side_effect = [pipeline, workflow] with pytest.raises(ValueError, match="Dataset not found"): - rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template( + "p1", {}, _make_account(), "t1", session=session + ) def test_get_recommended_plugins_skips_manifest_when_missing( @@ -2165,7 +2169,7 @@ def test_get_datasource_plugins_returns_empty_for_non_datasource_nodes( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = SimpleNamespace( graph_dict={"nodes": [{"id": "n1", "data": {"type": "start"}}]}, rag_pipeline_variables=[] ) @@ -2360,7 +2364,7 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() - pipeline = _make_pipeline() + pipeline = _make_pipeline(workflow_id="wf-1") workflow = SimpleNamespace( graph_dict={ "nodes": [ diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py index cee8e55f8cc..4ee1a5831a0 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py @@ -1,31 +1,16 @@ import logging -from collections.abc import Iterator from datetime import UTC, datetime from types import SimpleNamespace from typing import cast import pytest from pytest_mock import MockerFixture -from sqlalchemy import create_engine, select -from sqlalchemy.orm import Session, sessionmaker -from models.dataset import Dataset, Pipeline -from models.enums import DatasetRuntimeMode +from models.dataset import Dataset from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService -@pytest.fixture -def pipeline_session() -> Iterator[Session]: - engine = create_engine("sqlite:///:memory:") - Dataset.__table__.create(engine) - Pipeline.__table__.create(engine) - session_factory = sessionmaker(bind=engine, expire_on_commit=False) - with session_factory() as session: - yield session - engine.dispose() - - @pytest.mark.parametrize( ("doc_form", "datasource_type", "indexing_technique"), [ @@ -107,37 +92,47 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi install_mock.assert_called_once_with("tenant-1", ["missing-plugin:1.0.0"]) -def test_transform_to_empty_pipeline_updates_dataset_and_flushes( - mocker: MockerFixture, pipeline_session: Session -) -> None: +def test_transform_to_empty_pipeline_updates_dataset_and_commits(mocker: MockerFixture) -> None: service = RagPipelineTransformService() mocker.patch( "services.rag_pipeline.rag_pipeline_transform_service.current_user", SimpleNamespace(id="user-1"), ) - session = pipeline_session - dataset = Dataset( + class FakePipeline: + def __init__(self, **kwargs): + self.id = "pipeline-1" + self.tenant_id = kwargs["tenant_id"] + self.name = kwargs["name"] + self.description = kwargs["description"] + self.created_by = kwargs["created_by"] + + mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.Pipeline", FakePipeline) + session_mock = mocker.Mock() + add_mock = session_mock.add + flush_mock = session_mock.flush + commit_mock = session_mock.commit + + dataset = SimpleNamespace( + id="dataset-1", tenant_id="tenant-1", name="Dataset", description="desc", - created_by="user-1", + pipeline_id=None, + runtime_mode="general", + updated_by=None, + updated_at=None, ) - session.add(dataset) - session.commit() - flush_spy = mocker.spy(session, "flush") - commit_spy = mocker.spy(session, "commit") - result = service._transform_to_empty_pipeline(dataset, session) + result = service._transform_to_empty_pipeline(cast(Dataset, dataset), session=session_mock) - assert flush_spy.call_count == 2 - commit_spy.assert_not_called() - pipeline = session.scalar(select(Pipeline).where(Pipeline.id == dataset.pipeline_id)) - assert pipeline is not None - assert result == {"pipeline_id": pipeline.id, "dataset_id": dataset.id, "status": "success"} - assert dataset.pipeline_id == pipeline.id - assert dataset.runtime_mode == DatasetRuntimeMode.RAG_PIPELINE + assert result == {"pipeline_id": "pipeline-1", "dataset_id": "dataset-1", "status": "success"} + assert dataset.pipeline_id == "pipeline-1" + assert dataset.runtime_mode == "rag_pipeline" assert dataset.updated_by == "user-1" + add_mock.assert_called() + flush_mock.assert_called_once() + commit_mock.assert_called_once() # --- transform_dataset --- @@ -373,6 +368,7 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_deal_document_data") + session_mock.commit = mocker.Mock() # Mock current_user to have the same tenant_id as dataset mock_current_user = SimpleNamespace(current_tenant_id="t1") @@ -386,8 +382,6 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: assert result["pipeline_id"] == "p-new" assert dataset.runtime_mode == "rag_pipeline" assert dataset.chunk_structure == "text_model" - session_mock.flush.assert_called_once_with() - session_mock.commit.assert_not_called() def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(mocker: MockerFixture) -> None: @@ -439,11 +433,12 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: M service.transform_dataset("d1", session_mock) -def test_create_pipeline_raises_when_workflow_data_missing(pipeline_session: Session) -> None: +def test_create_pipeline_raises_when_workflow_data_missing(mocker: MockerFixture) -> None: service = RagPipelineTransformService() + session = mocker.Mock() with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service._create_pipeline({"rag_pipeline": {"name": "N"}}, pipeline_session) + service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=session) def test_deal_document_data_upload_file_with_existing_file(mocker: MockerFixture) -> None: diff --git a/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py b/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py index c86aaf1db22..cafac0656d1 100644 --- a/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py +++ b/api/tests/unit_tests/services/recommend_app/test_buildin_retrieval.py @@ -39,7 +39,7 @@ class TestBuildInRecommendAppRetrieval: return_value={"apps": []}, ) as mock_fetch: retrieval = BuildInRecommendAppRetrieval() - result = retrieval.get_recommended_apps_and_categories("en-US") + result = retrieval.get_recommended_apps_and_categories("en-US", session=MagicMock()) mock_fetch.assert_called_once_with("en-US") assert result == {"apps": []} @@ -47,11 +47,12 @@ class TestBuildInRecommendAppRetrieval: def test_get_learn_dify_apps_delegates_to_database(self, mock_database_retrieval): expected = {"recommended_apps": [{"id": "learn-dify-app"}]} mock_database_retrieval.fetch_learn_dify_apps_from_db.return_value = expected + session = MagicMock() - result = BuildInRecommendAppRetrieval().get_learn_dify_apps("en-US") + result = BuildInRecommendAppRetrieval().get_learn_dify_apps("en-US", session=session) assert result == expected - mock_database_retrieval.fetch_learn_dify_apps_from_db.assert_called_once_with("en-US") + mock_database_retrieval.fetch_learn_dify_apps_from_db.assert_called_once_with("en-US", session=session) def test_get_recommend_app_detail_delegates(self): with patch.object( @@ -60,7 +61,7 @@ class TestBuildInRecommendAppRetrieval: return_value={"id": "app-1"}, ) as mock_fetch: retrieval = BuildInRecommendAppRetrieval() - result = retrieval.get_recommend_app_detail("app-1") + result = retrieval.get_recommend_app_detail("app-1", session=MagicMock()) mock_fetch.assert_called_once_with("app-1") assert result == {"id": "app-1"} diff --git a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py index 55165deec25..9575aa9f52e 100644 --- a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py @@ -17,7 +17,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"id": "app-1"}, ) def test_get_recommend_app_detail_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1") + result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) assert result == {"id": "app-1"} mock_fetch.assert_called_once_with("app-1") @@ -32,7 +32,7 @@ class TestRemoteRecommendAppRetrieval: side_effect=ConnectionError("timeout"), ) def test_get_recommend_app_detail_falls_back_on_error(self, mock_fetch, mock_builtin): - result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1") + result = RemoteRecommendAppRetrieval().get_recommend_app_detail("app-1", session=MagicMock()) assert result == {"id": "fallback"} mock_builtin.assert_called_once_with("app-1") @@ -42,7 +42,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"recommended_apps": [], "categories": []}, ) def test_get_recommended_apps_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") + result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) assert result == {"recommended_apps": [], "categories": []} @patch( @@ -56,7 +56,7 @@ class TestRemoteRecommendAppRetrieval: side_effect=ValueError("server error"), ) def test_get_recommended_apps_falls_back_on_error(self, mock_fetch, mock_builtin): - result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US") + result = RemoteRecommendAppRetrieval().get_recommended_apps_and_categories("en-US", session=MagicMock()) assert result == {"recommended_apps": [{"id": "builtin"}]} @patch.object( @@ -65,7 +65,7 @@ class TestRemoteRecommendAppRetrieval: return_value={"recommended_apps": [{"id": "learn-dify-app"}]}, ) def test_get_learn_dify_apps_success(self, mock_fetch): - result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US") + result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US", session=MagicMock()) assert result == {"recommended_apps": [{"id": "learn-dify-app"}]} mock_fetch.assert_called_once_with("en-US") @@ -80,10 +80,12 @@ class TestRemoteRecommendAppRetrieval: side_effect=ValueError("server error"), ) def test_get_learn_dify_apps_falls_back_to_database_on_error(self, mock_fetch, mock_database): - result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US") + session = MagicMock() + + result = RemoteRecommendAppRetrieval().get_learn_dify_apps("en-US", session=session) assert result == {"recommended_apps": [{"id": "db-fallback"}]} - mock_database.assert_called_once_with("en-US") + mock_database.assert_called_once_with("en-US", session=session) class TestFetchFromDifyOfficial: diff --git a/api/tests/unit_tests/services/test_account_service.py b/api/tests/unit_tests/services/test_account_service.py index 233191ca0b7..b73fa112003 100644 --- a/api/tests/unit_tests/services/test_account_service.py +++ b/api/tests/unit_tests/services/test_account_service.py @@ -681,7 +681,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = TenantAccountRole.ADMIN - role = TenantService.get_account_role_in_tenant(mock_session, "account-1", "tenant-1") + role = TenantService.get_account_role_in_tenant("account-1", "tenant-1", session=mock_session) assert role == TenantAccountRole.ADMIN @@ -690,7 +690,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = None - role = TenantService.get_account_role_in_tenant(mock_session, "account-1", "tenant-1") + role = TenantService.get_account_role_in_tenant("account-1", "tenant-1", session=mock_session) assert role is None @@ -699,7 +699,7 @@ class TestTenantService: without ever touching the session.""" mock_session = MagicMock() - assert TenantService.get_account_role_in_tenant(mock_session, None, "tenant-1") is None + assert TenantService.get_account_role_in_tenant(None, "tenant-1", session=mock_session) is None mock_session.execute.assert_not_called() def test_get_account_role_in_tenant_query_is_scoped(self): @@ -711,7 +711,7 @@ class TestTenantService: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = TenantAccountRole.NORMAL - TenantService.get_account_role_in_tenant(mock_session, account_id, tenant_id) + TenantService.get_account_role_in_tenant(account_id, tenant_id, session=mock_session) stmt = mock_session.execute.call_args.args[0] compiled = str(stmt.compile(compile_kwargs={"literal_binds": True})) @@ -760,11 +760,7 @@ class TestTenantService: mock_tenant_instance.name = "Test User's Workspace" mock_tenant_class.return_value = mock_tenant_instance - # Mock the db import in CreditPoolService to avoid database connection - with patch("services.credit_pool_service.db") as mock_credit_pool_db: - mock_credit_pool_db.session.add = MagicMock() - mock_credit_pool_db.session.commit = MagicMock() - + with patch("services.credit_pool_service.CreditPoolService.create_default_pool"): # Execute test TenantService.create_owner_tenant_if_not_exist( mock_account, session=mock_db_dependencies["db"].session @@ -1052,6 +1048,7 @@ class TestTenantService: account_id="user-rbac", member_account_id="user-rbac", role_ids=["rbac-owner-id"], + session=mock_db_dependencies["db"].session, ) def test_admin_can_update_admin_member_role(self): @@ -1191,7 +1188,11 @@ class TestTenantService: with pytest.raises(NoPermissionError): TenantService.check_member_permission( - mock_tenant, mock_operator, mock_member, "remove", session=MagicMock() + mock_tenant, + mock_operator, + mock_member, + "remove", + session=mock_db_dependencies["db"].session, ) def test_rbac_member_can_remove_non_owner_member(self): @@ -1265,7 +1266,9 @@ class TestTenantService: ), patch("services.account_service.RBACService.Roles", mock_rbac_roles), ): - owner_account_id = AccountService.get_rbac_workspace_owner_account_id("tenant-1", "acct-1") + owner_account_id = AccountService.get_rbac_workspace_owner_account_id( + "tenant-1", "acct-1", session=MagicMock() + ) assert owner_account_id == "owner-account" call = mock_rbac_roles.members.call_args @@ -1912,7 +1915,9 @@ class TestRegisterService: is_setup=True, session=mock_db_dependencies["db"].session, ) - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, "newuser@example.com") + mock_lookup.assert_called_once_with( + "newuser@example.com", session=mock_db_dependencies["db"].session + ) def test_invite_new_member_normalizes_new_account_email( self, mock_db_dependencies, mock_redis_dependencies, mock_task_dependencies @@ -1958,7 +1963,7 @@ class TestRegisterService: is_setup=True, session=mock_db_dependencies["db"].session, ) - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, mixed_email) + mock_lookup.assert_called_once_with(mixed_email, session=mock_db_dependencies["db"].session) mock_check_permission.assert_called_once_with( mock_tenant, mock_inviter, @@ -2025,7 +2030,7 @@ class TestRegisterService: mock_tenant, mock_existing_account, "normal", requires_setup=True ) mock_task_dependencies.delay.assert_called_once() - mock_lookup.assert_called_once_with(mock_db_dependencies["db"].session, "existing@example.com") + mock_lookup.assert_called_once_with("existing@example.com", session=mock_db_dependencies["db"].session) def test_invite_existing_active_account_requires_acceptance_before_joining( self, mock_db_dependencies, mock_redis_dependencies, mock_task_dependencies @@ -2171,6 +2176,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_new_account.id, role_ids=["rbac-role-id-123"], + session=mock_db_dependencies["db"].session, ) def test_invite_new_member_rbac_enabled_existing_account( @@ -2220,6 +2226,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_existing_account.id, role_ids=["rbac-role-id-456"], + session=mock_db_dependencies["db"].session, ) def test_invite_new_member_rbac_enabled_existing_active_account_adds_role_before_signin_response( @@ -2268,6 +2275,7 @@ class TestRegisterService: account_id=mock_inviter.id, member_account_id=mock_existing_account.id, role_ids=["rbac-role-id-456"], + session=mock_db_dependencies["db"].session, ) mock_task_dependencies.delay.assert_not_called() @@ -2615,7 +2623,7 @@ class TestSessionInjectedGetters: sentinel_account = MagicMock(spec=Account) mock_session.get.return_value = sentinel_account - result = AccountService.get_account_by_id(mock_session, "user-123") + result = AccountService.get_account_by_id("user-123", session=mock_session) assert result is sentinel_account mock_session.get.assert_called_once_with(Account, "user-123") @@ -2625,7 +2633,7 @@ class TestSessionInjectedGetters: mock_session = MagicMock() mock_session.get.return_value = None - assert AccountService.get_account_by_id(mock_session, "missing") is None + assert AccountService.get_account_by_id("missing", session=mock_session) is None @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) def test_get_account_by_email_returns_scalar_or_none(self, sqlite_session: Session): @@ -2637,9 +2645,9 @@ class TestSessionInjectedGetters: sqlite_session.add(account) sqlite_session.commit() - assert AccountService.get_account_by_email(sqlite_session, "alice@example.com") == account - assert AccountService.get_account_by_email(sqlite_session, "ALICE@example.com") is None - assert AccountService.get_account_by_email(sqlite_session, "ghost@example.com") is None + assert AccountService.get_account_by_email("alice@example.com", session=sqlite_session) == account + assert AccountService.get_account_by_email("ALICE@example.com", session=sqlite_session) is None + assert AccountService.get_account_by_email("ghost@example.com", session=sqlite_session) is None def test_account_belongs_to_tenant_short_circuits_on_falsy_account_id(self): """SSO bearers with no ``account_id`` (and any other falsy id) @@ -2648,22 +2656,22 @@ class TestSessionInjectedGetters: """ mock_session = MagicMock() - assert TenantService.account_belongs_to_tenant(mock_session, None, "tenant-1") is False - assert TenantService.account_belongs_to_tenant(mock_session, "", "tenant-1") is False + assert TenantService.account_belongs_to_tenant(None, "tenant-1", session=mock_session) is False + assert TenantService.account_belongs_to_tenant("", "tenant-1", session=mock_session) is False mock_session.execute.assert_not_called() def test_account_belongs_to_tenant_true_when_join_row_exists(self): mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = "join-id" - assert TenantService.account_belongs_to_tenant(mock_session, "user-1", "tenant-1") is True + assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=mock_session) is True mock_session.execute.assert_called_once() def test_account_belongs_to_tenant_false_when_no_join(self): mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = None - assert TenantService.account_belongs_to_tenant(mock_session, "user-1", "tenant-1") is False + assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=mock_session) is False def test_get_account_memberships_returns_join_tenant_pairs(self): """Returns whatever ``session.query(...).join(...).filter(...).all()`` @@ -2674,7 +2682,7 @@ class TestSessionInjectedGetters: rows = [(MagicMock(), MagicMock()), (MagicMock(), MagicMock())] mock_session.query.return_value.join.return_value.filter.return_value.all.return_value = rows - out = TenantService.get_account_memberships(mock_session, "user-123") + out = TenantService.get_account_memberships("user-123", session=mock_session) assert out == rows # No fall-through to the global db.session proxy. @@ -2688,7 +2696,7 @@ class TestSessionInjectedGetters: rows = [(MagicMock(), MagicMock())] mock_session.execute.return_value.all.return_value = rows - out = TenantService.get_workspaces_for_account(mock_session, "user-123") + out = TenantService.get_workspaces_for_account("user-123", session=mock_session) assert out == rows assert mock_session.execute.called @@ -2704,20 +2712,20 @@ class TestSessionInjectedGetters: sentinel = MagicMock(spec=Tenant) mock_session.get.return_value = sentinel - assert TenantService.get_tenant_by_id(mock_session, "tenant-1") is sentinel + assert TenantService.get_tenant_by_id("tenant-1", session=mock_session) is sentinel mock_session.get.assert_called_once_with(Tenant, "tenant-1") def test_get_tenant_by_id_returns_none_when_missing(self): mock_session = MagicMock() mock_session.get.return_value = None - assert TenantService.get_tenant_by_id(mock_session, "missing") is None + assert TenantService.get_tenant_by_id("missing", session=mock_session) is None def test_get_tenants_by_ids_short_circuits_on_empty_input(self): """Empty id list must not emit ``WHERE id IN ()``.""" mock_session = MagicMock() - assert TenantService.get_tenants_by_ids(mock_session, []) == [] + assert TenantService.get_tenants_by_ids([], session=mock_session) == [] mock_session.execute.assert_not_called() def test_get_tenants_by_ids_returns_scalars(self): @@ -2725,7 +2733,7 @@ class TestSessionInjectedGetters: tenants = [MagicMock(), MagicMock()] mock_session.execute.return_value.scalars.return_value.all.return_value = tenants - assert TenantService.get_tenants_by_ids(mock_session, ["t1", "t2"]) == tenants + assert TenantService.get_tenants_by_ids(["t1", "t2"], session=mock_session) == tenants mock_session.execute.assert_called_once() def test_get_tenant_name_returns_scalar_or_none(self): @@ -2736,10 +2744,10 @@ class TestSessionInjectedGetters: mock_session = MagicMock() mock_session.execute.return_value.scalar_one_or_none.return_value = "Acme Inc." - assert TenantService.get_tenant_name(mock_session, "tenant-1") == "Acme Inc." + assert TenantService.get_tenant_name("tenant-1", session=mock_session) == "Acme Inc." mock_session.execute.return_value.scalar_one_or_none.return_value = None - assert TenantService.get_tenant_name(mock_session, "missing") is None + assert TenantService.get_tenant_name("missing", session=mock_session) is None def test_find_workspace_for_account_returns_first_row_or_none(self): """Per-id read returns ``session.execute(...).first()`` directly; @@ -2750,7 +2758,7 @@ class TestSessionInjectedGetters: sentinel_row = (MagicMock(), MagicMock()) mock_session.execute.return_value.first.return_value = sentinel_row - assert TenantService.find_workspace_for_account(mock_session, "user-123", "ws-1") is sentinel_row + assert TenantService.find_workspace_for_account("user-123", "ws-1", session=mock_session) is sentinel_row mock_session.execute.return_value.first.return_value = None - assert TenantService.find_workspace_for_account(mock_session, "user-123", "ws-1") is None + assert TenantService.find_workspace_for_account("user-123", "ws-1", session=mock_session) is None diff --git a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py index a9ed82413bb..c36980f5829 100644 --- a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py +++ b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py @@ -258,6 +258,7 @@ def test_workflow_sandbox_service_resolves_locator_and_returns_download_url( node_id="node-1", node_execution_id="node-exec-1", path="report.txt", + session=session_factory.create_session(), ) assert result.url == "https://files.example/report.txt?token=1&as_attachment=true" @@ -376,6 +377,7 @@ def test_workflow_sandbox_service_filters_by_node_execution_id() -> None: node_id="node-1", node_execution_id="node-exec-2", path="out.txt", + session=session_factory.create_session(), ) assert result.text == "hello" @@ -409,6 +411,7 @@ def test_workflow_sandbox_service_uses_latest_active_session_when_execution_id_o node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert result.path == "." @@ -428,6 +431,7 @@ def test_workflow_sandbox_service_raises_when_no_active_session() -> None: node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert exc_info.value.code == "no_active_session" @@ -447,6 +451,7 @@ def test_workflow_sandbox_service_raises_when_runtime_specs_missing() -> None: node_id="node-1", node_execution_id=None, path=".", + session=session_factory.create_session(), ) assert exc_info.value.code == "no_sandbox" diff --git a/api/tests/unit_tests/services/test_agent_drive_service.py b/api/tests/unit_tests/services/test_agent_drive_service.py index ad72e142522..6371b197325 100644 --- a/api/tests/unit_tests/services/test_agent_drive_service.py +++ b/api/tests/unit_tests/services/test_agent_drive_service.py @@ -125,6 +125,7 @@ def _commit(key: str, tool_file_id: str, *, owned: bool = True): value_owned_by_drive=owned, ) ], + session=session_factory.create_session(), ) @@ -132,15 +133,25 @@ def test_commit_then_manifest_lists_the_entry(): tf = _seed_tool_file() _commit("data/report.txt", tf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert [i["key"] for i in items] == ["data/report.txt"] assert items[0]["file_kind"] == "tool_file" assert items[0]["file_id"] == tf assert items[0]["mime_type"] == "text/plain" # prefix filter - assert AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, prefix="data/") != [] - assert AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, prefix="other/") == [] + assert ( + AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, prefix="data/", session=session_factory.create_session() + ) + != [] + ) + assert ( + AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, prefix="other/", session=session_factory.create_session() + ) + == [] + ) def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: @@ -157,6 +168,7 @@ def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description="Parses RFPs."), ) ], + session=session_factory.create_session(), ) with session_factory.create_session() as session: @@ -165,7 +177,7 @@ def test_commit_skill_row_persists_metadata_and_lists_catalog() -> None: assert row.is_skill is True assert row.skill_metadata == '{"description":"Parses RFPs.","name":"Tender Analyzer"}' - skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert len(skills) == 1 assert skills[0]["path"] == "tender-analyzer" assert skills[0]["skill_md_key"] == "tender-analyzer/SKILL.md" @@ -191,6 +203,7 @@ def test_commit_rejects_skill_row_without_skill_metadata() -> None: is_skill=True, ) ], + session=session_factory.create_session(), ) assert exc_info.value.code == "invalid_skill_metadata" @@ -220,7 +233,7 @@ def test_list_skills_raises_controlled_error_for_invalid_stored_metadata(raw_met session.commit() with pytest.raises(AgentDriveError) as exc_info: - AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert exc_info.value.code == "invalid_skill_metadata" @@ -239,6 +252,7 @@ def test_commit_rejects_non_skill_row_with_skill_metadata() -> None: skill_metadata=DriveSkillMetadata(name="Bad", description=""), ) ], + session=session_factory.create_session(), ) @@ -257,6 +271,7 @@ def test_commit_rejects_non_canonical_skill_key() -> None: skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description=""), ) ], + session=session_factory.create_session(), ) @@ -282,6 +297,7 @@ def test_commit_rejects_agent_from_another_tenant(): value_owned_by_drive=True, ) ], + session=session_factory.create_session(), ) assert exc_info.value.status_code == 404 assert exc_info.value.code == "agent_not_found" @@ -311,24 +327,27 @@ def test_batch_failure_does_not_delete_old_storage_before_commit(): _commit("doc.txt", tf1, owned=True) with patch("services.agent_drive_service.storage") as storage_mock: - with pytest.raises(AgentDriveError): - AgentDriveService().commit( - tenant_id=TENANT, - user_id=USER, - agent_id=AGENT, - items=[ - DriveCommitItem( - key="doc.txt", - file_ref={"kind": "tool_file", "id": tf2}, - value_owned_by_drive=True, - ), - DriveCommitItem( - key="bad.txt", - file_ref={"kind": "tool_file", "id": "44444444-4444-4444-4444-444444444444"}, - value_owned_by_drive=True, - ), - ], - ) + with session_factory.create_session() as session: + with pytest.raises(AgentDriveError): + AgentDriveService().commit( + tenant_id=TENANT, + user_id=USER, + agent_id=AGENT, + items=[ + DriveCommitItem( + key="doc.txt", + file_ref={"kind": "tool_file", "id": tf2}, + value_owned_by_drive=True, + ), + DriveCommitItem( + key="bad.txt", + file_ref={"kind": "tool_file", "id": "44444444-4444-4444-4444-444444444444"}, + value_owned_by_drive=True, + ), + ], + session=session, + ) + session.rollback() storage_mock.delete.assert_not_called() with session_factory.create_session() as session: @@ -389,6 +408,7 @@ def test_recommit_same_skill_value_updates_metadata_without_cleaning_backing_fil skill_metadata=DriveSkillMetadata(name="Tender Analyzer", description="v1"), ) ], + session=session_factory.create_session(), ) with patch("services.agent_drive_service.storage") as storage_mock: @@ -405,6 +425,7 @@ def test_recommit_same_skill_value_updates_metadata_without_cleaning_backing_fil skill_metadata=DriveSkillMetadata(name="Tender Analyzer v2", description="v2"), ) ], + session=session_factory.create_session(), ) storage_mock.delete.assert_not_called() @@ -449,6 +470,7 @@ def _commit_upload(key: str, upload_file_id: str, *, owned: bool = True): value_owned_by_drive=owned, ) ], + session=session_factory.create_session(), ) @@ -456,7 +478,7 @@ def test_commit_upload_file_source_and_manifest(): uf = _seed_upload_file() _commit_upload("docs/u.txt", uf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert items[0]["file_kind"] == "upload_file" assert items[0]["file_id"] == uf assert items[0]["mime_type"] == "text/plain" @@ -492,7 +514,9 @@ def test_manifest_includes_internal_download_url(): patch("core.app.workflow.file_runtime.DifyWorkflowFileRuntime") as runtime_cls, ): runtime_cls.return_value.resolve_file_url.return_value = "http://internal/files/x?sign=1" - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, include_download_url=True) + items = AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, include_download_url=True, session=session_factory.create_session() + ) assert items[0]["download_url"] == "http://internal/files/x?sign=1" # drive-owned resolution: internal URL (for_external=False) @@ -507,7 +531,9 @@ def test_manifest_download_url_none_when_unresolvable(): "services.agent_drive_service.file_factory.build_from_mapping", side_effect=ValueError("not found"), ): - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, include_download_url=True) + items = AgentDriveService().manifest( + tenant_id=TENANT, agent_id=AGENT, include_download_url=True, session=session_factory.create_session() + ) assert items[0]["download_url"] is None @@ -524,6 +550,7 @@ def test_delete_by_key_cleans_drive_owned_value(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/doomed.txt", file_ref=None)], + session=session_factory.create_session(), ) storage_mock.delete.assert_called_once() @@ -560,6 +587,7 @@ def test_commit_null_batch_removes_multiple_skill_keys(): DriveCommitItem(key="tender-analyzer/SKILL.md", file_ref=None), DriveCommitItem(key="tender-analyzer/.DIFY-SKILL-FULL.zip", file_ref=None), ], + session=session_factory.create_session(), ) assert sorted(item["key"] for item in removed) == [ @@ -581,6 +609,7 @@ def test_commit_null_is_idempotent_for_missing_keys(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/never-there.txt", file_ref=None)], + session=session_factory.create_session(), ) assert removed == [{"key": "files/never-there.txt", "removed": True, "noop": True}] @@ -595,6 +624,7 @@ def test_commit_null_keeps_shared_value_records(): user_id=USER, agent_id=AGENT, items=[DriveCommitItem(key="files/shared.txt", file_ref=None)], + session=session_factory.create_session(), ) storage_mock.delete.assert_not_called() @@ -639,7 +669,9 @@ def test_preview_returns_text_with_truncation_flags(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\nUse responsibly.\n"]) - result = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/SKILL.md") + result = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/SKILL.md", session=session_factory.create_session() + ) assert result == { "key": "pdf-toolkit/SKILL.md", @@ -656,13 +688,17 @@ def test_preview_marks_binary_and_oversized_content(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"\x00\x01\x02"]) - binary = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin") + binary = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin", session=session_factory.create_session() + ) assert binary["binary"] is True assert binary["text"] is None with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"x" * (AgentDriveService.PREVIEW_MAX_BYTES + 10)]) - oversized = AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin") + oversized = AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="files/blob.bin", session=session_factory.create_session() + ) assert oversized["truncated"] is True assert oversized["binary"] is False assert len(oversized["text"]) == AgentDriveService.PREVIEW_MAX_BYTES @@ -670,7 +706,9 @@ def test_preview_marks_binary_and_oversized_content(): def test_preview_unknown_key_is_404(): with pytest.raises(AgentDriveError) as exc_info: - AgentDriveService().preview(tenant_id=TENANT, agent_id=AGENT, key="ghost/SKILL.md") + AgentDriveService().preview( + tenant_id=TENANT, agent_id=AGENT, key="ghost/SKILL.md", session=session_factory.create_session() + ) assert exc_info.value.code == "drive_key_not_found" assert exc_info.value.status_code == 404 @@ -678,7 +716,10 @@ def test_preview_unknown_key_is_404(): def test_preview_rejects_cross_tenant_agent(): with pytest.raises(AgentDriveError) as exc_info: AgentDriveService().preview( - tenant_id="99999999-9999-9999-9999-999999999999", agent_id=AGENT, key="pdf-toolkit/SKILL.md" + tenant_id="99999999-9999-9999-9999-999999999999", + agent_id=AGENT, + key="pdf-toolkit/SKILL.md", + session=session_factory.create_session(), ) assert exc_info.value.code == "agent_not_found" @@ -688,7 +729,12 @@ def test_download_url_signs_external_audience(): _commit("pdf-toolkit/.DIFY-SKILL-FULL.zip", tf) with patch.object(AgentDriveService, "_resolve_download_url", return_value="https://signed.example/x") as resolver: - url = AgentDriveService().download_url(tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/.DIFY-SKILL-FULL.zip") + url = AgentDriveService().download_url( + tenant_id=TENANT, + agent_id=AGENT, + key="pdf-toolkit/.DIFY-SKILL-FULL.zip", + session=session_factory.create_session(), + ) assert url == "https://signed.example/x" # console downloads are for browsers: external signing, never the internal URL @@ -702,7 +748,9 @@ def test_upload_file_download_url_uses_attachment_filename(): with patch("core.app.workflow.file_runtime.DifyWorkflowFileRuntime") as runtime_cls: runtime_cls.return_value.resolve_upload_file_url.return_value = "https://files.example/report.pdf" - url = AgentDriveService().download_url(tenant_id=TENANT, agent_id=AGENT, key="files/report.pdf") + url = AgentDriveService().download_url( + tenant_id=TENANT, agent_id=AGENT, key="files/report.pdf", session=session_factory.create_session() + ) assert url == "https://files.example/report.pdf" assert runtime_cls.return_value.resolve_upload_file_url.call_args.kwargs["for_external"] is True @@ -712,7 +760,7 @@ def test_upload_file_download_url_uses_attachment_filename(): def test_manifest_items_carry_created_at_for_inspector(): tf = _seed_tool_file() _commit("files/x.txt", tf) - items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT) + items = AgentDriveService().manifest(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) assert items[0]["created_at"] is None or isinstance(items[0]["created_at"], int) @@ -744,13 +792,14 @@ def _commit_skill(*, manifest_files: list[str] | None = None) -> None: value_owned_by_drive=True, ), ], + session=session_factory.create_session(), ) def test_list_skills_uses_canonical_skill_rows(): _commit_skill(manifest_files=["SKILL.md", "scripts/run.py"]) - skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT) + skills = AgentDriveService().list_skills(tenant_id=TENANT, agent_id=AGENT, session=session_factory.create_session()) created_at = skills[0].pop("created_at") assert skills == [ @@ -773,7 +822,9 @@ def test_inspect_skill_returns_manifest_files_and_file_tree(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\n"]) - result = AgentDriveService().inspect_skill(tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit") + result = AgentDriveService().inspect_skill( + tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit", session=session_factory.create_session() + ) assert result["source"] == "skill_md" assert result["warnings"] == [] @@ -792,7 +843,9 @@ def test_inspect_skill_falls_back_to_drive_keys_when_manifest_missing(): with patch("services.agent_drive_service.storage") as storage_mock: storage_mock.load_stream.return_value = iter([b"# PDF Toolkit\n"]) - result = AgentDriveService().inspect_skill(tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit") + result = AgentDriveService().inspect_skill( + tenant_id=TENANT, agent_id=AGENT, skill_path="pdf-toolkit", session=session_factory.create_session() + ) assert result["warnings"] == ["manifest_files_unavailable"] assert [file["path"] for file in result["files"]] == ["SKILL.md"] @@ -808,6 +861,7 @@ def test_preview_skill_archive_member_from_manifest_without_drive_row(): tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/references/guide.md", + session=session_factory.create_session(), ) assert result == { @@ -831,6 +885,7 @@ def test_download_url_signs_skill_archive_member_from_manifest_without_drive_row tenant_id=TENANT, agent_id=AGENT, key="pdf-toolkit/references/guide.md", + session=session_factory.create_session(), ) assert url == "https://signed.example/member" @@ -856,5 +911,6 @@ def test_skill_metadata_rejects_non_canonical_rows(): skill_metadata=DriveSkillMetadata(name="Bad"), ) ], + session=session_factory.create_session(), ) assert exc_info.value.code == "invalid_skill_key" diff --git a/api/tests/unit_tests/services/test_agent_tool_inner_service.py b/api/tests/unit_tests/services/test_agent_tool_inner_service.py index 93e20d267da..61049d29e9e 100644 --- a/api/tests/unit_tests/services/test_agent_tool_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_tool_inner_service.py @@ -70,7 +70,7 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None: side_effect=lambda messages, **_kwargs: messages, ), ): - response = AgentToolInnerService().invoke(session, _request()) + response = AgentToolInnerService().invoke(_request(), session=session) assert response.observation == "ok" assert response.metadata == { @@ -89,7 +89,7 @@ def test_invoke_raises_app_not_found_when_session_has_no_app() -> None: session.get.return_value = None with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -102,7 +102,7 @@ def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant() -> session.get.return_value = fake_app with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_tenant_mismatch" assert exc_info.value.status_code == 403 @@ -120,7 +120,7 @@ def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_co patch("services.agent_tool_inner_service.ToolEngine.generic_invoke", side_effect=ValueError("app not found")), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -141,7 +141,7 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N ), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == "agent_tool_invoke_failed" @@ -161,6 +161,6 @@ def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(error: Excepti with patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", side_effect=error): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(session, _request()) + AgentToolInnerService().invoke(_request(), session=session) assert exc_info.value.error_code == expected_code diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index 2975c4df14c..79bbb5873ac 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -103,7 +103,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1") + AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1", session=mock_db.session) def test_up_insert_app_annotation_from_message_should_raise_value_error_when_answer_missing(self) -> None: """Test missing answer and content raises ValueError.""" @@ -121,7 +121,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_raise_not_found_when_message_missing(self) -> None: """Test missing message raises NotFound.""" @@ -139,7 +139,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_update_existing_annotation_when_found(self) -> None: """Test existing annotation is updated and indexed.""" @@ -161,7 +161,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, message, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation @@ -199,7 +199,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, message, None] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -231,7 +231,7 @@ class TestAppAnnotationServiceUpInsert: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) def test_up_insert_app_annotation_from_message_should_create_annotation_when_message_missing(self) -> None: """Test annotation is created when message_id is not provided.""" @@ -252,7 +252,7 @@ class TestAppAnnotationServiceUpInsert: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -383,7 +383,7 @@ class TestAppAnnotationServiceListAndExport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "") + AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "", session=mock_db.session) def test_get_annotation_list_by_app_id_should_return_items_with_keyword(self) -> None: """Test keyword search returns items and total.""" @@ -402,7 +402,9 @@ class TestAppAnnotationServiceListAndExport: mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "keyword") + items, total = AppAnnotationService.get_annotation_list_by_app_id( + app.id, 1, 10, "keyword", session=mock_db.session + ) # Assert assert items == ["a1"] @@ -424,7 +426,9 @@ class TestAppAnnotationServiceListAndExport: mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "") + items, total = AppAnnotationService.get_annotation_list_by_app_id( + app.id, 1, 10, "", session=mock_db.session + ) # Assert assert items == ["a1", "a2"] @@ -451,7 +455,7 @@ class TestAppAnnotationServiceListAndExport: mock_db.session.scalars.return_value.all.return_value = [annotation1, annotation2] # Act - result = AppAnnotationService.export_annotation_list_by_app_id(app.id) + result = AppAnnotationService.export_annotation_list_by_app_id(app.id, session=mock_db.session) # Assert assert result == [annotation1, annotation2] @@ -473,7 +477,7 @@ class TestAppAnnotationServiceListAndExport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.export_annotation_list_by_app_id("app-1") + AppAnnotationService.export_annotation_list_by_app_id("app-1", session=mock_db.session) class TestAppAnnotationServiceDirectManipulation: @@ -493,7 +497,7 @@ class TestAppAnnotationServiceDirectManipulation: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.insert_app_annotation_directly(args, "app-1") + AppAnnotationService.insert_app_annotation_directly(args, "app-1", session=mock_db.session) def test_insert_app_annotation_directly_should_raise_value_error_when_question_missing(self) -> None: """Test missing question raises ValueError.""" @@ -510,7 +514,7 @@ class TestAppAnnotationServiceDirectManipulation: # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly(args, app.id) + AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) def test_insert_app_annotation_directly_should_create_annotation_and_index(self) -> None: """Test insert creates annotation and triggers index task.""" @@ -531,7 +535,7 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.insert_app_annotation_directly(args, app.id) + result = AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) # Assert assert result == annotation_instance @@ -696,7 +700,9 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.execute.return_value.all.return_value = [] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1"]) + result = AppAnnotationService.delete_app_annotations_in_batch( + _make_app_ref(app), ["ann-1"], session=mock_db.session + ) # Assert assert result == {"deleted_count": 0} @@ -723,7 +729,9 @@ class TestAppAnnotationServiceDirectManipulation: mock_db.session.execute.side_effect = [execute_result_multi, MagicMock(), execute_result_delete] # Act - result = AppAnnotationService.delete_app_annotations_in_batch(_make_app_ref(app), ["ann-1", "ann-2"]) + result = AppAnnotationService.delete_app_annotations_in_batch( + _make_app_ref(app), ["ann-1", "ann-2"], session=mock_db.session + ) # Assert assert result == {"deleted_count": 2} @@ -755,7 +763,7 @@ class TestAppAnnotationServiceBatchImport: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.batch_import_app_annotations("app-1", file) + AppAnnotationService.batch_import_app_annotations("app-1", file, session=mock_db.session) def test_batch_import_app_annotations_should_return_error_when_columns_invalid(self) -> None: """Test invalid column count returns error message.""" @@ -777,7 +785,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -801,7 +809,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -829,7 +837,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -855,7 +863,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -885,7 +893,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -911,7 +919,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -937,7 +945,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -963,7 +971,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -994,7 +1002,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -1027,7 +1035,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert assert result == {"job_id": "uuid-3", "job_status": "waiting", "record_count": 1} @@ -1067,7 +1075,7 @@ class TestAppAnnotationServiceBatchImport: # Act with caplog.at_level(logging.DEBUG): - result = AppAnnotationService.batch_import_app_annotations(app.id, file) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) # Assert assert result["error_msg"] == "An error occurred while processing the file: boom" @@ -1090,7 +1098,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_hit_histories(_make_annotation_ref(app, "ann-1"), 1, 10) + AppAnnotationService.get_annotation_hit_histories( + _make_annotation_ref(app, "ann-1"), 1, 10, session=mock_db.session + ) def test_get_annotation_hit_histories_should_return_items_and_total(self) -> None: """Test hit histories pagination returns items and total.""" @@ -1114,6 +1124,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: _make_annotation_ref(app, annotation.id), 1, 10, + session=mock_db.session, ) # Assert @@ -1129,7 +1140,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.get.return_value = None # Act - result = AppAnnotationService.get_annotation_by_id("ann-1") + result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) # Assert assert result is None @@ -1142,7 +1153,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.get.return_value = annotation # Act - result = AppAnnotationService.get_annotation_by_id("ann-1") + result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) # Assert assert result == annotation @@ -1165,6 +1176,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: message_id="msg-1", from_source="chat", score=0.8, + session=mock_db.session, ) # Assert @@ -1187,7 +1199,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result["enabled"] is True @@ -1208,7 +1220,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_app_annotation_setting_by_app_id("app-1") + AppAnnotationService.get_app_annotation_setting_by_app_id("app-1", session=mock_db.session) def test_get_app_annotation_setting_by_app_id_should_return_empty_embedding_model_when_no_detail(self) -> None: """Test setting without detail returns empty embedding model.""" @@ -1224,7 +1236,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result["enabled"] is True @@ -1243,7 +1255,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, None] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) # Assert assert result == {"enabled": False} @@ -1265,7 +1277,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, setting.id, args, session=mock_db.session + ) # Assert assert result["enabled"] is True @@ -1292,7 +1306,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args) + result = AppAnnotationService.update_app_annotation_setting( + app.id, setting.id, args, session=mock_db.session + ) # Assert assert result["enabled"] is True @@ -1312,7 +1328,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_setting("app-1", "setting-1", {"score_threshold": 0.5}) + AppAnnotationService.update_app_annotation_setting( + "app-1", "setting-1", {"score_threshold": 0.5}, session=mock_db.session + ) def test_update_app_annotation_setting_should_raise_not_found_when_setting_missing(self) -> None: """Test update raises NotFound when setting is missing.""" @@ -1328,7 +1346,9 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.update_app_annotation_setting(app.id, "setting-1", {"score_threshold": 0.5}) + AppAnnotationService.update_app_annotation_setting( + app.id, "setting-1", {"score_threshold": 0.5}, session=mock_db.session + ) class TestAppAnnotationServiceClearAll: @@ -1361,7 +1381,7 @@ class TestAppAnnotationServiceClearAll: mock_db.session.scalars.side_effect = [annotations_scalars, histories_scalars_1, histories_scalars_2] # Act - result = AppAnnotationService.clear_all_annotations(app.id) + result = AppAnnotationService.clear_all_annotations(app.id, session=mock_db.session) # Assert assert result == {"result": "success"} @@ -1385,4 +1405,4 @@ class TestAppAnnotationServiceClearAll: # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.clear_all_annotations("app-1") + AppAnnotationService.clear_all_annotations("app-1", session=mock_db.session) diff --git a/api/tests/unit_tests/services/test_app_generate_service.py b/api/tests/unit_tests/services/test_app_generate_service.py index 865410fddf1..22c7514a522 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -235,12 +235,12 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "ok"} gen_spy.assert_called_once() @@ -256,12 +256,12 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.AGENT_CHAT), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "agent"} gen_spy.assert_called_once() @@ -278,12 +278,12 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=True) result = AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "agent-via-flag"} gen_spy.assert_called_once() @@ -300,12 +300,12 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=False) result = AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "chat"} gen_spy.assert_called_once() @@ -342,12 +342,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "advanced-blocking"} call_kwargs = gen_spy.call_args.kwargs @@ -375,12 +375,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) # In streaming mode it should go through retrieve_events, not generate gen_instance.retrieve_events.assert_called_once() @@ -401,12 +401,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) assert result == {"result": "workflow-blocking"} call_kwargs = gen_spy.call_args.kwargs @@ -435,12 +435,12 @@ class TestGenerate: ) result = AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) retrieve_spy.assert_called_once() # The inner on_subscribe closure was invoked by _build_streaming_task_on_subscribe @@ -451,12 +451,12 @@ class TestGenerate: app = _make_app("invalid-mode", is_agent=False) with pytest.raises(ValueError, match="Invalid app mode"): AppGenerateService.generate( - MagicMock(), app_model=app, user=_make_user(), args={}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) @@ -489,12 +489,12 @@ class TestGenerateBilling: ) AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) reserve_mock.assert_called_once_with(QuotaType.WORKFLOW, "tenant-id") quota_charge.commit.assert_called_once() @@ -513,12 +513,12 @@ class TestGenerateBilling: with pytest.raises(InvokeRateLimitError): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) def test_exception_refunds_quota_and_exits_rate_limit(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): @@ -539,12 +539,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -571,12 +571,12 @@ class TestGenerateBilling: ) AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) # exit is called in finally block for non-streaming assert exit_calls == ["dummy-request-id"] @@ -669,12 +669,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -701,12 +701,12 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=True, + session=MagicMock(), ) quota_charge.refund.assert_called_once() @@ -723,7 +723,7 @@ class TestGetWorkflow: ws.get_draft_workflow.return_value = draft_wf mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) - result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER) + result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) assert result is draft_wf ws.get_draft_workflow.assert_called_once() @@ -733,7 +733,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not initialized"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.DEBUGGER, session=MagicMock()) def test_non_debugger_fetches_published(self, mocker: MockerFixture): pub_wf = _make_workflow() @@ -741,7 +741,9 @@ class TestGetWorkflow: ws.get_published_workflow.return_value = pub_wf mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) - result = AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API) + result = AppGenerateService._get_workflow( + _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock() + ) assert result is pub_wf ws.get_published_workflow.assert_called_once() @@ -751,7 +753,7 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) with pytest.raises(ValueError, match="Workflow not published"): - AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API) + AppGenerateService._get_workflow(_make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, session=MagicMock()) def test_specific_workflow_id_valid_uuid(self, mocker: MockerFixture): valid_uuid = str(uuid.uuid4()) @@ -761,7 +763,10 @@ class TestGetWorkflow: mocker.patch("services.app_generate_service.WorkflowService", return_value=ws) result = AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id=valid_uuid, + session=MagicMock(), ) assert result is specific_wf ws.get_published_workflow_by_id.assert_called_once() @@ -772,7 +777,10 @@ class TestGetWorkflow: with pytest.raises(WorkflowIdFormatError): AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id="not-a-uuid" + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id="not-a-uuid", + session=MagicMock(), ) def test_specific_workflow_id_not_found(self, mocker: MockerFixture): @@ -783,7 +791,10 @@ class TestGetWorkflow: with pytest.raises(WorkflowNotFoundError): AppGenerateService._get_workflow( - _make_app(AppMode.WORKFLOW), InvokeFrom.SERVICE_API, workflow_id=valid_uuid + _make_app(AppMode.WORKFLOW), + InvokeFrom.SERVICE_API, + workflow_id=valid_uuid, + session=MagicMock(), ) @@ -804,7 +815,11 @@ class TestGenerateSingleIteration: ) app = _make_app(AppMode.ADVANCED_CHAT) result = AppGenerateService.generate_single_iteration( - app_model=app, user=_make_user(), node_id="n1", args={"k": "v"} + app_model=app, + user=_make_user(), + node_id="n1", + args={"k": "v"}, + session=MagicMock(), ) iter_spy.assert_called_once() assert result == {"event": "iteration"} @@ -822,7 +837,11 @@ class TestGenerateSingleIteration: ) app = _make_app(AppMode.WORKFLOW) result = AppGenerateService.generate_single_iteration( - app_model=app, user=_make_user(), node_id="n1", args={"k": "v"} + app_model=app, + user=_make_user(), + node_id="n1", + args={"k": "v"}, + session=MagicMock(), ) iter_spy.assert_called_once() assert result == {"event": "wf-iteration"} @@ -830,7 +849,9 @@ class TestGenerateSingleIteration: def test_invalid_mode_raises(self, mocker: MockerFixture): app = _make_app(AppMode.CHAT) with pytest.raises(ValueError, match="Invalid app mode"): - AppGenerateService.generate_single_iteration(app_model=app, user=_make_user(), node_id="n1", args={}) + AppGenerateService.generate_single_iteration( + app_model=app, user=_make_user(), node_id="n1", args={}, session=MagicMock() + ) # --------------------------------------------------------------------------- @@ -850,7 +871,11 @@ class TestGenerateSingleLoop: ) app = _make_app(AppMode.ADVANCED_CHAT) result = AppGenerateService.generate_single_loop( - app_model=app, user=_make_user(), node_id="n1", args=MagicMock() + app_model=app, + user=_make_user(), + node_id="n1", + args=MagicMock(), + session=MagicMock(), ) loop_spy.assert_called_once() assert result == {"event": "loop"} @@ -868,7 +893,11 @@ class TestGenerateSingleLoop: ) app = _make_app(AppMode.WORKFLOW) result = AppGenerateService.generate_single_loop( - app_model=app, user=_make_user(), node_id="n1", args=MagicMock() + app_model=app, + user=_make_user(), + node_id="n1", + args=MagicMock(), + session=MagicMock(), ) loop_spy.assert_called_once() assert result == {"event": "wf-loop"} @@ -876,7 +905,9 @@ class TestGenerateSingleLoop: def test_invalid_mode_raises(self, mocker: MockerFixture): app = _make_app(AppMode.COMPLETION) with pytest.raises(ValueError, match="Invalid app mode"): - AppGenerateService.generate_single_loop(app_model=app, user=_make_user(), node_id="n1", args=MagicMock()) + AppGenerateService.generate_single_loop( + app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), session=MagicMock() + ) # --------------------------------------------------------------------------- @@ -888,16 +919,18 @@ class TestGenerateMoreLikeThis: "services.app_generate_service.CompletionAppGenerator.generate_more_like_this", return_value={"result": "similar"}, ) + session = MagicMock() result = AppGenerateService.generate_more_like_this( - MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), message_id="msg-1", invoke_from=InvokeFrom.SERVICE_API, + session=session, streaming=True, ) assert result == {"result": "similar"} gen_spy.assert_called_once() + assert gen_spy.call_args.kwargs["session"] is session assert gen_spy.call_args.kwargs["stream"] is True diff --git a/api/tests/unit_tests/services/test_app_service.py b/api/tests/unit_tests/services/test_app_service.py index c57fb6ed775..36914679d3e 100644 --- a/api/tests/unit_tests/services/test_app_service.py +++ b/api/tests/unit_tests/services/test_app_service.py @@ -29,14 +29,14 @@ class TestOpenapiVisibilityHelpers: sentinel_app.status = "archived" # explicitly NOT "normal" mock_session.get.return_value = sentinel_app - assert AppService.get_app_by_id(mock_session, "app-uuid") is sentinel_app + assert AppService.get_app_by_id("app-uuid", session=mock_session) is sentinel_app mock_session.get.assert_called_once_with(App, "app-uuid") def test_get_app_by_id_returns_none_when_missing(self): mock_session = MagicMock() mock_session.get.return_value = None - assert AppService.get_app_by_id(mock_session, "missing") is None + assert AppService.get_app_by_id("missing", session=mock_session) is None def test_get_visible_app_by_id_returns_app_when_visible(self): mock_session = MagicMock() @@ -45,7 +45,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is app + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is app mock_session.get.assert_called_once_with(App, "app-uuid") @@ -53,7 +53,7 @@ class TestOpenapiVisibilityHelpers: mock_session = MagicMock() mock_session.get.return_value = None - assert AppService.get_visible_app_by_id(mock_session, "missing") is None + assert AppService.get_visible_app_by_id("missing", session=mock_session) is None def test_get_visible_app_by_id_returns_none_when_status_not_normal(self): """Soft-deleted/archived rows must not surface on the openapi @@ -65,7 +65,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is None + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is None def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self): """``is_openapi_visible`` is the per-row counterpart to @@ -78,7 +78,7 @@ class TestOpenapiVisibilityHelpers: mock_session.get.return_value = app with patch("services.app_service.is_openapi_visible", return_value=False): - assert AppService.get_visible_app_by_id(mock_session, "app-uuid") is None + assert AppService.get_visible_app_by_id("app-uuid", session=mock_session) is None def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self): """Tenant-scoped name lookup. The helper passes the SELECT through @@ -90,7 +90,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value = iter(rows) with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_name(mock_session, name="my-app", tenant_id="tenant-1") + out = AppService.find_visible_apps_by_name(name="my-app", tenant_id="tenant-1", session=mock_session) assert out == rows # Visibility gate must wrap the SELECT exactly once. @@ -102,7 +102,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value = iter([]) with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q): - out = AppService.find_visible_apps_by_name(mock_session, name="nope", tenant_id="tenant-1") + out = AppService.find_visible_apps_by_name(name="nope", tenant_id="tenant-1", session=mock_session) assert out == [] @@ -113,7 +113,7 @@ class TestOpenapiVisibilityHelpers: """ mock_session = MagicMock() - assert AppService.find_visible_apps_by_ids(mock_session, []) == [] + assert AppService.find_visible_apps_by_ids([], session=mock_session) == [] mock_session.execute.assert_not_called() def test_find_visible_apps_by_ids_passes_through_visibility_gate(self): @@ -127,7 +127,7 @@ class TestOpenapiVisibilityHelpers: mock_session.execute.return_value.scalars.return_value.all.return_value = rows with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_ids(mock_session, ["a", "b"]) + out = AppService.find_visible_apps_by_ids(["a", "b"], session=mock_session) assert out == rows gate.assert_called_once() @@ -208,6 +208,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert updated_app.name == "Iris" @@ -266,6 +267,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert backing_agent.role == "research assistant" @@ -317,6 +319,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) assert backing_agent.role == "" @@ -370,6 +373,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, + session=mock_db.session, ) mock_db.session.rollback.assert_called_once() @@ -392,7 +396,7 @@ class TestAgentAppType: patch("services.app_service.remove_app_and_related_data_task"), ): mock_db.session.scalar.return_value = backing_agent - AppService().delete_app(app) # type: ignore[arg-type] + AppService().delete_app(app, session=mock_db.session) # type: ignore[arg-type] assert backing_agent.status == AgentStatus.ARCHIVED assert backing_agent.archived_by == "account-2" diff --git a/api/tests/unit_tests/services/test_async_workflow_service.py b/api/tests/unit_tests/services/test_async_workflow_service.py index 1b9cc8a2ff6..567066845bf 100644 --- a/api/tests/unit_tests/services/test_async_workflow_service.py +++ b/api/tests/unit_tests/services/test_async_workflow_service.py @@ -331,7 +331,7 @@ class TestAsyncWorkflowService: assert trigger_log.triggered_at is not None repo.update.assert_called_once_with(trigger_log) session.commit.assert_called_once() - called_trigger_data = mock_trigger_workflow_async.call_args[0][2] + called_trigger_data = mock_trigger_workflow_async.call_args.args[1] assert isinstance(called_trigger_data, TriggerData) assert called_trigger_data.app_id == "app-123" @@ -465,11 +465,16 @@ class TestAsyncWorkflowServiceGetWorkflow: workflow_service.get_published_workflow_by_id.return_value = workflow # Act - result = AsyncWorkflowService._get_workflow(workflow_service, app_model, workflow_id="workflow-123") + session = MagicMock() + result = AsyncWorkflowService._get_workflow( + workflow_service, app_model, workflow_id="workflow-123", session=session + ) # Assert assert result == workflow - workflow_service.get_published_workflow_by_id.assert_called_once_with(app_model, "workflow-123", session=None) + workflow_service.get_published_workflow_by_id.assert_called_once_with( + app_model, "workflow-123", session=session + ) workflow_service.get_published_workflow.assert_not_called() def test_should_raise_when_specific_workflow_id_not_found(self): @@ -481,7 +486,9 @@ class TestAsyncWorkflowServiceGetWorkflow: # Act / Assert with pytest.raises(WorkflowNotFoundError, match="Published workflow not found: workflow-404"): - AsyncWorkflowService._get_workflow(workflow_service, app_model, workflow_id="workflow-404") + AsyncWorkflowService._get_workflow( + workflow_service, app_model, workflow_id="workflow-404", session=MagicMock() + ) def test_should_return_default_published_workflow_when_workflow_id_not_provided(self): """Test _get_workflow returns default published workflow when no id is provided.""" @@ -493,11 +500,12 @@ class TestAsyncWorkflowServiceGetWorkflow: workflow_service.get_published_workflow.return_value = workflow # Act - result = AsyncWorkflowService._get_workflow(workflow_service, app_model) + session = MagicMock() + result = AsyncWorkflowService._get_workflow(workflow_service, app_model, session=session) # Assert assert result == workflow - workflow_service.get_published_workflow.assert_called_once_with(app_model, session=None) + workflow_service.get_published_workflow.assert_called_once_with(app_model, session=session) workflow_service.get_published_workflow_by_id.assert_not_called() def test_should_raise_when_default_published_workflow_not_found(self): @@ -510,4 +518,4 @@ class TestAsyncWorkflowServiceGetWorkflow: # Act / Assert with pytest.raises(WorkflowNotFoundError, match="No published workflow found for app: app-123"): - AsyncWorkflowService._get_workflow(workflow_service, app_model) + AsyncWorkflowService._get_workflow(workflow_service, app_model, session=MagicMock()) diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index e5610545aa9..dc691176114 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -1115,7 +1115,7 @@ class TestBillingServiceAccountManagement: mock_db_session.scalar.return_value = mock_join # Act - should not raise exception - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) mock_db_session.scalar.assert_called_once() def test_is_tenant_owner_or_admin_admin(self, mock_db_session): @@ -1131,7 +1131,7 @@ class TestBillingServiceAccountManagement: mock_db_session.scalar.return_value = mock_join # Act - should not raise exception - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) mock_db_session.scalar.assert_called_once() def test_is_tenant_owner_or_admin_normal_user_raises_error(self, mock_db_session): @@ -1148,7 +1148,7 @@ class TestBillingServiceAccountManagement: # Act & Assert with pytest.raises(ValueError) as exc_info: - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) assert "Only team owner or team admin can perform this action" in str(exc_info.value) mock_db_session.scalar.assert_called_once() @@ -1163,7 +1163,7 @@ class TestBillingServiceAccountManagement: # Act & Assert with pytest.raises(ValueError) as exc_info: - BillingService.is_tenant_owner_or_admin(mock_db_session, current_user) + BillingService.is_tenant_owner_or_admin(current_user, session=mock_db_session) assert "Tenant account join not found" in str(exc_info.value) mock_db_session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/test_conversation_service.py b/api/tests/unit_tests/services/test_conversation_service.py index 2c7f13b79f3..e6f7b48f651 100644 --- a/api/tests/unit_tests/services/test_conversation_service.py +++ b/api/tests/unit_tests/services/test_conversation_service.py @@ -330,12 +330,9 @@ class TestConversationServiceHelpers: class TestConversationServiceConversationalVariable: """Test conversational variable operations.""" - @patch("services.conversation_service.session_factory") @patch("services.conversation_service.ConversationService.get_conversation") @patch("services.conversation_service.dify_config") - def test_get_conversational_variable_with_name_filter_mysql( - self, mock_config, mock_get_conversation, mock_session_factory - ): + def test_get_conversational_variable_with_name_filter_mysql(self, mock_config, mock_get_conversation): """ Test variable filtering by name for MySQL databases. @@ -351,7 +348,6 @@ class TestConversationServiceConversationalVariable: # Mock session mock_session = MagicMock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session mock_session.scalars.return_value.all.return_value = [] # Act @@ -362,6 +358,7 @@ class TestConversationServiceConversationalVariable: limit=10, last_id=None, variable_name="test_var", + session=mock_session, ) # Assert - JSON filter should be applied diff --git a/api/tests/unit_tests/services/test_credential_permission_service.py b/api/tests/unit_tests/services/test_credential_permission_service.py index e467e9c8c5e..cdcf4a6b00f 100644 --- a/api/tests/unit_tests/services/test_credential_permission_service.py +++ b/api/tests/unit_tests/services/test_credential_permission_service.py @@ -40,7 +40,7 @@ class TestGetPartialMemberList: session = MagicMock() session.scalars.return_value.all.return_value = [] result = CredentialPermissionService.get_partial_member_list( - session, credential_id, CredentialType.TRIGGER_SUBSCRIPTION + credential_id, CredentialType.TRIGGER_SUBSCRIPTION, session=session ) assert result == [] session.scalars.assert_called_once() @@ -49,7 +49,7 @@ class TestGetPartialMemberList: session = MagicMock() session.scalars.return_value.all.return_value = [user_id, other_user_id] result = CredentialPermissionService.get_partial_member_list( - session, credential_id, CredentialType.TRIGGER_SUBSCRIPTION + credential_id, CredentialType.TRIGGER_SUBSCRIPTION, session=session ) assert set(result) == {user_id, other_user_id} session.scalars.assert_called_once() diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index 5e589804c3d..f31d067525a 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -1,5 +1,3 @@ -from collections.abc import Generator -from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import MagicMock, patch from uuid import uuid4 @@ -7,7 +5,7 @@ from uuid import uuid4 import pytest from sqlalchemy import create_engine, select from sqlalchemy.engine import Engine -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.errors.error import QuotaExceededError from models import TenantCreditPool @@ -38,11 +36,8 @@ def _create_engine_with_pool(*, quota_limit: int, quota_used: int) -> tuple[Engi return engine, tenant_id, pool_id -@contextmanager -def _patched_session_factory(engine: Engine) -> Generator[None, None, None]: - session_maker = sessionmaker(bind=engine, expire_on_commit=False) - with patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker): - yield +def _make_session(engine: Engine) -> Session: + return sessionmaker(bind=engine, expire_on_commit=False)() def _get_quota_used(*, engine: Engine, pool_id: str) -> int | None: @@ -50,25 +45,17 @@ def _get_quota_used(*, engine: Engine, pool_id: str) -> int | None: return connection.scalar(select(TenantCreditPool.quota_used).where(TenantCreditPool.id == pool_id)) -def _make_session_maker(session: MagicMock) -> MagicMock: - session_maker = MagicMock() - transaction = session_maker.begin.return_value - transaction.__enter__.return_value = session - transaction.__exit__.return_value = None - return session_maker - - def _make_redis_lock() -> MagicMock: lock = MagicMock() lock.acquire.return_value = True return lock -def test_get_pool_uses_configured_session_factory_without_flask_app_context() -> None: +def test_get_pool_uses_provided_session() -> None: engine, tenant_id, _ = _create_engine_with_pool(quota_limit=10, quota_used=2) - with _patched_session_factory(engine): - pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL) + with _make_session(engine) as session: + pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=session) assert pool is not None assert pool.tenant_id == tenant_id @@ -78,36 +65,34 @@ def test_get_pool_uses_configured_session_factory_without_flask_app_context() -> def test_check_and_deduct_credits_deducts_exact_amount_when_sufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.check_and_deduct_credits( + tenant_id=tenant_id, credits_required=3, session=session + ) assert deducted_credits == 3 assert _get_quota_used(engine=engine, pool_id=pool_id) == 5 def test_check_and_deduct_credits_returns_zero_for_non_positive_request() -> None: - assert CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0) == 0 + assert ( + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 + ) def test_check_and_deduct_credits_raises_when_pool_is_missing() -> None: engine = create_engine("sqlite:///:memory:") TenantCreditPool.__table__.create(engine) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="Credit pool not found"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Credit pool not found"): + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1, session=session) def test_check_and_deduct_credits_raises_when_pool_is_empty() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="No credits remaining"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="No credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -115,11 +100,8 @@ def test_check_and_deduct_credits_raises_when_pool_is_empty() -> None: def test_check_and_deduct_credits_raises_without_partial_deduction_when_insufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) - with ( - _patched_session_factory(engine), - pytest.raises(QuotaExceededError, match="Insufficient credits remaining"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 9 @@ -128,25 +110,27 @@ def test_check_and_deduct_credits_wraps_unexpected_deduction_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1) + CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 def test_deduct_credits_capped_returns_zero_for_non_positive_request() -> None: - assert CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0) == 0 + assert CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 def test_deduct_credits_capped_returns_zero_when_pool_is_missing() -> None: engine = create_engine("sqlite:///:memory:") TenantCreditPool.__table__.create(engine) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=1) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=str(uuid4()), credits_required=1, session=session + ) assert deducted_credits == 0 @@ -154,8 +138,10 @@ def test_deduct_credits_capped_returns_zero_when_pool_is_missing() -> None: def test_deduct_credits_capped_returns_zero_when_pool_is_empty() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=1, session=session + ) assert deducted_credits == 0 assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -164,8 +150,10 @@ def test_deduct_credits_capped_returns_zero_when_pool_is_empty() -> None: def test_deduct_credits_capped_deducts_only_remaining_balance_when_insufficient() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) - with _patched_session_factory(engine): - deducted_credits = CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=3) + with _make_session(engine) as session: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=tenant_id, credits_required=3, session=session + ) assert deducted_credits == 1 assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 @@ -175,11 +163,11 @@ def test_deduct_credits_capped_wraps_unexpected_deduction_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 @@ -188,11 +176,11 @@ def test_deduct_credits_capped_reraises_quota_exceeded_errors() -> None: engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) with ( - _patched_session_factory(engine), + _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=QuotaExceededError("quota unavailable")), pytest.raises(QuotaExceededError, match="quota unavailable"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1) + CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 @@ -200,19 +188,18 @@ def test_deduct_credits_capped_reraises_quota_exceeded_errors() -> None: def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() -> None: tenant_id = "tenant-1" session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=10, quota_used=2) redis_lock = _make_redis_lock() with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock) as lock, - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool) as get_locked_pool, ): result = CreditPoolService.check_and_deduct_credits( tenant_id=tenant_id, credits_required=3, pool_type=ProviderQuotaType.TRIAL, + session=session, ) assert result == 3 @@ -230,19 +217,18 @@ def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() - def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> None: tenant_id = "tenant-1" session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=2, quota_used=8) redis_lock = _make_redis_lock() with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock) as lock, - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool) as get_locked_pool, ): result = CreditPoolService.deduct_credits_capped( tenant_id=tenant_id, credits_required=5, pool_type=ProviderQuotaType.PAID, + session=session, ) assert result == 2 @@ -266,38 +252,35 @@ def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> N ) def test_non_positive_credit_request_skips_tenant_redis_lock(deduct_method) -> None: with patch("services.credit_pool_service.redis_client.lock") as lock: - result = deduct_method(tenant_id="tenant-1", credits_required=0) + result = deduct_method(tenant_id="tenant-1", credits_required=0, session=MagicMock()) assert result == 0 lock.assert_not_called() def test_check_and_deduct_credits_wraps_redis_lock_errors_without_querying_db() -> None: - session_maker = MagicMock() + session = MagicMock() with ( patch("services.credit_pool_service.redis_client.lock", side_effect=RuntimeError("redis unavailable")), - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1) + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1, session=session) - session_maker.begin.assert_not_called() + session.scalar.assert_not_called() def test_deduct_credits_capped_ignores_release_errors_after_successful_deduction() -> None: session = MagicMock() - session_maker = _make_session_maker(session) pool = SimpleNamespace(remaining_credits=3, quota_used=7) redis_lock = _make_redis_lock() redis_lock.release.side_effect = RuntimeError("release failed") with ( patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock), - patch("services.credit_pool_service.session_factory.get_session_maker", return_value=session_maker), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool), ): - result = CreditPoolService.deduct_credits_capped(tenant_id="tenant-1", credits_required=2) + result = CreditPoolService.deduct_credits_capped(tenant_id="tenant-1", credits_required=2, session=session) assert result == 2 assert pool.quota_used == 9 diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 02d965f4bd2..dcb6250a8f0 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -344,7 +344,9 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session.scalar.return_value = object() with pytest.raises(DatasetNameDuplicateError, match="Dataset with name Dataset already exists"): - DatasetService.create_empty_dataset(mock_db.session, "tenant-1", "Dataset", None, "economy", account) + DatasetService.create_empty_dataset( + "tenant-1", "Dataset", None, "economy", account, session=mock_db.session + ) def test_create_empty_dataset_uses_default_embedding_model_for_high_quality_dataset(self): account = SimpleNamespace(id="user-1") @@ -512,7 +514,7 @@ class TestDatasetServiceCreationAndUpdate: session = MagicMock() with patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(ValueError, match="Dataset not found"): - DatasetService.update_dataset(session, "dataset-1", {}, SimpleNamespace(id="user-1")) + DatasetService.update_dataset("dataset-1", {}, SimpleNamespace(id="user-1"), session=session) def test_update_dataset_raises_when_new_name_conflicts(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -524,10 +526,7 @@ class TestDatasetServiceCreationAndUpdate: ): with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( - MagicMock(), - "dataset-1", - {"name": "New Dataset"}, - SimpleNamespace(id="user-1"), + "dataset-1", {"name": "New Dataset"}, SimpleNamespace(id="user-1"), session=MagicMock() ) def test_update_dataset_routes_external_datasets_to_external_helper(self): @@ -541,7 +540,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_external_dataset", return_value="updated") as update_external, ): session = MagicMock() - result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session=session) assert result == "updated" check_permission.assert_called_once() @@ -560,7 +559,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_internal_dataset", return_value="updated") as update_internal, ): session = MagicMock() - result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) + result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session=session) assert result == "updated" check_permission.assert_called_once() @@ -612,7 +611,7 @@ class TestDatasetServiceCreationAndUpdate: assert dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM assert dataset.updated_by == "user-1" assert dataset.updated_at is now - get_external_knowledge_api.assert_called_once_with(mock_db.session, "api-1", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with("api-1", dataset.tenant_id, session=mock_db.session) update_binding.assert_called_once_with("dataset-1", "knowledge-1", "api-1", mock_db.session) mock_db.session.add.assert_called_once_with(dataset) mock_db.session.commit.assert_called_once() @@ -652,7 +651,7 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session, ) - get_external_knowledge_api.assert_called_once_with(mock_db.session, "foreign-api", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id, session=mock_db.session) update_binding.assert_not_called() mock_db.session.commit.assert_not_called() @@ -1165,7 +1164,7 @@ class TestDatasetServiceRagPipelineSettings: with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id=None)): with pytest.raises(ValueError, match="Current user or current tenant not found"): - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) def test_update_rag_pipeline_dataset_settings_without_published_high_quality_updates_embedding_settings(self): session = MagicMock() @@ -1185,7 +1184,7 @@ class TestDatasetServiceRagPipelineSettings: ): model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) assert dataset.chunk_structure == "paragraph" assert dataset.indexing_technique == "high_quality" @@ -1211,7 +1210,7 @@ class TestDatasetServiceRagPipelineSettings: ) with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id="tenant-1")): - DatasetService.update_rag_pipeline_dataset_settings(session, dataset, knowledge_configuration) + DatasetService.update_rag_pipeline_dataset_settings(dataset, knowledge_configuration, session=session) assert dataset.indexing_technique == "economy" assert dataset.keyword_number == 12 @@ -1228,10 +1227,7 @@ class TestDatasetServiceRagPipelineSettings: with patch("services.dataset_service.current_user", SimpleNamespace(current_tenant_id="tenant-1")): with pytest.raises(ValueError, match="Chunk structure is not allowed to be updated"): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) def test_update_rag_pipeline_dataset_settings_with_published_rejects_switch_to_economy(self): @@ -1252,10 +1248,7 @@ class TestDatasetServiceRagPipelineSettings: match="Knowledge base indexing technique is not allowed to be updated to economy", ): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) def test_update_rag_pipeline_dataset_settings_with_published_adds_high_quality_index(self): @@ -1280,10 +1273,7 @@ class TestDatasetServiceRagPipelineSettings: model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.indexing_technique == "high_quality" @@ -1326,10 +1316,7 @@ class TestDatasetServiceRagPipelineSettings: ) DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.embedding_model_provider == "provider-two" @@ -1364,10 +1351,7 @@ class TestDatasetServiceRagPipelineSettings: ) DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.embedding_model_provider == "provider" @@ -1396,10 +1380,7 @@ class TestDatasetServiceRagPipelineSettings: patch("services.dataset_service.deal_dataset_index_update_task") as update_task, ): DatasetService.update_rag_pipeline_dataset_settings( - session, - dataset, - knowledge_configuration, - has_published=True, + dataset, knowledge_configuration, has_published=True, session=session ) assert dataset.keyword_number == 9 @@ -1457,7 +1438,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="does not have permission"): - DatasetPermissionService.check_permission(session, user, dataset, "all_team", []) + DatasetPermissionService.check_permission(user, dataset, "all_team", [], session=session) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1465,7 +1446,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): - DatasetPermissionService.check_permission(session, user, dataset, "only_me", []) + DatasetPermissionService.check_permission(user, dataset, "only_me", [], session=session) def test_check_permission_requires_partial_member_list_for_partial_members_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1473,7 +1454,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(ValueError, match="Partial member list is required"): - DatasetPermissionService.check_permission(session, user, dataset, "partial_members", []) + DatasetPermissionService.check_permission(user, dataset, "partial_members", [], session=session) def test_check_permission_rejects_dataset_operator_member_list_changes(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1485,11 +1466,7 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - session, - user, - dataset, - "partial_members", - [{"user_id": "user-2"}], + user, dataset, "partial_members", [{"user_id": "user-2"}], session=session ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged(self): @@ -1501,11 +1478,7 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( - session, - user, - dataset, - "partial_members", - [{"user_id": "user-1"}], + user, dataset, "partial_members", [{"user_id": "user-1"}], session=session ) def test_clear_partial_member_list_rolls_back_on_exception(self): diff --git a/api/tests/unit_tests/services/test_dataset_service_document.py b/api/tests/unit_tests/services/test_dataset_service_document.py index 02661fbe1f3..44619a29e83 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -1183,7 +1183,7 @@ class TestDocumentServiceTenantAndUpdateEdges: with patch("services.dataset_service.db") as mock_db: mock_db.session.scalar.return_value = 12 - result = DocumentService.get_tenant_documents_count(mock_db.session) + result = DocumentService.get_tenant_documents_count(session=mock_db.session) assert result == 12 diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index 34f3f947f96..c94093d59b7 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -306,13 +306,13 @@ class TestSegmentServiceQueries: def test_get_child_chunk_by_segment_ref_uses_full_ownership_chain(self): child_chunk = _make_child_chunk() segment_ref = _make_segment_ref() + session = MagicMock() + session.scalar.return_value = child_chunk - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = child_chunk - result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref) + result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref, session) assert result is child_chunk - stmt = mock_db.session.scalar.call_args.args[0] + stmt = session.scalar.call_args.args[0] sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) assert "child_chunks.id = 'child-a'" in sql assert "child_chunks.tenant_id = 'tenant-1'" in sql @@ -381,13 +381,13 @@ class TestSegmentServiceQueries: ) segment.id = "segment-1" segment_ref = _make_segment_ref() + session = MagicMock() + session.scalar.return_value = segment - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = segment - result = SegmentService.get_segment_by_ref(segment_ref) + result = SegmentService.get_segment_by_ref(segment_ref, session) assert result is segment - stmt = mock_db.session.scalar.call_args.args[0] + stmt = session.scalar.call_args.args[0] sql = str(stmt.compile(compile_kwargs={"literal_binds": True})) assert "document_segments.id = 'segment-1'" in sql assert "document_segments.tenant_id = 'tenant-1'" in sql @@ -566,7 +566,7 @@ class TestSegmentServiceMutations: assert all(segment.error == "vector failed" for segment in result) assert document.word_count == 5 + sum(len(item["content"]) + len(item["answer"]) for item in segments) vector_service.create_segments_vector.assert_called_once_with( - [["k1"], None], result, dataset, document.doc_form + [["k1"], None], result, dataset, document.doc_form, mock_db.session ) mock_db.session.commit.assert_called_once() @@ -641,7 +641,7 @@ class TestSegmentServiceMutations: assert result is refreshed_segment assert segment.keywords == ["new"] vector_service.update_segment_vector.assert_called_once_with(["new"], segment, dataset) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_regenerates_child_chunks_and_updates_manual_summary(self, account_context): segment = _make_segment(content="same content", word_count=len("same content")) @@ -684,10 +684,11 @@ class TestSegmentServiceMutations: dataset, embedding_model_instance, processing_rule, + mock_db.session, True, ) - update_summary.assert_called_once_with(segment, dataset, "new summary") - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_auto_regenerates_summary_after_content_change(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -725,8 +726,8 @@ class TestSegmentServiceMutations: assert segment.tokens == 9 assert document.word_count == 18 vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset) - generate_summary.assert_called_once_with(segment, dataset, {"enable": True}) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_regenerates_summary_when_manual_summary_is_unchanged(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -760,9 +761,9 @@ class TestSegmentServiceMutations: result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) assert result is refreshed_segment - generate_summary.assert_called_once_with(segment, dataset, {"enable": True}) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=mock_db.session) update_summary.assert_not_called() - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_delete_segment_removes_index_and_updates_document_word_count(self): segment = _make_segment(word_count=4, index_node_id="parent-node") @@ -972,7 +973,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: assert segment.word_count == len("question") + len("new answer") assert document.word_count == 20 + (len("question") + len("new answer") - 8) vector_service.update_segment_vector.assert_not_called() - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_content_change_uses_answer_when_counting_tokens_for_qa_segments(self, account_context): segment = _make_segment(content="old", word_count=3) @@ -1009,7 +1010,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: assert segment.tokens == 21 assert segment.word_count == len("new question") + len("new answer") vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_content_change_parent_child_uses_default_embedding_and_ignores_summary_failures( self, account_context @@ -1063,10 +1064,11 @@ class TestSegmentServiceAdditionalRegenerationBranches: dataset, embedding_model_instance, processing_rule, + mock_db.session, True, ) - update_summary.assert_called_once_with(segment, dataset, "new summary") - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) def test_update_segment_same_content_parent_child_marks_segment_error_for_non_high_quality_dataset( self, account_context diff --git a/api/tests/unit_tests/services/test_datasource_provider_service.py b/api/tests/unit_tests/services/test_datasource_provider_service.py index f374a294825..bd6891d846a 100644 --- a/api/tests/unit_tests/services/test_datasource_provider_service.py +++ b/api/tests/unit_tests/services/test_datasource_provider_service.py @@ -177,11 +177,11 @@ class TestDatasourceProviderService: def test_should_return_true_when_tenant_oauth_params_enabled(self, service, mock_db_session): mock_db_session.scalar.return_value = 1 - assert service.is_tenant_oauth_params_enabled("t1", make_id()) is True + assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is True def test_should_return_false_when_tenant_oauth_params_disabled(self, service, mock_db_session): mock_db_session.scalar.return_value = 0 - assert service.is_tenant_oauth_params_enabled("t1", make_id()) is False + assert service.is_tenant_oauth_params_enabled("t1", make_id(), session=mock_db_session) is False # ----------------------------------------------------------------------- # remove_oauth_custom_client_params (lines 55-61) @@ -453,7 +453,7 @@ class TestDatasourceProviderService: tenant_params.client_params = {"k": "v"} mock_db_session.scalar.return_value = tenant_params with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)): - result = service.get_tenant_oauth_client("t1", make_id(), mask=True) + result = service.get_tenant_oauth_client("t1", make_id(), mask=True, session=mock_db_session) assert result == {"k": "mask"} def test_should_return_decrypted_credentials_when_mask_is_false(self, service, mock_db_session): @@ -461,12 +461,12 @@ class TestDatasourceProviderService: tenant_params.client_params = {"k": "v"} mock_db_session.scalar.return_value = tenant_params with patch.object(service, "get_oauth_encrypter", return_value=(self._enc, None)): - result = service.get_tenant_oauth_client("t1", make_id(), mask=False) + result = service.get_tenant_oauth_client("t1", make_id(), mask=False, session=mock_db_session) assert result == {"k": "dec"} def test_should_return_none_when_no_tenant_oauth_config_exists(self, service, mock_db_session): mock_db_session.scalar.return_value = None - assert service.get_tenant_oauth_client("t1", make_id()) is None + assert service.get_tenant_oauth_client("t1", make_id(), session=mock_db_session) is None # ----------------------------------------------------------------------- # get_oauth_client (lines 423-457) @@ -657,7 +657,7 @@ class TestDatasourceProviderService: def test_should_return_empty_list_when_no_credentials_stored(self, service, mock_db_session): mock_db_session.scalars.return_value.all.return_value = [] - assert service.list_datasource_credentials("t1", "prov", "org/plug") == [] + assert service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == [] def test_should_return_masked_credentials_list_when_credentials_exist(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) @@ -666,7 +666,7 @@ class TestDatasourceProviderService: p.is_default = False mock_db_session.scalars.return_value.all.return_value = [p] with patch.object(service, "extract_secret_variables", return_value=["sk"]): - result = service.list_datasource_credentials("t1", "prov", "org/plug") + result = service.list_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) assert len(result) == 1 # ----------------------------------------------------------------------- @@ -682,7 +682,9 @@ class TestDatasourceProviderService: mock_mgr.return_value.fetch_installed_datasource_providers.return_value = [ds] cred = {"credential": {"k": "v"}, "is_default": True} with patch.object(service, "list_datasource_credentials", return_value=[cred]): - results = service.get_all_datasource_credentials("t1") + session = MagicMock() + session.scalar.return_value = 0 + results = service.get_all_datasource_credentials("t1", session=session) assert len(results) == 1 def test_should_include_oauth_schema_for_hardcoded_plugin_ids(self, service, mock_db_session): @@ -707,7 +709,7 @@ class TestDatasourceProviderService: patch.object(service, "is_tenant_oauth_params_enabled", return_value=False), patch.object(service, "is_system_oauth_params_exist", return_value=False), ): - results = service.get_all_datasource_credentials("t1") + results = service.get_all_datasource_credentials("t1", session=mock_db_session) assert len(results) == 1 assert results[0]["oauth_schema"] is not None @@ -717,7 +719,7 @@ class TestDatasourceProviderService: def test_should_return_empty_list_when_no_real_credentials_exist(self, service, mock_db_session): mock_db_session.scalars.return_value.all.return_value = [] - assert service.get_real_datasource_credentials("t1", "prov", "org/plug") == [] + assert service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) == [] def test_should_return_decrypted_credential_list_when_credentials_exist(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) @@ -725,7 +727,7 @@ class TestDatasourceProviderService: p.encrypted_credentials = {"sk": "v"} mock_db_session.scalars.return_value.all.return_value = [p] with patch.object(service, "extract_secret_variables", return_value=["sk"]): - result = service.get_real_datasource_credentials("t1", "prov", "org/plug") + result = service.get_real_datasource_credentials("t1", "prov", "org/plug", session=mock_db_session) assert len(result) == 1 # ----------------------------------------------------------------------- @@ -788,11 +790,11 @@ class TestDatasourceProviderService: def test_should_delete_provider_and_commit_when_found(self, service, mock_db_session): p = MagicMock(spec=DatasourceProvider) mock_db_session.scalar.return_value = p - service.remove_datasource_credentials("t1", "id", "prov", "org/plug") + service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session) mock_db_session.delete.assert_called_once_with(p) def test_should_do_nothing_when_credential_not_found_on_remove(self, service, mock_db_session): """No error raised; no delete called when record doesn't exist (lines 994 branch).""" mock_db_session.scalar.return_value = None - service.remove_datasource_credentials("t1", "id", "prov", "org/plug") + service.remove_datasource_credentials("t1", "id", "prov", "org/plug", session=mock_db_session) mock_db_session.delete.assert_not_called() diff --git a/api/tests/unit_tests/services/test_external_dataset_service.py b/api/tests/unit_tests/services/test_external_dataset_service.py index dbb4627759c..9dff74f8dd5 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -145,7 +145,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_success_basic( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test successful retrieval of external knowledge APIs with pagination.""" # Arrange @@ -158,7 +158,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 5 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -170,11 +170,11 @@ class TestExternalDatasetServiceGetAPIs: assert result_total == 5 assert result_items[0].id == "api-0" assert result_items[4].id == "api-4" - mock_paginate.assert_called_once() + mock_paginate_query.assert_called_once() @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_with_search_filter( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with search filter.""" # Arrange @@ -186,7 +186,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -200,14 +200,14 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_empty_results( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with no results.""" # Arrange mock_pagination = MagicMock() mock_pagination.items = [] mock_pagination.total = 0 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -220,7 +220,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_large_result_set( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with large result set.""" # Arrange @@ -229,7 +229,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[:10] mock_pagination.total = 100 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -242,7 +242,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_pagination_last_page( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test last page pagination with partial results.""" # Arrange @@ -251,7 +251,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 100 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -264,7 +264,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_case_insensitive_search( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test case-insensitive search functionality.""" # Arrange @@ -276,7 +276,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 2 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -289,7 +289,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_special_characters_search( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test search with special characters.""" # Arrange @@ -298,7 +298,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -310,7 +310,7 @@ class TestExternalDatasetServiceGetAPIs: @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_max_per_page_limit( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test that max_per_page limit is enforced.""" # Arrange @@ -319,7 +319,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1000 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -327,12 +327,12 @@ class TestExternalDatasetServiceGetAPIs: ) # Assert - call_args = mock_paginate.call_args + call_args = mock_paginate_query.call_args assert call_args.kwargs["max_per_page"] == 100 @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_ordered_by_created_at_desc( - self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate_query, factory: ExternalDatasetServiceTestDataFactory ): """Test that results are ordered by created_at descending.""" # Arrange @@ -344,7 +344,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[::-1] # Reversed to simulate DESC order mock_pagination.total = 5 - mock_paginate.return_value = mock_pagination + mock_paginate_query.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -437,13 +437,13 @@ class TestExternalDatasetServiceValidateAPIList: class TestExternalDatasetServiceCreateAPI: """Test create_external_knowledge_api operations.""" + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_success_full( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test successful creation with all fields.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -453,7 +453,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, session=mock_db.session) # Assert assert result.name == "Test API" @@ -462,55 +462,63 @@ class TestExternalDatasetServiceCreateAPI: assert result.created_by == user_id assert result.updated_by == user_id mock_check.assert_called_once_with(args["settings"]) - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() + mock_db.session.add.assert_called_once() + mock_db.session.commit.assert_called_once() + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_minimal_fields( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with minimal required fields.""" # Arrange - mock_session = MagicMock() args = { "name": "Minimal API", "settings": {"endpoint": "https://api.example.com", "api_key": "key"}, } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.name == "Minimal API" assert result.description == "" - def test_create_external_knowledge_api_missing_settings(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_knowledge_api_missing_settings( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test creation fails when settings are missing.""" # Arrange - mock_session = MagicMock() args = {"name": "Test API", "description": "Test"} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) - def test_create_external_knowledge_api_none_settings(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_knowledge_api_none_settings(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test creation fails when settings are explicitly None.""" # Arrange - mock_session = MagicMock() args = {"name": "Test API", "settings": None} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_settings_json_serialization( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that settings are properly JSON serialized.""" # Arrange - mock_session = MagicMock() settings = { "endpoint": "https://api.example.com", "api_key": "test-key", @@ -519,20 +527,22 @@ class TestExternalDatasetServiceCreateAPI: args = {"name": "Test API", "settings": settings} # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert isinstance(result.settings, str) parsed_settings = json.loads(result.settings) assert parsed_settings == settings + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_unicode_handling( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test proper handling of Unicode characters in name and description.""" # Arrange - mock_session = MagicMock() args = { "name": "测试API", "description": "テストの説明", @@ -540,19 +550,21 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.name == "测试API" assert result.description == "テストの説明" + @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_long_description( - self, mock_check, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with very long description.""" # Arrange - mock_session = MagicMock() long_description = "A" * 1000 args = { "name": "Test API", @@ -561,7 +573,9 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) + result = ExternalDatasetService.create_external_knowledge_api( + "tenant-123", "user-123", args, session=mock_db.session + ) # Assert assert result.description == long_description @@ -824,43 +838,43 @@ class TestExternalDatasetServiceCheckEndpoint: class TestExternalDatasetServiceGetAPI: """Test get_external_knowledge_api operations.""" - def test_get_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge API.""" # Arrange - mock_session = MagicMock() api_id = "api-123" expected_api = factory.create_external_knowledge_api_mock(api_id=api_id) - mock_session.scalar.return_value = expected_api + mock_db.session.scalar.return_value = expected_api # Act tenant_id = "tenant-123" - result = ExternalDatasetService.get_external_knowledge_api(mock_session, api_id, tenant_id) + result = ExternalDatasetService.get_external_knowledge_api(api_id, tenant_id, session=mock_db.session) # Assert assert result.id == api_id - def test_get_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.get_external_knowledge_api(mock_session, "nonexistent-id", "tenant-123") + ExternalDatasetService.get_external_knowledge_api("nonexistent-id", "tenant-123", session=mock_db.session) class TestExternalDatasetServiceUpdateAPI: """Test update_external_knowledge_api operations.""" @patch("services.external_knowledge_service.naive_utc_now") + @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_success_all_fields( - self, mock_now, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_now, factory: ExternalDatasetServiceTestDataFactory ): """Test successful update with all fields.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" user_id = "user-456" @@ -875,24 +889,26 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://new.example.com", "api_key": "new-key"}, } - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, user_id, api_id, args) + result = ExternalDatasetService.update_external_knowledge_api( + tenant_id, user_id, api_id, args, session=mock_db.session + ) # Assert assert result.name == "Updated API" assert result.description == "Updated description" assert result.updated_by == user_id assert result.updated_at == current_time - mock_session.commit.assert_called_once() + mock_db.session.commit.assert_called_once() + @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_preserve_hidden_api_key( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that hidden API key is preserved from existing settings.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" @@ -907,47 +923,51 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://api.example.com", "api_key": HIDDEN_VALUE}, } - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, "user-123", api_id, args) + result = ExternalDatasetService.update_external_knowledge_api( + tenant_id, "user-123", api_id, args, session=mock_db.session + ) # Assert settings = json.loads(result.settings) assert settings["api_key"] == "original-secret-key" - def test_update_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_session, "tenant-123", "user-123", "api-123", args + "tenant-123", "user-123", "api-123", args, session=mock_db.session ) - def test_update_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_tenant_mismatch( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when tenant ID doesn't match.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_session, "wrong-tenant", "user-123", "api-123", args + "wrong-tenant", "user-123", "api-123", args, session=mock_db.session ) - def test_update_external_knowledge_api_name_only(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_update_external_knowledge_api_name_only(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test updating only the name field.""" # Arrange - mock_session = MagicMock() existing_api = factory.create_external_knowledge_api_mock( description="Original description", settings={"endpoint": "https://api.example.com", "api_key": "key"}, @@ -955,11 +975,11 @@ class TestExternalDatasetServiceUpdateAPI: args = {"name": "New Name Only"} - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act result = ExternalDatasetService.update_external_knowledge_api( - mock_session, "tenant-123", "user-123", "api-123", args + "tenant-123", "user-123", "api-123", args, session=mock_db.session ) # Assert @@ -969,92 +989,104 @@ class TestExternalDatasetServiceUpdateAPI: class TestExternalDatasetServiceDeleteAPI: """Test delete_external_knowledge_api operations.""" - def test_delete_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful deletion of external knowledge API.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" existing_api = factory.create_external_knowledge_api_mock(api_id=api_id, tenant_id=tenant_id) - mock_session.scalar.return_value = existing_api + mock_db.session.scalar.return_value = existing_api # Act - ExternalDatasetService.delete_external_knowledge_api(mock_session, tenant_id, api_id) + ExternalDatasetService.delete_external_knowledge_api(tenant_id, api_id, session=mock_db.session) # Assert - mock_session.delete.assert_called_once_with(existing_api) - mock_session.commit.assert_called_once() + mock_db.session.delete.assert_called_once_with(existing_api) + mock_db.session.commit.assert_called_once() - def test_delete_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_session, "tenant-123", "api-123") + ExternalDatasetService.delete_external_knowledge_api("tenant-123", "api-123", session=mock_db.session) - def test_delete_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_delete_external_knowledge_api_tenant_mismatch( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when tenant ID doesn't match.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_session, "wrong-tenant", "api-123") + ExternalDatasetService.delete_external_knowledge_api("wrong-tenant", "api-123", session=mock_db.session) class TestExternalDatasetServiceAPIUseCheck: """Test external_knowledge_api_use_check operations.""" - def test_external_knowledge_api_use_check_in_use_single(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_in_use_single( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test API use check when API has one binding.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 1 + mock_db.session.scalar.return_value = 1 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is True assert count == 1 - assert "tenant_id" in str(mock_session.scalar.call_args.args[0]) + assert "tenant_id" in str(mock_db.session.scalar.call_args.args[0]) - def test_external_knowledge_api_use_check_in_use_multiple(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_in_use_multiple( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test API use check with multiple bindings.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 10 + mock_db.session.scalar.return_value = 10 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is True assert count == 10 - def test_external_knowledge_api_use_check_not_in_use(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_external_knowledge_api_use_check_not_in_use(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test API use check when API is not in use.""" # Arrange - mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_session.scalar.return_value = 0 + mock_db.session.scalar.return_value = 0 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check( + api_id, tenant_id, session=mock_db.session + ) # Assert assert in_use is False @@ -1064,46 +1096,48 @@ class TestExternalDatasetServiceAPIUseCheck: class TestExternalDatasetServiceGetBinding: """Test get_external_knowledge_binding_with_dataset_id operations.""" - def test_get_external_knowledge_binding_success(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_binding_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge binding.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" expected_binding = factory.create_external_knowledge_binding_mock(tenant_id=tenant_id, dataset_id=dataset_id) - mock_session.scalar.return_value = expected_binding + mock_db.session.scalar.return_value = expected_binding # Act result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_session, tenant_id, dataset_id + tenant_id, dataset_id, session=mock_db.session ) # Assert assert result.dataset_id == dataset_id assert result.tenant_id == tenant_id - def test_get_external_knowledge_binding_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_get_external_knowledge_binding_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when binding is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="external knowledge binding not found"): ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_session, "tenant-123", "dataset-123" + "tenant-123", "dataset-123", session=mock_db.session ) class TestExternalDatasetServiceDocumentValidate: """Test document_create_args_validate operations.""" - def test_document_create_args_validate_success_all_params(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_success_all_params( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test successful validation with all required parameters.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1117,17 +1151,21 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {"param1": "value1", "param2": "value2"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate( + tenant_id, api_id, process_parameter, session=mock_db.session + ) - def test_document_create_args_validate_missing_required_param(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_missing_required_param( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test validation fails when required parameter is missing.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1135,42 +1173,46 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {} # Act & Assert with pytest.raises(ValueError, match="required_param is required"): - ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate( + tenant_id, api_id, process_parameter, session=mock_db.session + ) - def test_document_create_args_validate_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test validation fails when API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) - def test_document_create_args_validate_no_custom_parameters(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_document_create_args_validate_no_custom_parameters( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test validation succeeds when no custom parameters defined.""" # Arrange - mock_session = MagicMock() settings = {} api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) + @patch("services.external_knowledge_service.db") def test_document_create_args_validate_optional_params_not_required( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test that optional parameters don't cause validation failure.""" # Arrange - mock_session = MagicMock() settings = { "document_process_setting": [ {"name": "required_param", "required": True}, @@ -1180,12 +1222,14 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_session.scalar.return_value = api + mock_db.session.scalar.return_value = api process_parameter = {"required_param": "value"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", process_parameter) + ExternalDatasetService.document_create_args_validate( + "tenant-123", "api-123", process_parameter, session=mock_db.session + ) class TestExternalDatasetServiceProcessAPI: @@ -1475,10 +1519,10 @@ class TestExternalDatasetServiceGetSettings: class TestExternalDatasetServiceCreateDataset: """Test create_external_dataset operations.""" - def test_create_external_dataset_success_full(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_success_full(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test successful creation of external dataset with all fields.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -1491,84 +1535,90 @@ class TestExternalDatasetServiceCreateDataset: api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] # Act - result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, mock_session) + result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, session=mock_db.session) # Assert assert result.name == "Test External Dataset" assert result.description == "Comprehensive test description" assert result.provider == "external" assert result.created_by == user_id - mock_session.add.assert_called() - mock_session.commit.assert_called_once() + mock_db.session.add.assert_called() + mock_db.session.commit.assert_called_once() - def test_create_external_dataset_duplicate_name_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_duplicate_name_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when dataset name already exists.""" # Arrange - mock_session = MagicMock() existing_dataset = factory.create_dataset_mock(name="Duplicate Dataset") - mock_session.scalar.return_value = existing_dataset + mock_db.session.scalar.return_value = existing_dataset args = {"name": "Duplicate Dataset"} # Act & Assert with pytest.raises(DatasetNameDuplicateError): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_api_not_found_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_api_not_found_error(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): """Test error when external knowledge API is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.side_effect = [None, None] + mock_db.session.scalar.side_effect = [None, None] args = {"name": "Test Dataset", "external_knowledge_api_id": "nonexistent-api"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_missing_knowledge_id_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_missing_knowledge_id_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when external_knowledge_id is missing.""" # Arrange - mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_api_id": "api-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) - def test_create_external_dataset_missing_api_id_error(self, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.db") + def test_create_external_dataset_missing_api_id_error( + self, mock_db, factory: ExternalDatasetServiceTestDataFactory + ): """Test error when external_knowledge_api_id is missing.""" # Arrange - mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [None, api] + mock_db.session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_id": "knowledge-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_api_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) class TestExternalDatasetServiceFetchRetrieval: """Test fetch_external_knowledge_retrieval operations.""" @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_success_with_results( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test successful external knowledge retrieval with results.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" query = "test query for retrieval" @@ -1578,7 +1628,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1594,7 +1644,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, tenant_id, dataset_id, query, external_retrieval_parameters + tenant_id, + dataset_id, + query, + external_retrieval_parameters, + session=mock_db.session, ) # Assert @@ -1602,46 +1656,46 @@ class TestExternalDatasetServiceFetchRetrieval: assert result[0]["content"] == "result 1" assert result[1]["score"] == 0.8 + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_binding_not_found_error( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test error when external knowledge binding is not found.""" # Arrange - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external knowledge binding not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {} + "tenant-123", "dataset-123", "query", {}, session=mock_db.session ) + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_cross_tenant_api_template_error( - self, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, factory: ExternalDatasetServiceTestDataFactory ): """Test error when a binding points to an API template outside the dataset tenant.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() - mock_session.scalar.side_effect = [binding, None] + mock_db.session.scalar.side_effect = [binding, None] # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external api template not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {} + "tenant-123", "dataset-123", "query", {}, session=mock_db.session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_results( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with empty results.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1650,23 +1704,27 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) # Assert assert len(result) == 0 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_with_score_threshold( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with score threshold enabled.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1681,7 +1739,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", external_retrieval_parameters + "tenant-123", + "dataset-123", + "query", + external_retrieval_parameters, + session=mock_db.session, ) # Assert @@ -1691,16 +1753,16 @@ class TestExternalDatasetServiceFetchRetrieval: assert call_args.params["retrieval_setting"]["score_threshold"] == 0.75 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_non_200_status_raises_exception( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test that non-200 status code raises Exception with response text.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 500 @@ -1710,7 +1772,11 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="Internal Server Error: Database connection failed"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @pytest.mark.parametrize( @@ -1727,12 +1793,12 @@ class TestExternalDatasetServiceFetchRetrieval: ], ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_various_error_status_codes( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message ): """Test that various error status codes raise exceptions with response text.""" # Arrange - mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" @@ -1741,7 +1807,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = status_code @@ -1751,20 +1817,20 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match=re.escape(error_message)): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, tenant_id, dataset_id, "query", {"top_k": 5} + tenant_id, dataset_id, "query", {"top_k": 5}, session=mock_db.session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") + @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_response_text( - self, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test exception with empty response text.""" # Arrange - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 503 @@ -1774,17 +1840,21 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_db, mock_process, factory): """Test malformed JSON success responses are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1793,17 +1863,21 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_db, mock_process, factory): """Test malformed success payload shapes are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1812,17 +1886,21 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_db, mock_process, factory): """Test non-list records payloads are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1831,20 +1909,28 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_process, factory): + @patch("services.external_knowledge_service.db") + def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_db, mock_process, factory): """Test transport/runtime failures are normalized to external retrieval errors.""" - mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_session.scalar.side_effect = [binding, api] + mock_db.session.scalar.side_effect = [binding, api] mock_process.side_effect = RuntimeError("connection reset by peer") with pytest.raises(ExternalKnowledgeRetrievalError, match="connection reset by peer"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} + "tenant-123", + "dataset-123", + "query", + {"top_k": 5}, + session=mock_db.session, ) diff --git a/api/tests/unit_tests/services/test_file_service.py b/api/tests/unit_tests/services/test_file_service.py index b81fb823949..41b86fda0cb 100644 --- a/api/tests/unit_tests/services/test_file_service.py +++ b/api/tests/unit_tests/services/test_file_service.py @@ -377,7 +377,7 @@ class TestFileService: def test_get_upload_files_by_ids_empty(self): session = MagicMock() - result = FileService.get_upload_files_by_ids(session, "tenant_id", []) + result = FileService.get_upload_files_by_ids("tenant_id", [], session=session) assert result == {} def test_get_upload_files_by_ids(self): @@ -387,7 +387,9 @@ class TestFileService: session = MagicMock() session.scalars().all.return_value = [upload_file] - result = FileService.get_upload_files_by_ids(session, "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"]) + result = FileService.get_upload_files_by_ids( + "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"], session=session + ) assert result["550e8400-e29b-41d4-a716-446655440000"] == upload_file def test_sanitize_zip_entry_name(self): diff --git a/api/tests/unit_tests/services/test_message_service.py b/api/tests/unit_tests/services/test_message_service.py index 6588c8a8de6..13f340e9f4a 100644 --- a/api/tests/unit_tests/services/test_message_service.py +++ b/api/tests/unit_tests/services/test_message_service.py @@ -102,6 +102,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -124,6 +125,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="", first_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -166,6 +168,7 @@ class TestMessageServicePaginationByFirstId: first_id=None, limit=10, order="desc", + session=mock_db.session, ) # Assert @@ -209,6 +212,7 @@ class TestMessageServicePaginationByFirstId: first_id=None, limit=10, order="asc", + session=mock_db.session, ) # Assert @@ -258,6 +262,7 @@ class TestMessageServicePaginationByFirstId: first_id="msg-005", limit=10, order="desc", + session=mock_db.session, ) # Assert @@ -288,6 +293,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id="nonexistent-msg", limit=10, + session=mock_db.session, ) # Test 07: Has_more flag when results exceed limit @@ -323,6 +329,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -353,6 +360,7 @@ class TestMessageServicePaginationByFirstId: conversation_id="conv-001", first_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -389,6 +397,7 @@ class TestMessageServicePaginationByLastId: user=None, last_id=None, limit=10, + session=MagicMock(), ) # Assert @@ -421,6 +430,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -459,6 +469,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id="msg-005", limit=10, + session=mock_db.session, ) # Assert @@ -482,6 +493,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id="nonexistent-msg", limit=10, + session=mock_db.session, ) # Test 13: Pagination with conversation_id filter @@ -516,6 +528,7 @@ class TestMessageServicePaginationByLastId: last_id=None, limit=10, conversation_id="conv-001", + session=mock_db.session, ) # Assert @@ -546,6 +559,7 @@ class TestMessageServicePaginationByLastId: last_id=None, limit=10, include_ids=["msg-001", "msg-003"], + session=mock_db.session, ) # Assert @@ -578,6 +592,7 @@ class TestMessageServicePaginationByLastId: user=user, last_id=None, limit=10, + session=mock_db.session, ) # Assert @@ -680,8 +695,8 @@ class TestMessageServiceGetMessage: mock_db.session.scalar.return_value = message - # Act - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123") + # Act, + result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) # Assert assert result == message @@ -700,8 +715,8 @@ class TestMessageServiceGetMessage: mock_db.session.scalar.return_value = message - # Act - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123") + # Act, + result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) # Assert assert result == message @@ -718,7 +733,7 @@ class TestMessageServiceGetMessage: # Act & Assert with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=user, message_id="msg-123") + MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) class TestMessageServiceFeedback: @@ -748,6 +763,7 @@ class TestMessageServiceFeedback: user=user, rating=FeedbackRating.LIKE, content="Good answer", + session=mock_db.session, ) # Assert @@ -780,6 +796,7 @@ class TestMessageServiceFeedback: user=user, rating=FeedbackRating.DISLIKE, content="Bad answer", + session=mock_db.session, ) # Assert @@ -808,6 +825,7 @@ class TestMessageServiceFeedback: user=user, rating=None, content=None, + session=mock_db.session, ) # Assert @@ -826,8 +844,8 @@ class TestMessageServiceFeedback: mock_db.session.scalars.return_value.all.return_value = [feedback] - # Act - result = MessageService.get_all_messages_feedbacks(app_model=app, page=1, limit=10) + # Act, + result = MessageService.get_all_messages_feedbacks(app_model=app, page=1, limit=10, session=mock_db.session) # Assert assert result == [{"id": "fb-1"}] @@ -846,7 +864,11 @@ class TestMessageServiceSuggestedQuestions: app = factory.create_app_mock() with pytest.raises(ValueError, match="user cannot be None"): MessageService.get_suggested_questions_after_answer( - app_model=app, user=None, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=None, + message_id="msg-123", + invoke_from=MagicMock(), + session=MagicMock(), ) # Test 28: get_suggested_questions_after_answer - Advanced Chat success @@ -890,7 +912,11 @@ class TestMessageServiceSuggestedQuestions: # Act result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=InvokeFrom.WEB_APP + app_model=app, + user=user, + message_id="msg-123", + invoke_from=InvokeFrom.WEB_APP, + session=MagicMock(), ) # Assert @@ -938,7 +964,11 @@ class TestMessageServiceSuggestedQuestions: # Act result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) # Assert @@ -996,6 +1026,7 @@ class TestMessageServiceSuggestedQuestions: user=user, message_id="msg-123", invoke_from=InvokeFrom.WEB_APP, + session=mock_db.session, ) assert result == ["Q1?"] @@ -1059,7 +1090,11 @@ class TestMessageServiceSuggestedQuestions: mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) assert result == ["Q1?"] @@ -1168,7 +1203,11 @@ class TestMessageServiceSuggestedQuestions: mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] result = MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=mock_db.session, ) assert result == ["Q1?"] @@ -1209,5 +1248,9 @@ class TestMessageServiceSuggestedQuestions: # Act & Assert with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=app, user=user, message_id="msg-123", invoke_from=MagicMock() + app_model=app, + user=user, + message_id="msg-123", + invoke_from=MagicMock(), + session=MagicMock(), ) diff --git a/api/tests/unit_tests/services/test_metadata_bug_complete.py b/api/tests/unit_tests/services/test_metadata_bug_complete.py index 6792243e9d0..00f16f75ac0 100644 --- a/api/tests/unit_tests/services/test_metadata_bug_complete.py +++ b/api/tests/unit_tests/services/test_metadata_bug_complete.py @@ -48,14 +48,14 @@ class TestMetadataBugCompleteValidation: account = _make_account() # Should crash with TypeError with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) # Test update method as well account = _make_account() none_name = cast(str, None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): MetadataService.update_metadata_name( - Mock(), "dataset-123", "metadata-456", none_name, account, "tenant-123" + "dataset-123", "metadata-456", none_name, account, "tenant-123", session=Mock() ) def test_3_database_constraints_verification(self) -> None: @@ -99,7 +99,7 @@ class TestMetadataBugCompleteValidation: account = _make_account() with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) def test_7_end_to_end_validation_layers(self) -> None: """Test all validation layers work together correctly.""" diff --git a/api/tests/unit_tests/services/test_metadata_nullable_bug.py b/api/tests/unit_tests/services/test_metadata_nullable_bug.py index ae93fe5ef51..cfd3d034df2 100644 --- a/api/tests/unit_tests/services/test_metadata_nullable_bug.py +++ b/api/tests/unit_tests/services/test_metadata_nullable_bug.py @@ -37,7 +37,7 @@ class TestMetadataNullableBug: account = _make_account() # This should crash with TypeError when calling len(None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata(Mock(), "dataset-123", mock_metadata_args, account, "tenant-123") + MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) def test_metadata_service_update_with_none_name_crashes(self) -> None: """Test that MetadataService.update_metadata_name crashes when name is None.""" @@ -46,7 +46,7 @@ class TestMetadataNullableBug: # This should crash with TypeError when calling len(None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): MetadataService.update_metadata_name( - Mock(), "dataset-123", "metadata-456", none_name, account, "tenant-123" + "dataset-123", "metadata-456", none_name, account, "tenant-123", session=Mock() ) def test_api_layer_now_uses_pydantic_validation(self) -> None: diff --git a/api/tests/unit_tests/services/test_model_load_balancing_service.py b/api/tests/unit_tests/services/test_model_load_balancing_service.py index 827567f1afe..743e6e797a3 100644 --- a/api/tests/unit_tests/services/test_model_load_balancing_service.py +++ b/api/tests/unit_tests/services/test_model_load_balancing_service.py @@ -80,9 +80,9 @@ def service(mocker: MockerFixture) -> ModelLoadBalancingService: @pytest.fixture -def mock_db(mocker: MockerFixture) -> MagicMock: +def mock_db() -> MagicMock: # Arrange - mocked_db = mocker.patch("services.model_load_balancing_service.db") + mocked_db = MagicMock() mocked_db.session = MagicMock() return mocked_db @@ -159,7 +159,7 @@ def test_get_load_balancing_configs_should_raise_value_error_when_provider_missi # Act + Assert with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_configs("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM) + service.get_load_balancing_configs("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, session=MagicMock()) def test_get_load_balancing_configs_should_insert_inherit_config_when_missing_for_custom_provider( @@ -201,6 +201,7 @@ def test_get_load_balancing_configs_should_insert_inherit_config_when_missing_fo "openai", "gpt-4o-mini", ModelType.LLM, + session=mock_db.session, ) # Assert @@ -263,6 +264,7 @@ def test_get_load_balancing_configs_should_reorder_existing_inherit_and_tolerate "gpt-4o-mini", ModelType.LLM, config_from="predefined-model", + session=mock_db.session, ) # Assert @@ -282,7 +284,9 @@ def test_get_load_balancing_config_should_raise_value_error_when_provider_missin # Act + Assert with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=MagicMock() + ) def test_get_load_balancing_config_should_return_none_when_config_not_found( @@ -295,7 +299,9 @@ def test_get_load_balancing_config_should_return_none_when_config_not_found( mock_db.session.scalar.return_value = None # Act - result = service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + result = service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session + ) # Assert assert result is None @@ -315,7 +321,9 @@ def test_get_load_balancing_config_should_return_obfuscated_payload_when_config_ mock_db.session.scalar.return_value = config # Act - result = service.get_load_balancing_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1") + result = service.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session + ) # Assert assert result == { @@ -334,7 +342,9 @@ def test_init_inherit_config_should_create_and_persist_inherit_configuration( model_type = ModelType.LLM # Act - inherit_config = service._init_inherit_config("tenant-1", "openai", "gpt-4o-mini", model_type) + inherit_config = service._init_inherit_config( + "tenant-1", "openai", "gpt-4o-mini", model_type, session=mock_db.session + ) # Assert assert inherit_config.tenant_id == "tenant-1" @@ -361,6 +371,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_provider_mi ModelType.LLM, [], "custom-model", + session=MagicMock(), ) @@ -380,6 +391,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_configs_is_ ModelType.LLM, cast(list[dict[str, object]], "invalid-configs"), "custom-model", + session=MagicMock(), ) @@ -401,6 +413,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_config_item ModelType.LLM, cast(list[dict[str, object]], ["bad-item"]), "custom-model", + session=mock_db.session, ) @@ -423,6 +436,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credential_ ModelType.LLM, [{"credential_id": "cred-1", "enabled": True}], "predefined-model", + session=mock_db.session, ) @@ -444,6 +458,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_name_or_ena ModelType.LLM, [{"enabled": True}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config enabled"): @@ -454,6 +469,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_name_or_ena ModelType.LLM, [{"name": "cfg-without-enabled"}], "custom-model", + session=mock_db.session, ) @@ -476,6 +492,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_existing_co ModelType.LLM, [{"id": "cfg-2", "name": "invalid", "enabled": True}], "custom-model", + session=mock_db.session, ) @@ -498,6 +515,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credentials ModelType.LLM, [{"id": "cfg-1", "name": "new", "enabled": True, "credentials": "bad"}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config credentials"): @@ -508,6 +526,7 @@ def test_update_load_balancing_configs_should_raise_value_error_when_credentials ModelType.LLM, [{"name": "new-config", "enabled": True, "credentials": "bad"}], "custom-model", + session=mock_db.session, ) @@ -548,6 +567,7 @@ def test_update_load_balancing_configs_should_update_existing_create_new_and_del {"name": "new-config", "enabled": True, "credentials": {"api_key": "plain"}}, ], "custom-model", + session=mock_db.session, ) # Assert @@ -579,6 +599,7 @@ def test_update_load_balancing_configs_should_raise_value_error_for_invalid_new_ ModelType.LLM, [{"name": "__inherit__", "enabled": True, "credentials": {"api_key": "x"}}], "custom-model", + session=mock_db.session, ) with pytest.raises(ValueError, match="Invalid load balancing config credentials"): @@ -589,6 +610,7 @@ def test_update_load_balancing_configs_should_raise_value_error_for_invalid_new_ ModelType.LLM, [{"name": "new", "enabled": True}], "custom-model", + session=mock_db.session, ) @@ -611,6 +633,7 @@ def test_update_load_balancing_configs_should_create_from_existing_provider_cred ModelType.LLM, [{"credential_id": "cred-1", "enabled": True}], "predefined-model", + session=mock_db.session, ) # Assert @@ -636,6 +659,7 @@ def test_validate_load_balancing_credentials_should_raise_value_error_when_provi "gpt-4o-mini", ModelType.LLM, {"api_key": "plain"}, + session=MagicMock(), ) @@ -657,6 +681,7 @@ def test_validate_load_balancing_credentials_should_raise_value_error_when_confi ModelType.LLM, {"api_key": "plain"}, config_id="cfg-1", + session=mock_db.session, ) @@ -680,6 +705,7 @@ def test_validate_load_balancing_credentials_should_delegate_to_custom_validate_ ModelType.LLM, {"api_key": "plain"}, config_id="cfg-1", + session=mock_db.session, ) service.validate_load_balancing_credentials( "tenant-1", @@ -687,6 +713,7 @@ def test_validate_load_balancing_credentials_should_delegate_to_custom_validate_ "gpt-4o-mini", ModelType.LLM, {"api_key": "plain"}, + session=mock_db.session, ) # Assert diff --git a/api/tests/unit_tests/services/test_oauth_device_flow.py b/api/tests/unit_tests/services/test_oauth_device_flow.py index fcb3f29a76f..00b2919240d 100644 --- a/api/tests/unit_tests/services/test_oauth_device_flow.py +++ b/api/tests/unit_tests/services/test_oauth_device_flow.py @@ -83,7 +83,7 @@ def test_revoke_oauth_token_invalidates_redis_cache_when_live_hash_seen(): redis = MagicMock() - revoke_oauth_token(session, redis, "token-id") + revoke_oauth_token(redis, "token-id", session=session) assert session.execute.called # UPDATE ... WHERE revoked_at IS NULL assert session.commit.called @@ -101,7 +101,7 @@ def test_revoke_oauth_token_is_idempotent_when_already_revoked(): redis = MagicMock() - revoke_oauth_token(session, redis, "token-id") + revoke_oauth_token(redis, "token-id", session=session) assert session.execute.called assert session.commit.called @@ -126,7 +126,7 @@ def test_list_active_sessions_returns_session_execute_rows(): fake_rows = [MagicMock(), MagicMock()] session.execute.return_value.scalars.return_value.all.return_value = fake_rows - out = list_active_sessions(session, _account_ctx(), datetime.now(UTC)) + out = list_active_sessions(_account_ctx(), datetime.now(UTC), session=session) assert out == fake_rows assert session.execute.called @@ -136,11 +136,11 @@ def test_token_belongs_to_subject_true_when_row_present(): session = MagicMock() session.execute.return_value.first.return_value = ("some-id",) - assert token_belongs_to_subject(session, "token-id", _account_ctx()) is True + assert token_belongs_to_subject("token-id", _account_ctx(), session=session) is True def test_token_belongs_to_subject_false_when_no_row(): session = MagicMock() session.execute.return_value.first.return_value = None - assert token_belongs_to_subject(session, "token-id", _account_ctx()) is False + assert token_belongs_to_subject("token-id", _account_ctx(), session=session) is False diff --git a/api/tests/unit_tests/services/test_summary_index_service.py b/api/tests/unit_tests/services/test_summary_index_service.py index 19418c43926..7ece6204ce3 100644 --- a/api/tests/unit_tests/services/test_summary_index_service.py +++ b/api/tests/unit_tests/services/test_summary_index_service.py @@ -118,7 +118,7 @@ def test_generate_summary_for_segment_raises_when_empty(monkeypatch: pytest.Monk SummaryIndexService.generate_summary_for_segment(_segment(), _dataset(), {"a": 1}) -def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_summary_record_updates_existing_and_reenables() -> None: existing = _summary_record(summary_content="old", node_id="n1") existing.enabled = False existing.disabled_at = datetime(2024, 1, 1) @@ -127,13 +127,12 @@ def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytes session = MagicMock(name="session") session.scalar.return_value = existing - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - segment = _segment() dataset = _dataset() - result = SummaryIndexService.create_summary_record(segment, dataset, "new", status=SummaryStatus.GENERATING) + result = SummaryIndexService.create_summary_record( + segment, dataset, "new", status=SummaryStatus.GENERATING, session=session + ) assert result is existing assert existing.summary_content == "new" assert existing.status == SummaryStatus.GENERATING @@ -145,14 +144,13 @@ def test_create_summary_record_updates_existing_and_reenables(monkeypatch: pytes session.flush.assert_called_once() -def test_create_summary_record_creates_new(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_summary_record_creates_new() -> None: session = MagicMock(name="session") session.scalar.return_value = None - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - - record = SummaryIndexService.create_summary_record(_segment(), _dataset(), "new", status=SummaryStatus.GENERATING) + record = SummaryIndexService.create_summary_record( + _segment(), _dataset(), "new", status=SummaryStatus.GENERATING, session=session + ) assert record.dataset_id == "dataset-1" assert record.chunk_id == "seg-1" assert record.summary_content == "new" @@ -331,17 +329,12 @@ def test_generate_and_vectorize_summary_success(monkeypatch: pytest.MonkeyPatch) session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0))) ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) - out = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + out = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert out is record session.refresh.assert_called_once_with(record) session.commit.assert_called() @@ -355,18 +348,13 @@ def test_generate_and_vectorize_summary_vectorize_failure_sets_error(monkeypatch session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0))) ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) with pytest.raises(RuntimeError, match="boom"): - SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert record.status == SummaryStatus.ERROR # Outer exception handler overwrites the error with the raw exception message. assert record.error == "boom" @@ -562,18 +550,12 @@ def test_generate_and_vectorize_summary_creates_missing_record_and_logs_usage( session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - usage = MagicMock(total_tokens=4, prompt_tokens=1, completion_tokens=3) monkeypatch.setattr(SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", usage))) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) with caplog.at_level(logging.INFO, logger="services.summary_index_service"): - result = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}) + result = SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) assert result.status in {SummaryStatus.GENERATING, SummaryStatus.COMPLETED} assert any(r.levelno >= logging.INFO for r in caplog.records) @@ -833,11 +815,12 @@ def test_delete_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.Mon def test_update_summary_for_segment_skip_conditions() -> None: + session = MagicMock() economy_dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY) - assert SummaryIndexService.update_summary_for_segment(_segment(), economy_dataset, "x") is None + assert SummaryIndexService.update_summary_for_segment(_segment(), economy_dataset, "x", session=session) is None seg = _segment(has_document=True) seg.document.doc_form = IndexStructureType.QA_INDEX - assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x") is None + assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x", session=session) is None def test_update_summary_for_segment_empty_content_deletes_existing(monkeypatch: pytest.MonkeyPatch) -> None: @@ -850,13 +833,7 @@ def test_update_summary_for_segment_empty_content_deletes_existing(monkeypatch: vector_instance = MagicMock() monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None vector_instance.delete_by_ids.assert_called_once_with(["n1"]) session.delete.assert_called_once_with(record) session.commit.assert_called_once() @@ -872,18 +849,12 @@ def test_update_summary_for_segment_empty_content_delete_vector_warns( session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vector_instance = MagicMock() vector_instance.delete_by_ids.side_effect = RuntimeError("boom") monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) with caplog.at_level(logging.WARNING, logger="services.summary_index_service"): - assert SummaryIndexService.update_summary_for_segment(segment, dataset, "") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, "", session=session) is None assert any(r.levelno >= logging.WARNING for r in caplog.records) @@ -893,13 +864,7 @@ def test_update_summary_for_segment_empty_content_no_record_noop(monkeypatch: py session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ") is None + assert SummaryIndexService.update_summary_for_segment(segment, dataset, " ", session=session) is None def test_update_summary_for_segment_updates_existing_and_vectorizes(monkeypatch: pytest.MonkeyPatch) -> None: @@ -912,16 +877,10 @@ def test_update_summary_for_segment_updates_existing_and_vectorizes(monkeypatch: vector_instance = MagicMock() monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vectorize_mock = MagicMock() monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new summary") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new summary", session=session) assert out is record vectorize_mock.assert_called_once() session.refresh.assert_called_once_with(record) @@ -938,19 +897,13 @@ def test_update_summary_for_segment_existing_vector_delete_warns( session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - vector_instance = MagicMock() vector_instance.delete_by_ids.side_effect = RuntimeError("boom") monkeypatch.setattr(summary_module, "Vector", MagicMock(return_value=vector_instance)) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) with caplog.at_level(logging.WARNING, logger="services.summary_index_service"): - SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert any(r.levelno >= logging.WARNING for r in caplog.records) @@ -963,14 +916,9 @@ def test_update_summary_for_segment_existing_vectorize_failure_returns_error_rec session = MagicMock() session.scalar.return_value = record - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out is record assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") @@ -982,18 +930,11 @@ def test_update_summary_for_segment_new_record_success(monkeypatch: pytest.Monke session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - created = _summary_record(summary_content="new", node_id=None) monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created)) - session.merge.return_value = created monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out is created session.refresh.assert_called() session.commit.assert_called() @@ -1007,81 +948,60 @@ def test_update_summary_for_segment_outer_exception_sets_error_and_reraises(monk session = MagicMock() session.scalar.return_value = record session.flush.side_effect = RuntimeError("flush boom") - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - with pytest.raises(RuntimeError, match="flush boom"): - SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert record.status == SummaryStatus.ERROR assert record.error == "flush boom" session.commit.assert_called() -def test_get_segment_summary_and_document_summaries(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_segment_summary_and_document_summaries() -> None: record = _summary_record(summary_content="sum", node_id="n1") session = MagicMock() session.scalar.return_value = record session.scalars.return_value.all.return_value = [record] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - - assert SummaryIndexService.get_segment_summary("seg-1", "dataset-1") is record - assert SummaryIndexService.get_document_summaries("doc-1", "dataset-1", segment_ids=["seg-1"]) == [record] + assert SummaryIndexService.get_segment_summary("seg-1", "dataset-1", session=session) is record + assert SummaryIndexService.get_document_summaries("doc-1", "dataset-1", segment_ids=["seg-1"], session=session) == [ + record + ] -def test_get_segments_summaries_non_empty(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_segments_summaries_non_empty() -> None: record1 = _summary_record() record1.chunk_id = "seg-1" record2 = _summary_record() record2.chunk_id = "seg-2" session = MagicMock() session.scalars.return_value.all.return_value = [record1, record2] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - out = SummaryIndexService.get_segments_summaries(["seg-1", "seg-2"], "dataset-1") + out = SummaryIndexService.get_segments_summaries(["seg-1", "seg-2"], "dataset-1", session=session) assert set(out.keys()) == {"seg-1", "seg-2"} -def test_get_document_summary_index_status_no_segments_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_document_summary_index_status_no_segments_returns_none() -> None: session = MagicMock() session.scalars.return_value.all.return_value = [] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), + assert ( + SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session) is None ) - assert SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1") is None -def test_get_documents_summary_index_status_empty_input(monkeypatch: pytest.MonkeyPatch) -> None: - assert SummaryIndexService.get_documents_summary_index_status([], "dataset-1", "tenant-1") == {} +def test_get_documents_summary_index_status_empty_input() -> None: + assert ( + SummaryIndexService.get_documents_summary_index_status([], "dataset-1", "tenant-1", session=MagicMock()) == {} + ) def test_get_documents_summary_index_status_no_pending_sets_none(monkeypatch: pytest.MonkeyPatch) -> None: session = MagicMock() session.execute.return_value.all.return_value = [SimpleNamespace(id="seg-1", document_id="doc-1")] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.COMPLETED)}), ) - result = SummaryIndexService.get_documents_summary_index_status(["doc-1"], "dataset-1", "tenant-1") + result = SummaryIndexService.get_documents_summary_index_status(["doc-1"], "dataset-1", "tenant-1", session=session) assert result["doc-1"] is None @@ -1094,26 +1014,19 @@ def test_update_summary_for_segment_creates_new_and_vectorize_fails_returns_erro session = MagicMock() session.scalar.return_value = None - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - created = _summary_record(summary_content="new", node_id=None) monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created)) - session.merge.return_value = created vectorize_mock = MagicMock(side_effect=RuntimeError("boom")) monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock) - out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new") + out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") def test_get_segments_summaries_empty_list() -> None: - assert SummaryIndexService.get_segments_summaries([], "dataset-1") == {} + assert SummaryIndexService.get_segments_summaries([], "dataset-1", session=MagicMock()) == {} def test_get_document_summary_index_status_and_documents_status(monkeypatch: pytest.MonkeyPatch) -> None: @@ -1121,30 +1034,27 @@ def test_get_document_summary_index_status_and_documents_status(monkeypatch: pyt session = MagicMock() session.scalars.return_value.all.return_value = ["seg-1"] # get_document_summary_index_status returns IDs - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) - monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.GENERATING)}), ) - assert SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1") == "SUMMARIZING" + assert ( + SummaryIndexService.get_document_summary_index_status("doc-1", "dataset-1", "tenant-1", session=session) + == "SUMMARIZING" + ) # Multiple docs session2 = MagicMock() session2.execute.return_value.all.return_value = [seg_row] # get_documents_summary_index_status uses execute - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session2))), - ) monkeypatch.setattr( SummaryIndexService, "get_segments_summaries", MagicMock(return_value={"seg-1": SimpleNamespace(status=SummaryStatus.NOT_STARTED)}), ) - result = SummaryIndexService.get_documents_summary_index_status(["doc-1", "doc-2"], "dataset-1", "tenant-1") + result = SummaryIndexService.get_documents_summary_index_status( + ["doc-1", "doc-2"], "dataset-1", "tenant-1", session=session2 + ) assert result["doc-1"] == "SUMMARIZING" assert result["doc-2"] is None diff --git a/api/tests/unit_tests/services/test_trigger_provider_service.py b/api/tests/unit_tests/services/test_trigger_provider_service.py index 0a4452cf478..ff11bbb3035 100644 --- a/api/tests/unit_tests/services/test_trigger_provider_service.py +++ b/api/tests/unit_tests/services/test_trigger_provider_service.py @@ -444,7 +444,7 @@ def test_delete_trigger_provider_should_raise_error_when_subscription_missing( # Act + Assert with pytest.raises(ValueError, match="not found"): - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-1") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) def test_delete_trigger_provider_should_delete_and_clear_cache_even_if_unsubscribe_fails( @@ -476,7 +476,7 @@ def test_delete_trigger_provider_should_delete_and_clear_cache_even_if_unsubscri mock_delete_cache = mocker.patch("services.trigger.trigger_provider_service.delete_cache_for_subscription") # Act - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-1") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) # Assert mock_session.delete.assert_called_once_with(subscription) @@ -507,7 +507,7 @@ def test_delete_trigger_provider_should_skip_unsubscribe_for_unauthorized( ) # Act - TriggerProviderService.delete_trigger_provider(mock_session, "tenant-1", "sub-2") + TriggerProviderService.delete_trigger_provider("tenant-1", "sub-2", session=mock_session) # Assert mock_unsubscribe.assert_not_called() diff --git a/api/tests/unit_tests/services/test_vector_service.py b/api/tests/unit_tests/services/test_vector_service.py index e7ebada6bea..3659b85228b 100644 --- a/api/tests/unit_tests/services/test_vector_service.py +++ b/api/tests/unit_tests/services/test_vector_service.py @@ -98,7 +98,7 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) index_processor.load.assert_called_once() args, kwargs = index_processor.load.call_args @@ -123,7 +123,7 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector([["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) assert index_processor.load.call_count == 2 first_args, first_kwargs = index_processor.load.call_args_list[0] @@ -145,7 +145,7 @@ def test_create_segments_vector_with_no_segments_does_not_load(monkeypatch: pyte factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX) + VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, MagicMock()) index_processor.load.assert_not_called() @@ -189,11 +189,7 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex processing_rule = MagicMock(name="processing_rule") processing_rule.to_dict.return_value = {"rules": {}} - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) embedding_model_instance = MagicMock(name="embedding_model_instance") model_manager_instance = MagicMock(name="model_manager_instance") @@ -211,12 +207,22 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) model_manager_instance.get_model_instance.assert_called_once() generate_child_chunks_mock.assert_called_once_with( - segment, dataset_document, dataset, embedding_model_instance, processing_rule, False + segment, + dataset_document, + dataset, + embedding_model_instance, + processing_rule, + db_mock.session, + False, ) index_processor.load.assert_not_called() @@ -239,11 +245,7 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p processing_rule = MagicMock() processing_rule.to_dict.return_value = {"rules": {}} - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) embedding_model_instance = MagicMock() model_manager_instance = MagicMock() @@ -261,7 +263,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) model_manager_instance.get_default_model_instance.assert_called_once() @@ -276,11 +282,7 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c segment = _make_segment() processing_rule = MagicMock() - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule) index_processor = MagicMock() factory_instance = MagicMock() @@ -289,7 +291,11 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c with caplog.at_level(logging.WARNING, logger="services.vector_service"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) assert any(r.levelno >= logging.WARNING for r in caplog.records) index_processor.load.assert_not_called() @@ -301,15 +307,15 @@ def test_create_segments_vector_parent_child_missing_processing_rule_raises(monk dataset_document = MagicMock() dataset_document.dataset_process_rule_id = "rule-1" - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None) with pytest.raises(ValueError, match="No processing rule found"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) @@ -322,15 +328,15 @@ def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch dataset_document = MagicMock() dataset_document.dataset_process_rule_id = "rule-1" processing_rule = MagicMock() - monkeypatch.setattr( - vector_service_module, - "db", - _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule), - ) + db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) with pytest.raises(ValueError, match="not high quality"): VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + db_mock.session, ) @@ -404,10 +410,7 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch child_chunk_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) monkeypatch.setattr(vector_service_module, "ChildChunk", child_chunk_ctor) - db_mock = MagicMock() - db_mock.session.add = MagicMock() - db_mock.session.commit = MagicMock() - monkeypatch.setattr(vector_service_module, "db", db_mock) + session = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -415,6 +418,7 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch dataset=dataset, embedding_model_instance=MagicMock(), processing_rule=processing_rule, + session=session, regenerate=True, ) @@ -422,8 +426,8 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch _, transform_kwargs = index_processor.transform.call_args assert transform_kwargs["process_rule"]["rules"]["parent_mode"] == vector_service_module.ParentMode.FULL_DOC index_processor.load.assert_called_once() - assert db_mock.session.add.call_count == 2 - db_mock.session.commit.assert_called_once() + assert session.add.call_count == 2 + session.commit.assert_called_once() def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest.MonkeyPatch) -> None: @@ -442,8 +446,7 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - db_mock = MagicMock() - monkeypatch.setattr(vector_service_module, "db", db_mock) + session = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -451,12 +454,13 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest dataset=dataset, embedding_model_instance=MagicMock(), processing_rule=processing_rule, + session=session, regenerate=False, ) index_processor.load.assert_not_called() - db_mock.session.add.assert_not_called() - db_mock.session.commit.assert_called_once() + session.add.assert_not_called() + session.commit.assert_called_once() def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.MonkeyPatch) -> None: @@ -554,9 +558,10 @@ def test_update_multimodel_vector_returns_when_not_high_quality(monkeypatch: pyt vector_cls = MagicMock() db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["a"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["a"], dataset=dataset, session=db_mock.session + ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -568,9 +573,10 @@ def test_update_multimodel_vector_returns_when_no_actual_change(monkeypatch: pyt vector_cls = MagicMock() db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["b", "a"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["b", "a"], dataset=dataset, session=db_mock.session + ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -586,9 +592,8 @@ def test_update_multimodel_vector_deletes_bindings_and_commits_on_empty_new_ids( db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset) + VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=db_mock.session) vector_cls.assert_called_once_with(dataset=dataset) vector_instance.delete_by_ids.assert_called_once_with(["old-1", "old-2"]) @@ -605,9 +610,10 @@ def test_update_multimodel_vector_commits_when_no_upload_files_found(monkeypatch vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) - monkeypatch.setattr(vector_service_module, "db", db_mock) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["new-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["new-1"], dataset=dataset, session=db_mock.session + ) db_mock.session.commit.assert_called_once() db_mock.session.add_all.assert_not_called() @@ -624,7 +630,6 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - monkeypatch.setattr(vector_service_module, "db", db_mock) binding_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) monkeypatch.setattr(vector_service_module, "SegmentAttachmentBinding", binding_ctor) @@ -632,7 +637,12 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up monkeypatch.setattr(vector_service_module, "select", MagicMock()) with caplog.at_level(logging.WARNING, logger="services.vector_service"): - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, + attachment_ids=["file-1", "missing"], + dataset=dataset, + session=db_mock.session, + ) assert any(r.levelno >= logging.WARNING for r in caplog.records) db_mock.session.add_all.assert_called_once() bindings = db_mock.session.add_all.call_args.args[0] @@ -656,14 +666,15 @@ def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - monkeypatch.setattr(vector_service_module, "db", db_mock) monkeypatch.setattr( vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) ) monkeypatch.setattr(vector_service_module, "delete", MagicMock()) monkeypatch.setattr(vector_service_module, "select", MagicMock()) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["file-1"], dataset=dataset, session=db_mock.session + ) vector_instance.delete_by_ids.assert_not_called() vector_instance.add_texts.assert_not_called() @@ -682,7 +693,6 @@ def test_update_multimodel_vector_rolls_back_and_reraises_on_error( monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) db_mock.session.commit.side_effect = RuntimeError("boom") - monkeypatch.setattr(vector_service_module, "db", db_mock) monkeypatch.setattr( vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) ) @@ -691,7 +701,9 @@ def test_update_multimodel_vector_rolls_back_and_reraises_on_error( with caplog.at_level(logging.ERROR, logger="services.vector_service"): with pytest.raises(RuntimeError, match="boom"): - VectorService.update_multimodel_vector(segment=segment, attachment_ids=["file-1"], dataset=dataset) + VectorService.update_multimodel_vector( + segment=segment, attachment_ids=["file-1"], dataset=dataset, session=db_mock.session + ) assert any(r.levelno >= logging.ERROR for r in caplog.records) db_mock.session.rollback.assert_called_once() diff --git a/api/tests/unit_tests/services/test_workflow_collaboration_service.py b/api/tests/unit_tests/services/test_workflow_collaboration_service.py index a61e49c02fa..6b269443fa2 100644 --- a/api/tests/unit_tests/services/test_workflow_collaboration_service.py +++ b/api/tests/unit_tests/services/test_workflow_collaboration_service.py @@ -32,7 +32,7 @@ class TestWorkflowCollaborationService: patch.object(collaboration_service, "broadcast_online_users"), ): # Act - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) # Assert assert result == ("u-1", True) @@ -52,7 +52,7 @@ class TestWorkflowCollaborationService: socketio.get_session.return_value = {} # Act - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) # Assert assert result is None @@ -63,7 +63,7 @@ class TestWorkflowCollaborationService: collaboration_service, repository, socketio = service socketio.get_session.return_value = {"user_id": "u-1", "username": "Jane", "avatar": None} - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) assert result is None repository.set_session_info.assert_not_called() @@ -82,7 +82,7 @@ class TestWorkflowCollaborationService: } with patch.object(collaboration_service, "_can_access_workflow", return_value=False): - result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1") + result = collaboration_service.authorize_and_join_workflow_room("wf-1", "sid-1", session=Mock()) assert result is None repository.set_session_info.assert_not_called() @@ -106,21 +106,12 @@ class TestWorkflowCollaborationService: {"user_id": "u-1", "username": "Jane", "avatar": "avatar.png", "tenant_id": "t-1"}, ) - def test_can_access_workflow_uses_session_factory( - self, service: tuple[WorkflowCollaborationService, Mock, Mock] - ) -> None: + def test_can_access_workflow_uses_session(self, service: tuple[WorkflowCollaborationService, Mock, Mock]) -> None: collaboration_service, _repository, _socketio = service session = Mock() session.scalar.return_value = "wf-1" - session_context = Mock() - session_context.__enter__ = Mock(return_value=session) - session_context.__exit__ = Mock(return_value=False) - with patch( - "services.workflow_collaboration_service.session_factory.create_session", - return_value=session_context, - ): - result = collaboration_service._can_access_workflow("wf-1", "tenant-1") + result = collaboration_service._can_access_workflow("wf-1", "tenant-1", session=session) assert result is True session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/test_workflow_service.py b/api/tests/unit_tests/services/test_workflow_service.py index 67b3e80da6b..0a75a0a8788 100644 --- a/api/tests/unit_tests/services/test_workflow_service.py +++ b/api/tests/unit_tests/services/test_workflow_service.py @@ -312,7 +312,7 @@ class TestWorkflowService: # Mock the database query to return True mock_db_session.session.execute.return_value.scalar_one.return_value = True - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=mock_db_session.session) assert result is True @@ -323,7 +323,7 @@ class TestWorkflowService: # Mock the database query to return False mock_db_session.session.execute.return_value.scalar_one.return_value = False - result = workflow_service.is_workflow_exist(app) + result = workflow_service.is_workflow_exist(app, session=mock_db_session.session) assert result is False @@ -343,7 +343,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_draft_workflow mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=mock_db_session.session) assert result == mock_workflow @@ -367,7 +367,7 @@ class TestWorkflowService: # Mock db.session.scalar() to return None mock_db_session.session.scalar.return_value = None - result = workflow_service.get_draft_workflow(app) + result = workflow_service.get_draft_workflow(app, session=mock_db_session.session) assert result is None @@ -380,7 +380,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow_by_id mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_draft_workflow(app, workflow_id=workflow_id) + result = workflow_service.get_draft_workflow(app, workflow_id=workflow_id, session=mock_db_session.session) assert result == mock_workflow @@ -411,7 +411,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow_by_id mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_published_workflow_by_id(app, workflow_id) + result = workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) assert result == mock_workflow @@ -432,7 +432,7 @@ class TestWorkflowService: mock_db_session.session.scalar.return_value = mock_workflow with pytest.raises(IsDraftWorkflowError): - workflow_service.get_published_workflow_by_id(app, workflow_id) + workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) def test_get_published_workflow_by_id_returns_none(self, workflow_service, mock_db_session): """Test get_published_workflow_by_id returns None when workflow not found.""" @@ -442,7 +442,7 @@ class TestWorkflowService: # Mock db.session.scalar() to return None mock_db_session.session.scalar.return_value = None - result = workflow_service.get_published_workflow_by_id(app, workflow_id) + result = workflow_service.get_published_workflow_by_id(app, workflow_id, session=mock_db_session.session) assert result is None @@ -455,7 +455,7 @@ class TestWorkflowService: # Mock db.session.scalar() used by get_published_workflow mock_db_session.session.scalar.return_value = mock_workflow - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=mock_db_session.session) assert result == mock_workflow @@ -463,7 +463,7 @@ class TestWorkflowService: """Test get_published_workflow returns None when app has no workflow_id.""" app = TestWorkflowAssociatedDataFactory.create_app_mock(workflow_id=None) - result = workflow_service.get_published_workflow(app) + result = workflow_service.get_published_workflow(app, session=MagicMock()) assert result is None @@ -499,6 +499,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) # Verify workflow was added to session @@ -536,6 +537,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) # Verify workflow was updated @@ -571,6 +573,7 @@ class TestWorkflowService: account=account, environment_variables=[], conversation_variables=[], + session=mock_db_session.session, ) def test_restore_published_workflow_to_draft_keeps_source_features_unmodified( @@ -648,6 +651,7 @@ class TestWorkflowService: app_model=app, workflow_id=source_workflow.id, account=account, + session=mock_db_session.session, ) mock_validate_features.assert_called_once_with(app_model=app, features=normalized_features) @@ -761,6 +765,7 @@ class TestWorkflowService: app_model=app, environment_variables=variables, account=account, + session=mock_db_session.session, ) assert workflow.environment_variables == variables @@ -779,6 +784,7 @@ class TestWorkflowService: app_model=app, environment_variables=[], account=account, + session=MagicMock(), ) def test_update_draft_workflow_conversation_variables_updates_workflow(self, workflow_service, mock_db_session): @@ -796,6 +802,7 @@ class TestWorkflowService: app_model=app, conversation_variables=variables, account=account, + session=mock_db_session.session, ) assert workflow.conversation_variables == variables @@ -814,6 +821,7 @@ class TestWorkflowService: app_model=app, conversation_variables=[], account=account, + session=MagicMock(), ) # ==================== Publish Workflow Tests ==================== @@ -1429,7 +1437,7 @@ class TestWorkflowService: mock_new_app = TestWorkflowAssociatedDataFactory.create_app_mock(mode=AppMode.WORKFLOW) mock_converter.convert_to_workflow.return_value = mock_new_app - result = workflow_service.convert_to_workflow(app, account, args) + result = workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) assert result == mock_new_app mock_converter.convert_to_workflow.assert_called_once() @@ -1451,7 +1459,7 @@ class TestWorkflowService: mock_new_app = TestWorkflowAssociatedDataFactory.create_app_mock(mode=AppMode.WORKFLOW) mock_converter.convert_to_workflow.return_value = mock_new_app - result = workflow_service.convert_to_workflow(app, account, args) + result = workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) assert result == mock_new_app @@ -1467,7 +1475,7 @@ class TestWorkflowService: args = {} with pytest.raises(ValueError, match="not supported convert to workflow"): - workflow_service.convert_to_workflow(app, account, args) + workflow_service.convert_to_workflow(app, account, args, session=MagicMock()) # =========================================================================== @@ -1520,7 +1528,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check: # Should not raise; mock allows the call - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) mock_check.assert_called_once() def test_validate_workflow_credentials_should_check_default_credential_when_no_credential_id( @@ -1541,10 +1549,11 @@ class TestWorkflowServiceCredentialValidation: # Act with patch.object(service, "_check_default_tool_credential") as mock_default: - service._validate_workflow_credentials(workflow) + session = MagicMock() + service._validate_workflow_credentials(workflow, session=session) # Assert - mock_default.assert_called_once_with("tenant-1", "my-provider") + mock_default.assert_called_once_with("tenant-1", "my-provider", session=session) def test_validate_workflow_credentials_should_skip_tool_node_without_provider( self, service: WorkflowService @@ -1556,7 +1565,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert (no error raised) with patch.object(service, "_check_default_tool_credential") as mock_default: - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) mock_default.assert_not_called() def test_validate_workflow_credentials_should_validate_llm_node_with_model_config( @@ -1579,7 +1588,7 @@ class TestWorkflowServiceCredentialValidation: patch.object(service, "_validate_llm_model_config") as mock_llm, patch.object(service, "_validate_load_balancing_credentials"), ): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) # Assert mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4") @@ -1599,7 +1608,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with pytest.raises(ValueError, match="Missing provider or model configuration"): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) def test_validate_workflow_credentials_should_wrap_unexpected_exception_in_value_error( self, service: WorkflowService @@ -1620,7 +1629,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert with patch.object(service, "_validate_llm_model_config", side_effect=RuntimeError("boom")): with pytest.raises(ValueError, match="boom"): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) def test_validate_workflow_credentials_should_validate_agent_node_model(self, service: WorkflowService) -> None: # Arrange @@ -1643,7 +1652,7 @@ class TestWorkflowServiceCredentialValidation: patch.object(service, "_validate_llm_model_config") as mock_llm, patch.object(service, "_validate_load_balancing_credentials"), ): - service._validate_workflow_credentials(workflow) + service._validate_workflow_credentials(workflow, session=MagicMock()) # Assert mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4") @@ -1675,11 +1684,12 @@ class TestWorkflowServiceCredentialValidation: patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check, patch.object(service, "_check_default_tool_credential") as mock_default, ): - service._validate_workflow_credentials(workflow) + session = MagicMock() + service._validate_workflow_credentials(workflow, session=session) # Assert mock_check.assert_called_once() # provider-a has credential_id - mock_default.assert_called_once_with("tenant-1", "provider-b") + mock_default.assert_called_once_with("tenant-1", "provider-b", session=session) # --- _validate_llm_model_config --- @@ -1739,7 +1749,7 @@ class TestWorkflowServiceCredentialValidation: # Arrange with patch("services.workflow_service.db") as mock_db: # Act + Assert (should NOT raise) - service._check_default_tool_credential("tenant-1", "some-provider") + service._check_default_tool_credential("tenant-1", "some-provider", session=MagicMock()) def test_check_default_tool_credential_should_raise_when_compliance_fails(self, service: WorkflowService) -> None: # Arrange @@ -1751,7 +1761,7 @@ class TestWorkflowServiceCredentialValidation: ): # Act + Assert with pytest.raises(ValueError, match="Failed to validate default credential"): - service._check_default_tool_credential("tenant-1", "some-provider") + service._check_default_tool_credential("tenant-1", "some-provider", session=MagicMock()) # --- _is_load_balancing_enabled --- @@ -1811,7 +1821,7 @@ class TestWorkflowServiceCredentialValidation: side_effect=RuntimeError("fail"), ): # Act - result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4") + result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock()) # Assert assert result == [] @@ -1828,7 +1838,7 @@ class TestWorkflowServiceCredentialValidation: ], ): # Act - result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4") + result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock()) # Assert — only entries with a credential_id should be returned assert len(result) == 2 @@ -1845,7 +1855,7 @@ class TestWorkflowServiceCredentialValidation: node_data: dict[str, Any] = {} # no model key # Act + Assert (no error expected) - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) def test_validate_load_balancing_credentials_should_skip_when_lb_not_enabled( self, service: WorkflowService @@ -1856,7 +1866,7 @@ class TestWorkflowServiceCredentialValidation: # Act + Assert (no error expected) with patch.object(service, "_is_load_balancing_enabled", return_value=False): - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) def test_validate_load_balancing_credentials_should_raise_when_compliance_fails( self, service: WorkflowService @@ -1876,7 +1886,7 @@ class TestWorkflowServiceCredentialValidation: ), ): with pytest.raises(ValueError, match="Invalid load balancing credentials"): - service._validate_load_balancing_credentials(workflow, node_data, "node-1") + service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock()) # =========================================================================== @@ -2673,7 +2683,9 @@ class TestWorkflowServiceHumanInputOperations: def test_get_human_input_form_preview_should_raise_if_workflow_not_init(self, service: WorkflowService) -> None: service.get_draft_workflow = MagicMock(return_value=None) with pytest.raises(ValueError, match="Workflow not initialized"): - service.get_human_input_form_preview(app_model=MagicMock(), account=MagicMock(), node_id="node-1") + service.get_human_input_form_preview( + app_model=MagicMock(), account=MagicMock(), node_id="node-1", session=MagicMock() + ) def test_get_human_input_form_preview_should_raise_if_wrong_node_type(self, service: WorkflowService) -> None: draft = MagicMock() @@ -2681,7 +2693,9 @@ class TestWorkflowServiceHumanInputOperations: service.get_draft_workflow = MagicMock(return_value=draft) with patch("models.workflow.Workflow.get_node_type_from_node_config", return_value=BuiltinNodeTypes.LLM): with pytest.raises(ValueError, match="Node type must be human-input"): - service.get_human_input_form_preview(app_model=MagicMock(), account=MagicMock(), node_id="node-1") + service.get_human_input_form_preview( + app_model=MagicMock(), account=MagicMock(), node_id="node-1", session=MagicMock() + ) def test_get_human_input_form_preview_success(self, service: WorkflowService) -> None: app_model = MagicMock(spec=App) @@ -2716,7 +2730,9 @@ class TestWorkflowServiceHumanInputOperations: patch("services.workflow_service.HumanInputNode", return_value=mock_node), patch("services.workflow_service.HumanInputRequired") as mock_required_cls, ): - service.get_human_input_form_preview(app_model=app_model, account=account, node_id="node-1") + service.get_human_input_form_preview( + app_model=app_model, account=account, node_id="node-1", session=MagicMock() + ) mock_node.render_form_content_before_submission.assert_called_once() mock_required_cls.return_value.model_dump.assert_called_once() @@ -2760,7 +2776,12 @@ class TestWorkflowServiceHumanInputOperations: patch("services.workflow_service.DraftVariableSaver") as mock_saver_cls, ): result = service.submit_human_input_form_preview( - app_model=app_model, account=account, node_id="node-1", form_inputs={"field1": "val1"}, action="submit" + app_model=app_model, + account=account, + node_id="node-1", + form_inputs={"field1": "val1"}, + action="submit", + session=MagicMock(), ) assert result["__action_id"] == "submit" mock_validate.assert_called_once() @@ -2785,7 +2806,11 @@ class TestWorkflowServiceHumanInputOperations: ): mock_resolve.return_value = MagicMock() service.test_human_input_delivery( - app_model=MagicMock(), account=MagicMock(), node_id="node-1", delivery_method_id="method-1" + app_model=MagicMock(), + account=MagicMock(), + node_id="node-1", + delivery_method_id="method-1", + session=MagicMock(), ) mock_test_srv.return_value.send_test.assert_called_once() @@ -2801,7 +2826,11 @@ class TestWorkflowServiceHumanInputOperations: ): with pytest.raises(ValueError, match="Delivery method not found"): service.test_human_input_delivery( - app_model=MagicMock(), account=MagicMock(), node_id="node-1", delivery_method_id="none" + app_model=MagicMock(), + account=MagicMock(), + node_id="node-1", + delivery_method_id="none", + session=MagicMock(), ) def test_load_email_recipients_parsing_failure(self, service: WorkflowService) -> None: diff --git a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py index c210db580e0..549f50cb370 100644 --- a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py +++ b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py @@ -354,7 +354,7 @@ class TestGetBuiltinToolProviderCredentialInfo: def test_returns_credential_info(self, mock_tm, mock_creds, mock_oauth): mock_tm.get_builtin_provider.return_value.get_supported_credential_types.return_value = ["api-key"] - result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google", session=MagicMock()) assert result.credentials == [] assert result.supported_credential_types == ["api-key"] @@ -368,7 +368,7 @@ class TestGetBuiltinToolProviderCredentials: mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False) mock_db.session.scalars.return_value.all.return_value = [] - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) assert result == [] @@ -391,7 +391,7 @@ class TestGetBuiltinToolProviderCredentials: credential_entity = MagicMock() mock_transform.convert_builtin_provider_to_credential_entity.return_value = credential_entity - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google") + result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) assert len(result) == 1 assert result[0] is credential_entity diff --git a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py index 6f6c56fd67f..7f720575154 100644 --- a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py +++ b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py @@ -1,8 +1,8 @@ """Unit tests for NodeOutputInspectorService (Stage 4 §8). The service reads from postgres and resolves agent v2 bindings; this suite -mocks ``session_factory`` and the binding resolver so we exercise the -view-construction logic without DB / network access. +mocks the DB session and binding resolver so we exercise the view-construction +logic without DB / network access. """ from __future__ import annotations @@ -100,26 +100,17 @@ def _non_agent_node(*, node_id: str = "tool-node-1", node_type: str = "tool", ti } -def _patch_session( +def _mock_session( *, workflow_run: SimpleNamespace | None, executions: list[SimpleNamespace] | None = None, ): - """Patch ``session_factory.create_session`` to return the configured rows. - - Returns a context manager that the test uses with ``with``. - """ + """Build a mock DB session with the configured rows.""" executions = executions or [] - mock_session = MagicMock() - mock_session.scalar.return_value = workflow_run - mock_session.scalars.return_value.all.return_value = executions - cm = MagicMock() - cm.__enter__.return_value = mock_session - cm.__exit__.return_value = False - return patch( - "services.workflow.node_output_inspector_service.session_factory.create_session", - return_value=cm, - ) + session = MagicMock() + session.scalar.return_value = workflow_run + session.scalars.return_value.all.return_value = executions + return session def _stub_binding_resolver(*, declared_outputs: list[DeclaredOutputConfig]): @@ -149,9 +140,9 @@ def _make_service(declared_outputs: list[DeclaredOutputConfig] | None = None) -> def test_snapshot_404_when_workflow_run_missing(): service = _make_service() - with _patch_session(workflow_run=None): - with pytest.raises(NodeOutputInspectorError) as exc: - service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="missing") + session = _mock_session(workflow_run=None) + with pytest.raises(NodeOutputInspectorError) as exc: + service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="missing", session=session) assert exc.value.code == "workflow_run_not_found" @@ -162,8 +153,8 @@ def test_snapshot_accepts_published_run_d1_lifted(): nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" assert [n.node_id for n in snapshot.node_outputs] == ["agent-1"] @@ -175,17 +166,17 @@ def test_snapshot_accepts_webhook_triggered_run(): nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.WEBHOOK, ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" def test_node_detail_404_when_node_id_absent_from_graph(): service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.node_detail(app_model=_app_model(), workflow_run_id="run-1", node_id="ghost") + session = _mock_session(workflow_run=run, executions=[]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.node_detail(app_model=_app_model(), workflow_run_id="run-1", node_id="ghost", session=session) assert exc.value.code == "node_not_in_workflow_run" @@ -195,28 +186,30 @@ def test_output_preview_404_when_output_name_unknown(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - with _patch_session(workflow_run=run, executions=[ex]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.output_preview( - app_model=_app_model(), - workflow_run_id="run-1", - node_id="agent-1", - output_name="missing", - ) + session = _mock_session(workflow_run=run, executions=[ex]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.output_preview( + app_model=_app_model(), + workflow_run_id="run-1", + node_id="agent-1", + output_name="missing", + session=session, + ) assert exc.value.code == "node_output_not_declared" def test_output_preview_404_when_node_id_absent_from_graph(): service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - with pytest.raises(NodeOutputInspectorError) as exc: - service.output_preview( - app_model=_app_model(), - workflow_run_id="run-1", - node_id="ghost", - output_name="report", - ) + session = _mock_session(workflow_run=run, executions=[]) + with pytest.raises(NodeOutputInspectorError) as exc: + service.output_preview( + app_model=_app_model(), + workflow_run_id="run-1", + node_id="ghost", + output_name="report", + session=session, + ) assert exc.value.code == "node_not_in_workflow_run" @@ -230,8 +223,8 @@ def test_snapshot_status_pending_when_node_has_no_execution(): declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert len(snapshot.node_outputs) == 1 node = snapshot.node_outputs[0] @@ -245,8 +238,8 @@ def test_snapshot_status_running(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.RUNNING) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].node_status == NodeStatus.RUNNING assert snapshot.node_outputs[0].outputs[0].status == NodeOutputStatus.RUNNING @@ -260,8 +253,8 @@ def test_snapshot_status_failed_node_marks_all_outputs_failed(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.FAILED) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"a": NodeOutputStatus.FAILED, "b": NodeOutputStatus.FAILED} @@ -272,8 +265,8 @@ def test_snapshot_status_ready_when_outputs_present_and_no_failure_metadata(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.READY assert output.value_preview == "hello" @@ -294,8 +287,8 @@ def test_snapshot_marks_type_check_failure(): } }, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.TYPE_CHECK_FAILED assert output.type_check is not None @@ -324,14 +317,12 @@ def test_snapshot_marks_output_check_failure_when_type_check_passed(): }, }, ) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", - return_value="https://signed.example/x", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", + return_value="https://signed.example/x", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.OUTPUT_CHECK_FAILED assert output.output_check is not None @@ -348,8 +339,8 @@ def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload() ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hi"}) # optional_meta missing - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"text": NodeOutputStatus.READY, "optional_meta": NodeOutputStatus.NOT_PRODUCED} @@ -367,8 +358,8 @@ def test_non_agent_node_outputs_inferred_from_payload_keys(): node_type="tool", outputs={"message": "sent", "thread_ts": "1234"}, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output_names = sorted(o.name for o in snapshot.node_outputs[0].outputs) assert output_names == ["message", "thread_ts"] # All inferred outputs should have ``type=None`` since we don't know the @@ -393,14 +384,12 @@ def test_file_output_preview_includes_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview_value = snapshot.node_outputs[0].outputs[0].value_preview assert isinstance(preview_value, dict) assert preview_value["preview_url"] == "https://signed.example/x.pdf" @@ -419,18 +408,17 @@ def test_file_output_preview_endpoint_returns_full_value_with_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="report", + session=session, ) assert preview.output_name == "report" assert preview.status == NodeOutputStatus.READY @@ -484,26 +472,25 @@ def test_array_file_output_preview_includes_signed_urls_for_each_item(): }, ] ex = _execution(node_id="agent-1", outputs={"files": file_payloads}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - side_effect=[ - "https://signed.example/1.pdf", - "https://signed.example/2.pdf", - "https://signed.example/1-detail.pdf", - "https://signed.example/2-detail.pdf", - "https://signed.example/1-full.pdf", - "https://signed.example/2-full.pdf", - ], - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + side_effect=[ + "https://signed.example/1.pdf", + "https://signed.example/2.pdf", + "https://signed.example/1-detail.pdf", + "https://signed.example/2-detail.pdf", + "https://signed.example/1-full.pdf", + "https://signed.example/2-full.pdf", + ], ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="files", + session=session, ) snapshot_value = snapshot.node_outputs[0].outputs[0].value_preview @@ -531,14 +518,12 @@ def test_file_output_preview_uses_none_when_signed_url_resolution_fails(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - side_effect=RuntimeError("boom"), - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + side_effect=RuntimeError("boom"), ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview_value = snapshot.node_outputs[0].outputs[0].value_preview assert isinstance(preview_value, dict) @@ -557,19 +542,18 @@ def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"meta": raw_value}) - with ( - _patch_session(workflow_run=run, executions=[ex]), - patch( - "services.workflow.node_output_inspector_service._resolve_preview_url", - return_value="https://signed.example/x.pdf", - ), + session = _mock_session(workflow_run=run, executions=[ex]) + with patch( + "services.workflow.node_output_inspector_service._resolve_preview_url", + return_value="https://signed.example/x.pdf", ): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) preview = service.output_preview( app_model=_app_model(), workflow_run_id="run-1", node_id="agent-1", output_name="meta", + session=session, ) assert snapshot.node_outputs[0].outputs[0].value_preview == raw_value @@ -591,8 +575,8 @@ def test_retried_count_pulled_from_attempt_metadata(): outputs={"text": "ok"}, execution_metadata={"attempt": 2}, ) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].retried == 2 @@ -610,8 +594,8 @@ def test_keeps_latest_execution_per_node_by_index(): run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) older = _execution(node_id="agent-1", outputs={"text": "old"}, index=1) newer = _execution(node_id="agent-1", outputs={"text": "new"}, index=5) - with _patch_session(workflow_run=run, executions=[older, newer]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[older, newer]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].value_preview == "new" @@ -632,8 +616,8 @@ def test_array_typed_output_with_array_item_renders_correctly(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"files": []}) - with _patch_session(workflow_run=run, executions=[ex]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[ex]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.type == DeclaredOutputType.ARRAY @@ -654,6 +638,6 @@ def test_unparseable_graph_blob_yields_empty_snapshot_not_500(): status=WorkflowExecutionStatus.RUNNING, graph="{not valid json", ) - with _patch_session(workflow_run=run, executions=[]): - snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1") + session = _mock_session(workflow_run=run, executions=[]) + snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs == [] diff --git a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py index 2aaf3bdf1d5..f471e4aeb56 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py @@ -118,6 +118,7 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, + session=MagicMock(), ) assert len(nodes) == 2 @@ -160,6 +161,7 @@ def test__convert_to_http_request_node_for_workflow_app(default_variables: list[ app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, + session=MagicMock(), ) body = json.loads(nodes[0]["data"]["body"]["data"]) @@ -364,6 +366,7 @@ def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(conve icon_type="emoji", icon="robot", icon_background="#fff", + session=MagicMock(), ) @@ -389,7 +392,6 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( monkeypatch.setattr(converter_module, "App", FakeApp) db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) send_mock = MagicMock() monkeypatch.setattr(converter_module.app_was_created, "send", send_mock) @@ -417,6 +419,7 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( icon_type="", icon="", icon_background="", + session=db_session, ) assert new_app.name == "Source App(workflow)" @@ -501,12 +504,12 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", + session=db_session, ) graph = json.loads(workflow.graph) @@ -568,12 +571,12 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", + session=db_session, ) graph = json.loads(workflow.graph) @@ -644,6 +647,7 @@ def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_ app_model=app_model, variables=[], external_data_variables=external_data_variables, + session=MagicMock(), ) assert nodes == [] @@ -810,10 +814,9 @@ def test_get_api_based_extension_should_raise_when_extension_not_found( monkeypatch: pytest.MonkeyPatch, ) -> None: db_session = SimpleNamespace(scalar=MagicMock(return_value=None)) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) with pytest.raises(ValueError, match="API Based Extension not found"): - converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1") + converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1", session=db_session) db_session.scalar.assert_called_once() @@ -823,9 +826,10 @@ def test_get_api_based_extension_should_return_entity_when_found( ) -> None: extension = SimpleNamespace(id="ext-1") db_session = SimpleNamespace(scalar=MagicMock(return_value=extension)) - monkeypatch.setattr(converter_module, "db", SimpleNamespace(session=db_session)) - result = converter._get_api_based_extension(tenant_id="tenant-1", api_based_extension_id="ext-1") + result = converter._get_api_based_extension( + tenant_id="tenant-1", api_based_extension_id="ext-1", session=db_session + ) assert result is extension db_session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py b/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py index 5bcb13c360c..cd97fa2e53b 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_human_input_delivery.py @@ -64,6 +64,7 @@ def test_human_input_delivery_requires_draft_workflow(): account=account, node_id="node-1", delivery_method_id="delivery-1", + session=MagicMock(), ) @@ -98,6 +99,7 @@ def test_human_input_delivery_allows_disabled_method(monkeypatch: pytest.MonkeyP account=account, node_id="node-1", delivery_method_id=str(delivery_method.id), + session=MagicMock(), ) test_service_instance.send_test.assert_called_once() @@ -135,6 +137,7 @@ def test_human_input_delivery_dispatches_to_test_service(monkeypatch: pytest.Mon node_id="node-1", delivery_method_id=str(delivery_method.id), inputs={"#node-1.output#": "value"}, + session=MagicMock(), ) pool_args = service._build_human_input_variable_pool.call_args.kwargs @@ -173,6 +176,7 @@ def test_human_input_delivery_debug_mode_overrides_recipients(monkeypatch: pytes account=account, node_id="node-1", delivery_method_id=str(delivery_method.id), + session=MagicMock(), ) test_service_instance.send_test.assert_called_once()