diff --git a/api/commands/plugin.py b/api/commands/plugin.py index 3695c742921..e5be0cd1b10 100644 --- a/api/commands/plugin.py +++ b/api/commands/plugin.py @@ -9,6 +9,7 @@ from sqlalchemy import delete, func, select from sqlalchemy.engine import CursorResult from configs import dify_config +from core.db.session_factory import session_factory from core.helper import encrypter from core.plugin.entities.plugin_daemon import CredentialType from core.plugin.impl.plugin import PluginInstaller @@ -578,9 +579,6 @@ def install_rag_pipeline_plugins(input_file, output_file, workers): """ click.echo(click.style("Installing rag pipeline plugins", fg="yellow")) plugin_migration = PluginMigration() - plugin_migration.install_rag_pipeline_plugins( - input_file, - output_file, - workers, - ) + with session_factory.create_session() as session: + plugin_migration.install_rag_pipeline_plugins(input_file, output_file, workers, session=session) click.echo(click.style("Installing rag pipeline plugins successfully", fg="green")) diff --git a/api/commands/system.py b/api/commands/system.py index 7755d3b5bcd..c1f91b43869 100644 --- a/api/commands/system.py +++ b/api/commands/system.py @@ -188,23 +188,26 @@ where sites.id is null limit 1000""" if app_id in failed_app_ids: continue + session = db.session() try: - app = db.session.scalar(select(App).where(App.id == app_id)) + app = session.scalar(select(App).where(App.id == app_id)) if not app: logger.info("App %s not found", app_id) continue - tenant = app.tenant + tenant = session.get(Tenant, app.tenant_id) if tenant: - accounts = tenant.get_accounts() + accounts = tenant.get_accounts(session=session) if not accounts: logger.info("Fix failed for app %s", app.id) continue account = accounts[0] logger.info("Fixing missing site for app %s", app.id) - app_was_created.send(app, account=account) + app_was_created.send(app, account=account, session=session) + session.commit() except Exception: + session.rollback() failed_app_ids.append(app_id) click.echo(click.style(f"Failed to fix missing site for app {app_id}", fg="red")) logger.exception("Failed to fix app related site missing issue, app_id: %s", app_id) diff --git a/api/commands/vector.py b/api/commands/vector.py index 39884418b23..437ea40f529 100644 --- a/api/commands/vector.py +++ b/api/commands/vector.py @@ -1,10 +1,11 @@ import json +from typing import cast import click from flask import current_app from sqlalchemy import select from sqlalchemy.exc import SQLAlchemyError -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from configs import dify_config from core.rag.datasource.vdb.vector_factory import Vector @@ -101,7 +102,8 @@ def migrate_annotation_vector_database(): ) documents.append(document) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + with Session(db.engine) as session: + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) click.echo(f"Migrating annotations for app: {app.id}.") try: @@ -176,6 +178,7 @@ def migrate_knowledge_vector_database(): VectorType.OCEANBASE, } page = 1 + db_session = db.session() while True: try: stmt = ( @@ -184,7 +187,7 @@ def migrate_knowledge_vector_database(): .order_by(Dataset.created_at.desc()) ) - datasets = paginate_query(stmt, page=page, per_page=50, max_per_page=50) + datasets = paginate_query(stmt, page=page, per_page=50, max_per_page=50, session=db_session) if not datasets.items: break except SQLAlchemyError: @@ -227,7 +230,8 @@ def migrate_knowledge_vector_database(): index_struct_dict = {"type": vector_type, "vector_store": {"class_prefix": collection_name}} dataset.index_struct = json.dumps(index_struct_dict) - vector = Vector(dataset) + with Session(db.engine) as session: + vector = Vector(dataset, session=session) click.echo(f"Migrating dataset {dataset.id}.") try: @@ -274,7 +278,7 @@ def migrate_knowledge_vector_database(): }, ) if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=db_session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -410,7 +414,9 @@ def old_metadata_migration(): .where(DatasetDocument.doc_metadata.is_not(None)) .order_by(DatasetDocument.created_at.desc()) ) - documents = paginate_query(stmt, page=page, per_page=50, max_per_page=50) + documents = paginate_query( + stmt, page=page, per_page=50, max_per_page=50, session=cast(Session, db.session()) + ) except SQLAlchemyError: raise if not documents: diff --git a/api/controllers/common/agent_app_parameters.py b/api/controllers/common/agent_app_parameters.py index 8c2fbccd513..c1c9fcdb23a 100644 --- a/api/controllers/common/agent_app_parameters.py +++ b/api/controllers/common/agent_app_parameters.py @@ -1,27 +1,29 @@ from typing import Any from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features from core.app.apps.agent_app.app_variable_projection import agent_app_variables_to_user_input_form from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError -from extensions.ext_database import db from models.agent import Agent, AgentConfigSnapshot, AgentStatus from models.agent_config_entities import AgentSoulConfig -from models.model import App +from models.model import App, load_annotation_reply_config def get_published_agent_app_feature_dict_and_user_input_form( app_model: App, + *, + session: Session, ) -> tuple[dict[str, Any], list[dict[str, Any]]]: """Return public Agent App parameters backed by the published Agent Soul.""" - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) agent_id = app_model.bound_agent_id if not agent_id: raise AgentAppGeneratorError("Agent App has no bound Agent") - agent = db.session.scalar( + agent = session.scalar( select(Agent) .where( Agent.tenant_id == app_model.tenant_id, @@ -37,7 +39,7 @@ def get_published_agent_app_feature_dict_and_user_input_form( if not agent.active_config_snapshot_id: raise AgentAppNotPublishedError("Agent has not been published") - snapshot = db.session.scalar( + snapshot = session.scalar( select(AgentConfigSnapshot) .where( AgentConfigSnapshot.tenant_id == app_model.tenant_id, @@ -50,5 +52,10 @@ def get_published_agent_app_feature_dict_and_user_input_form( raise AgentAppGeneratorError("Agent published version not found") agent_soul = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict) - features_dict = merge_agent_app_features(agent_soul=agent_soul, app_model_config=app_model_config) + annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None + features_dict = merge_agent_app_features( + agent_soul=agent_soul, + app_model_config=app_model_config, + annotation_reply=annotation_reply, + ) return features_dict, agent_app_variables_to_user_input_form(agent_soul.app_variables) diff --git a/api/controllers/common/app_access.py b/api/controllers/common/app_access.py index 214d2de71b4..0644b09ca73 100644 --- a/api/controllers/common/app_access.py +++ b/api/controllers/common/app_access.py @@ -4,7 +4,8 @@ from collections.abc import Sequence from dataclasses import dataclass from typing import TYPE_CHECKING -from extensions.ext_database import db +from sqlalchemy.orm import Session + from services.enterprise import rbac_service as enterprise_rbac_service if TYPE_CHECKING: @@ -68,6 +69,7 @@ def resolve_app_access_filter( tenant_id: str, account_id: str, *, + session: Session, permissions: MyPermissionsResponse | None = None, ) -> AppAccessFilter: """Compute the RBAC app-access filter for ``account_id`` in ``tenant_id``. @@ -77,7 +79,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, session=db.session()) + permissions = enterprise_rbac_service.RBACService.MyPermissions.get(tenant_id, account_id, session=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/common/session.py b/api/controllers/common/session.py index fac2bec6767..24b1a8729d3 100644 --- a/api/controllers/common/session.py +++ b/api/controllers/common/session.py @@ -2,8 +2,10 @@ `with_session` is an HTTP controller helper: it opens one SQLAlchemy session for a Resource handler and injects it as the first argument after `self`. -Handlers use a transaction by default so migrated write paths keep -commit/rollback handling; pure read handlers may opt out with `write=False`. +Write handlers commit on success and roll back on failure. They use a regular +Session context so existing services may commit an intermediate unit and keep +using the same Session through SQLAlchemy's autobegin behavior. Pure read +handlers may opt out with `write=False`. """ from collections.abc import Callable @@ -38,14 +40,20 @@ def with_session[T, **P, R]( ) -> ( Callable[Concatenate[T, P], R] | Callable[[Callable[Concatenate[T, Session, P], R]], Callable[Concatenate[T, P], R]] ): - """Inject a request-scoped session, using a transaction only for write handlers.""" + """Inject a request-scoped session and finalize write handlers.""" def decorator(view: Callable[Concatenate[T, Session, P], R]) -> Callable[Concatenate[T, P], R]: @wraps(view) def wrapper(self: T, *args: P.args, **kwargs: P.kwargs) -> R: if write: - with session_factory.get_session_maker().begin() as session: - return view(self, session, *args, **kwargs) + with session_factory.create_session() as session: + try: + result = view(self, session, *args, **kwargs) + session.commit() + return result + except Exception: + session.rollback() # noqa: no-new-controller-sqlalchemy decorator owns transaction rollback + raise with session_factory.create_session() as session: return view(self, session, *args, **kwargs) diff --git a/api/controllers/console/agent/app_helpers.py b/api/controllers/console/agent/app_helpers.py index 7af38b0164d..0de65d5a89f 100644 --- a/api/controllers/console/agent/app_helpers.py +++ b/api/controllers/console/agent/app_helpers.py @@ -1,20 +1,21 @@ from uuid import UUID -from extensions.ext_database import db +from sqlalchemy.orm import Session + from models.model import App from services.agent.roster_service import AgentRosterService -def resolve_agent_app_model(*, tenant_id: str, agent_id: UUID) -> App: +def resolve_agent_app_model(*, session: Session, tenant_id: str, agent_id: UUID) -> App: """Resolve a roster Agent's public Agent App.""" - return AgentRosterService(db.session).get_agent_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) + return AgentRosterService(session).get_agent_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) -def resolve_agent_runtime_app_model(*, tenant_id: str, agent_id: UUID) -> App: +def resolve_agent_runtime_app_model(*, session: Session, tenant_id: str, agent_id: UUID) -> App: """Resolve the App that backs an Agent runtime surface. This accepts both roster Agent Apps and workflow-only inline Agents with a hidden backing App. """ - return AgentRosterService(db.session).get_agent_runtime_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) + return AgentRosterService(session).get_agent_runtime_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) diff --git a/api/controllers/console/agent/composer.py b/api/controllers/console/agent/composer.py index b5be73f2596..ef860d86b8f 100644 --- a/api/controllers/console/agent/composer.py +++ b/api/controllers/console/agent/composer.py @@ -2,9 +2,11 @@ from uuid import UUID from flask import request from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.app.wraps import get_app_model from controllers.console.wraps import ( @@ -17,7 +19,6 @@ 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, @@ -59,20 +60,21 @@ class WorkflowAgentComposerApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, account_id: str, app_model: App, node_id: str): + @with_session + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def get(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): query = WorkflowAgentComposerQuery.model_validate(request.args.to_dict(flat=True)) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.load_workflow_composer( + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, account_id=account_id, snapshot_id=query.snapshot_id, - session=db.session(), ), ) @@ -85,20 +87,21 @@ class WorkflowAgentComposerApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_user_id @with_current_tenant_id - def put(self, tenant_id: str, account_id: str, app_model: App, node_id: str): + @with_session + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def put(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, account_id=account_id, payload=payload, - session=db.session(), ), ) @@ -116,14 +119,16 @@ class WorkflowAgentComposerCopyFromRosterApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, account_id: str, app_model: App, node_id: str): + @with_session + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def post(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): payload = WorkflowComposerCopyFromRosterPayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, @@ -131,7 +136,6 @@ 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(), ), ) @@ -145,19 +149,22 @@ class WorkflowAgentComposerValidateApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_tenant_id - def post(self, tenant_id: str, app_model: App, node_id: str): + @with_session(write=False) + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def post(self, session: Session, tenant_id: str, app_model: App, node_id: str): 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) + AgentComposerService.validate_knowledge_datasets( + session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + ) findings = AgentComposerService.collect_validation_findings( + session=session, 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, session=db.session() + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id ), - session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -170,18 +177,19 @@ class WorkflowAgentComposerCandidatesApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, current_user_id: str, app_model: App, node_id: str): + @with_session(write=False) + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def get(self, session: Session, tenant_id: str, current_user_id: str, app_model: App, node_id: str): return dump_response( AgentComposerCandidatesResponse, AgentComposerService.get_workflow_candidates( + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, user_id=current_user_id, - session=db.session(), ), ) @@ -193,9 +201,10 @@ class WorkflowAgentComposerImpactApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_tenant_id - def post(self, tenant_id: str, app_model: App, node_id: str): + @with_session(write=False) + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def post(self, session: Session, tenant_id: str, app_model: App, node_id: str): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) current_snapshot_id = payload.binding.current_snapshot_id if payload.binding else None if not current_snapshot_id: @@ -205,7 +214,7 @@ class WorkflowAgentComposerImpactApi(Resource): return dump_response( AgentComposerImpactResponse, AgentComposerService.calculate_impact( - tenant_id=tenant_id, current_snapshot_id=current_snapshot_id, session=db.session() + session=session, tenant_id=tenant_id, current_snapshot_id=current_snapshot_id ), ) @@ -221,26 +230,27 @@ class WorkflowAgentComposerSaveToRosterApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, account_id: str, app_model: App, node_id: str): + @with_session + @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) + def post(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( + session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id, account_id=account_id, payload=payload, - session=db.session(), ), ) -def _require_snippet_app_id(*, tenant_id: str, snippet_id: UUID) -> str: - snippet = SnippetService(session=db.session()).get_snippet_by_id( +def _require_snippet_app_id(*, session: Session, tenant_id: str, snippet_id: UUID) -> str: + snippet = SnippetService(session=session).get_snippet_by_id( snippet_id=str(snippet_id), tenant_id=tenant_id, ) @@ -258,17 +268,18 @@ class SnippetAgentComposerApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): + @with_session + def get(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): query = WorkflowAgentComposerQuery.model_validate(request.args.to_dict(flat=True)) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.load_workflow_composer( + session=session, tenant_id=tenant_id, - app_id=_require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id), + app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, snapshot_id=query.snapshot_id, - session=db.session(), ), ) @@ -283,17 +294,18 @@ class SnippetAgentComposerApi(Resource): ) @with_current_user_id @with_current_tenant_id - def put(self, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): + @with_session + def put(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( + session=session, tenant_id=tenant_id, - app_id=_require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id), + app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, payload=payload, - session=db.session(), ), ) @@ -313,19 +325,20 @@ class SnippetAgentComposerCopyFromRosterApi(Resource): ) @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): + @with_session + def post(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): payload = WorkflowComposerCopyFromRosterPayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id=tenant_id, - app_id=_require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id), + app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, source_agent_id=payload.source_agent_id, source_snapshot_id=payload.source_snapshot_id, idempotency_key=payload.idempotency_key, - session=db.session(), ), ) @@ -340,21 +353,24 @@ class SnippetAgentComposerValidateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, snippet_id: UUID, node_id: str): - app_id = _require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id) + @with_session(write=False) + def post(self, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): + app_id = _require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id) 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) + AgentComposerService.validate_knowledge_datasets( + session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + ) findings = AgentComposerService.collect_validation_findings( + session=session, tenant_id=tenant_id, payload=payload, agent_id=AgentComposerService.resolve_workflow_node_agent_id( + session=session, tenant_id=tenant_id, app_id=app_id, node_id=node_id, - session=db.session(), ), - session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -369,15 +385,16 @@ class SnippetAgentComposerCandidatesApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, current_user_id: str, snippet_id: UUID, node_id: str): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user_id: str, snippet_id: UUID, node_id: str): return dump_response( AgentComposerCandidatesResponse, AgentComposerService.get_workflow_candidates( + session=session, tenant_id=tenant_id, - app_id=_require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id), + app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, user_id=current_user_id, - session=db.session(), ), ) @@ -390,8 +407,9 @@ class SnippetAgentComposerImpactApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, snippet_id: UUID, node_id: str): - _require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id) + @with_session(write=False) + def post(self, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): + _require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id) payload = ComposerSavePayload.model_validate(console_ns.payload or {}) current_snapshot_id = payload.binding.current_snapshot_id if payload.binding else None if not current_snapshot_id: @@ -401,9 +419,9 @@ class SnippetAgentComposerImpactApi(Resource): return dump_response( AgentComposerImpactResponse, AgentComposerService.calculate_impact( + session=session, tenant_id=tenant_id, current_snapshot_id=current_snapshot_id, - session=db.session(), ), ) @@ -423,17 +441,18 @@ class SnippetAgentComposerSaveToRosterApi(Resource): ) @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): + @with_session + def post(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( + session=session, tenant_id=tenant_id, - app_id=_require_snippet_app_id(tenant_id=tenant_id, snippet_id=snippet_id), + app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, payload=payload, - session=db.session(), ), ) @@ -445,10 +464,11 @@ class AgentComposerApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, agent_id: UUID): return dump_response( AgentAppComposerResponse, - AgentComposerService.load_agent_composer(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()), + AgentComposerService.load_agent_composer(session=session, tenant_id=tenant_id, agent_id=str(agent_id)), ) @console_ns.expect(console_ns.models[ComposerSavePayload.__name__]) @@ -460,16 +480,17 @@ class AgentComposerApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_user_id @with_current_tenant_id - def put(self, tenant_id: str, account_id: str, agent_id: UUID): + @with_session + def put(self, session: Session, tenant_id: str, account_id: str, agent_id: UUID): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return dump_response( AgentAppComposerResponse, AgentComposerService.save_agent_composer( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=account_id, payload=payload, - session=db.session(), ), ) @@ -484,16 +505,19 @@ class AgentComposerValidateApi(Resource): @login_required @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), session=db.session()) + @with_session + def post(self, session: Session, tenant_id: str, agent_id: UUID): + AgentComposerService.load_agent_composer(session=session, tenant_id=tenant_id, agent_id=str(agent_id)) 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) + AgentComposerService.validate_knowledge_datasets( + session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + ) findings = AgentComposerService.collect_validation_findings( + session=session, tenant_id=tenant_id, payload=payload, agent_id=str(agent_id), - session=db.session(), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) @@ -508,13 +532,14 @@ class AgentComposerCandidatesApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, current_user_id: str, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user_id: str, agent_id: UUID): return dump_response( AgentComposerCandidatesResponse, AgentComposerService.get_agent_app_candidates( + session=session, 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 1467cc0c246..c8d2134ed83 100644 --- a/api/controllers/console/agent/roster.py +++ b/api/controllers/console/agent/roster.py @@ -4,6 +4,7 @@ from flask import abort, request from flask_restx import Resource from pydantic import AliasChoices, BaseModel, Field, field_validator from sqlalchemy import func, select +from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, @@ -11,8 +12,8 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session from controllers.console import console_ns -from controllers.console.agent.app_helpers import resolve_agent_app_model, resolve_agent_runtime_app_model from controllers.console.apikey import ApiKeyItem, ApiKeyList, BaseApiKeyListResource, BaseApiKeyResource from controllers.console.app.app import ( APP_LIST_QUERY_ARRAY_FIELDS, @@ -42,7 +43,6 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) -from extensions.ext_database import db from fields.agent_fields import ( AgentConfigDraftSummaryResponse, AgentConfigSnapshotDetailResponse, @@ -339,11 +339,13 @@ register_response_schema_models( ) -def _agent_roster_service() -> AgentRosterService: - return AgentRosterService(db.session) +def _agent_roster_service(session: Session) -> AgentRosterService: + return AgentRosterService(session) -def _serialize_agent_app_detail(app_model, *, current_user: Account, agent_id: str | None = None) -> dict: +def _serialize_agent_app_detail( + session: Session, app_model, *, current_user: Account, agent_id: str | None = None +) -> dict: """Serialize an Agent App detail using roster-only DTOs. `/agent` responses are roster-shaped rather than raw app-shaped: `id` @@ -353,15 +355,19 @@ def _serialize_agent_app_detail(app_model, *, current_user: Account, agent_id: s roster persona fields without widening the shared /apps detail schema. """ - app_model = AppService().get_app(app_model) + app_model = AppService().get_app(app_model, session=session) if FeatureService.get_system_features().webapp_auth.enabled: app_setting = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=str(app_model.id)) app_model.access_mode = app_setting.access_mode # type: ignore[attr-defined] - roster_service = _agent_roster_service() - payload = AgentAppDetailWithSite.model_validate(app_model, from_attributes=True).model_dump(mode="json") + roster_service = _agent_roster_service(session) + payload = AgentAppDetailWithSite.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") agent = ( - db.session.scalar( + session.scalar( select(Agent).where( Agent.tenant_id == app_model.tenant_id, Agent.id == agent_id, @@ -382,6 +388,7 @@ def _serialize_agent_app_detail(app_model, *, current_user: Account, agent_id: s tenant_id=app_model.tenant_id, agent_id=agent.id, account_id=current_user.id, + commit=False, ) message_count = roster_service.count_agent_app_debug_conversation_messages( conversation_id=debug_conversation_id, @@ -397,7 +404,7 @@ def _serialize_agent_app_detail(app_model, *, current_user: Account, agent_id: s return payload -def _serialize_agent_app_pagination(app_pagination, *, tenant_id: str, current_user: Account) -> dict: +def _serialize_agent_app_pagination(session: Session, app_pagination, *, tenant_id: str, current_user: Account) -> dict: """Serialize Agent App lists with roster-shaped items. Each item starts from the shared App list shape, then drops @@ -407,7 +414,7 @@ def _serialize_agent_app_pagination(app_pagination, *, tenant_id: str, current_u """ app_ids = [str(app.id) for app in app_pagination.items] - roster_service = _agent_roster_service() + roster_service = _agent_roster_service(session) agents_by_app_id = roster_service.load_app_backing_agents_by_app_id( tenant_id=tenant_id, app_ids=app_ids, @@ -425,7 +432,11 @@ def _serialize_agent_app_pagination(app_pagination, *, tenant_id: str, current_u agents=list(agents_by_app_id.values()), account_id=current_user.id, ) - payload = AgentAppPagination.model_validate(app_pagination, from_attributes=True).model_dump(mode="json") + payload = AgentAppPagination.model_validate( + app_pagination, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") for item in payload["data"]: app_id = item["id"] item.pop("bound_agent_id", None) @@ -456,13 +467,17 @@ def _serialize_agent_app_pagination(app_pagination, *, tenant_id: str, current_u ) -def _resolve_agent_app_model(*, tenant_id: str, agent_id: UUID): - return resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) +def _resolve_agent_app_model(session: Session, *, tenant_id: str, agent_id: UUID) -> App: + return _agent_roster_service(session).get_agent_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) -def _agent_api_key_count(app_id: str) -> int: +def _resolve_agent_runtime_app_model(session: Session, *, tenant_id: str, agent_id: UUID) -> App: + return _agent_roster_service(session).get_agent_runtime_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) + + +def _agent_api_key_count(session: Session, app_id: str) -> int: return ( - db.session.scalar( + session.scalar( select(func.count(ApiToken.id)).where( ApiToken.type == ApiTokenType.APP, ApiToken.app_id == app_id, @@ -472,7 +487,7 @@ def _agent_api_key_count(app_id: str) -> int: ) -def _serialize_agent_api_access(app_model: App) -> dict: +def _serialize_agent_api_access(session: Session, app_model: App) -> dict: base_url = app_model.api_base_url response = AgentApiAccessResponse( enabled=bool(app_model.enable_api), @@ -487,13 +502,13 @@ def _serialize_agent_api_access(app_model: App) -> dict: meta_endpoint=f"{base_url}/meta", api_rpm=app_model.api_rpm or 0, api_rph=app_model.api_rph or 0, - api_key_count=_agent_api_key_count(str(app_model.id)), + api_key_count=_agent_api_key_count(session, str(app_model.id)), ) return response.model_dump(mode="json") -def _agent_observability_service() -> AgentObservabilityService: - return AgentObservabilityService(db.session) +def _agent_observability_service(session: Session) -> AgentObservabilityService: + return AgentObservabilityService(session) def _parse_observability_time_range(start: str | None, end: str | None, account: Account): @@ -520,7 +535,8 @@ class AgentAppListApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): + @with_session + def get(self, session: Session, current_tenant_id: str, current_user: Account): args = query_params_from_request(AppListQuery, list_fields=APP_LIST_QUERY_ARRAY_FIELDS) params = AppListParams( page=args.page, @@ -534,12 +550,13 @@ 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, 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") return _serialize_agent_app_pagination( + session, app_pagination, tenant_id=current_tenant_id, current_user=current_user, @@ -555,7 +572,8 @@ class AgentAppListApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account): + @with_session + def post(self, session: Session, current_tenant_id: str, current_user: Account): args = AgentAppCreatePayload.model_validate(console_ns.payload) params = CreateAppParams( name=args.name, @@ -567,8 +585,8 @@ class AgentAppListApi(Resource): icon_background=args.icon_background, ) - app = AppService().create_app(current_tenant_id, params, current_user, session=db.session()) - return _serialize_agent_app_detail(app, current_user=current_user), 201 + app = AppService().create_app(current_tenant_id, params, current_user, session=session) + return _serialize_agent_app_detail(session, app, current_user=current_user), 201 @console_ns.route("/agent/") @@ -580,9 +598,10 @@ class AgentAppApi(Resource): @enterprise_license_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _serialize_agent_app_detail(app_model, current_user=current_user, agent_id=str(agent_id)) + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + return _serialize_agent_app_detail(session, app_model, current_user=current_user, agent_id=str(agent_id)) @console_ns.expect(console_ns.models[AgentAppUpdatePayload.__name__]) @console_ns.response(200, "Agent app updated successfully", console_ns.models[AgentAppDetailWithSite.__name__]) @@ -594,8 +613,9 @@ class AgentAppApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def put(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session + def put(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) args = AgentAppUpdatePayload.model_validate(console_ns.payload) args_dict: AppService.ArgsDict = { "name": args.name, @@ -607,8 +627,8 @@ class AgentAppApi(Resource): "max_active_requests": args.max_active_requests or 0, "role": args.role, } - updated = AppService().update_app(app_model, args_dict, session=db.session()) - return _serialize_agent_app_detail(updated, current_user=current_user) + updated = AppService().update_app(app_model, args_dict, session=session) + return _serialize_agent_app_detail(session, updated, current_user=current_user) @console_ns.response(204, "Agent app deleted successfully") @console_ns.response(403, "Insufficient permissions") @@ -617,9 +637,10 @@ class AgentAppApi(Resource): @account_initialization_required @edit_permission_required @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, session=db.session()) + @with_session + def delete(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + AppService().delete_app(app_model, session=session) return "", 204 @@ -637,8 +658,9 @@ class AgentDebugConversationRefreshApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): - debug_conversation_id = _agent_roster_service().refresh_agent_app_debug_conversation_id( + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + debug_conversation_id = _agent_roster_service(session).refresh_agent_app_debug_conversation_id( tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, @@ -661,14 +683,15 @@ class AgentPublishApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): args = AgentPublishPayload.model_validate(console_ns.payload or {}) return AgentComposerService.publish_agent_app_draft( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, version_note=args.version_note, - session=db.session(), ) @@ -682,14 +705,15 @@ class AgentBuildDraftCheckoutApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): args = AgentBuildDraftCheckoutPayload.model_validate(console_ns.payload or {}) return AgentComposerService.checkout_agent_app_build_draft( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, force=args.force, - session=db.session(), ) @@ -702,12 +726,13 @@ class AgentBuildDraftApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return AgentComposerService.load_agent_app_build_draft( + session=session, 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__]) @@ -718,14 +743,15 @@ class AgentBuildDraftApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def put(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def put(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): payload = ComposerSavePayload.model_validate(console_ns.payload or {}) return AgentComposerService.save_agent_app_build_draft( + session=session, tenant_id=tenant_id, 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__]) @@ -735,12 +761,13 @@ class AgentBuildDraftApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return AgentComposerService.discard_agent_app_build_draft( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, - session=db.session(), ) @@ -753,12 +780,13 @@ class AgentBuildDraftApplyApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return AgentComposerService.apply_agent_app_build_draft( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, - session=db.session(), ) @@ -774,9 +802,10 @@ class AgentAppCopyApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): args = AgentAppCopyPayload.model_validate(console_ns.payload or {}) - copied_app = _agent_roster_service().duplicate_agent_app( + copied_app = _agent_roster_service(session).duplicate_agent_app( tenant_id=tenant_id, agent_id=str(agent_id), account=current_user, @@ -787,7 +816,7 @@ class AgentAppCopyApi(Resource): icon=args.icon, icon_background=args.icon_background, ) - return _serialize_agent_app_detail(copied_app, current_user=current_user), 201 + return _serialize_agent_app_detail(session, copied_app, current_user=current_user), 201 @console_ns.route("/agent//api-access") @@ -797,9 +826,10 @@ class AgentApiAccessApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _serialize_agent_api_access(app_model) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + return _serialize_agent_api_access(session, app_model) @console_ns.route("/agent//api-enable") @@ -813,11 +843,12 @@ class AgentApiStatusApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @with_current_tenant_id - def post(self, tenant_id: str, agent_id: UUID): - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session + def post(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = _resolve_agent_app_model(session, 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, session=db.session()) - return _serialize_agent_api_access(app_model) + app_model = AppService().update_app_api_status(app_model, args.enable_api, session=session) + return _serialize_agent_api_access(session, app_model) @console_ns.route("/agent//api-keys") @@ -829,18 +860,23 @@ class AgentApiKeyListApi(BaseApiKeyListResource): @console_ns.response(200, "Agent service API keys", console_ns.models[ApiKeyList.__name__]) @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID) -> dict[str, object]: - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - return dump_response(ApiKeyList, self._get_api_key_list(str(app_model.id), tenant_id)) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID) -> dict[str, object]: + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + return dump_response(ApiKeyList, self._get_api_key_list(str(app_model.id), tenant_id, session=session)) @console_ns.response(201, "Agent service API key created", console_ns.models[ApiKeyItem.__name__]) @console_ns.response(400, "Maximum keys exceeded") @with_current_tenant_id @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) - def post(self, tenant_id: str, agent_id: UUID) -> tuple[dict[str, object], int]: - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - return dump_response(ApiKeyItem, self._create_api_key(str(app_model.id), tenant_id)), 201 + @with_session + def post(self, session: Session, tenant_id: str, agent_id: UUID) -> tuple[dict[str, object], int]: + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + return dump_response( + ApiKeyItem, + self._create_api_key(str(app_model.id), tenant_id, session=session), + ), 201 @console_ns.route("/agent//api-keys/") @@ -853,9 +889,17 @@ class AgentApiKeyApi(BaseApiKeyResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID, api_key_id: UUID) -> tuple[str, int]: - app_model = _resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - self._delete_api_key(str(app_model.id), str(api_key_id), tenant_id, current_user) + @with_session + def delete( + self, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + api_key_id: UUID, + ) -> tuple[str, int]: + app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + self._delete_api_key(str(app_model.id), str(api_key_id), tenant_id, current_user, session=session) return "", 204 @@ -867,11 +911,12 @@ class AgentInviteOptionsApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): + @with_session(write=False) + def get(self, session: Session, tenant_id: str): query = AgentInviteOptionsQuery.model_validate(request.args.to_dict(flat=True)) return dump_response( AgentInviteOptionsResponse, - _agent_roster_service().list_invite_options( + _agent_roster_service(session).list_invite_options( tenant_id=tenant_id, page=query.page, limit=query.limit, @@ -890,15 +935,16 @@ class AgentLogsApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) query_data: dict[str, object] = dict(request.args.to_dict(flat=True)) query_data["sources"] = _query_values("sources", "source") query_data["statuses"] = _query_values("statuses", "status") query = AgentLogsQuery.model_validate(query_data) start, end = _parse_observability_time_range(query.start, query.end, current_user) try: - payload = _agent_observability_service().list_logs( + payload = _agent_observability_service(session).list_logs( app=app_model, agent_id=str(agent_id), params=AgentLogQueryParams( @@ -927,15 +973,16 @@ class AgentLogMessagesApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, conversation_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, conversation_id: UUID): + app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) query_data: dict[str, object] = dict(request.args.to_dict(flat=True)) query_data["sources"] = _query_values("sources", "source") query_data["statuses"] = _query_values("statuses", "status") query = AgentLogsQuery.model_validate(query_data) start, end = _parse_observability_time_range(query.start, query.end, current_user) try: - payload = _agent_observability_service().list_log_messages( + payload = _agent_observability_service(session).list_log_messages( app=app_model, agent_id=str(agent_id), conversation_id=str(conversation_id), @@ -964,9 +1011,10 @@ class AgentLogSourcesApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - payload = _agent_observability_service().list_log_sources(app=app_model, agent_id=str(agent_id)) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) + payload = _agent_observability_service(session).list_log_sources(app=app_model, agent_id=str(agent_id)) return dump_response(AgentLogSourceListResponse, payload) @@ -983,13 +1031,14 @@ class AgentStatisticsSummaryApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) query = AgentStatisticsQuery.model_validate(request.args.to_dict(flat=True)) timezone = current_user.timezone or "UTC" start, end = _parse_observability_time_range(query.start, query.end, current_user) try: - payload = _agent_observability_service().get_statistics_summary( + payload = _agent_observability_service(session).get_statistics_summary( app=app_model, agent_id=str(agent_id), params=AgentStatisticsQueryParams(source=query.source, start=start, end=end, timezone=timezone), @@ -1006,10 +1055,11 @@ class AgentRosterVersionsApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): return dump_response( AgentConfigSnapshotListResponse, - {"data": _agent_roster_service().list_agent_versions(tenant_id=tenant_id, agent_id=str(agent_id))}, + {"data": _agent_roster_service(session).list_agent_versions(tenant_id=tenant_id, agent_id=str(agent_id))}, ) @@ -1020,10 +1070,11 @@ class AgentRosterVersionDetailApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID, version_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID, version_id: UUID): return dump_response( AgentConfigSnapshotDetailResponse, - _agent_roster_service().get_agent_version_detail( + _agent_roster_service(session).get_agent_version_detail( tenant_id=tenant_id, agent_id=str(agent_id), version_id=str(version_id), @@ -1040,10 +1091,11 @@ class AgentRosterVersionRestoreApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID, version_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, version_id: UUID): return dump_response( AgentConfigSnapshotRestoreResponse, - _agent_roster_service().restore_agent_version( + _agent_roster_service(session).restore_agent_version( tenant_id=tenant_id, agent_id=str(agent_id), version_id=str(version_id), diff --git a/api/controllers/console/apikey.py b/api/controllers/console/apikey.py index dcea303d7cc..41a81267ff9 100644 --- a/api/controllers/console/apikey.py +++ b/api/controllers/console/apikey.py @@ -6,12 +6,12 @@ from flask_restx import Resource from flask_restx._http import HTTPStatus from pydantic import field_validator from sqlalchemy import delete, func, select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden from configs import dify_config from controllers.common.schema import register_response_schema_models -from extensions.ext_database import db +from controllers.common.session import with_session from fields.base import ResponseModel from libs.helper import dump_response, to_timestamp from libs.login import login_required @@ -54,11 +54,10 @@ class ApiKeyList(ResponseModel): register_response_schema_models(console_ns, ApiKeyItem, ApiKeyList) -def _get_resource(resource_id, tenant_id, resource_model): - with sessionmaker(db.engine).begin() as session: - resource = session.execute( - select(resource_model).filter_by(id=resource_id, tenant_id=tenant_id) - ).scalar_one_or_none() +def _get_resource(resource_id, tenant_id, resource_model, *, session: Session): + resource = session.execute( + select(resource_model).filter_by(id=resource_id, tenant_id=tenant_id) + ).scalar_one_or_none() if resource is None: flask_restx.abort(HTTPStatus.NOT_FOUND, message=f"{resource_model.__name__} not found.") @@ -75,14 +74,18 @@ class BaseApiKeyListResource(Resource): token_prefix: str | None = None max_keys = 10 - def get(self, resource_id: str, current_tenant_id: str) -> dict[str, object]: - return dump_response(ApiKeyList, self._get_api_key_list(resource_id, current_tenant_id)) + @with_session(write=False) + def get(self, session: Session, resource_id: str, current_tenant_id: str) -> dict[str, object]: + return dump_response( + ApiKeyList, + self._get_api_key_list(resource_id, current_tenant_id, session=session), + ) - def _get_api_key_list(self, resource_id: str, current_tenant_id: str) -> ApiKeyList: + def _get_api_key_list(self, resource_id: str, current_tenant_id: str, *, session: Session) -> ApiKeyList: assert self.resource_id_field is not None, "resource_id_field must be set" - _get_resource(resource_id, current_tenant_id, self.resource_model) - keys = db.session.scalars( + _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) + keys = session.scalars( select(ApiToken).where( ApiToken.type == self.resource_type, getattr(ApiToken, self.resource_id_field) == resource_id ) @@ -90,14 +93,18 @@ class BaseApiKeyListResource(Resource): return ApiKeyList.model_validate({"data": keys}, from_attributes=True) @edit_permission_required - def post(self, resource_id: str, current_tenant_id: str) -> tuple[dict[str, object], int]: - return dump_response(ApiKeyItem, self._create_api_key(resource_id, current_tenant_id)), 201 + @with_session + def post(self, session: Session, resource_id: str, current_tenant_id: str) -> tuple[dict[str, object], int]: + return dump_response( + ApiKeyItem, + self._create_api_key(resource_id, current_tenant_id, session=session), + ), 201 - def _create_api_key(self, resource_id: str, current_tenant_id: str) -> ApiToken: + def _create_api_key(self, resource_id: str, current_tenant_id: str, *, session: Session) -> ApiToken: assert self.resource_id_field is not None, "resource_id_field must be set" - _get_resource(resource_id, current_tenant_id, self.resource_model) + _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) current_key_count: int = ( - db.session.scalar( + session.scalar( select(func.count(ApiToken.id)).where( ApiToken.type == self.resource_type, getattr(ApiToken, self.resource_id_field) == resource_id ) @@ -112,15 +119,15 @@ class BaseApiKeyListResource(Resource): custom="max_keys_exceeded", ) - key = ApiToken.generate_api_key(self.token_prefix or "", 24) + key = ApiToken.generate_api_key(self.token_prefix or "", 24, session=session) assert self.resource_type is not None, "resource_type must be set" api_token = ApiToken() setattr(api_token, self.resource_id_field, resource_id) api_token.tenant_id = current_tenant_id api_token.token = key api_token.type = self.resource_type - db.session.add(api_token) - db.session.commit() + session.add(api_token) + session.commit() return api_token @@ -131,10 +138,16 @@ class BaseApiKeyResource(Resource): resource_model: type | None = None resource_id_field: str | None = None + @with_session def delete( - self, resource_id: str, api_key_id: str, current_tenant_id: str, current_user: Account + self, + session: Session, + resource_id: str, + api_key_id: str, + current_tenant_id: str, + current_user: Account, ) -> tuple[str, int]: - self._delete_api_key(resource_id, api_key_id, current_tenant_id, current_user) + self._delete_api_key(resource_id, api_key_id, current_tenant_id, current_user, session=session) return "", 204 def _delete_api_key( @@ -143,14 +156,16 @@ class BaseApiKeyResource(Resource): api_key_id: str, current_tenant_id: str, current_user: Account, + *, + session: Session, ) -> None: assert self.resource_id_field is not None, "resource_id_field must be set" - _get_resource(resource_id, current_tenant_id, self.resource_model) + _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) if not dify_config.RBAC_ENABLED and not current_user.is_admin_or_owner: raise Forbidden() - key = db.session.scalar( + key = session.scalar( select(ApiToken) .where( getattr(ApiToken, self.resource_id_field) == resource_id, @@ -168,8 +183,8 @@ class BaseApiKeyResource(Resource): assert key is not None # nosec - for type checker only ApiTokenCache.delete(key.token, key.type) - db.session.execute(delete(ApiToken).where(ApiToken.id == api_key_id)) - db.session.commit() + session.execute(delete(ApiToken).where(ApiToken.id == api_key_id)) + session.commit() @console_ns.route("/apps//api-keys") @@ -179,9 +194,13 @@ class AppApiKeyListResource(BaseApiKeyListResource): @console_ns.doc(params={"resource_id": "App ID"}) @console_ns.response(200, "API keys retrieved successfully", console_ns.models[ApiKeyList.__name__]) @with_current_tenant_id - def get(self, current_tenant_id: str, resource_id: UUID) -> dict[str, object]: + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, resource_id: UUID) -> dict[str, object]: """Get all API keys for an app""" - return dump_response(ApiKeyList, self._get_api_key_list(str(resource_id), current_tenant_id)) + return dump_response( + ApiKeyList, + self._get_api_key_list(str(resource_id), current_tenant_id, session=session), + ) @console_ns.doc("create_app_api_key") @console_ns.doc(description="Create a new API key for an app") @@ -191,9 +210,13 @@ class AppApiKeyListResource(BaseApiKeyListResource): @with_current_tenant_id @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) - def post(self, current_tenant_id: str, resource_id: UUID) -> tuple[dict[str, object], int]: + @with_session + def post(self, session: Session, current_tenant_id: str, resource_id: UUID) -> tuple[dict[str, object], int]: """Create a new API key for an app""" - return dump_response(ApiKeyItem, self._create_api_key(str(resource_id), current_tenant_id)), 201 + return dump_response( + ApiKeyItem, + self._create_api_key(str(resource_id), current_tenant_id, session=session), + ), 201 resource_type = ApiTokenType.APP resource_model = App @@ -210,11 +233,23 @@ class AppApiKeyResource(BaseApiKeyResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) + @with_session def delete( - self, current_tenant_id: str, current_user: Account, resource_id: UUID, api_key_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + resource_id: UUID, + api_key_id: UUID, ) -> tuple[str, int]: """Delete an API key for an app""" - self._delete_api_key(str(resource_id), str(api_key_id), current_tenant_id, current_user) + self._delete_api_key( + str(resource_id), + str(api_key_id), + current_tenant_id, + current_user, + session=session, + ) return "", 204 resource_type = ApiTokenType.APP @@ -229,9 +264,13 @@ class DatasetApiKeyListResource(BaseApiKeyListResource): @console_ns.doc(params={"resource_id": "Dataset ID"}) @console_ns.response(200, "API keys retrieved successfully", console_ns.models[ApiKeyList.__name__]) @with_current_tenant_id - def get(self, current_tenant_id: str, resource_id: UUID) -> dict[str, object]: + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, resource_id: UUID) -> dict[str, object]: """Get all API keys for a dataset""" - return dump_response(ApiKeyList, self._get_api_key_list(str(resource_id), current_tenant_id)) + return dump_response( + ApiKeyList, + self._get_api_key_list(str(resource_id), current_tenant_id, session=session), + ) @console_ns.doc("create_dataset_api_key") @console_ns.doc(description="Create a new API key for a dataset") @@ -241,9 +280,13 @@ class DatasetApiKeyListResource(BaseApiKeyListResource): @with_current_tenant_id @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE) - def post(self, current_tenant_id: str, resource_id: UUID) -> tuple[dict[str, object], int]: + @with_session + def post(self, session: Session, current_tenant_id: str, resource_id: UUID) -> tuple[dict[str, object], int]: """Create a new API key for a dataset""" - return dump_response(ApiKeyItem, self._create_api_key(str(resource_id), current_tenant_id)), 201 + return dump_response( + ApiKeyItem, + self._create_api_key(str(resource_id), current_tenant_id, session=session), + ), 201 resource_type = ApiTokenType.DATASET resource_model = Dataset @@ -260,11 +303,23 @@ class DatasetApiKeyResource(BaseApiKeyResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE) + @with_session def delete( - self, current_tenant_id: str, current_user: Account, resource_id: UUID, api_key_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + resource_id: UUID, + api_key_id: UUID, ) -> tuple[str, int]: """Delete an API key for a dataset""" - self._delete_api_key(str(resource_id), str(api_key_id), current_tenant_id, current_user) + self._delete_api_key( + str(resource_id), + str(api_key_id), + current_tenant_id, + current_user, + session=session, + ) return "", 204 resource_type = ApiTokenType.DATASET diff --git a/api/controllers/console/app/agent.py b/api/controllers/console/app/agent.py index 81d17ace37a..8be24acc8fb 100644 --- a/api/controllers/console/app/agent.py +++ b/api/controllers/console/app/agent.py @@ -5,6 +5,7 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select +from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, @@ -12,6 +13,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session 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 @@ -24,7 +26,6 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) -from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import uuid_value from libs.login import login_required @@ -169,23 +170,23 @@ register_response_schema_models( ) -def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: +def _resolve_agent_id(session: Session, 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, session=db.session() + session=session, tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id ) - return app_model.bound_agent_id + return app_model.bound_agent_id_with_session(session=session) def _agent_not_bound() -> tuple[dict[str, str], int]: return {"code": "agent_not_bound", "message": "no agent is bound for this app/node"}, 400 -def _upload_skill_for_app(*, current_user: Account, app_model: App): +def _upload_skill_for_app(*, session: Session, current_user: Account, app_model: App): """Upload one skill package and commit its normalized files into the agent drive.""" query = query_params_from_request(AgentDriveMutationQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) if not agent_id: return _agent_not_bound() if "file" not in request.files: @@ -202,22 +203,22 @@ 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(), + session=session, ) except (SkillPackageError, AgentDriveError) as exc: return {"code": exc.code, "message": exc.message}, exc.status_code return result, 201 -def _commit_drive_file_for_app(*, current_user: Account, app_model: App, allow_node_id: bool = True): +def _commit_drive_file_for_app(*, session: Session, current_user: Account, app_model: App, allow_node_id: bool = True): query = query_params_from_request(AgentDriveMutationQuery) node_id = query.node_id if allow_node_id else None - agent_id = _resolve_agent_id(app_model, node_id) + agent_id = _resolve_agent_id(session, app_model, node_id) if not agent_id: return _agent_not_bound() payload = AgentDriveFilePayload.model_validate(console_ns.payload or {}) - upload_file = db.session.scalar( + upload_file = session.scalar( select(UploadFile).where( UploadFile.id == payload.upload_file_id, UploadFile.tenant_id == app_model.tenant_id, @@ -241,7 +242,7 @@ def _commit_drive_file_for_app(*, current_user: Account, app_model: App, allow_n value_owned_by_drive=True, ) ], - session=db.session(), + session=session, ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -258,10 +259,10 @@ def _commit_drive_file_for_app(*, current_user: Account, app_model: App, allow_n }, 201 -def _delete_drive_file_for_app(*, current_user: Account, app_model: App, allow_node_id: bool = True): +def _delete_drive_file_for_app(*, session: Session, current_user: Account, app_model: App, allow_node_id: bool = True): query = query_params_from_request(AgentDriveDeleteFileQuery) node_id = query.node_id if allow_node_id else None - agent_id = _resolve_agent_id(app_model, node_id) + agent_id = _resolve_agent_id(session, app_model, node_id) if not agent_id: return _agent_not_bound() try: @@ -275,7 +276,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(), + session=session, ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -283,10 +284,12 @@ def _delete_drive_file_for_app(*, current_user: Account, app_model: App, allow_n return {"result": "success", "removed_keys": removed_keys} -def _delete_skill_for_app(*, current_user: Account, app_model: App, slug: str, allow_node_id: bool = True): +def _delete_skill_for_app( + *, session: Session, current_user: Account, app_model: App, slug: str, allow_node_id: bool = True +): query = query_params_from_request(AgentDriveMutationQuery) node_id = query.node_id if allow_node_id else None - agent_id = _resolve_agent_id(app_model, node_id) + agent_id = _resolve_agent_id(session, app_model, node_id) if not agent_id: return _agent_not_bound() if "/" in slug or not slug.strip(): @@ -301,7 +304,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(), + session=session, ) except AgentDriveError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -309,16 +312,16 @@ def _delete_skill_for_app(*, current_user: Account, app_model: App, slug: str, a return {"result": "success", "removed_keys": removed_keys} -def _infer_skill_tools_for_app(*, app_model: App, slug: str): +def _infer_skill_tools_for_app(*, session: Session, app_model: App, slug: str): query = query_params_from_request(AgentDriveMutationQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) if not agent_id: return _agent_not_bound() 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, session=db.session() + tenant_id=app_model.tenant_id, agent_id=agent_id, slug=slug, session=session ) except SkillToolInferenceError as exc: return {"code": exc.code, "message": exc.message}, exc.status_code @@ -336,12 +339,13 @@ class AgentLogApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=[AppMode.AGENT_CHAT]) - def get(self, app_model: App): + def get(self, session: Session, app_model: App): """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, db.session()) + return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id, session) @console_ns.route("/agent//skills/upload") @@ -356,9 +360,10 @@ class AgentSkillUploadByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _upload_skill_for_app(current_user=current_user, app_model=app_model) + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + return _upload_skill_for_app(session=session, current_user=current_user, app_model=app_model) @console_ns.route("/apps//agent/skills/upload") @@ -378,11 +383,12 @@ class AgentSkillUploadApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) @with_current_user - def post(self, current_user: Account, app_model: App): + @with_session + @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) + def post(self, session: Session, current_user: Account, app_model: App): """Upload a Skill, validate it, and commit drive-backed skill files.""" - return _upload_skill_for_app(current_user=current_user, app_model=app_model) + return _upload_skill_for_app(session=session, current_user=current_user, app_model=app_model) @console_ns.route("/agent//files") @@ -399,9 +405,12 @@ class AgentDriveFilesByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _commit_drive_file_for_app(current_user=current_user, app_model=app_model, allow_node_id=False) + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + return _commit_drive_file_for_app( + session=session, current_user=current_user, app_model=app_model, allow_node_id=False + ) @console_ns.doc("delete_agent_drive_file_by_agent") @console_ns.doc(description="Delete one Agent App drive file by key") @@ -412,9 +421,12 @@ class AgentDriveFilesByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _delete_drive_file_for_app(current_user=current_user, app_model=app_model, allow_node_id=False) + @with_session + def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + return _delete_drive_file_for_app( + session=session, current_user=current_user, app_model=app_model, allow_node_id=False + ) @console_ns.route("/apps//agent/files") @@ -429,11 +441,12 @@ class AgentDriveFilesApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) @with_current_user - def post(self, current_user: Account, app_model: App): + @with_session + @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) + def post(self, session: Session, current_user: Account, app_model: App): """ADD FILE: commit one uploaded file into the bound agent's drive.""" - return _commit_drive_file_for_app(current_user=current_user, app_model=app_model) + return _commit_drive_file_for_app(session=session, current_user=current_user, app_model=app_model) @console_ns.doc("delete_agent_drive_file") @console_ns.doc(description="Delete one drive file by key via drive commit-null semantics") @@ -442,10 +455,11 @@ class AgentDriveFilesApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) @with_current_user - def delete(self, current_user: Account, app_model: App): - return _delete_drive_file_for_app(current_user=current_user, app_model=app_model) + @with_session + @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) + def delete(self, session: Session, current_user: Account, app_model: App): + return _delete_drive_file_for_app(session=session, current_user=current_user, app_model=app_model) @console_ns.route("/agent//skills/") @@ -459,9 +473,12 @@ class AgentSkillByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID, slug: str): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _delete_skill_for_app(current_user=current_user, app_model=app_model, slug=slug, allow_node_id=False) + @with_session + def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, slug: str): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + return _delete_skill_for_app( + session=session, current_user=current_user, app_model=app_model, slug=slug, allow_node_id=False + ) @console_ns.route("/apps//agent/skills/") @@ -479,10 +496,11 @@ class AgentSkillApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) @with_current_user - def delete(self, current_user: Account, app_model: App, slug: str): - return _delete_skill_for_app(current_user=current_user, app_model=app_model, slug=slug) + @with_session + @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) + def delete(self, session: Session, current_user: Account, app_model: App, slug: str): + return _delete_skill_for_app(session=session, current_user=current_user, app_model=app_model, slug=slug) @console_ns.route("/agent//skills//infer-tools") @@ -499,9 +517,10 @@ class AgentSkillInferToolsByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, agent_id: UUID, slug: str): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) - return _infer_skill_tools_for_app(app_model=app_model, slug=slug) + @with_session(write=False) + def post(self, session: Session, tenant_id: str, agent_id: UUID, slug: str): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + return _infer_skill_tools_for_app(session=session, app_model=app_model, slug=slug) @console_ns.route("/apps//agent/skills//infer-tools") @@ -525,7 +544,8 @@ class AgentSkillInferToolsApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_AGENT_DRIVE_APP_MODES) - def post(self, app_model: App, slug: str): + def post(self, session: Session, app_model: App, slug: str): """Suggest CLI tools/env for a skill. Saving still goes through composer validation.""" - return _infer_skill_tools_for_app(app_model=app_model, slug=slug) + return _infer_skill_tools_for_app(session=session, app_model=app_model, slug=slug) diff --git a/api/controllers/console/app/agent_app_access.py b/api/controllers/console/app/agent_app_access.py index 4e79beac594..14c049d2f89 100644 --- a/api/controllers/console/app/agent_app_access.py +++ b/api/controllers/console/app/agent_app_access.py @@ -9,12 +9,13 @@ from uuid import UUID from flask_restx import Resource from pydantic import Field +from sqlalchemy.orm import Session from controllers.common.schema import register_response_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_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 services.agent.roster_service import AgentRosterService @@ -55,9 +56,10 @@ class AgentAppReferencingWorkflowsResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_app_model(tenant_id=tenant_id, agent_id=agent_id) - workflows = AgentRosterService(db.session).list_workflows_referencing_app_agent( + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = resolve_agent_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + workflows = AgentRosterService(session).list_workflows_referencing_app_agent( tenant_id=tenant_id, app_id=app_model.id ) return AgentReferencingWorkflowsResponse( diff --git a/api/controllers/console/app/agent_app_feature.py b/api/controllers/console/app/agent_app_feature.py index edd2f31f75f..99925727335 100644 --- a/api/controllers/console/app/agent_app_feature.py +++ b/api/controllers/console/app/agent_app_feature.py @@ -13,9 +13,11 @@ from uuid import UUID from flask_restx import Resource from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from controllers.common.fields import SimpleResultResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.wraps import ( @@ -29,7 +31,6 @@ from controllers.console.wraps import ( with_current_user, ) from events.app_event import app_model_config_was_updated -from extensions.ext_database import db from libs.login import login_required from models import Account from models.agent_config_entities import ( @@ -85,17 +86,23 @@ class AgentAppFeatureConfigResource(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) args = AgentAppFeaturesPayload.model_validate(console_ns.payload or {}) new_app_model_config = AgentAppFeatureConfigService.update_features( app_model=app_model, account=current_user, config=args.model_dump(exclude_none=True), - session=db.session(), + session=session, ) - app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config) + app_model_config_was_updated.send( + app_model, + app_model_config=new_app_model_config, + session=session, + ) + session.commit() return SimpleResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/app/agent_app_sandbox.py b/api/controllers/console/app/agent_app_sandbox.py index 6f3811ccdc8..17b5bcd23ad 100644 --- a/api/controllers/console/app/agent_app_sandbox.py +++ b/api/controllers/console/app/agent_app_sandbox.py @@ -14,6 +14,7 @@ from dify_agent.client import DifyAgentClientError, DifyAgentHTTPError, DifyAgen from flask import request from flask_restx import Resource from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, @@ -21,6 +22,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session 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 @@ -153,8 +155,9 @@ class AgentAppSandboxInfoResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) query = query_params_from_request(AgentSandboxInfoQuery) try: result = AgentAppSandboxService().get_info( @@ -177,8 +180,9 @@ class AgentAppSandboxListResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) query = query_params_from_request(AgentSandboxListQuery) try: result = AgentAppSandboxService().list_files( @@ -202,8 +206,9 @@ class AgentAppSandboxReadResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) query = query_params_from_request(AgentSandboxFileQuery) try: result = AgentAppSandboxService().read_file( @@ -227,8 +232,9 @@ class AgentAppSandboxUploadResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def post(self, session: Session, tenant_id: str, agent_id: UUID): + app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) payload = AgentSandboxUploadPayload.model_validate(request.get_json(silent=True) or {}) try: result = AgentAppSandboxService().upload_file( diff --git a/api/controllers/console/app/agent_config_inspector.py b/api/controllers/console/app/agent_config_inspector.py index 0f7aa80ca78..170b3756b39 100644 --- a/api/controllers/console/app/agent_config_inspector.py +++ b/api/controllers/console/app/agent_config_inspector.py @@ -13,6 +13,7 @@ from uuid import UUID from flask import Response, request, send_file, url_for from flask_restx import Resource from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, @@ -20,6 +21,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session 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 @@ -33,7 +35,6 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) -from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required from models.account import Account @@ -247,15 +248,15 @@ def _service() -> AgentConfigService: return AgentConfigService() -def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: +def _resolve_agent_id(session: Session, app_model: App, node_id: str | None) -> str | None: if node_id: return AgentComposerService.resolve_workflow_node_agent_id( + session=session, tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id, - session=db.session(), ) - return app_model.bound_agent_id + return app_model.bound_agent_id_with_session(session=session) def _agent_not_bound() -> tuple[dict[str, object], int]: @@ -275,6 +276,7 @@ def _json_response(data: Mapping[str, Any]) -> Response: def _resolve_console_version( *, + session: Session, tenant_id: str, agent_id: str, account_id: str, @@ -286,26 +288,24 @@ def _resolve_console_version( try: if draft_type == "debug_build": state = AgentComposerService.load_agent_app_build_draft( + session=session, 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, session=db.session() - ) + state = AgentComposerService.load_agent_composer(session=session, tenant_id=tenant_id, agent_id=agent_id) draft = state.get("draft") or {} draft_id = draft.get("id") if isinstance(draft_id, str) and draft_id: # load_agent_composer creates the normal draft on first access. # Config asset services use their own SQLAlchemy session, so the # draft must be visible before we hand its id across that boundary. - db.session.commit() + session.commit() return draft_id, AgentConfigVersionKind.DRAFT except AgentVersionNotFoundError as exc: raise AgentConfigServiceError( @@ -322,6 +322,7 @@ def _resolve_console_version( def _resolve_target( *, + session: Session, tenant_id: str, agent_id: str, account_id: str, @@ -329,6 +330,7 @@ def _resolve_target( draft_type: str | None, ) -> _ResolvedConsoleTarget: resolved_version_id, version_kind = _resolve_console_version( + session=session, tenant_id=tenant_id, agent_id=agent_id, account_id=account_id, @@ -346,13 +348,15 @@ def _resolve_target( def _resolve_agent_route_target( *, + session: Session, tenant_id: str, agent_id: UUID, current_user: Account, query: AgentConfigByAgentQuery, ) -> _ResolvedConsoleTarget: - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) return _resolve_target( + session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, @@ -363,14 +367,16 @@ def _resolve_agent_route_target( def _resolve_app_route_target( *, + session: Session, app_model: App, current_user: Account, query: AgentConfigQuery, ) -> _ResolvedConsoleTarget | tuple[dict[str, object], int]: - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) if not agent_id: return _agent_not_bound() return _resolve_target( + session=session, tenant_id=app_model.tenant_id, agent_id=agent_id, account_id=current_user.id, @@ -381,6 +387,7 @@ def _resolve_app_route_target( def _with_agent_route_target( *, + session: Session, tenant_id: str, agent_id: UUID, current_user: Account, @@ -389,6 +396,7 @@ def _with_agent_route_target( query = query_params_from_request(AgentConfigByAgentQuery) try: target = _resolve_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -401,13 +409,14 @@ def _with_agent_route_target( def _with_app_route_target( *, + session: Session, app_model: App, current_user: Account, action: Callable[[_ResolvedConsoleTarget], Any], ) -> Any: query = query_params_from_request(AgentConfigQuery) try: - target = _resolve_app_route_target(app_model=app_model, current_user=current_user, query=query) + target = _resolve_app_route_target(session=session, app_model=app_model, current_user=current_user, query=query) if isinstance(target, tuple): return target return action(target) @@ -649,8 +658,10 @@ class AgentConfigManifestByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -666,10 +677,13 @@ class AgentConfigManifestApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App): - return _with_app_route_target(app_model=app_model, current_user=current_user, action=_manifest_response) + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App): + return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=_manifest_response + ) @console_ns.route("/agent//config/skills/upload") @@ -689,8 +703,10 @@ class AgentConfigSkillUploadByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -715,10 +731,13 @@ class AgentConfigSkillUploadApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def post(self, current_user: Account, app_model: App): - return _with_app_route_target(app_model=app_model, current_user=current_user, action=_skill_upload_response) + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def post(self, session: Session, current_user: Account, app_model: App): + return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=_skill_upload_response + ) @console_ns.route("/agent//config/skills") @@ -731,8 +750,10 @@ class AgentConfigSkillsByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -748,10 +769,13 @@ class AgentConfigSkillsApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App): - return _with_app_route_target(app_model=app_model, current_user=current_user, action=_skill_list_response) + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App): + return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=_skill_list_response + ) @console_ns.route("/agent//config/files") @@ -764,8 +788,10 @@ class AgentConfigFilesByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -781,9 +807,11 @@ class AgentConfigFilesByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account, agent_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): payload = AgentConfigFileUploadPayload.model_validate(console_ns.payload or {}) return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -799,10 +827,13 @@ class AgentConfigFilesApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App): - return _with_app_route_target(app_model=app_model, current_user=current_user, action=_file_list_response) + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App): + return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=_file_list_response + ) @console_ns.doc("upload_agent_config_file") @console_ns.doc(params={"app_id": "Application ID", **query_params_from_model(AgentConfigQuery)}) @@ -813,11 +844,13 @@ class AgentConfigFilesApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def post(self, current_user: Account, app_model: App): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def post(self, session: Session, current_user: Account, app_model: App): payload = AgentConfigFileUploadPayload.model_validate(console_ns.payload or {}) return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _file_upload_response(target, payload), @@ -836,8 +869,10 @@ class AgentConfigSkillInspectByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -855,10 +890,12 @@ class AgentConfigSkillInspectApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _skill_inspect_response(target, name), @@ -883,10 +920,12 @@ class AgentConfigSkillFilePreviewByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): query = query_params_from_request(AgentConfigSkillFileByAgentQuery) try: target = _resolve_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -913,12 +952,15 @@ class AgentConfigSkillFilePreviewApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): query = query_params_from_request(AgentConfigSkillFileQuery) try: - target = _resolve_app_route_target(app_model=app_model, current_user=current_user, query=query) + target = _resolve_app_route_target( + session=session, app_model=app_model, current_user=current_user, query=query + ) if isinstance(target, tuple): return target return _skill_file_preview_response(target, name, query.path) @@ -938,8 +980,10 @@ class AgentConfigSkillDownloadByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -957,10 +1001,12 @@ class AgentConfigSkillDownloadApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _skill_download_response(target, name), @@ -983,10 +1029,12 @@ class AgentConfigSkillFileDownloadByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): query = query_params_from_request(AgentConfigSkillFileByAgentQuery) try: target = _resolve_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1017,12 +1065,15 @@ class AgentConfigSkillFileDownloadApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): query = query_params_from_request(AgentConfigSkillFileQuery) try: - target = _resolve_app_route_target(app_model=app_model, current_user=current_user, query=query) + target = _resolve_app_route_target( + session=session, app_model=app_model, current_user=current_user, query=query + ) if isinstance(target, tuple): return target return _skill_file_download_response( @@ -1047,10 +1098,12 @@ class AgentConfigSkillFileDownloadContentByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): query = query_params_from_request(AgentConfigSkillFileByAgentQuery) try: target = _resolve_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1069,12 +1122,15 @@ class AgentConfigSkillFileDownloadContentApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): query = query_params_from_request(AgentConfigSkillFileQuery) try: - target = _resolve_app_route_target(app_model=app_model, current_user=current_user, query=query) + target = _resolve_app_route_target( + session=session, app_model=app_model, current_user=current_user, query=query + ) if isinstance(target, tuple): return target return _skill_file_raw_download_response(target, name, query.path) @@ -1094,8 +1150,10 @@ class AgentConfigSkillByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1115,10 +1173,12 @@ class AgentConfigSkillApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def delete(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def delete(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _skill_delete_response(target, name), @@ -1137,8 +1197,10 @@ class AgentConfigFilePreviewByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1156,10 +1218,12 @@ class AgentConfigFilePreviewApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _file_preview_response(target, name), @@ -1178,8 +1242,10 @@ class AgentConfigFileDownloadByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1197,10 +1263,12 @@ class AgentConfigFileDownloadApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def get(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def get(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _file_download_response(target, name), @@ -1219,8 +1287,10 @@ class AgentConfigFileByAgentApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def delete(self, tenant_id: str, current_user: Account, agent_id: UUID, name: str): + @with_session + def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID, name: str): return _with_agent_route_target( + session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, @@ -1240,10 +1310,12 @@ class AgentConfigFileApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - @get_app_model(mode=_WORKFLOW_APP_MODES) @with_current_user - def delete(self, current_user: Account, app_model: App, name: str): + @with_session + @get_app_model(mode=_WORKFLOW_APP_MODES) + def delete(self, session: Session, current_user: Account, app_model: App, name: str): return _with_app_route_target( + session=session, app_model=app_model, current_user=current_user, action=lambda target: _file_delete_response(target, name), diff --git a/api/controllers/console/app/agent_drive_inspector.py b/api/controllers/console/app/agent_drive_inspector.py index 5166393b3d9..e682953c015 100644 --- a/api/controllers/console/app/agent_drive_inspector.py +++ b/api/controllers/console/app/agent_drive_inspector.py @@ -18,17 +18,18 @@ from uuid import UUID from flask import Response from flask_restx import Resource from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, query_params_from_request, register_response_schema_models, ) +from controllers.common.session import with_session 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 @@ -144,13 +145,13 @@ register_response_schema_models( ) -def _resolve_agent_id(app_model: App, node_id: str | None) -> str | None: +def _resolve_agent_id(session: Session, 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, session=db.session() + session=session, tenant_id=app_model.tenant_id, app_id=app_model.id, node_id=node_id ) - return app_model.bound_agent_id + return app_model.bound_agent_id_with_session(session=session) def _agent_not_bound() -> tuple[dict[str, object], int]: @@ -181,12 +182,13 @@ class AgentDriveListByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): query = query_params_from_request(AgentDriveListByAgentQuery) - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) try: items = AgentDriveService().manifest( - tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix, session=db.session() + tenant_id=tenant_id, agent_id=str(agent_id), prefix=query.prefix, session=session ) except AgentDriveError as exc: return _handle(exc) @@ -203,10 +205,11 @@ class AgentDriveSkillListByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): + resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) try: - items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id), session=db.session()) + items = AgentDriveService().list_skills(tenant_id=tenant_id, agent_id=str(agent_id), session=session) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -222,15 +225,16 @@ class AgentDriveSkillInspectByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID, skill_path: str): - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID, skill_path: str): + resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) try: return _json_response( AgentDriveService().inspect_skill( tenant_id=tenant_id, agent_id=str(agent_id), skill_path=skill_path, - session=db.session(), + session=session, ) ) except AgentDriveError as exc: @@ -247,12 +251,13 @@ class AgentDrivePreviewByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): query = query_params_from_request(AgentDriveFileByAgentQuery) - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) try: return AgentDriveService().preview( - tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=db.session() + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=session ) except AgentDriveError as exc: return _handle(exc) @@ -268,12 +273,13 @@ class AgentDriveDownloadByAgentApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, agent_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id: str, agent_id: UUID): query = query_params_from_request(AgentDriveFileByAgentQuery) - resolve_agent_runtime_app_model(tenant_id=tenant_id, agent_id=agent_id) + resolve_agent_runtime_app_model(session=session, 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, session=db.session() + tenant_id=tenant_id, agent_id=str(agent_id), key=query.key, session=session ) except AgentDriveError as exc: return _handle(exc) @@ -289,15 +295,16 @@ class AgentDriveListApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_APP_MODES) - def get(self, app_model: App): + def get(self, session: Session, app_model: App): query = query_params_from_request(AgentDriveListQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) 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, session=db.session() + tenant_id=app_model.tenant_id, agent_id=agent_id, prefix=query.prefix, session=session ) except AgentDriveError as exc: return _handle(exc) @@ -315,16 +322,15 @@ class AgentDriveSkillListApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_APP_MODES) - def get(self, app_model: App): + def get(self, session: Session, app_model: App): query = query_params_from_request(AgentDriveListQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) if not agent_id: return _agent_not_bound() try: - items = AgentDriveService().list_skills( - tenant_id=app_model.tenant_id, agent_id=agent_id, session=db.session() - ) + items = AgentDriveService().list_skills(tenant_id=app_model.tenant_id, agent_id=agent_id, session=session) except AgentDriveError as exc: return _handle(exc) return {"items": items} @@ -345,10 +351,11 @@ class AgentDriveSkillInspectApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_APP_MODES) - def get(self, app_model: App, skill_path: str): + def get(self, session: Session, app_model: App, skill_path: str): query = query_params_from_request(AgentDriveSkillInspectQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) if not agent_id: return _agent_not_bound() try: @@ -357,7 +364,7 @@ class AgentDriveSkillInspectApi(Resource): tenant_id=app_model.tenant_id, agent_id=agent_id, skill_path=skill_path, - session=db.session(), + session=session, ) ) except AgentDriveError as exc: @@ -373,15 +380,16 @@ class AgentDrivePreviewApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_APP_MODES) - def get(self, app_model: App): + def get(self, session: Session, app_model: App): query = query_params_from_request(AgentDriveFileQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) 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, session=db.session() + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=session ) except AgentDriveError as exc: return _handle(exc) @@ -396,15 +404,16 @@ class AgentDriveDownloadApi(Resource): @setup_required @login_required @account_initialization_required + @with_session(write=False) @get_app_model(mode=_WORKFLOW_APP_MODES) - def get(self, app_model: App): + def get(self, session: Session, app_model: App): query = query_params_from_request(AgentDriveFileQuery) - agent_id = _resolve_agent_id(app_model, query.node_id) + agent_id = _resolve_agent_id(session, app_model, query.node_id) 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, session=db.session() + tenant_id=app_model.tenant_id, agent_id=agent_id, key=query.key, session=session ) except AgentDriveError as exc: return _handle(exc) diff --git a/api/controllers/console/app/annotation.py b/api/controllers/console/app/annotation.py index 961f9e2f1d8..6653a6e288c 100644 --- a/api/controllers/console/app/annotation.py +++ b/api/controllers/console/app/annotation.py @@ -5,10 +5,12 @@ from flask import abort, request from flask_restx import Resource from pydantic import BaseModel, Field, TypeAdapter, field_validator from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.common.errors import NoFileUploadedError, TooManyFilesError from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.wraps import ( RBACPermission, @@ -21,7 +23,6 @@ from controllers.console.wraps import ( rbac_permission_required, setup_required, ) -from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.annotation_fields import ( Annotation, @@ -46,9 +47,9 @@ from services.annotation_service import ( from services.app_ref_service import AppRef, AppRefService -def _get_app_ref(app_id: str) -> AppRef: +def _get_app_ref(session: Session, app_id: str) -> AppRef: _, 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 app is None: @@ -210,8 +211,9 @@ class AppAnnotationSettingDetailApi(Resource): @account_initialization_required @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), session=db.session()) + @with_session(write=False) + def get(self, session: Session, app_id: UUID): + result = AppAnnotationService.get_app_annotation_setting_by_app_id(str(app_id), session) return dump_response(AnnotationSettingResponse, result), 200 @@ -228,14 +230,15 @@ class AppAnnotationSettingUpdateApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - def post(self, app_id: UUID, annotation_setting_id: UUID): + @with_session + def post(self, session: Session, app_id: UUID, annotation_setting_id: UUID): annotation_setting_id_str = str(annotation_setting_id) args = AnnotationSettingUpdatePayload.model_validate(console_ns.payload) setting_args: UpdateAnnotationSettingArgs = {"score_threshold": args.score_threshold} result = AppAnnotationService.update_app_annotation_setting( - str(app_id), annotation_setting_id_str, setting_args, session=db.session() + str(app_id), annotation_setting_id_str, setting_args, session ) return dump_response(AnnotationSettingResponse, result), 200 @@ -286,14 +289,15 @@ class AnnotationApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) - def get(self, app_id: UUID): + @with_session(write=False) + def get(self, session: Session, app_id: UUID): args = AnnotationListQuery.model_validate(request.args.to_dict(flat=True)) page = args.page limit = args.limit keyword = args.keyword annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - str(app_id), page, limit, keyword, session=db.session() + str(app_id), page, limit, keyword, session ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return AnnotationList( @@ -312,7 +316,8 @@ class AnnotationApi(Resource): @cloud_edition_billing_resource_check("annotation") @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - def post(self, app_id: UUID): + @with_session + def post(self, session: Session, app_id: UUID): args = CreateAnnotationPayload.model_validate(console_ns.payload) upsert_args: UpsertAnnotationArgs = {} if args.answer is not None: @@ -323,9 +328,7 @@ 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), session=db.session() - ) + annotation = AppAnnotationService.up_insert_app_annotation_from_message(upsert_args, str(app_id), session) return dump_response(Annotation, annotation), 201 @setup_required @@ -334,7 +337,8 @@ class AnnotationApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @console_ns.response(204, "Annotations deleted successfully") - def delete(self, app_id: UUID): + @with_session + def delete(self, session: Session, app_id: UUID): # Use request.args.getlist to get annotation_ids array directly annotation_ids = request.args.getlist("annotation_id") @@ -348,12 +352,12 @@ class AnnotationApi(Resource): "message": "annotation_ids are required if the parameter is provided.", }, 400 - app_ref = _get_app_ref(str(app_id)) - AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids, session=db.session()) + app_ref = _get_app_ref(session, str(app_id)) + AppAnnotationService.delete_app_annotations_in_batch(app_ref, annotation_ids, session) return "", 204 # If no annotation_ids are provided, handle clearing all annotations else: - AppAnnotationService.clear_all_annotations(str(app_id), session=db.session()) + AppAnnotationService.clear_all_annotations(str(app_id), session) return "", 204 @@ -373,8 +377,9 @@ class AnnotationExportApi(Resource): @account_initialization_required @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), session=db.session()) + @with_session(write=False) + def get(self, session: Session, app_id: UUID): + annotation_list = AppAnnotationService.export_annotation_list_by_app_id(str(app_id), session) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return ( AnnotationExportList(data=annotation_models).model_dump(mode="json"), @@ -401,16 +406,17 @@ class AnnotationUpdateDeleteApi(Resource): @cloud_edition_billing_resource_check("annotation") @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - def post(self, app_id: UUID, annotation_id: UUID): + @with_session + def post(self, session: Session, app_id: UUID, annotation_id: UUID): args = UpdateAnnotationPayload.model_validate(console_ns.payload) update_args: UpdateAnnotationArgs = {} if args.answer is not None: update_args["answer"] = args.answer if args.question is not None: update_args["question"] = args.question - app_ref = _get_app_ref(str(app_id)) + app_ref = _get_app_ref(session, 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, session) return Annotation.model_validate(annotation, from_attributes=True).model_dump(mode="json") @setup_required @@ -419,10 +425,11 @@ class AnnotationUpdateDeleteApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @console_ns.response(204, "Annotation deleted successfully") - def delete(self, app_id: UUID, annotation_id: UUID): - app_ref = _get_app_ref(str(app_id)) + @with_session + def delete(self, session: Session, app_id: UUID, annotation_id: UUID): + app_ref = _get_app_ref(session, 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, session) return "", 204 @@ -446,7 +453,8 @@ class AnnotationBatchImportApi(Resource): @annotation_import_concurrency_limit @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - def post(self, app_id: UUID): + @with_session + def post(self, session: Session, app_id: UUID): from configs import dify_config # check file @@ -481,7 +489,7 @@ class AnnotationBatchImportApi(Resource): return dump_response( AnnotationBatchImportResponse, - AppAnnotationService.batch_import_app_annotations(str(app_id), file, session=db.session()), + AppAnnotationService.batch_import_app_annotations(str(app_id), file, session), ) @@ -533,16 +541,17 @@ class AnnotationHitHistoryListApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) - def get(self, app_id: UUID, annotation_id: UUID): + @with_session(write=False) + def get(self, session: Session, app_id: UUID, annotation_id: UUID): page = request.args.get("page", default=1, type=int) limit = request.args.get("limit", default=20, type=int) - app_ref = _get_app_ref(str(app_id)) + app_ref = _get_app_ref(session, str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) annotation_hit_history_list, total = AppAnnotationService.get_annotation_hit_histories( annotation_ref, page, limit, - session=db.session(), + 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 07820e0a31b..335aefdaf2d 100644 --- a/api/controllers/console/app/app.py +++ b/api/controllers/console/app/app.py @@ -6,7 +6,7 @@ from typing import Any, Literal from flask import request from flask_restx import Resource -from pydantic import AliasChoices, BaseModel, Field, computed_field, field_validator +from pydantic import AliasChoices, BaseModel, Field, ValidationInfo, computed_field, field_validator, model_validator from sqlalchemy import select from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, NotFound @@ -52,7 +52,14 @@ from libs.login import login_required from models import Account, App, DatasetPermissionEnum, Workflow from models.model import IconType from services.app_dsl_service import AppDslService -from services.app_service import AppListParams, AppListSortBy, AppService, CreateAppParams, StarredAppListParams +from services.app_service import ( + AppListParams, + AppListSortBy, + AppResponseView, + AppService, + CreateAppParams, + StarredAppListParams, +) from services.enterprise import rbac_service as enterprise_rbac_service from services.enterprise.enterprise_service import EnterpriseService from services.entities.dsl_entities import DslImportWarning, ImportMode, ImportStatus @@ -348,7 +355,18 @@ class DeletedTool(ResponseModel): provider_id: str -class AppPartial(ResponseModel): +class AppResponseModel(ResponseModel): + @model_validator(mode="before") + @classmethod + def _use_request_session(cls, value: Any, info: ValidationInfo) -> Any: + if not isinstance(value, App): + return value + if info.context is None or "session" not in info.context: + raise ValueError("session context is required to serialize an App") + return AppResponseView(value, session=info.context["session"]) + + +class AppPartial(AppResponseModel): id: str name: str max_active_requests: int | None = None @@ -392,7 +410,7 @@ class AppPartial(ResponseModel): return to_timestamp(value) -class AppDetail(ResponseModel): +class AppDetail(AppResponseModel): id: str name: str description: str | None = None @@ -587,12 +605,13 @@ class AppListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user_id, - session=db.session(), + session=session, ) if dify_config.RBAC_ENABLED: access_filter = resolve_app_access_filter( str(current_tenant_id), current_user_id, + session=session, permissions=permissions, ) access_filter.apply_to_params(params) @@ -608,7 +627,11 @@ class AppListApi(Resource): permission_keys_map = permissions.app.permission_keys_by_resource_ids(app_ids) _enrich_app_list_items(session, apps=app_pagination.items, tenant_id=current_tenant_id) - pagination_model = AppPagination.model_validate(app_pagination, from_attributes=True) + pagination_model = AppPagination.model_validate( + app_pagination, + from_attributes=True, + context={"session": session}, + ) if app_pagination.items: pagination_model = pagination_model.model_copy( update={ @@ -634,7 +657,8 @@ class AppListApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account): + @with_session + def post(self, session: Session, current_tenant_id: str, current_user: Account): """Create app""" args = CreateAppPayload.model_validate(console_ns.payload) params = CreateAppParams( @@ -647,7 +671,7 @@ class AppListApi(Resource): ) app_service = AppService() - app = app_service.create_app(current_tenant_id, params, current_user, session=db.session()) + app = app_service.create_app(current_tenant_id, params, current_user, session=session) if dify_config.RBAC_ENABLED: enterprise_rbac_service.RBACService.AppAccess.replace_whitelist( tenant_id=str(current_tenant_id), @@ -660,11 +684,13 @@ class AppListApi(Resource): 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), [])} + session=session, ) + app_detail = AppDetailWithSite.model_validate( + app, + from_attributes=True, + context={"session": session}, + ).model_copy(update={"permission_keys": permission_keys_map.get(str(app.id), [])}) return app_detail.model_dump(mode="json"), 201 @@ -700,7 +726,14 @@ class StarredAppListApi(Resource): return empty.model_dump(mode="json"), 200 _enrich_app_list_items(session, apps=app_pagination.items, tenant_id=current_tenant_id) - return AppPagination.model_validate(app_pagination, from_attributes=True).model_dump(mode="json"), 200 + return ( + AppPagination.model_validate( + app_pagination, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json"), + 200, + ) @console_ns.route("/apps//star") @@ -751,12 +784,13 @@ class AppApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=None) - def get(self, current_tenant_id: str, current_user: Account, app_model: App): + def get(self, session: Session, current_tenant_id: str, current_user: Account, app_model: App): """Get app detail""" app_service = AppService() - app_model = app_service.get_app(app_model) + app_model = app_service.get_app(app_model, session=session) if FeatureService.get_system_features().webapp_auth.enabled: app_setting = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(app_id=str(app_model.id)) @@ -766,13 +800,15 @@ class AppApi(Resource): str(current_tenant_id), current_user.id, app_id=str(app_model.id), - session=db.session(), + session=session, ) permission_keys_map = permissions.app.permission_keys_by_resource_ids([str(app_model.id)]) - response_model = AppDetailWithSite.model_validate(app_model, from_attributes=True).model_copy( - update={"permission_keys": permission_keys_map.get(str(app_model.id), [])} - ) + response_model = AppDetailWithSite.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_copy(update={"permission_keys": permission_keys_map.get(str(app_model.id), [])}) return response_model.model_dump(mode="json") @console_ns.doc("update_app") @@ -787,8 +823,9 @@ class AppApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) + @with_session @get_app_model(mode=None) - def put(self, app_model: App): + def put(self, session: Session, app_model: App): """Update app""" args = UpdateAppPayload.model_validate(console_ns.payload) @@ -803,8 +840,12 @@ 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, session=db.session()) - return dump_response(AppDetailWithSite, app_model) + app_model = app_service.update_app(app_model, args_dict, session=session) + return AppDetailWithSite.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") @console_ns.doc("delete_app") @console_ns.doc(description="Delete application") @@ -816,11 +857,12 @@ class AppApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_DELETE) + @with_session @get_app_model - def delete(self, app_model: App): + def delete(self, session: Session, app_model: App): """Delete app""" app_service = AppService() - app_service.delete_app(app_model, session=db.session()) + app_service.delete_app(app_model, session=session) return "", 204 @@ -883,20 +925,21 @@ class AppCopyApi(Resource): stmt = select(App).where(App.id == result.app_id) app = session.scalar(stmt) + if not app: + raise NotFound("App not found") - if not app: - raise NotFound("App not found") - - permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( - 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), [])} - ) - return response_model.model_dump(mode="json"), 201 + permission_keys_map = enterprise_rbac_service.RBACService.AppPermissions.batch_get( + str(current_tenant_id), + current_user.id, + [str(app.id)], + session=session, + ) + response_model = AppDetailWithSite.model_validate( + app, + from_attributes=True, + context={"session": session}, + ).model_copy(update={"permission_keys": permission_keys_map.get(str(app.id), [])}) + return response_model.model_dump(mode="json"), 201 @console_ns.route("/apps//export") @@ -966,13 +1009,18 @@ class AppNameApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) + @with_session @get_app_model(mode=None) - def post(self, app_model: App): + def post(self, session: Session, app_model: App): args = AppNamePayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_name(app_model, args.name, session=db.session()) - return dump_response(AppDetail, app_model) + app_model = app_service.update_app_name(app_model, args.name, session=session) + return AppDetail.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") @console_ns.route("/apps//icon") @@ -988,8 +1036,9 @@ class AppIconApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) + @with_session @get_app_model(mode=None) - def post(self, app_model: App): + def post(self, session: Session, app_model: App): args = AppIconPayload.model_validate(console_ns.payload or {}) app_service = AppService() @@ -998,9 +1047,13 @@ class AppIconApi(Resource): args.icon or "", args.icon_background or "", args.icon_type, - session=db.session(), + session=session, ) - return dump_response(AppDetail, app_model) + return AppDetail.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") @console_ns.route("/apps//site-enable") @@ -1016,13 +1069,18 @@ class AppSiteStatus(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) + @with_session @get_app_model(mode=None) - def post(self, app_model: App): + def post(self, session: Session, app_model: App): args = AppSiteStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_site_status(app_model, args.enable_site, session=db.session()) - return dump_response(AppDetail, app_model) + app_model = app_service.update_app_site_status(app_model, args.enable_site, session=session) + return AppDetail.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") @console_ns.route("/apps//api-enable") @@ -1038,13 +1096,18 @@ class AppApiStatus(Resource): @is_admin_or_owner_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) + @with_session @get_app_model(mode=None) - def post(self, app_model: App): + def post(self, session: Session, app_model: App): args = AppApiStatusPayload.model_validate(console_ns.payload) app_service = AppService() - app_model = app_service.update_app_api_status(app_model, args.enable_api, session=db.session()) - return dump_response(AppDetail, app_model) + app_model = app_service.update_app_api_status(app_model, args.enable_api, session=session) + return AppDetail.model_validate( + app_model, + from_attributes=True, + context={"session": session}, + ).model_dump(mode="json") @console_ns.route("/apps//trace") diff --git a/api/controllers/console/app/audio.py b/api/controllers/console/app/audio.py index 06ef0dc7acb..059cc96269d 100644 --- a/api/controllers/console/app/audio.py +++ b/api/controllers/console/app/audio.py @@ -127,6 +127,7 @@ def _transcribe_audio_to_text( *, app_model: App, file: FileStorage | None, + session: Session, agent_soul: AgentSoulConfig | None = None, ) -> dict[str, str]: try: @@ -134,6 +135,7 @@ def _transcribe_audio_to_text( response = AudioService.transcript_asr( app_model=app_model, file=file, + session=session, end_user=None, ) else: @@ -141,6 +143,7 @@ def _transcribe_audio_to_text( app_model=app_model, agent_soul=agent_soul, file=file, + session=session, end_user=None, ) return dump_response(AudioTranscriptResponse, response) @@ -194,7 +197,11 @@ class ChatMessageAudioApi(Resource): @account_initialization_required @get_app_model(mode=_CONSOLE_AUDIO_TRANSCRIPT_APP_MODES) def post(self, app_model: App): - return _transcribe_audio_to_text(app_model=app_model, file=request.files.get("file")) + return _transcribe_audio_to_text( + app_model=app_model, + file=request.files.get("file"), + session=db.session(), + ) @console_ns.route("/agent//audio-to-text") @@ -228,7 +235,11 @@ class AgentChatMessageAudioApi(Resource): agent_id: UUID, ): payload = AgentAudioTranscriptFormPayload.model_validate(request.form.to_dict(flat=True)) - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) # Agent routes expose Agent ids, while APP RBAC is keyed by the resolved runtime App id. enforce_rbac_access( tenant_id=current_tenant_id, @@ -248,6 +259,7 @@ class AgentChatMessageAudioApi(Resource): app_model=app_model, agent_soul=agent_soul, file=request.files.get("file"), + session=session, ) diff --git a/api/controllers/console/app/completion.py b/api/controllers/console/app/completion.py index 62cf38b86d8..ae76fee38d9 100644 --- a/api/controllers/console/app/completion.py +++ b/api/controllers/console/app/completion.py @@ -44,7 +44,6 @@ from core.errors.error import ( QuotaExceededError, ) from core.helper.trace_id_helper import get_external_trace_id -from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs import helper from libs.helper import uuid_value @@ -158,8 +157,8 @@ class CompletionMessageApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - @get_app_model(mode=AppMode.COMPLETION) @with_session + @get_app_model(mode=AppMode.COMPLETION) def post(self, session: Session, current_user: Account, app_model: App): args_model = CompletionMessagePayload.model_validate(console_ns.payload) args = args_model.model_dump(exclude_none=True, by_alias=True) @@ -240,8 +239,8 @@ class ChatMessageApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.AGENT]) @with_session + @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.AGENT]) def post(self, session: Session, current_tenant_id: str, current_user: Account, app_model: App): return _create_chat_message( session=session, current_tenant_id=current_tenant_id, current_user=current_user, app_model=app_model @@ -266,7 +265,9 @@ class AgentChatMessageApi(Resource): @with_current_tenant_id @with_session def post(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) + app_model = AgentRosterService(session).get_agent_runtime_app_model( + tenant_id=current_tenant_id, agent_id=str(agent_id) + ) return _create_chat_message( session=session, current_tenant_id=current_tenant_id, @@ -293,7 +294,9 @@ class AgentBuildChatFinalizeApi(Resource): @with_current_tenant_id @with_session def post(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) + app_model = AgentRosterService(session).get_agent_runtime_app_model( + tenant_id=current_tenant_id, agent_id=str(agent_id) + ) return _create_build_chat_finalization_message( session=session, current_tenant_id=current_tenant_id, @@ -329,15 +332,20 @@ class AgentChatMessageStopApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, current_tenant_id: str, current_user_id: str, agent_id: UUID, task_id: str): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) + @with_session(write=False) + def post(self, session: Session, current_tenant_id: str, current_user_id: str, agent_id: UUID, task_id: str): + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) return _stop_chat_message(current_user_id=current_user_id, app_model=app_model, task_id=task_id) def _resolve_current_user_agent_debug_conversation_id( - *, current_tenant_id: str, current_user: Account, app_model: App, agent_id: str | None + *, session: Session, current_tenant_id: str, current_user: Account, app_model: App, agent_id: str | None ) -> str: - roster_service = AgentRosterService(db.session) + roster_service = AgentRosterService(session) if agent_id: return roster_service.get_or_create_agent_app_debug_conversation_id( tenant_id=current_tenant_id, @@ -369,6 +377,7 @@ def _create_chat_message( if AppMode.value_of(app_model.mode) == AppMode.AGENT: debug_conversation_id = _resolve_current_user_agent_debug_conversation_id( + session=session, current_tenant_id=current_tenant_id or app_model.tenant_id, current_user=current_user, app_model=app_model, @@ -404,6 +413,7 @@ def _create_build_chat_finalization_message( *, session: Session, current_user: Account, app_model: App, current_tenant_id: str, agent_id: str ): debug_conversation_id = _resolve_current_user_agent_debug_conversation_id( + session=session, current_tenant_id=current_tenant_id, current_user=current_user, app_model=app_model, diff --git a/api/controllers/console/app/conversation.py b/api/controllers/console/app/conversation.py index b7d422d30b1..14ddb2446b8 100644 --- a/api/controllers/console/app/conversation.py +++ b/api/controllers/console/app/conversation.py @@ -6,10 +6,11 @@ from flask import abort, request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import func, or_ -from sqlalchemy.orm import selectinload +from sqlalchemy.orm import Session, selectinload from werkzeug.exceptions import NotFound from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.app.wraps import get_app_model from controllers.console.wraps import ( @@ -22,7 +23,6 @@ from controllers.console.wraps import ( with_current_user, ) from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from fields.conversation_fields import ( Conversation as ConversationResponse, ) @@ -35,6 +35,7 @@ from fields.conversation_fields import ( from fields.conversation_fields import ( ConversationPagination as ConversationPaginationResponse, ) +from fields.conversation_fields import ConversationResponseSource from fields.conversation_fields import ( ConversationWithSummaryPagination as ConversationWithSummaryPaginationResponse, ) @@ -105,8 +106,9 @@ class CompletionConversationApi(Resource): @edit_permission_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=AppMode.COMPLETION) - def get(self, current_user: Account, app_model: App): + def get(self, session: Session, current_user: Account, app_model: App): args = CompletionConversationQuery.model_validate(request.args.to_dict(flat=True)) query = sa.select(Conversation).where( @@ -157,9 +159,18 @@ class CompletionConversationApi(Resource): query = query.order_by(Conversation.created_at.desc()) - conversations = paginate_query(query, page=args.page, per_page=args.limit) + conversations = paginate_query(query, session=session, page=args.page, per_page=args.limit) - return dump_response(ConversationPaginationResponse, conversations) + return dump_response( + ConversationPaginationResponse, + { + "page": conversations.page, + "per_page": conversations.per_page, + "total": conversations.total, + "has_next": conversations.has_next, + "items": [ConversationResponseSource(item, session=session) for item in conversations.items], + }, + ) @console_ns.route("/apps//completion-conversations/") @@ -176,11 +187,15 @@ class CompletionConversationDetailApi(Resource): @edit_permission_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session @get_app_model(mode=AppMode.COMPLETION) - def get(self, current_user: Account, app_model: App, conversation_id: UUID): + def get(self, session: Session, current_user: Account, app_model: App, conversation_id: UUID): conversation_id_str = str(conversation_id) return dump_response( - ConversationMessageDetailResponse, _get_conversation(current_user, app_model, conversation_id_str) + ConversationMessageDetailResponse, + ConversationResponseSource( + _get_conversation(session, current_user, app_model, conversation_id_str), session=session + ), ) @console_ns.doc("delete_completion_conversation") @@ -195,12 +210,13 @@ class CompletionConversationDetailApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_user + @with_session @get_app_model(mode=AppMode.COMPLETION) - def delete(self, current_user: Account, app_model: App, conversation_id: UUID): + def delete(self, session: Session, current_user: Account, app_model: App, conversation_id: UUID): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) + ConversationService.delete(app_model, conversation_id_str, current_user, session=session) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") @@ -220,8 +236,9 @@ class ChatConversationApi(Resource): @edit_permission_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, current_user: Account, app_model: App): + def get(self, session: Session, current_user: Account, app_model: App): args = ChatConversationQuery.model_validate(request.args.to_dict(flat=True)) subquery = ( @@ -311,9 +328,18 @@ class ChatConversationApi(Resource): case _: query = query.order_by(Conversation.created_at.desc()) - conversations = paginate_query(query, page=args.page, per_page=args.limit) + conversations = paginate_query(query, session=session, page=args.page, per_page=args.limit) - return dump_response(ConversationWithSummaryPaginationResponse, conversations) + return dump_response( + ConversationWithSummaryPaginationResponse, + { + "page": conversations.page, + "per_page": conversations.per_page, + "total": conversations.total, + "has_next": conversations.has_next, + "items": [ConversationResponseSource(item, session=session) for item in conversations.items], + }, + ) @console_ns.route("/apps//chat-conversations/") @@ -330,11 +356,15 @@ class ChatConversationDetailApi(Resource): @edit_permission_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, current_user: Account, app_model: App, conversation_id: UUID): + def get(self, session: Session, current_user: Account, app_model: App, conversation_id: UUID): conversation_id_str = str(conversation_id) return dump_response( - ConversationDetailResponse, _get_conversation(current_user, app_model, conversation_id_str) + ConversationDetailResponse, + ConversationResponseSource( + _get_conversation(session, current_user, app_model, conversation_id_str), session=session + ), ) @console_ns.doc("delete_chat_conversation") @@ -349,27 +379,28 @@ class ChatConversationDetailApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_user + @with_session @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def delete(self, current_user: Account, app_model: App, conversation_id: UUID): + def delete(self, session: Session, current_user: Account, app_model: App, conversation_id: UUID): conversation_id_str = str(conversation_id) try: - ConversationService.delete(app_model, conversation_id_str, current_user, session=db.session()) + ConversationService.delete(app_model, conversation_id_str, current_user, session=session) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") return "", 204 -def _get_conversation(current_user: Account, app_model, conversation_id): - conversation = db.session.scalar( +def _get_conversation(session: Session, current_user: Account, app_model, conversation_id): + conversation = session.scalar( sa.select(Conversation).where(Conversation.id == conversation_id, Conversation.app_id == app_model.id).limit(1) ) if not conversation: raise NotFound("Conversation Not Exists.") - db.session.execute( + session.execute( sa.update(Conversation) .where(Conversation.id == conversation_id, Conversation.read_at.is_(None)) # Keep updated_at unchanged when only marking a conversation as read. @@ -379,7 +410,7 @@ def _get_conversation(current_user: Account, app_model, conversation_id): updated_at=Conversation.updated_at, ) ) - db.session.commit() - db.session.refresh(conversation) + session.flush() + session.refresh(conversation) return conversation diff --git a/api/controllers/console/app/message.py b/api/controllers/console/app/message.py index 958b356de94..080772107dc 100644 --- a/api/controllers/console/app/message.py +++ b/api/controllers/console/app/message.py @@ -6,11 +6,13 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import exists, func, select +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError, NotFound from controllers.common.controller_schemas import MessageFeedbackPayload as _MessageFeedbackPayloadBase from controllers.common.fields import SimpleResultResponse, TextFileResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model from controllers.console.app.error import ( @@ -39,6 +41,7 @@ from fields.base import ResponseModel from fields.conversation_fields import ( MessageDetail as BaseMessageDetailResponse, ) +from fields.conversation_fields import MessageResponseSource from graphon.model_runtime.errors.invoke import InvokeError from libs.helper import dump_response, uuid_value from libs.infinite_scroll_pagination import InfiniteScrollPagination @@ -152,9 +155,10 @@ class ChatMessageListApi(Resource): @edit_permission_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, current_user: Account, app_model: App): - return _list_chat_messages(app_model=app_model, current_user=current_user) + def get(self, session: Session, current_user: Account, app_model: App): + return _list_chat_messages(session=session, app_model=app_model, current_user=current_user) @console_ns.route("/agent//chat-messages") @@ -172,9 +176,14 @@ class AgentChatMessageListApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) - return _list_chat_messages(app_model=app_model, current_user=current_user) + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) + return _list_chat_messages(session=session, app_model=app_model, current_user=current_user) @console_ns.route("/apps//feedbacks") @@ -190,9 +199,10 @@ class MessageFeedbackApi(Resource): @login_required @account_initialization_required @with_current_user + @with_session @get_app_model - def post(self, current_user: Account, app_model: App): - return _update_message_feedback(current_user=current_user, app_model=app_model) + def post(self, session: Session, current_user: Account, app_model: App): + return _update_message_feedback(session=session, current_user=current_user, app_model=app_model) @console_ns.route("/agent//feedbacks") @@ -208,9 +218,14 @@ class AgentMessageFeedbackApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) - return _update_message_feedback(current_user=current_user, app_model=app_model) + @with_session + def post(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID): + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) + return _update_message_feedback(session=session, current_user=current_user, app_model=app_model) @console_ns.route("/apps//annotations/count") @@ -252,9 +267,12 @@ class MessageSuggestedQuestionApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, current_user: Account, app_model: App, message_id: UUID): - return _get_message_suggested_questions(current_user=current_user, app_model=app_model, message_id=message_id) + def get(self, session: Session, current_user: Account, app_model: App, message_id: UUID): + return _get_message_suggested_questions( + session=session, current_user=current_user, app_model=app_model, message_id=message_id + ) @console_ns.route("/agent//chat-messages//suggested-questions") @@ -273,9 +291,16 @@ class AgentMessageSuggestedQuestionApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account, agent_id: UUID, message_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) - return _get_message_suggested_questions(current_user=current_user, app_model=app_model, message_id=message_id) + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account, agent_id: UUID, message_id: UUID): + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) + return _get_message_suggested_questions( + session=session, current_user=current_user, app_model=app_model, message_id=message_id + ) @console_ns.route("/apps//feedbacks/export") @@ -333,9 +358,10 @@ class MessageApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) + @with_session(write=False) @get_app_model - def get(self, app_model: App, message_id: UUID): - return _get_message_detail(app_model=app_model, message_id=message_id) + def get(self, session: Session, app_model: App, message_id: UUID): + return _get_message_detail(session=session, app_model=app_model, message_id=message_id) @console_ns.route("/agent//messages/") @@ -349,12 +375,17 @@ class AgentMessageApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, current_tenant_id: str, agent_id: UUID, message_id: UUID): - app_model = resolve_agent_runtime_app_model(tenant_id=current_tenant_id, agent_id=agent_id) - return _get_message_detail(app_model=app_model, message_id=message_id) + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, agent_id: UUID, message_id: UUID): + app_model = resolve_agent_runtime_app_model( + session=session, + tenant_id=current_tenant_id, + agent_id=agent_id, + ) + return _get_message_detail(session=session, app_model=app_model, message_id=message_id) -def _list_chat_messages(*, app_model: App, current_user: Account | None = None): +def _list_chat_messages(*, session: Session, app_model: App, current_user: Account | None = None): args = ChatMessagesQuery.model_validate(request.args.to_dict()) if AppMode.value_of(app_model.mode) == AppMode.AGENT and current_user is not None: @@ -363,12 +394,12 @@ 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(), + session=session, ) except ConversationNotExistsError: raise NotFound("Conversation Not Exists.") else: - conversation = db.session.scalar( + conversation = session.scalar( select(Conversation) .where(Conversation.id == args.conversation_id, Conversation.app_id == app_model.id) .limit(1) @@ -378,14 +409,14 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): raise NotFound("Conversation Not Exists.") if args.first_id: - first_message = db.session.scalar( + first_message = session.scalar( select(Message).where(Message.conversation_id == conversation.id, Message.id == args.first_id).limit(1) ) if not first_message: raise NotFound("First message not found") - history_messages = db.session.scalars( + history_messages = session.scalars( select(Message) .where( Message.conversation_id == conversation.id, @@ -396,7 +427,7 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): .limit(args.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()) @@ -407,7 +438,7 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): if len(history_messages) == args.limit: current_page_first_message = history_messages[-1] # Check if there are more messages before the current page - has_more = db.session.scalar( + has_more = session.scalar( select( exists().where( Message.conversation_id == conversation.id, @@ -425,26 +456,28 @@ def _list_chat_messages(*, app_model: App, current_user: Account | None = None): return dump_response( MessageInfiniteScrollPaginationResponse, - InfiniteScrollPagination(data=history_messages, limit=args.limit, has_more=has_more), + InfiniteScrollPagination( + data=[MessageResponseSource(message, session=session) for message in history_messages], + limit=args.limit, + has_more=has_more, + ), ) -def _update_message_feedback(*, current_user: Account, app_model: App): +def _update_message_feedback(*, session: Session, current_user: Account, app_model: App): args = MessageFeedbackPayload.model_validate(console_ns.payload) message_id = args.message_id - message = db.session.scalar( - select(Message).where(Message.id == message_id, Message.app_id == app_model.id).limit(1) - ) + message = session.scalar(select(Message).where(Message.id == message_id, Message.app_id == app_model.id).limit(1)) if not message: raise NotFound("Message Not Exists.") - feedback = message.admin_feedback + feedback = message.admin_feedback_with_session(session=session) if not args.rating and feedback: - db.session.delete(feedback) + session.delete(feedback) elif args.rating and feedback: feedback.rating = FeedbackRating(args.rating) feedback.content = args.content @@ -463,14 +496,14 @@ def _update_message_feedback(*, current_user: Account, app_model: App): from_source=FeedbackFromSource.ADMIN, from_account_id=current_user.id, ) - db.session.add(feedback) + session.add(feedback) - db.session.commit() + session.commit() return SimpleResultResponse(result="success").model_dump(mode="json") -def _get_message_suggested_questions(*, current_user: Account, app_model: App, message_id: UUID): +def _get_message_suggested_questions(*, session: Session, current_user: Account, app_model: App, message_id: UUID): message_id_str = str(message_id) try: @@ -479,7 +512,7 @@ def _get_message_suggested_questions(*, current_user: Account, app_model: App, m message_id=message_id_str, user=current_user, invoke_from=InvokeFrom.DEBUGGER, - session=db.session(), + session=session, ) except MessageNotExistsError: raise NotFound("Message not found") @@ -502,10 +535,10 @@ def _get_message_suggested_questions(*, current_user: Account, app_model: App, m return dump_response(SuggestedQuestionsResponse, {"data": questions}) -def _get_message_detail(*, app_model: App, message_id: UUID): +def _get_message_detail(*, session: Session, app_model: App, message_id: UUID): message_id_str = str(message_id) - message = db.session.scalar( + message = session.scalar( select(Message).where(Message.id == message_id_str, Message.app_id == app_model.id).limit(1) ) @@ -513,4 +546,4 @@ def _get_message_detail(*, app_model: App, message_id: UUID): raise NotFound("Message Not Exists.") attach_message_extra_contents([message]) - return dump_response(MessageDetailResponse, message) + return dump_response(MessageDetailResponse, MessageResponseSource(message, session=session)) diff --git a/api/controllers/console/app/model_config.py b/api/controllers/console/app/model_config.py index 3a016e3b9b2..15298366414 100644 --- a/api/controllers/console/app/model_config.py +++ b/api/controllers/console/app/model_config.py @@ -4,9 +4,11 @@ from typing import Any, cast from flask import request from flask_restx import Resource from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from controllers.common.fields import SimpleResultResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.app.wraps import get_app_model from controllers.console.wraps import ( @@ -23,7 +25,6 @@ from core.agent.entities import AgentToolEntity from core.tools.tool_manager import ToolManager from core.tools.utils.configuration import ToolParameterConfigurationManager from events.app_event import app_model_config_was_updated -from extensions.ext_database import db from libs.datetime_utils import naive_utc_now from libs.login import login_required from models.model import App, AppMode, AppModelConfig @@ -93,14 +94,16 @@ class ModelConfigResource(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id + @with_session @get_app_model(mode=[AppMode.AGENT_CHAT, AppMode.CHAT, AppMode.COMPLETION]) - def post(self, current_tenant_id: str, current_user_id: str, app_model: App): - """Modify app model config""" + def post(self, session: Session, current_tenant_id: str, current_user_id: str, app_model: App): + """Modify the app model config and dataset joins in one request transaction.""" # validate config model_configuration = AppModelConfigService.validate_configuration( tenant_id=current_tenant_id, config=cast(dict, request.json), app_mode=AppMode.value_of(app_model.mode), + session=session, ) new_app_model_config = AppModelConfig( @@ -110,9 +113,8 @@ class ModelConfigResource(Resource): ) new_app_model_config = new_app_model_config.from_model_config_dict(model_configuration) - if app_model.mode == AppMode.AGENT_CHAT or app_model.is_agent: - # get original app model config - original_app_model_config = db.session.get(AppModelConfig, app_model.app_model_config_id) + if app_model.mode == AppMode.AGENT_CHAT or app_model.is_agent_with_session(session=session): + original_app_model_config = app_model.app_model_config_with_session(session=session) if original_app_model_config is None: raise ValueError("Original app model config not found") agent_mode = original_app_model_config.agent_mode_dict @@ -204,14 +206,17 @@ class ModelConfigResource(Resource): # update app model config new_app_model_config.agent_mode = json.dumps(agent_mode) - db.session.add(new_app_model_config) - db.session.flush() + session.add(new_app_model_config) + session.flush() app_model.app_model_config_id = new_app_model_config.id app_model.updated_by = current_user_id app_model.updated_at = naive_utc_now() - db.session.commit() - app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config) + app_model_config_was_updated.send( + app_model, + app_model_config=new_app_model_config, + session=session, + ) return {"result": "success"} diff --git a/api/controllers/console/app/site.py b/api/controllers/console/app/site.py index da24c03b830..64228c247d5 100644 --- a/api/controllers/console/app/site.py +++ b/api/controllers/console/app/site.py @@ -3,10 +3,12 @@ from typing import Literal from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from constants.languages import supported_language from controllers.common.schema import register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.app.wraps import get_app_model from controllers.console.wraps import ( @@ -19,7 +21,6 @@ from controllers.console.wraps import ( setup_required, with_current_user, ) -from extensions.ext_database import db from fields.base import ResponseModel from libs.datetime_utils import naive_utc_now from libs.helper import dump_response @@ -94,10 +95,11 @@ class AppSite(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @account_initialization_required @with_current_user + @with_session @get_app_model - def post(self, current_user: Account, app_model: App): + def post(self, session: Session, current_user: Account, app_model: App): args = AppSiteUpdatePayload.model_validate(console_ns.payload or {}) - site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) + site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) if not site: raise NotFound @@ -126,7 +128,7 @@ class AppSite(Resource): site.updated_by = current_user.id site.updated_at = naive_utc_now() - db.session.commit() + session.flush() return dump_response(AppSiteResponse, site) @@ -145,16 +147,17 @@ class AppSiteAccessTokenReset(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @account_initialization_required @with_current_user + @with_session @get_app_model - def post(self, current_user: Account, app_model: App): - site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) + def post(self, session: Session, current_user: Account, app_model: App): + site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) if not site: raise NotFound - site.code = Site.generate_code(16) + site.code = Site.generate_code(16, session=session) site.updated_by = current_user.id site.updated_at = naive_utc_now() - db.session.commit() + session.flush() return dump_response(AppSiteResponse, site) diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 53c7c6ea788..7e8bd598d32 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -324,6 +324,27 @@ class WorkflowResponse(ResponseModel): return [_serialize_environment_variable(item) for item in value] +class _WorkflowResponseSource: + def __init__(self, workflow: Workflow, *, session: Session) -> None: + self._workflow = workflow + self._session = session + + def __getattr__(self, name: str) -> object: + return getattr(self._workflow, name) # noqa: no-new-getattr response adapter delegates model fields + + @property + def created_by_account(self) -> Account | None: + return self._workflow.get_created_by_account(session=self._session) + + @property + def updated_by_account(self) -> Account | None: + return self._workflow.get_updated_by_account(session=self._session) + + @property + def tool_published(self) -> bool: + return self._workflow.get_tool_published(session=self._session) + + class WorkflowPaginationResponse(ResponseModel): items: list[WorkflowResponse] page: int @@ -629,10 +650,10 @@ class AdvancedChatDraftWorkflowRunApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - @get_app_model(mode=[AppMode.ADVANCED_CHAT]) @with_current_user @edit_permission_required @with_session + @get_app_model(mode=[AppMode.ADVANCED_CHAT]) def post(self, session: Session, current_user: Account, app_model: App): """ Run draft workflow @@ -1074,10 +1095,10 @@ class DraftWorkflowRunApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - @get_app_model(mode=[AppMode.WORKFLOW]) @with_current_user @edit_permission_required @with_session + @get_app_model(mode=[AppMode.WORKFLOW]) def post(self, session: Session, current_user: Account, app_model: App): """ Run draft workflow @@ -1438,7 +1459,7 @@ class PublishedAllWorkflowApi(Resource): ) return WorkflowPaginationResponse.model_validate( { - "items": workflows, + "items": [_WorkflowResponseSource(workflow, session=session) for workflow in workflows], "page": page, "limit": limit, "has_more": has_more, @@ -1532,7 +1553,9 @@ class WorkflowByIdApi(Resource): if not workflow: raise NotFound("Workflow not found") - return dump_response(WorkflowResponse, workflow) + response = dump_response(WorkflowResponse, _WorkflowResponseSource(workflow, session=session)) + + return response @setup_required @login_required @@ -1626,10 +1649,10 @@ class DraftWorkflowTriggerRunApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - @get_app_model(mode=[AppMode.WORKFLOW]) @with_current_user @edit_permission_required @with_session + @get_app_model(mode=[AppMode.WORKFLOW]) def post(self, session: Session, current_user: Account, app_model: App): """ Poll for trigger events and execute full workflow when event arrives @@ -1637,7 +1660,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, session=db.session()) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=session) if not draft_workflow: raise ValueError("Workflow not found") @@ -1777,11 +1800,11 @@ class DraftWorkflowTriggerRunAllApi(Resource): @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.WORKFLOW]) @with_current_user @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) @with_session + @get_app_model(mode=[AppMode.WORKFLOW]) def post(self, session: Session, current_user: Account, app_model: App): """ Full workflow debug when the start node is a trigger @@ -1790,7 +1813,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, session=db.session()) + draft_workflow = workflow_service.get_draft_workflow(app_model, session=session) if not draft_workflow: raise ValueError("Workflow not found") diff --git a/api/controllers/console/app/wraps.py b/api/controllers/console/app/wraps.py index 097a41e64c3..48021ba18f6 100644 --- a/api/controllers/console/app/wraps.py +++ b/api/controllers/console/app/wraps.py @@ -1,9 +1,8 @@ """Controller decorators for console app resources. -App-loading decorators prefer a session injected by -`controllers.common.session.with_session` when present, while still supporting -existing handlers that have not been migrated yet and still rely on -Flask-SQLAlchemy's scoped `db.session`. +`get_app_model` still supports legacy handlers backed by Flask-SQLAlchemy's +scoped session. Trial app handlers compose `get_app_model_with_trial` under +`controllers.common.session.with_session` and always reuse that request session. """ from collections.abc import Callable @@ -41,9 +40,9 @@ def _load_app_model_from_scoped_session(app_id: str) -> App | None: return app_model -def _load_app_model_with_trial(app_id: str) -> App | None: +def _load_app_model_with_trial(session: Session, app_id: str) -> App | None: """Load a normal app through its trial registration without applying current-tenant scope.""" - app_model = db.session.scalar( + app_model = session.scalar( select(App).join(TrialApp, TrialApp.app_id == App.id).where(App.id == app_id, App.status == "normal").limit(1) ) return app_model @@ -157,7 +156,7 @@ def get_app_model_with_trial[**P, R]( *, mode: AppMode | list[AppMode] | None = None, ) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]: - """Inject an app registered for trial or available from the recommended catalog.""" + """Inject a trial-registered or recommended App using the Session supplied by `with_session`.""" def decorator(view_func: Callable[P, R]) -> Callable[P, R]: @wraps(view_func) @@ -170,9 +169,12 @@ def get_app_model_with_trial[**P, R]( del kwargs["app_id"] - app_model = _load_app_model_with_trial(app_id) + session = _get_injected_session(args) + if session is None: + raise RuntimeError("get_app_model_with_trial requires @with_session") + app_model = _load_app_model_with_trial(session, app_id) if app_model is None: - app_model = RecommendedAppService.get_app(app_id, session=db.session()) + app_model = RecommendedAppService.get_app(app_id, session=session) if not app_model: raise AppNotFoundError() diff --git a/api/controllers/console/auth/forgot_password.py b/api/controllers/console/auth/forgot_password.py index 6456bb480f4..8ea15e1ee5a 100644 --- a/api/controllers/console/auth/forgot_password.py +++ b/api/controllers/console/auth/forgot_password.py @@ -203,5 +203,5 @@ class ForgotPasswordResetApi(Resource): ): 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 + account.set_current_tenant_with_session(tenant, session=db.session()) tenant_was_created.send(tenant) diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index ab92fc0db74..0497cfd03cf 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -4,6 +4,7 @@ import flask_login from flask import make_response, request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from werkzeug.exceptions import Unauthorized import services @@ -16,6 +17,7 @@ from controllers.common.fields import ( SimpleResultResponse, ) from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.auth.error import ( AuthenticationFailedError, @@ -317,7 +319,7 @@ class EmailCodeLoginApi(Resource): 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") - account.current_tenant = new_tenant + account.set_current_tenant_with_session(new_tenant, session=db.session()) tenant_was_created.send(new_tenant) if account is None: @@ -356,7 +358,8 @@ class EmailCodeLoginApi(Resource): class RefreshTokenApi(Resource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @console_ns.response(401, "Unauthorized", console_ns.models[SimpleResultMessageResponse.__name__]) - def post(self): + @with_session(write=False) + def post(self, session: Session): # Get refresh token from cookie instead of request body refresh_token = extract_refresh_token(request) @@ -366,7 +369,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=session) except Unauthorized as exc: return SimpleResultMessageResponse(result="fail", message=exc.description or "Unauthorized.").model_dump( mode="json" diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index c4c3c5642d5..46d4eff01ec 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -284,7 +284,7 @@ def _generate_account( 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") - account.current_tenant = new_tenant + account.set_current_tenant_with_session(new_tenant, session=db.session()) tenant_was_created.send(new_tenant) if not account: diff --git a/api/controllers/console/auth/oauth_server.py b/api/controllers/console/auth/oauth_server.py index d068fb0785e..096426972c9 100644 --- a/api/controllers/console/auth/oauth_server.py +++ b/api/controllers/console/auth/oauth_server.py @@ -10,7 +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 core.db.session_factory import session_factory from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.login import login_required from models import Account @@ -132,9 +132,10 @@ 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, db.session() - ) + with session_factory.create_session() as session: + account = OAuthServerService.validate_oauth_access_token( + oauth_provider_app.client_id, access_token, session + ) if not account: response = jsonify({"error": "access_token or client_id is invalid"}) response.status_code = 401 diff --git a/api/controllers/console/datasets/data_source.py b/api/controllers/console/datasets/data_source.py index 17f027df9b3..590fbdbb87d 100644 --- a/api/controllers/console/datasets/data_source.py +++ b/api/controllers/console/datasets/data_source.py @@ -8,11 +8,12 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_serializer from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.common.fields import SimpleResultResponse, TextContentResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from core.datasource.entities.datasource_entities import DatasourceProviderType, OnlineDocumentPagesMessage from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin from core.entities.knowledge_entities import IndexingEstimate @@ -188,16 +189,16 @@ class DataSourceApi(Resource): @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @with_current_tenant_id + @with_session def patch( - self, current_tenant_id: str, binding_id: UUID, action: Literal["enable", "disable"] + self, session: Session, current_tenant_id: str, binding_id: UUID, action: Literal["enable", "disable"] ) -> tuple[dict[str, str], int]: binding_id_str = str(binding_id) - with sessionmaker(db.engine, expire_on_commit=False).begin() as session: - data_source_binding = session.execute( - select(DataSourceOauthBinding).where( - DataSourceOauthBinding.id == binding_id_str, DataSourceOauthBinding.tenant_id == current_tenant_id - ) - ).scalar_one_or_none() + data_source_binding = session.scalar( + select(DataSourceOauthBinding).where( + DataSourceOauthBinding.id == binding_id_str, DataSourceOauthBinding.tenant_id == current_tenant_id + ) + ) if data_source_binding is None: raise NotFound("Data source binding not found.") # enable binding @@ -206,8 +207,6 @@ class DataSourceApi(Resource): if data_source_binding.disabled: data_source_binding.disabled = False data_source_binding.updated_at = naive_utc_now() - db.session.add(data_source_binding) - db.session.commit() else: raise ValueError("Data source is not disabled.") # disable binding @@ -215,8 +214,6 @@ class DataSourceApi(Resource): if not data_source_binding.disabled: data_source_binding.disabled = True data_source_binding.updated_at = naive_utc_now() - db.session.add(data_source_binding) - db.session.commit() else: raise ValueError("Data source is disabled.") return {"result": "success"}, 200 @@ -231,7 +228,8 @@ class DataSourceNotionListApi(Resource): @console_ns.response(200, "Success", console_ns.models[NotionIntegrateInfoListResponse.__name__]) @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account) -> tuple[dict[str, Any], int]: + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account) -> tuple[dict[str, Any], int]: query = DataSourceNotionListQuery.model_validate(request.args.to_dict(flat=True)) datasource_provider_service = DatasourceProviderService() credential = datasource_provider_service.get_datasource_credentials( @@ -245,13 +243,13 @@ 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, session) if not dataset: raise NotFound("Dataset not found.") if dataset.data_source_type != "notion_import": raise ValueError("Dataset is not notion type.") - documents = db.session.scalars( + documents = session.scalars( select(Document).where( Document.dataset_id == query.dataset_id, Document.tenant_id == current_tenant_id, @@ -355,7 +353,8 @@ class DataSourceNotionIndexingEstimateApi(Resource): @console_ns.expect(console_ns.models[NotionEstimatePayload.__name__]) @console_ns.response(200, "Success", console_ns.models[IndexingEstimate.__name__]) @with_current_tenant_id - def post(self, current_tenant_id: str) -> tuple[dict[str, Any], int]: + @with_session + def post(self, session: Session, current_tenant_id: str) -> tuple[dict[str, Any], int]: payload = NotionEstimatePayload.model_validate(console_ns.payload or {}) args = payload.model_dump() # validate args @@ -382,11 +381,12 @@ class DataSourceNotionIndexingEstimateApi(Resource): extract_settings.append(extract_setting) indexing_runner = IndexingRunner() response = indexing_runner.indexing_estimate( - current_tenant_id, - extract_settings, - args["process_rule"], - args["doc_form"], - args["doc_language"], + tenant_id=current_tenant_id, + extract_settings=extract_settings, + tmp_processing_rule=args["process_rule"], + doc_form=args["doc_form"], + doc_language=args["doc_language"], + session=session, ) return dump_response(IndexingEstimate, response), 200 @@ -398,13 +398,14 @@ class DataSourceNotionDatasetSyncApi(Resource): @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, dataset_id: UUID) -> tuple[dict[str, str], int]: + @with_session(write=False) + def get(self, session: Session, 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, 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, session) for document in documents: document_indexing_sync_task.delay(dataset_id_str, document.id) return {"result": "success"}, 200 @@ -417,14 +418,15 @@ class DataSourceNotionDocumentSyncApi(Resource): @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, dataset_id: UUID, document_id: UUID) -> tuple[dict[str, str], int]: + @with_session(write=False) + def get(self, session: Session, 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, 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=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 0bf249f026b..ea6368075cb 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -1,3 +1,4 @@ +from dataclasses import dataclass from datetime import datetime from typing import Any from uuid import UUID @@ -13,10 +14,10 @@ import services from configs import dify_config from controllers.common.fields import ApiBaseUrlResponse, SimpleResultResponse, UsageCheckResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.apikey import ApiKeyItem, ApiKeyList from controllers.console.app.error import ProviderNotInitializeError -from controllers.console.app.wraps import with_session from controllers.console.datasets.error import DatasetInUseError, DatasetNameDuplicateError, IndexingEstimateError from controllers.console.wraps import ( RBACPermission, @@ -39,18 +40,18 @@ from core.rag.extractor.entity.datasource_type import DatasourceType from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.retrieval.retrieval_methods import RetrievalMethod -from extensions.ext_database import db from fields.base import ResponseModel -from fields.dataset_fields import DatasetDetailResponse +from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source from graphon.model_runtime.entities.model_entities import ModelType from libs.helper import build_icon_url, dump_response, to_timestamp from libs.login import login_required from libs.url_utils import normalize_api_base_url -from models import Account, ApiToken, Dataset, Document, DocumentSegment, UploadFile -from models.dataset import DatasetPermission, DatasetPermissionEnum +from models import Account, ApiToken, App, Dataset, Document, DocumentSegment, UploadFile +from models.dataset import DatasetPermission, DatasetPermissionEnum, DatasetQuery from models.enums import ApiTokenType, SegmentStatus from models.provider_ids import ModelProviderID from services.api_token_service import ApiTokenCache +from services.app_service import AppService from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService from services.enterprise import rbac_service as enterprise_rbac_service from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings @@ -205,6 +206,21 @@ class DatasetQueryDetailResponse(ResponseModel): return to_timestamp(value) +@dataclass(frozen=True) +class _DatasetQueryResponseSource: + """Expose query content through the request's database session.""" + + query: DatasetQuery + session: Session + + @property + def queries(self) -> list[dict[str, Any]]: + return self.query.get_queries(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.query, name) # noqa: no-new-getattr response adapter delegates model fields + + class DatasetQueryListResponse(ResponseModel): data: list[DatasetQueryDetailResponse] has_more: bool @@ -229,6 +245,21 @@ class RelatedAppResponse(ResponseModel): return self +@dataclass(frozen=True) +class _RelatedAppResponseSource: + """Expose the compatible app mode through the request's database session.""" + + app: App + session: Session + + @property + def mode_compatible_with_agent(self) -> str: + return self.app.mode_compatible_with_agent_with_session(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.app, name) # noqa: no-new-getattr response adapter delegates model fields + + class RelatedAppListResponse(ResponseModel): data: list[RelatedAppResponse] total: int @@ -397,7 +428,8 @@ class DatasetListApi(Resource): @enterprise_license_required @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account): # Convert query parameters to dict, handling list parameters correctly query_params: dict[str, str | list[str]] = dict(request.args.to_dict()) # Handle ids and tag_ids as lists (Flask request.args.getlist returns list even for single value) @@ -410,7 +442,7 @@ class DatasetListApi(Resource): permissions = enterprise_rbac_service.RBACService.MyPermissions.get( str(current_tenant_id), current_user.id, - session=db.session(), + session=session, ) accessible_dataset_ids: list[str] | None = None @@ -449,12 +481,13 @@ class DatasetListApi(Resource): user=current_user, accessible_dataset_ids=accessible_dataset_ids, include_own_datasets=include_own_datasets, + session=session, ) else: datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session(), + session, current_tenant_id, current_user, query.keyword, @@ -479,11 +512,14 @@ class DatasetListApi(Resource): for embedding_model in embedding_models: model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}") - data = [dump_response(DatasetDetailResponse, dataset) for dataset in datasets] + data = [ + dump_response(DatasetDetailResponse, dataset_detail_response_source(dataset, session=session)) + for dataset in datasets + ] dataset_ids = [item["id"] for item in data if item.get("permission") == "partial_members"] partial_members_map: dict[str, list[str]] = {} if dataset_ids: - partial_member_rows = db.session.execute( + partial_member_rows = session.execute( select(DatasetPermission.dataset_id, DatasetPermission.account_id).where( DatasetPermission.dataset_id.in_(dataset_ids) ) @@ -579,9 +615,9 @@ class DatasetListApi(Resource): session=session, ) - item = DatasetDetailWithPartialMembersResponse.model_validate(dataset, from_attributes=True).model_dump( - mode="json" - ) + item = DatasetDetailWithPartialMembersResponse.model_validate( + dataset_detail_response_source(dataset, session=session), from_attributes=True + ).model_dump(mode="json") item["permission_keys"] = permission_keys_map.get(dataset.id, []) return item, 201 @@ -604,30 +640,31 @@ class DatasetApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, 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(), + session=session, ) permission_keys_map = permissions.dataset.permission_keys_by_resource_ids([dataset_id_str]) - data = dump_response(DatasetDetailResponse, dataset) + data = dump_response(DatasetDetailResponse, dataset_detail_response_source(dataset, session=session)) data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: if dataset.embedding_model_provider: 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, session) data.update({"partial_member_list": part_users_list}) # check embedding setting @@ -671,7 +708,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, session) if dataset is None: raise NotFound("Dataset not found.") @@ -704,19 +741,19 @@ class DatasetApi(Resource): [dataset_id_str], session=session, ) - result_data = dump_response(DatasetDetailResponse, dataset) + result_data = dump_response(DatasetDetailResponse, dataset_detail_response_source(dataset, session=session)) result_data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) tenant_id = current_tenant_id 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, 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, 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, session) result_data.update({"partial_member_list": partial_member_list}) return dump_response(DatasetDetailWithPartialMembersResponse, result_data), 200 @@ -728,15 +765,16 @@ class DatasetApi(Resource): @console_ns.response(204, "Dataset deleted successfully") @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def delete(self, current_user: Account, dataset_id: UUID): + @with_session + def delete(self, session: Session, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) if not (current_user.has_edit_permission or current_user.is_dataset_operator): 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, session): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, session) return "", 204 else: raise NotFound("Dataset not found.") @@ -758,10 +796,11 @@ class DatasetUseCheckApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, dataset_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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, session) return UsageCheckResponse(is_using=dataset_is_using).model_dump(mode="json"), 200 @@ -780,24 +819,27 @@ class DatasetQueryApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) page = request.args.get("page", default=1, type=int) limit = request.args.get("limit", default=20, type=int) - dataset_queries, total = DatasetService.get_dataset_queries(dataset_id=dataset.id, page=page, per_page=limit) + dataset_queries, total = DatasetService.get_dataset_queries( + dataset_id=dataset.id, page=page, per_page=limit, session=session + ) response = { - "data": dataset_queries, + "data": [_DatasetQueryResponseSource(query=query, session=session) for query in dataset_queries], "has_more": len(dataset_queries) == limit, "limit": limit, "total": total, @@ -820,7 +862,8 @@ class DatasetIndexingEstimateApi(Resource): @account_initialization_required @console_ns.expect(console_ns.models[IndexingEstimatePayload.__name__]) @with_current_tenant_id - def post(self, current_tenant_id: str): + @with_session + def post(self, session: Session, current_tenant_id: str): payload = IndexingEstimatePayload.model_validate(console_ns.payload or {}) args = payload.model_dump() # validate args @@ -829,7 +872,7 @@ class DatasetIndexingEstimateApi(Resource): match args["info_list"]["data_source_type"]: case "upload_file": file_ids = args["info_list"]["file_info_list"]["file_ids"] - file_details = db.session.scalars( + file_details = session.scalars( select(UploadFile).where(UploadFile.tenant_id == current_tenant_id, UploadFile.id.in_(file_ids)) ).all() if file_details is None: @@ -886,13 +929,14 @@ class DatasetIndexingEstimateApi(Resource): indexing_runner = IndexingRunner() try: response = indexing_runner.indexing_estimate( - current_tenant_id, - extract_settings, - args["process_rule"], - args["doc_form"], - args["doc_language"], - args["dataset_id"], - args["indexing_technique"], + tenant_id=current_tenant_id, + extract_settings=extract_settings, + tmp_processing_rule=args["process_rule"], + doc_form=args["doc_form"], + doc_language=args["doc_language"], + dataset_id=args["dataset_id"], + indexing_technique=args["indexing_technique"], + session=session, ) except LLMBadRequestError: raise ProviderNotInitializeError( @@ -931,24 +975,25 @@ class DatasetRelatedAppListApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, 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, session) related_apps = [] for app_dataset_join in app_dataset_joins: - app_model = app_dataset_join.app + app_model = AppService.get_app_by_id(app_dataset_join.app_id, session) if app_model: - related_apps.append(app_model) + related_apps.append(_RelatedAppResponseSource(app=app_model, session=session)) return dump_response(RelatedAppListResponse, {"data": related_apps, "total": len(related_apps)}), 200 @@ -968,15 +1013,16 @@ class DatasetIndexingStatusApi(Resource): @account_initialization_required @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_tenant_id: str, dataset_id: UUID): + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, dataset_id: UUID): dataset_id_str = str(dataset_id) - documents = db.session.scalars( + documents = session.scalars( select(Document).where(Document.dataset_id == dataset_id_str, Document.tenant_id == current_tenant_id) ).all() documents_status = [] for document in documents: completed_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), DocumentSegment.document_id == str(document.id), @@ -986,7 +1032,7 @@ class DatasetIndexingStatusApi(Resource): or 0 ) total_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, @@ -1026,8 +1072,9 @@ class DatasetApiKeyApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, current_tenant_id: str): - keys = db.session.scalars( + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str): + keys = session.scalars( select(ApiToken).where(ApiToken.type == self.resource_type, ApiToken.tenant_id == current_tenant_id) ).all() return dump_response(ApiKeyList, {"data": keys}) @@ -1040,9 +1087,10 @@ class DatasetApiKeyApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): + @with_session + def post(self, session: Session, current_tenant_id: str): current_key_count = ( - db.session.scalar( + session.scalar( select(func.count(ApiToken.id)).where( ApiToken.type == self.resource_type, ApiToken.tenant_id == current_tenant_id ) @@ -1057,13 +1105,13 @@ class DatasetApiKeyApi(Resource): custom="max_keys_exceeded", ) - key = ApiToken.generate_api_key(self.token_prefix, 24) + key = ApiToken.generate_api_key(self.token_prefix, 24, session=session) api_token = ApiToken() api_token.tenant_id = current_tenant_id api_token.token = key api_token.type = self.resource_type - db.session.add(api_token) - db.session.commit() + session.add(api_token) + session.flush() return dump_response(ApiKeyItem, api_token), 200 @@ -1081,9 +1129,10 @@ class DatasetApiDeleteApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def delete(self, current_tenant_id: str, api_key_id: UUID): + @with_session + def delete(self, session: Session, current_tenant_id: str, api_key_id: UUID): api_key_id_str = str(api_key_id) - key = db.session.scalar( + key = session.scalar( select(ApiToken) .where( ApiToken.tenant_id == current_tenant_id, @@ -1101,8 +1150,7 @@ class DatasetApiDeleteApi(Resource): assert key is not None # nosec - for type checker only ApiTokenCache.delete(key.token, key.type) - db.session.delete(key) - db.session.commit() + session.delete(key) return "", 204 @@ -1114,10 +1162,11 @@ class DatasetEnableApiApi(Resource): @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, dataset_id: UUID, status: str): + @with_session + def post(self, session: Session, 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", session) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -1184,12 +1233,13 @@ class DatasetErrorDocs(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, session) return dump_response(ErrorDocsResponse, {"data": results, "total": len(results)}), 200 @@ -1211,17 +1261,18 @@ class DatasetPermissionUserListApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, 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, session) return dump_response(PartialMemberListResponse, {"data": partial_members_list}), 200 @@ -1241,10 +1292,11 @@ class DatasetAutoDisableLogApi(Resource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, 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 df4bb1bb66a..3d79af3ffc9 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -1,7 +1,7 @@ import json import logging from argparse import ArgumentTypeError -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from contextlib import ExitStack from datetime import datetime from typing import Any, Literal, cast @@ -12,12 +12,14 @@ from flask import request, send_file from flask_restx import Resource from pydantic import BaseModel, Field, JsonValue, field_validator from sqlalchemy import asc, desc, func, select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services from controllers.common.controller_schemas import DocumentBatchDownloadZipPayload from controllers.common.fields import SimpleResultMessageResponse, SimpleResultResponse, UrlResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required from core.entities.knowledge_entities import IndexingEstimate @@ -34,13 +36,15 @@ from core.rag.entities import Rule from core.rag.extractor.entity.datasource_type import DatasourceType from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo from core.rag.index_processor.constant.index_type import IndexTechniqueType -from extensions.ext_database import db from fields.base import ResponseModel from fields.document_fields import ( DocumentMetadataResponse, DocumentResponse, DocumentStatusListResponse, DocumentStatusResponse, + DocumentWithSession, + document_response, + document_responses, normalize_enum, ) from graphon.model_runtime.entities.model_entities import ModelType @@ -49,7 +53,7 @@ from libs.datetime_utils import naive_utc_now from libs.helper import dump_response, to_timestamp from libs.login import login_required from libs.pagination import paginate_query -from models import Account, DatasetProcessRule, Document, DocumentSegment, UploadFile +from models import Account, Document, DocumentSegment, UploadFile from models.dataset import DocumentPipelineExecutionLog from models.enums import IndexingStatus, ProcessRuleMode, SegmentStatus from services.dataset_ref_service import DatasetRefService @@ -110,6 +114,24 @@ class DocumentWithSegmentsResponse(DocumentResponse): total_segments: int | None = Field(default=None, exclude_if=lambda value: value is None) +class DocumentWithSegmentsSession(DocumentWithSession): + @property + def process_rule_dict(self) -> Any: + process_rule = self.document.get_dataset_process_rule(session=self.session) + return process_rule.to_dict() if process_rule else None + + +def document_with_segments_responses( + documents: Sequence[Document], *, session: Session +) -> list[DocumentWithSegmentsResponse]: + return [ + DocumentWithSegmentsResponse.model_validate( + DocumentWithSegmentsSession(document=document, session=session), from_attributes=True + ) + for document in documents + ] + + class DatasetAndDocumentResponse(ResponseModel): dataset: DatasetResponse documents: list[DocumentResponse] @@ -269,18 +291,18 @@ register_response_schema_models( class DocumentResource(Resource): def get_document( - self, dataset_id: str, document_id: str, current_user: Account, current_tenant_id: str + self, session: Session, 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, 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, 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=session) if not document: raise NotFound("Document not found.") @@ -290,17 +312,19 @@ 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()) + def get_batch_documents( + self, session: Session, dataset_id: str, batch: str, current_user: Account + ) -> Sequence[Document]: + dataset = DatasetService.get_dataset(dataset_id, 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, 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, session) if not documents: raise NotFound("Documents not found.") @@ -318,7 +342,8 @@ class GetProcessRuleApi(Resource): @login_required @account_initialization_required @with_current_user - def get(self, current_user: Account): + @with_session(write=False) + def get(self, session: Session, current_user: Account): req_data = request.args document_id = req_data.get("document_id") @@ -328,26 +353,21 @@ class GetProcessRuleApi(Resource): rules = DocumentService.DEFAULT_RULES["rules"] limits = DocumentService.DEFAULT_RULES["limits"] if document_id: - # get the latest process rule - document = db.get_or_404(Document, document_id) + document = DocumentService.get_document_by_id(document_id, session) + if document is None: + raise NotFound("Document not found.") - dataset = DatasetService.get_dataset(document.dataset_id, db.session()) + dataset = DatasetService.get_dataset(document.dataset_id, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - # get the latest process rule - dataset_process_rule = db.session.scalar( - select(DatasetProcessRule) - .where(DatasetProcessRule.dataset_id == document.dataset_id) - .order_by(DatasetProcessRule.created_at.desc()) - .limit(1) - ) + dataset_process_rule = dataset.get_latest_process_rule(session=session) if dataset_process_rule: mode = dataset_process_rule.mode rules = dataset_process_rule.rules_dict @@ -381,7 +401,8 @@ class DatasetDocumentListApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) raw_args = request.args.to_dict() param = DocumentDatasetListParam.model_validate(raw_args) @@ -407,12 +428,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, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -457,20 +478,20 @@ class DatasetDocumentListApi(Resource): desc(Document.position), ) - paginated_documents = paginate_query(query, page=page, per_page=limit, max_per_page=100) + paginated_documents = paginate_query(query, session=session, page=page, per_page=limit, max_per_page=100) documents = paginated_documents.items DocumentService.enrich_documents_with_summary_index_status( documents=documents, dataset=dataset, tenant_id=current_tenant_id, - session=db.session(), + session=session, ) if fetch: for document in documents: completed_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), DocumentSegment.document_id == str(document.id), @@ -480,7 +501,7 @@ class DatasetDocumentListApi(Resource): or 0 ) total_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, @@ -491,7 +512,7 @@ class DatasetDocumentListApi(Resource): document.completed_segments = completed_segments document.total_segments = total_segments response = { - "data": documents, + "data": document_with_segments_responses(documents, session=session), "has_more": len(documents) == limit, "limit": limit, "total": paginated_documents.total, @@ -509,10 +530,11 @@ class DatasetDocumentListApi(Resource): @console_ns.response(200, "Documents created successfully", console_ns.models[DatasetAndDocumentResponse.__name__]) @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, dataset_id: UUID): + @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, session) if not dataset: raise NotFound("Dataset not found.") @@ -522,7 +544,7 @@ class DatasetDocumentListApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session()) + DatasetService.check_dataset_permission(dataset, current_user, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -536,9 +558,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=session ) - dataset = DatasetService.get_dataset(dataset_id_str, db.session()) + dataset = DatasetService.get_dataset(dataset_id_str, session) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -547,7 +569,10 @@ class DatasetDocumentListApi(Resource): except ModelCurrentlyNotSupportError: raise ProviderModelCurrentlyNotSupportError() - return dump_response(DatasetAndDocumentResponse, {"dataset": dataset, "documents": documents, "batch": batch}) + return dump_response( + DatasetAndDocumentResponse, + {"dataset": dataset, "documents": document_responses(documents, session=session), "batch": batch}, + ) @setup_required @login_required @@ -555,9 +580,10 @@ class DatasetDocumentListApi(Resource): @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Documents deleted successfully") @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def delete(self, dataset_id: UUID): + @with_session + def delete(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, session) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting @@ -566,7 +592,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.get_doc_form(session=session), session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -589,7 +615,8 @@ class DatasetInitApi(Resource): @cloud_edition_billing_rate_limit_check("knowledge") @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account): + @with_session + def post(self, session: Session, current_tenant_id: str, current_user: Account): # 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() @@ -625,7 +652,7 @@ class DatasetInitApi(Resource): tenant_id=current_tenant_id, knowledge_config=knowledge_config, account=current_user, - session=db.session(), + session=session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -634,7 +661,10 @@ class DatasetInitApi(Resource): except ModelCurrentlyNotSupportError: raise ProviderModelCurrentlyNotSupportError() - return dump_response(DatasetAndDocumentResponse, {"dataset": dataset, "documents": documents, "batch": batch}) + return dump_response( + DatasetAndDocumentResponse, + {"dataset": dataset, "documents": document_responses(documents, session=session), "batch": batch}, + ) @console_ns.route("/datasets//documents//indexing-estimate") @@ -655,23 +685,24 @@ class DocumentIndexingEstimateApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def get(self, session: Session, 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) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) if document.indexing_status in {IndexingStatus.COMPLETED, IndexingStatus.ERROR}: raise DocumentAlreadyFinishedError() - data_process_rule = document.dataset_process_rule - data_process_rule_dict = data_process_rule.to_dict() if data_process_rule else {} + data_process_rule = document.get_dataset_process_rule(session=session) + data_process_rule_dict: Mapping[str, Any] = data_process_rule.to_dict() if data_process_rule else {} if document.data_source_type == "upload_file": data_source_info = document.data_source_info_dict if data_source_info and "upload_file_id" in data_source_info: file_id = data_source_info["upload_file_id"] - file = db.session.scalar( + file = session.scalar( select(UploadFile) .where(UploadFile.tenant_id == document.tenant_id, UploadFile.id == file_id) .limit(1) @@ -689,12 +720,13 @@ class DocumentIndexingEstimateApi(DocumentResource): try: estimate_response = indexing_runner.indexing_estimate( - current_tenant_id, - [extract_setting], - data_process_rule_dict, - document.doc_form, - "English", - dataset_id_str, + tenant_id=current_tenant_id, + extract_settings=[extract_setting], + tmp_processing_rule=data_process_rule_dict, + doc_form=document.doc_form, + doc_language="English", + dataset_id=dataset_id_str, + session=session, ) return ( # TODO: why using zero here? the same for the below endpoint @@ -745,9 +777,10 @@ class DocumentBatchIndexingEstimateApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, batch: str): + @with_session + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, batch: str): dataset_id_str = str(dataset_id) - documents = self.get_batch_documents(dataset_id_str, batch, current_user) + documents = self.get_batch_documents(session, dataset_id_str, batch, current_user) if not documents: return ( IndexingEstimateResponse( @@ -759,8 +792,8 @@ class DocumentBatchIndexingEstimateApi(DocumentResource): ).model_dump(mode="json", exclude_none=True), 200, ) - data_process_rule = documents[0].dataset_process_rule - data_process_rule_dict = data_process_rule.to_dict() if data_process_rule else {} + data_process_rule = documents[0].get_dataset_process_rule(session=session) + data_process_rule_dict: Mapping[str, Any] = data_process_rule.to_dict() if data_process_rule else {} extract_settings = [] for document in documents: if document.indexing_status in {IndexingStatus.COMPLETED, IndexingStatus.ERROR}: @@ -771,7 +804,7 @@ class DocumentBatchIndexingEstimateApi(DocumentResource): if not data_source_info: continue file_id = data_source_info["upload_file_id"] - file_detail = db.session.scalar( + file_detail = session.scalar( select(UploadFile) .where(UploadFile.tenant_id == current_tenant_id, UploadFile.id == file_id) .limit(1) @@ -825,12 +858,13 @@ class DocumentBatchIndexingEstimateApi(DocumentResource): indexing_runner = IndexingRunner() try: response = indexing_runner.indexing_estimate( - current_tenant_id, - extract_settings, - data_process_rule_dict, - document.doc_form, - "English", - dataset_id_str, + tenant_id=current_tenant_id, + extract_settings=extract_settings, + tmp_processing_rule=data_process_rule_dict, + doc_form=document.doc_form, + doc_language="English", + dataset_id=dataset_id_str, + session=session, ) return ( IndexingEstimateResponse( @@ -865,13 +899,14 @@ class DocumentBatchIndexingStatusApi(DocumentResource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_user: Account, dataset_id: UUID, batch: str): + @with_session(write=False) + def get(self, session: Session, current_user: Account, dataset_id: UUID, batch: str): dataset_id_str = str(dataset_id) - documents = self.get_batch_documents(dataset_id_str, batch, current_user) + documents = self.get_batch_documents(session, dataset_id_str, batch, current_user) documents_status = [] for document in documents: completed_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), DocumentSegment.document_id == str(document.id), @@ -881,7 +916,7 @@ class DocumentBatchIndexingStatusApi(DocumentResource): or 0 ) total_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, @@ -923,13 +958,14 @@ class DocumentIndexingStatusApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) completed_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), DocumentSegment.document_id == document_id_str, @@ -939,7 +975,7 @@ class DocumentIndexingStatusApi(DocumentResource): or 0 ) total_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.document_id == document_id_str, DocumentSegment.status != SegmentStatus.RE_SEGMENT, @@ -987,10 +1023,11 @@ class DocumentApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) metadata = request.args.get("metadata", "all") if metadata not in self.METADATA_CHOICES: @@ -1002,20 +1039,22 @@ class DocumentApi(DocumentResource): { "id": document.id, "doc_type": document.doc_type, - "doc_metadata": document.doc_metadata_details, + "doc_metadata": document.get_doc_metadata_details(session=session), } ) return response.model_dump(mode="json", include={"id", *metadata_fields}, exclude_unset=True), 200 - 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 {} + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, session) + document_process_rule = document.get_dataset_process_rule(session=session) + document_process_rules: Mapping[str, Any] = document_process_rule.to_dict() if document_process_rule else {} + segment_count = document.get_segment_count(session=session) response = DocumentDetailResponse.model_validate( { "id": document.id, "position": document.position, "data_source_type": document.data_source_type, "data_source_info": document.data_source_info_dict, - "data_source_detail_dict": document.data_source_detail_dict, + "data_source_detail_dict": document.get_data_source_detail_dict(session=session), "dataset_process_rule_id": document.dataset_process_rule_id, "dataset_process_rule": dataset_process_rules, "document_process_rule": document_process_rules, @@ -1034,10 +1073,10 @@ class DocumentApi(DocumentResource): "disabled_by": document.disabled_by, "archived": document.archived, "doc_type": document.doc_type, - "doc_metadata": document.doc_metadata_details, - "segment_count": document.segment_count, - "average_segment_length": document.average_segment_length, - "hit_count": document.hit_count, + "doc_metadata": document.get_doc_metadata_details(session=session), + "segment_count": segment_count, + "average_segment_length": (document.word_count or 0) // segment_count if segment_count else 0, + "hit_count": document.get_hit_count(session=session), "display_status": document.display_status, "doc_form": document.doc_form, "doc_language": document.doc_language, @@ -1055,19 +1094,22 @@ class DocumentApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def delete(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def delete( + self, session: Session, 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, session) if dataset is None: raise NotFound("Dataset not found.") # check user's model setting DatasetService.check_dataset_model_setting(dataset) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) try: - DocumentService.delete_document(document, db.session()) + DocumentService.delete_document(document, session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") @@ -1088,12 +1130,13 @@ class DocumentDownloadApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_DOCUMENT_DOWNLOAD) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID) -> dict[str, Any]: + @with_session(write=False) + def get( + self, session: Session, 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 UrlResponse(url=DocumentService.get_document_download_url(document, db.session())).model_dump( - mode="json" - ) + document = self.get_document(session, str(dataset_id), str(document_id), current_user, current_tenant_id) + return UrlResponse(url=DocumentService.get_document_download_url(document, session)).model_dump(mode="json") @console_ns.route("/datasets//documents/download-zip") @@ -1111,7 +1154,8 @@ class DocumentBatchDownloadZipApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): + @with_session(write=False) + def post(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): """Stream a ZIP archive containing the requested uploaded documents.""" # Parse and validate request payload. payload = DocumentBatchDownloadZipPayload.model_validate(console_ns.payload or {}) @@ -1123,7 +1167,7 @@ class DocumentBatchDownloadZipApi(DocumentResource): document_ids=document_ids, tenant_id=current_tenant_id, current_user=current_user, - session=db.session(), + session=session, ) # Delegate ZIP packing to FileService, but keep Flask response+cleanup in the route. @@ -1162,8 +1206,10 @@ class DocumentProcessingApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( self, + session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, @@ -1172,7 +1218,7 @@ class DocumentProcessingApi(DocumentResource): ): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) # 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: @@ -1186,7 +1232,6 @@ class DocumentProcessingApi(DocumentResource): document.paused_by = current_user.id document.paused_at = naive_utc_now() document.is_paused = True - db.session.commit() case "resume": if document.indexing_status not in {IndexingStatus.PAUSED, IndexingStatus.ERROR}: @@ -1195,7 +1240,6 @@ class DocumentProcessingApi(DocumentResource): document.paused_by = None document.paused_at = None document.is_paused = False - db.session.commit() return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -1219,10 +1263,11 @@ class DocumentMetadataApi(DocumentResource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def put(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def put(self, session: Session, 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) - document = self.get_document(dataset_id_str, document_id_str, current_user, current_tenant_id) + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) req_data = DocumentMetadataUpdatePayload.model_validate(request.get_json() or {}) @@ -1254,7 +1299,6 @@ class DocumentMetadataApi(DocumentResource): document.doc_type = doc_type document.updated_at = naive_utc_now() - db.session.commit() return SimpleResultMessageResponse(result="success", message="Document metadata updated.").model_dump( mode="json" @@ -1271,11 +1315,16 @@ class DocumentStatusApi(DocumentResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( - self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable", "archive", "un_archive"] + self, + session: Session, + 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, session) if dataset is None: raise NotFound("Dataset not found.") @@ -1287,12 +1336,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, 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, session) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -1311,16 +1360,17 @@ class DocumentPauseApi(DocumentResource): @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Document paused successfully") @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def patch(self, dataset_id: UUID, document_id: UUID): + @with_session + def patch(self, session: Session, dataset_id: UUID, document_id: UUID): """pause 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, 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=session) # 404 if document not found if document is None: @@ -1332,7 +1382,7 @@ class DocumentPauseApi(DocumentResource): try: # pause document - DocumentService.pause_document(document, db.session()) + DocumentService.pause_document(document, session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot pause completed document.") @@ -1347,14 +1397,15 @@ class DocumentRecoverApi(DocumentResource): @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Document resumed successfully") @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def patch(self, dataset_id: UUID, document_id: UUID): + @with_session + def patch(self, session: Session, dataset_id: UUID, document_id: UUID): """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, 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=session) # 404 if document not found if document is None: @@ -1365,7 +1416,7 @@ class DocumentRecoverApi(DocumentResource): raise ArchivedDocumentImmutableError() try: # pause document - DocumentService.recover_document(document, db.session()) + DocumentService.recover_document(document, session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Document is not in paused status.") @@ -1381,17 +1432,18 @@ class DocumentRetryApi(DocumentResource): @console_ns.expect(console_ns.models[DocumentRetryPayload.__name__]) @console_ns.response(204, "Documents retry started successfully") @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, dataset_id: UUID): + @with_session + def post(self, session: Session, dataset_id: UUID): """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, 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=session) # 404 if document not found if document is None: @@ -1409,7 +1461,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, session) return "", 204 @@ -1423,22 +1475,23 @@ class DocumentRenameApi(DocumentResource): @console_ns.expect(console_ns.models[DocumentRenamePayload.__name__]) @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def post(self, session: Session, current_user: Account, dataset_id: UUID, document_id: UUID): # 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(str(dataset_id), db.session()) + dataset = DatasetService.get_dataset(dataset_id, 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=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, session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") - return dump_response(DocumentResponse, document) + return dump_response(DocumentResponse, document_response(document, session=session)) @console_ns.route("/datasets//documents//website-sync") @@ -1449,14 +1502,15 @@ class WebsiteDocumentSyncApi(DocumentResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def get(self, session: Session, 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, 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=session) if not document: raise NotFound("Document not found.") if document.tenant_id != current_tenant_id: @@ -1467,7 +1521,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, session) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -1483,17 +1537,18 @@ class DocumentPipelineExecutionLogApi(DocumentResource): @login_required @account_initialization_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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, 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=session) if not document: raise NotFound("Document not found.") - log = db.session.scalar( + log = session.scalar( select(DocumentPipelineExecutionLog) .where(DocumentPipelineExecutionLog.document_id == document_id_str) .order_by(DocumentPipelineExecutionLog.created_at.desc()) @@ -1532,7 +1587,8 @@ class DocumentGenerateSummaryApi(Resource): @cloud_edition_billing_rate_limit_check("knowledge") @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, dataset_id: UUID): + @with_session + def post(self, session: Session, current_user: Account, dataset_id: UUID): """ Generate summary index for specified documents. @@ -1543,7 +1599,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, session) if not dataset: raise NotFound("Dataset not found.") @@ -1552,7 +1608,7 @@ class DocumentGenerateSummaryApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session()) + DatasetService.check_dataset_permission(dataset, current_user, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1577,7 +1633,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, session) if len(documents) != len(document_list): found_ids = {doc.id for doc in documents} @@ -1593,7 +1649,7 @@ class DocumentGenerateSummaryApi(Resource): DocumentService.update_documents_need_summary( dataset_id=dataset_id_str, document_ids=document_ids_to_update, - session=db.session(), + session=session, need_summary=True, ) @@ -1631,7 +1687,8 @@ class DocumentSummaryStatusApi(DocumentResource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, current_user: Account, dataset_id: UUID, document_id: UUID): """ Get summary index generation status for a document. @@ -1649,13 +1706,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, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1665,7 +1722,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=session, ) return dump_response(DocumentSummaryStatusResponse, result), 200 diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 19e5670f239..33a5a1f752a 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -8,6 +8,7 @@ from flask_restx import Resource from pydantic import BaseModel, Field from sqlalchemy import String, case, cast, func, literal, or_, select from sqlalchemy.dialects.postgresql import JSONB +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -20,6 +21,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.app.error import ProviderNotInitializeError from controllers.console.datasets.error import ( @@ -42,7 +44,6 @@ from controllers.console.wraps import ( from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.model_manager import ModelManager from core.rag.index_processor.constant.index_type import IndexTechniqueType -from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.base import ResponseModel from fields.segment_fields import ( @@ -165,7 +166,7 @@ register_response_schema_models( def _get_segment_for_document( - dataset: Dataset, document: Document, segment_id: str + session: Session, dataset: Dataset, document: Document, segment_id: str ) -> tuple[SegmentRef, DocumentSegment]: dataset_ref = DatasetRefService.create_dataset_ref(dataset) document_ref = DatasetRefService.create_document_ref(dataset_ref, document) @@ -173,7 +174,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, db.session()) + segment = SegmentService.get_segment_by_ref(segment_ref, session=session) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -190,19 +191,20 @@ class DatasetDocumentSegmentListApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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, 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, 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=session) if not document: raise NotFound("Document not found.") @@ -271,19 +273,19 @@ class DatasetDocumentSegmentListApi(Resource): elif args.enabled.lower() == "false": query = query.where(DocumentSegment.enabled == False) - segments = paginate_query(query, page=page, per_page=limit, max_per_page=100) + segments = paginate_query(query, session=session, page=page, per_page=limit, max_per_page=100) segment_list = list(segments.items) segment_ids = [segment.id for segment in segment_list] summaries: dict[str, str | None] = {} if segment_ids: summary_records = SummaryIndexService.get_segments_summaries( - segment_ids=segment_ids, dataset_id=dataset_id_str, session=db.session() + segment_ids=segment_ids, dataset_id=dataset_id_str, session=session ) summaries = {chunk_id: summary.summary_content for chunk_id, summary in summary_records.items()} response = { - "data": segment_responses_with_summaries(segment_list, summaries), + "data": segment_responses_with_summaries(segment_list, summaries, session=session), "limit": limit, "total": segments.total, "total_pages": segments.pages, @@ -300,17 +302,18 @@ class DatasetDocumentSegmentListApi(Resource): @console_ns.response(204, "Segments deleted successfully") @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def delete(self, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def delete(self, session: Session, 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, 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=session) if not document: raise NotFound("Document not found.") segment_ids = request.args.getlist("segment_id") @@ -319,10 +322,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, 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, session) return "", 204 @@ -339,8 +342,10 @@ class DatasetDocumentSegmentApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( self, + session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, @@ -348,11 +353,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, 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=session) if not document: raise NotFound("Document not found.") # check user's model setting @@ -362,7 +367,7 @@ class DatasetDocumentSegmentApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user, db.session()) + DatasetService.check_dataset_permission(dataset, current_user, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -388,7 +393,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, session) except Exception as e: raise InvalidActionError(str(e)) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -408,15 +413,23 @@ class DatasetDocumentSegmentAddApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def post( + self, + session: Session, + 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, 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=session) if not document: raise NotFound("Document not found.") if not current_user.is_dataset_editor: @@ -438,22 +451,21 @@ 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, 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()), - ) + segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset, session)) summary = SummaryIndexService.get_segment_summary( - segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + segment_id=segment.id, dataset_id=dataset_id_str, session=session ) response = { - "data": segment_response_with_summary(segment, summary.summary_content if summary else None), + "data": segment_response_with_summary( + segment, summary.summary_content if summary else None, session=session + ), "doc_form": document.doc_form, } return dump_response(SegmentDetailResponse, response), 200 @@ -472,26 +484,33 @@ class DatasetDocumentSegmentUpdateApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( - self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID, segment_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + 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, 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=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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -511,7 +530,7 @@ class DatasetDocumentSegmentUpdateApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) segment_id_str = str(segment_id) - _, segment = _get_segment_for_document(dataset, document, segment_id_str) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # validate args payload = SegmentUpdatePayload.model_validate(console_ns.payload or {}) payload_dict = payload.model_dump(exclude_none=True) @@ -523,13 +542,15 @@ class DatasetDocumentSegmentUpdateApi(Resource): segment, document, dataset, - db.session(), + session, ) summary = SummaryIndexService.get_segment_summary( - segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + segment_id=segment.id, dataset_id=dataset_id_str, session=session ) response = { - "data": segment_response_with_summary(segment, summary.summary_content if summary else None), + "data": segment_response_with_summary( + segment, summary.summary_content if summary else None, session=session + ), "doc_form": document.doc_form, } return dump_response(SegmentDetailResponse, response), 200 @@ -543,31 +564,38 @@ class DatasetDocumentSegmentUpdateApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def delete( - self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID, segment_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + 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, 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=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, 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()) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) + SegmentService.delete_segment(segment, document, dataset, session) return "", 204 @@ -587,22 +615,30 @@ class DatasetDocumentSegmentBatchImportApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @with_session + def post( + self, + session: Session, + 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, 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=session) if not document: raise NotFound("Document not found.") payload = BatchImportPayload.model_validate(console_ns.payload or {}) upload_file_id = payload.upload_file_id - upload_file = db.session.scalar(select(UploadFile).where(UploadFile.id == upload_file_id).limit(1)) + upload_file = session.scalar(select(UploadFile).where(UploadFile.id == upload_file_id).limit(1)) if not upload_file: raise NotFound("UploadFile not found.") @@ -660,23 +696,30 @@ class ChildChunkAddApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def post( - self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID, segment_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + 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, 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=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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # check embedding model setting @@ -696,11 +739,11 @@ class ChildChunkAddApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) segment_id_str = str(segment_id) - _, segment = _get_segment_for_document(dataset, document, segment_id_str) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # 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, session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 @@ -713,21 +756,22 @@ class ChildChunkAddApi(Resource): @account_initialization_required @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) - def get(self, current_tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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, 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=session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - _get_segment_for_document(dataset, document, segment_id_str) + _get_segment_for_document(session, dataset, document, segment_id_str) args = query_params_from_request(ChildChunkListQuery, use_defaults_for_malformed_ints=True) page = args.page @@ -735,7 +779,13 @@ class ChildChunkAddApi(Resource): keyword = args.keyword child_chunks = SegmentService.get_child_chunks( - segment_id_str, document_id_str, dataset_id_str, page, limit, keyword + segment_id_str, + document_id_str, + dataset_id_str, + page, + limit, + keyword, + session=session, ) response = { "data": child_chunks.items, @@ -761,34 +811,41 @@ class ChildChunkAddApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( - self, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID, segment_id: UUID + self, + session: Session, + current_tenant_id: str, + current_user: Account, + 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, 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=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, 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) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # 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, session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkBatchUpdateResponse, {"data": child_chunks}), 200 @@ -807,8 +864,10 @@ class ChildChunkUpdateApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def delete( self, + session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, @@ -818,31 +877,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, 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=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, 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) + segment_ref, _ = _get_segment_for_document(session, 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, db.session()) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, session=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, session) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) return "", 204 @@ -858,8 +917,10 @@ class ChildChunkUpdateApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) + @with_session def patch( self, + session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, @@ -869,34 +930,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, 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=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, 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) + segment_ref, segment = _get_segment_for_document(session, 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, db.session()) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, session=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, 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 cff0589f9eb..94efe388561 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -1,3 +1,4 @@ +from dataclasses import dataclass from datetime import datetime from typing import Any from uuid import UUID @@ -10,9 +11,13 @@ from werkzeug.exceptions import Forbidden, InternalServerError, NotFound import services from controllers.common.fields import UsageCountResponse -from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.schema import ( + query_params_from_model, + register_response_schema_models, + register_schema_models, +) +from controllers.common.session import with_session from controllers.console import console_ns -from controllers.console.app.wraps import with_session from controllers.console.datasets.error import DatasetNameDuplicateError from controllers.console.wraps import ( RBACPermission, @@ -25,10 +30,11 @@ from controllers.console.wraps import ( with_current_user, ) from fields.base import ResponseModel -from fields.dataset_fields import DatasetDetailResponse +from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source from libs.helper import dump_response from libs.login import login_required from models import Account +from models.dataset import ExternalKnowledgeApis from services.dataset_service import DatasetService from services.enterprise import rbac_service as enterprise_rbac_service from services.external_knowledge_service import ExternalDatasetService @@ -90,6 +96,28 @@ class ExternalKnowledgeApiResponse(ResponseModel): return value +@dataclass(frozen=True) +class ExternalKnowledgeApiResponseSource: + external_knowledge_api: ExternalKnowledgeApis + session: Session + + @property + def dataset_bindings(self) -> Any: + return self.external_knowledge_api.get_dataset_bindings(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.external_knowledge_api, name) # noqa: no-new-getattr response adapter delegates model fields + + +def external_knowledge_api_response( + external_knowledge_api: ExternalKnowledgeApis, *, session: Session +) -> ExternalKnowledgeApiResponse: + return ExternalKnowledgeApiResponse.model_validate( + ExternalKnowledgeApiResponseSource(external_knowledge_api=external_knowledge_api, session=session), + from_attributes=True, + ) + + class ExternalKnowledgeApiListResponse(ResponseModel): data: list[ExternalKnowledgeApiResponse] has_more: bool @@ -162,14 +190,15 @@ class ExternalApiTemplateListApi(Resource): @login_required @with_current_tenant_id @account_initialization_required - def get(self, current_tenant_id: str): + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str): query = ExternalApiTemplateListQuery.model_validate(request.args.to_dict()) external_knowledge_apis, total = ExternalDatasetService.get_external_knowledge_apis( - query.page, query.limit, current_tenant_id, query.keyword + query.page, query.limit, current_tenant_id, query.keyword, session=session ) return ExternalKnowledgeApiListResponse( - data=[ExternalKnowledgeApiResponse.model_validate(item) for item in external_knowledge_apis], + data=[external_knowledge_api_response(item, session=session) for item in external_knowledge_apis], has_more=len(external_knowledge_apis) == query.limit, limit=query.limit, total=total, @@ -210,7 +239,7 @@ class ExternalApiTemplateListApi(Resource): except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() - return dump_response(ExternalKnowledgeApiResponse, external_knowledge_api), 201 + return external_knowledge_api_response(external_knowledge_api, session=session).model_dump(mode="json"), 201 @console_ns.route("/datasets/external-knowledge-api/") @@ -237,7 +266,7 @@ class ExternalApiTemplateApi(Resource): if external_knowledge_api is None: raise NotFound("API template not found.") - return dump_response(ExternalKnowledgeApiResponse, external_knowledge_api), 200 + return external_knowledge_api_response(external_knowledge_api, session=session).model_dump(mode="json"), 200 @console_ns.doc("update_external_api_template") @console_ns.doc(description="Update external knowledge API template") @@ -269,7 +298,7 @@ class ExternalApiTemplateApi(Resource): session=session, ) - return dump_response(ExternalKnowledgeApiResponse, external_knowledge_api), 200 + return external_knowledge_api_response(external_knowledge_api, session=session).model_dump(mode="json"), 200 @setup_required @login_required @@ -354,7 +383,9 @@ class ExternalDatasetCreateApi(Resource): [dataset_id_str], session=session, ) - data = DatasetDetailResponse.model_validate(dataset).model_dump(mode="json") + data = DatasetDetailResponse.model_validate( + dataset_detail_response_source(dataset, session=session) + ).model_dump(mode="json") data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) return data, 201 diff --git a/api/controllers/console/datasets/hit_testing.py b/api/controllers/console/datasets/hit_testing.py index b30aadfbbc7..8018343b01f 100644 --- a/api/controllers/console/datasets/hit_testing.py +++ b/api/controllers/console/datasets/hit_testing.py @@ -53,7 +53,7 @@ class HitTestingApi(Resource, DatasetsHitTestingBase): ) -> dict[str, object]: dataset_id_str = str(dataset_id) - dataset = self.get_and_validate_dataset(dataset_id_str, current_user, current_tenant_id) + dataset = self.get_and_validate_dataset(session, dataset_id_str, current_user, current_tenant_id) args = self.parse_args(console_ns.payload) self.hit_testing_args_check(args) diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index 656a426c125..512f44ed9a3 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -19,7 +19,6 @@ from core.errors.error import ( ProviderTokenNotInitError, QuotaExceededError, ) -from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs.login import resolve_account_fallback from models.account import Account @@ -83,15 +82,18 @@ class DatasetsHitTestingBase: @staticmethod def get_and_validate_dataset( - dataset_id: str, current_user: Account | None = None, current_tenant_id: str | None = None + session: Session, + 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, 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, 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 42ae4903673..32c4151a017 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -2,10 +2,12 @@ from typing import Literal from uuid import UUID from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.common.controller_schemas import MetadataUpdatePayload from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.wraps import ( RBACPermission, @@ -17,7 +19,6 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) -from extensions.ext_database import db from fields.dataset_fields import ( DatasetMetadataBuiltInFieldsResponse, DatasetMetadataListResponse, @@ -57,17 +58,18 @@ class DatasetMetadataCreateApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_tenant_id: str, current_user: Account, dataset_id: UUID): + @with_session + def post(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): 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, 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, session) metadata = MetadataService.create_metadata( - dataset_id_str, metadata_args, current_user, current_tenant_id, session=db.session() + dataset_id_str, metadata_args, current_user, current_tenant_id, session=session ) return dump_response(DatasetMetadataResponse, metadata), 201 @@ -79,12 +81,13 @@ class DatasetMetadataCreateApi(Resource): 200, "Metadata retrieved successfully", console_ns.models[DatasetMetadataListResponse.__name__] ) @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) - def get(self, dataset_id: UUID): + @with_session(write=False) + def get(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, session) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) + metadata = MetadataService.get_dataset_metadatas(dataset, session) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -99,19 +102,27 @@ class DatasetMetadataApi(Resource): @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def patch(self, current_tenant_id: str, current_user: Account, dataset_id: UUID, metadata_id: UUID): + @with_session + def patch( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + metadata_id: UUID, + ): payload = MetadataUpdatePayload.model_validate(console_ns.payload or {}) name = payload.name 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, 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, session) metadata = MetadataService.update_metadata_name( - dataset_id_str, metadata_id_str, name, current_user, current_tenant_id, session=db.session() + dataset_id_str, metadata_id_str, name, current_user, current_tenant_id, session=session ) return dump_response(DatasetMetadataResponse, metadata), 200 @@ -122,15 +133,16 @@ class DatasetMetadataApi(Resource): @console_ns.response(204, "Metadata deleted successfully") @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def delete(self, current_user: Account, dataset_id: UUID, metadata_id: UUID): + @with_session + def delete(self, session: Session, 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, 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, session) - MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -160,18 +172,19 @@ class DatasetMetadataBuiltInFieldActionApi(Resource): @console_ns.response(204, "Action completed successfully") @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, dataset_id: UUID, action: Literal["enable", "disable"]): + @with_session + def post(self, session: Session, 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, 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, session) match action: case "enable": - MetadataService.enable_built_in_field(dataset, session=db.session()) + MetadataService.enable_built_in_field(dataset, session) case "disable": - MetadataService.disable_built_in_field(dataset, session=db.session()) + MetadataService.disable_built_in_field(dataset, session) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -189,16 +202,17 @@ class DocumentMetadataEditApi(Resource): ) @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, dataset_id: UUID): + @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, 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, session) metadata_args = MetadataOperationData.model_validate(console_ns.payload or {}) - MetadataService.update_documents_metadata(dataset, metadata_args, current_user, session=db.session()) + MetadataService.update_documents_metadata(dataset, metadata_args, current_user, session=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/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index a61fc2639db..a264450078c 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -27,7 +27,7 @@ from controllers.console.app.workflow import ( WorkflowResponse, ) from controllers.console.app.wraps import with_session -from controllers.console.datasets.wraps import get_rag_pipeline +from controllers.console.datasets.wraps import get_rag_pipeline, load_rag_pipeline from controllers.console.wraps import ( RBACPermission, RBACResourceScope, @@ -344,11 +344,11 @@ class DraftRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @with_session - @get_rag_pipeline - def post(self, session: Session, current_user: Account, pipeline: Pipeline): + def post(self, session: Session, current_user: Account, pipeline_id: UUID): """ Run draft workflow """ + pipeline = load_rag_pipeline(session, str(pipeline_id)) payload = DraftWorkflowRunPayload.model_validate(console_ns.payload or {}) args = payload.model_dump() @@ -378,11 +378,11 @@ class PublishedRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @with_session - @get_rag_pipeline - def post(self, session: Session, current_user: Account, pipeline: Pipeline): + def post(self, session: Session, current_user: Account, pipeline_id: UUID): """ Run published workflow """ + pipeline = load_rag_pipeline(session, str(pipeline_id)) payload = PublishedWorkflowRunPayload.model_validate(console_ns.payload or {}) args = payload.model_dump(exclude_none=True) streaming = payload.response_mode == "streaming" diff --git a/api/controllers/console/datasets/wraps.py b/api/controllers/console/datasets/wraps.py index b5a9cd753ff..ea55cccb976 100644 --- a/api/controllers/console/datasets/wraps.py +++ b/api/controllers/console/datasets/wraps.py @@ -1,13 +1,21 @@ 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 from libs.login import current_account_with_tenant from models.dataset import Pipeline +from services.rag_pipeline.rag_pipeline import RagPipelineService + + +def load_rag_pipeline(session: Session, pipeline_id: str) -> Pipeline: + _, current_tenant_id = current_account_with_tenant() + pipeline = RagPipelineService.get_pipeline_by_id(pipeline_id, current_tenant_id, session=session) + if not pipeline: + raise PipelineNotFoundError() + return pipeline def get_rag_pipeline[**P, R](view_func: Callable[P, R]) -> Callable[P, R]: @@ -16,22 +24,11 @@ def get_rag_pipeline[**P, R](view_func: Callable[P, R]) -> Callable[P, R]: if not kwargs.get("pipeline_id"): raise ValueError("missing pipeline_id in path parameters") - _, current_tenant_id = current_account_with_tenant() - pipeline_id = kwargs.get("pipeline_id") pipeline_id = str(pipeline_id) del kwargs["pipeline_id"] - - 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() - - kwargs["pipeline"] = pipeline + kwargs["pipeline"] = load_rag_pipeline(db.session(), pipeline_id) return view_func(*args, **kwargs) diff --git a/api/controllers/console/explore/audio.py b/api/controllers/console/explore/audio.py index 3ad170ebccf..6219571b2d6 100644 --- a/api/controllers/console/explore/audio.py +++ b/api/controllers/console/explore/audio.py @@ -50,14 +50,19 @@ register_response_schema_models(console_ns, AudioBinaryResponse, AudioTranscript class ChatAudioApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[AudioTranscriptResponse.__name__]) def post(self, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() file = request.files["file"] try: - response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=None) + response = AudioService.transcript_asr( + app_model=app_model, + file=file, + session=db.session(), + end_user=None, + ) return response except services.errors.app_model_config.AppModelConfigBrokenError: @@ -96,7 +101,7 @@ class ChatTextApi(InstalledAppResource): @console_ns.expect(console_ns.models[TextToAudioPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[AudioBinaryResponse.__name__]) def post(self, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() try: diff --git a/api/controllers/console/explore/completion.py b/api/controllers/console/explore/completion.py index 6b546d183d4..5f034fabb81 100644 --- a/api/controllers/console/explore/completion.py +++ b/api/controllers/console/explore/completion.py @@ -90,7 +90,7 @@ class CompletionApi(InstalledAppResource): @with_current_user @with_session def post(self, session: Session, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode != AppMode.COMPLETION: @@ -146,8 +146,9 @@ class CompletionApi(InstalledAppResource): class CompletionStopApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @with_current_user_id - def post(self, current_user_id: str, installed_app: InstalledApp, task_id: str): - app_model = installed_app.app + @with_session(write=False) + def post(self, session: Session, current_user_id: str, installed_app: InstalledApp, task_id: str): + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode != AppMode.COMPLETION: @@ -173,7 +174,7 @@ class ChatApi(InstalledAppResource): @with_current_user @with_session def post(self, session: Session, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) @@ -195,7 +196,7 @@ class ChatApi(InstalledAppResource): app_model=app_model, conversation_id=payload.conversation_id, user=current_user, - session=db.session(), + session=session, ) response = AppGenerateService.generate( @@ -240,8 +241,9 @@ class ChatApi(InstalledAppResource): class ChatStopApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @with_current_user_id - def post(self, current_user_id: str, installed_app: InstalledApp, task_id: str): - app_model = installed_app.app + @with_session(write=False) + def post(self, session: Session, current_user_id: str, installed_app: InstalledApp, task_id: str): + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) diff --git a/api/controllers/console/explore/conversation.py b/api/controllers/console/explore/conversation.py index 25239203d8d..9e21fd496af 100644 --- a/api/controllers/console/explore/conversation.py +++ b/api/controllers/console/explore/conversation.py @@ -16,6 +16,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom from extensions.ext_database import db from fields.conversation_fields import ( ConversationInfiniteScrollPagination, + ConversationResponseSource, ResultResponse, SimpleConversation, ) @@ -53,7 +54,7 @@ class ConversationListApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ConversationInfiniteScrollPagination.__name__]) @with_current_user def get(self, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) @@ -84,7 +85,13 @@ class ConversationListApi(InstalledAppResource): pinned=args.pinned, ) adapter = TypeAdapter(SimpleConversation) - conversations = [adapter.validate_python(item, from_attributes=True) for item in pagination.data] + conversations = [ + adapter.validate_python( + ConversationResponseSource(item, session=session), + from_attributes=True, + ) + for item in pagination.data + ] return ConversationInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, @@ -102,7 +109,7 @@ class ConversationApi(InstalledAppResource): @console_ns.response(204, "Conversation deleted successfully") @with_current_user def delete(self, current_user: Account, installed_app: InstalledApp, c_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) @@ -127,7 +134,7 @@ class ConversationRenameApi(InstalledAppResource): @console_ns.response(200, "Conversation renamed successfully", console_ns.models[SimpleConversation.__name__]) @with_current_user def post(self, current_user: Account, installed_app: InstalledApp, c_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) @@ -139,12 +146,13 @@ class ConversationRenameApi(InstalledAppResource): payload = ConversationRenamePayload.model_validate(console_ns.payload or {}) try: + session = db.session() conversation = ConversationService.rename( - app_model, conversation_id, current_user, payload.name, payload.auto_generate, session=db.session() + app_model, conversation_id, current_user, payload.name, payload.auto_generate, session=session ) return ( TypeAdapter(SimpleConversation) - .validate_python(conversation, from_attributes=True) + .validate_python(ConversationResponseSource(conversation, session=session), from_attributes=True) .model_dump(mode="json") ) except ConversationNotExistsError: @@ -159,7 +167,7 @@ class ConversationPinApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ResultResponse.__name__]) @with_current_user def patch(self, current_user: Account, installed_app: InstalledApp, c_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) @@ -184,7 +192,7 @@ class ConversationUnPinApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ResultResponse.__name__]) @with_current_user def patch(self, current_user: Account, installed_app: InstalledApp, c_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) diff --git a/api/controllers/console/explore/message.py b/api/controllers/console/explore/message.py index 7b316b0382d..fb78ddede70 100644 --- a/api/controllers/console/explore/message.py +++ b/api/controllers/console/explore/message.py @@ -28,7 +28,7 @@ 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.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import ( ExploreMessageInfiniteScrollPagination, ExploreMessageListItem, @@ -76,7 +76,8 @@ class MessageListApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ExploreMessageInfiniteScrollPagination.__name__]) @with_current_user def get(self, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + session = db.session() + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() @@ -92,10 +93,13 @@ class MessageListApi(InstalledAppResource): args.conversation_id, args.first_id or None, args.limit, - session=db.session(), + session=session, ) adapter = TypeAdapter(ExploreMessageListItem) - items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] + items = [ + adapter.validate_python(MessageResponseSource(message, session=session), from_attributes=True) + for message in pagination.data + ] return ExploreMessageInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, @@ -116,7 +120,7 @@ class MessageFeedbackApi(InstalledAppResource): @console_ns.response(200, "Feedback submitted successfully", console_ns.models[ResultResponse.__name__]) @with_current_user def post(self, current_user: Account, installed_app: InstalledApp, message_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() @@ -149,7 +153,7 @@ class MessageMoreLikeThisApi(InstalledAppResource): @with_current_user @with_session def get(self, session: Session, current_user: Account, installed_app: InstalledApp, message_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode != "completion": @@ -199,7 +203,7 @@ class MessageSuggestedQuestionApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SuggestedQuestionsResponse.__name__]) @with_current_user def get(self, current_user: Account, installed_app: InstalledApp, message_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() app_mode = AppMode.value_of(app_model.mode) diff --git a/api/controllers/console/explore/parameter.py b/api/controllers/console/explore/parameter.py index 680885f9bd3..9208bb63695 100644 --- a/api/controllers/console/explore/parameter.py +++ b/api/controllers/console/explore/parameter.py @@ -9,7 +9,7 @@ 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 models.model import AppMode, InstalledApp, load_annotation_reply_config from services.app_service import AppService @@ -32,24 +32,29 @@ class AppParameterApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[fields.Parameters.__name__]) def get(self, installed_app: InstalledApp): """Retrieve app parameters.""" - app_model = installed_app.app + session = db.session() + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() features_dict: dict[str, Any] = workflow.features_dict user_input_form = workflow.user_input_form(to_old_structure=True) else: - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if app_model_config is None: raise AppUnavailableError() - features_dict = cast(dict[str, Any], app_model_config.to_dict()) + annotation_reply = load_annotation_reply_config(session, app_model.id) + features_dict = cast( + dict[str, Any], + app_model_config.to_dict(annotation_reply=annotation_reply), + ) user_input_form = features_dict.get("user_input_form", []) @@ -62,7 +67,7 @@ class ExploreAppMetaApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ExploreAppMetaResponse.__name__]) def get(self, installed_app: InstalledApp): """Get app meta""" - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if not app_model: raise ValueError("App not found") return AppService().get_app_meta(app_model, session=db.session()) diff --git a/api/controllers/console/explore/saved_message.py b/api/controllers/console/explore/saved_message.py index e3fd730a3cc..d057c49bff7 100644 --- a/api/controllers/console/explore/saved_message.py +++ b/api/controllers/console/explore/saved_message.py @@ -12,7 +12,7 @@ from controllers.console.explore.error import NotCompletionAppError from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import with_current_user from extensions.ext_database import db -from fields.conversation_fields import ResultResponse +from fields.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import SavedMessageInfiniteScrollPagination, SavedMessageItem from models import Account from models.model import InstalledApp @@ -29,7 +29,8 @@ class SavedMessageListApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SavedMessageInfiniteScrollPagination.__name__]) @with_current_user def get(self, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + session = db.session() + app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode != "completion": @@ -38,10 +39,13 @@ class SavedMessageListApi(InstalledAppResource): args = SavedMessageListQuery.model_validate(request.args.to_dict()) pagination = SavedMessageService.pagination_by_last_id( - app_model, current_user, str(args.last_id) if args.last_id else None, args.limit, session=db.session() + app_model, current_user, str(args.last_id) if args.last_id else None, args.limit, session=session ) adapter = TypeAdapter(SavedMessageItem) - items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] + items = [ + adapter.validate_python(MessageResponseSource(message, session=session), from_attributes=True) + for message in pagination.data + ] return SavedMessageInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, @@ -52,7 +56,7 @@ class SavedMessageListApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[ResultResponse.__name__]) @with_current_user def post(self, current_user: Account, installed_app: InstalledApp): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() if app_model.mode != "completion": @@ -75,7 +79,7 @@ class SavedMessageApi(InstalledAppResource): @console_ns.response(204, "Saved message deleted successfully") @with_current_user def delete(self, current_user: Account, installed_app: InstalledApp, message_id: UUID): - app_model = installed_app.app + app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index 47c4b625512..f15577ea728 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -1,4 +1,6 @@ import logging +from collections.abc import Mapping +from dataclasses import dataclass from datetime import datetime from typing import Any, Literal @@ -65,11 +67,12 @@ from libs import helper from libs.helper import dump_response, to_timestamp, uuid_value from models import Account from models.account import TenantStatus -from models.model import AppMode, Site +from models.model import AppMode, Site, load_annotation_reply_config from models.workflow import Workflow +from services.account_service import TenantService from services.app_generate_service import AppGenerateService from services.app_ref_service import AppRefService -from services.app_service import AppService +from services.app_service import AppResponseView, AppService from services.audio_service import AudioService from services.dataset_service import DatasetService from services.errors.audio import ( @@ -379,6 +382,27 @@ class TrialWorkflowResponse(ResponseModel): return to_timestamp(value) +@dataclass(frozen=True) +class TrialWorkflowResponseSource: + workflow: Workflow + session: Session + + @property + def created_by_account(self) -> Account | None: + return self.workflow.get_created_by_account(session=self.session) + + @property + def updated_by_account(self) -> Account | None: + return self.workflow.get_updated_by_account(session=self.session) + + @property + def tool_published(self) -> bool: + return self.workflow.get_tool_published(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.workflow, name) # noqa: no-new-getattr response adapter delegates model fields + + register_schema_models( console_ns, WorkflowRunRequest, @@ -594,7 +618,12 @@ class TrialChatAudioApi(TrialAppResource): app_id = app_model.id user_id = current_user.id - response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=None) + response = AudioService.transcript_asr( + app_model=app_model, + file=file, + session=db.session(), + end_user=None, + ) RecommendedAppService.add_trial_app_record(app_id, user_id, session=db.session()) return response except services.errors.app_model_config.AppModelConfigBrokenError: @@ -746,19 +775,21 @@ class TrialSitApi(Resource): """Resource for trial app sites.""" @console_ns.response(200, "Success", console_ns.models[SiteResponse.__name__]) + @with_session(write=False) @get_app_model_with_trial(None) - def get(self, app_model): + def get(self, session: Session, app_model): """Retrieve app site info. Returns the site configuration for the application including theme, icons, and text. """ - site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) + site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) if not site: raise Forbidden() - assert app_model.tenant - if app_model.tenant.status == TenantStatus.ARCHIVE: + tenant = TenantService.get_tenant_by_id(app_model.tenant_id, session=session) + assert tenant + if tenant.status == TenantStatus.ARCHIVE: raise Forbidden() return SiteResponse.model_validate(site).model_dump(mode="json") @@ -768,26 +799,29 @@ class TrialAppParameterApi(Resource): """Resource for app variables.""" @console_ns.response(200, "Success", console_ns.models[ParametersResponse.__name__]) + @with_session(write=False) @get_app_model_with_trial(None) - def get(self, app_model): + def get(self, session: Session, app_model): """Retrieve app parameters.""" if app_model is None: raise AppUnavailableError() + features_dict: Mapping[str, Any] if app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() features_dict = workflow.features_dict user_input_form = workflow.user_input_form(to_old_structure=True) else: - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if app_model_config is None: raise AppUnavailableError() - features_dict = app_model_config.to_dict() + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + features_dict = app_model_config.to_dict(annotation_reply=annotation_reply) user_input_form = features_dict.get("user_input_form", []) @@ -797,43 +831,52 @@ class TrialAppParameterApi(Resource): class AppApi(Resource): @console_ns.response(200, "Success", console_ns.models[TrialAppDetailResponse.__name__]) + @with_session(write=False) @get_app_model_with_trial(None) - def get(self, app_model): + def get(self, session: Session, app_model): """Get app detail""" app_service = AppService() - app_model = app_service.get_app(app_model) + app_model = app_service.get_app(app_model, session=session) - return dump_response(TrialAppDetailResponse, app_model) + return TrialAppDetailResponse.model_validate( + AppResponseView(app_model, session=session), + from_attributes=True, + ).model_dump(mode="json") class AppWorkflowApi(Resource): @console_ns.response(200, "Success", console_ns.models[TrialWorkflowResponse.__name__]) + @with_session(write=False) @get_app_model_with_trial(None) - def get(self, app_model): + def get(self, session: Session, app_model): """Get workflow detail""" if not app_model.workflow_id: raise AppUnavailableError() - workflow = db.session.get(Workflow, app_model.workflow_id) + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() - return dump_response(TrialWorkflowResponse, workflow) + return TrialWorkflowResponse.model_validate( + TrialWorkflowResponseSource(workflow=workflow, session=session), + from_attributes=True, + ).model_dump(mode="json") class DatasetListApi(Resource): @console_ns.doc(params=query_params_from_model(TrialDatasetListQuery)) @console_ns.response(200, "Success", console_ns.models[TrialDatasetListResponse.__name__]) + @with_session(write=False) @get_app_model_with_trial(None) - def get(self, app_model): + def get(self, session: Session, app_model): page = request.args.get("page", default=1, type=int) limit = request.args.get("limit", default=20, type=int) ids = request.args.getlist("ids") tenant_id = app_model.tenant_id if ids: - datasets, total = DatasetService.get_datasets_by_ids(ids, tenant_id) + datasets, total = DatasetService.get_datasets_by_ids(ids, tenant_id, session=session) else: raise NeedAddIdsError() diff --git a/api/controllers/console/explore/workflow.py b/api/controllers/console/explore/workflow.py index 72c9085921f..a8c176c6778 100644 --- a/api/controllers/console/explore/workflow.py +++ b/api/controllers/console/explore/workflow.py @@ -51,7 +51,7 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): """ Run workflow """ - app_model = installed_app.app + app_model = installed_app.app_with_session(session=session) if not app_model: raise NotWorkflowAppError() app_mode = AppMode.value_of(app_model.mode) @@ -92,11 +92,12 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): @console_ns.route("/installed-apps//workflows/tasks//stop") class InstalledAppWorkflowTaskStopApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) - def post(self, installed_app: InstalledApp, task_id: str): + @with_session(write=False) + def post(self, session: Session, installed_app: InstalledApp, task_id: str): """ Stop workflow task """ - app_model = installed_app.app + app_model = installed_app.app_with_session(session=session) if not app_model: raise NotWorkflowAppError() app_mode = AppMode.value_of(app_model.mode) diff --git a/api/controllers/console/explore/wraps.py b/api/controllers/console/explore/wraps.py index 9f7e829ae8a..d67f3e18d53 100644 --- a/api/controllers/console/explore/wraps.py +++ b/api/controllers/console/explore/wraps.py @@ -30,7 +30,7 @@ def installed_app_required[**P, R](view: Callable[Concatenate[InstalledApp, P], if installed_app is None: raise NotFound("Installed app not found") - if not installed_app.app: + if not installed_app.app_with_session(session=db.session()): db.session.delete(installed_app) db.session.commit() @@ -74,17 +74,18 @@ def trial_app_required[**P, R](view: Callable[Concatenate[App, P], R] | None = N @wraps(view) def decorated(app_id: str, *args: P.args, **kwargs: P.kwargs): current_user, _ = current_account_with_tenant() + session = db.session() - trial_app = db.session.scalar(select(TrialApp).where(TrialApp.app_id == str(app_id)).limit(1)) + trial_app = session.scalar(select(TrialApp).where(TrialApp.app_id == str(app_id)).limit(1)) if trial_app is None: raise TrialAppNotAllowed() - app = trial_app.app + app = trial_app.app_with_session(session=session) if app is None: raise TrialAppNotAllowed() - account_trial_app_record = db.session.scalar( + account_trial_app_record = session.scalar( select(AccountTrialAppRecord) .where(AccountTrialAppRecord.account_id == current_user.id, AccountTrialAppRecord.app_id == app_id) .limit(1) diff --git a/api/controllers/console/socketio/workflow.py b/api/controllers/console/socketio/workflow.py index 6d8d316ad90..04178aeff18 100644 --- a/api/controllers/console/socketio/workflow.py +++ b/api/controllers/console/socketio/workflow.py @@ -4,7 +4,7 @@ from typing import cast from flask import Request as FlaskRequest -from extensions.ext_database import db +from core.db.session_factory import session_factory from extensions.ext_socketio import sio from libs.passport import PassportService from libs.token import extract_access_token @@ -43,8 +43,8 @@ def socket_connect(sid, environ, auth): logging.warning("Socket connect rejected: missing user_id (sid=%s)", sid) return False - with sio.app.app_context(): - user = AccountService.load_logged_in_account(account_id=user_id, session=db.session()) + with sio.app.app_context(), session_factory.create_session() as session: + user = AccountService.load_logged_in_account(account_id=user_id, session=session) if not user: logging.warning("Socket connect rejected: user not found (user_id=%s, sid=%s)", user_id, sid) return False @@ -69,8 +69,8 @@ def handle_user_connect(sid, data): if not workflow_id: return {"msg": "workflow_id is required"}, 400 - with sio.app.app_context(): - result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid, session=db.session()) + with sio.app.app_context(), session_factory.create_session() as session: + result = collaboration_service.authorize_and_join_workflow_room(workflow_id, sid, session=session) if not result: return {"msg": "unauthorized"}, 401 diff --git a/api/controllers/console/workspace/workspace.py b/api/controllers/console/workspace/workspace.py index 23ce116b349..a15663be557 100644 --- a/api/controllers/console/workspace/workspace.py +++ b/api/controllers/console/workspace/workspace.py @@ -6,7 +6,8 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select -from werkzeug.exceptions import Unauthorized +from sqlalchemy.orm import Session +from werkzeug.exceptions import NotFound, Unauthorized import services from configs import dify_config @@ -23,6 +24,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.admin import admin_required from controllers.console.error import AccountNotLinkTenantError @@ -220,10 +222,11 @@ class TenantListApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str, current_user: Account): tenant_rows: list[tuple[Tenant, TenantAccountJoin]] = [ (tenant, membership) - for tenant, membership in TenantService.get_workspaces_for_account(current_user.id, session=db.session()) + for tenant, membership in TenantService.get_workspaces_for_account(current_user.id, session=session) if tenant.status == TenantStatus.NORMAL ] tenants = [tenant for tenant, _ in tenant_rows] @@ -274,11 +277,12 @@ class WorkspaceListApi(Resource): @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[WorkspacePaginationResponse.__name__]) @setup_required @admin_required - def get(self): + @with_session(write=False) + def get(self, session: Session): args = query_params_from_request(WorkspaceListQuery) stmt = select(Tenant).order_by(Tenant.created_at.desc()) - tenants = paginate_query(stmt, page=args.page, per_page=args.limit) + tenants = paginate_query(stmt, session=session, page=args.page, per_page=args.limit) has_more = False if tenants.has_next: @@ -297,7 +301,8 @@ class TenantApi(Resource): @account_initialization_required @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[TenantInfoResponse.__name__]) @with_current_user - def post(self, current_user: Account): + @with_session + def post(self, session: Session, current_user: Account): if request.path == "/info": logger.warning("Deprecated URL /info was used.") @@ -306,17 +311,17 @@ 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=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=session) tenant = tenants[0] # else, raise Unauthorized else: raise Unauthorized("workspace is archived") return ( - dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant, session=db.session())), + dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant, session=session)), HTTPStatus.OK, ) @@ -329,22 +334,23 @@ class SwitchWorkspaceApi(Resource): @login_required @account_initialization_required @with_current_user - def post(self, current_user: Account): + @with_session + def post(self, session: Session, current_user: Account): payload = console_ns.payload or {} args = SwitchWorkspacePayload.model_validate(payload) # 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=session) except Exception: raise AccountNotLinkTenantError("Account not link tenant") - new_tenant = db.session.get(Tenant, args.tenant_id) # Get new tenant + new_tenant = TenantService.get_tenant_by_id(args.tenant_id, session=session) if new_tenant is None: raise ValueError("Tenant not found") return SwitchWorkspaceResponse( - result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant, session=db.session()) + result="success", new_tenant=WorkspaceService.get_tenant_info(new_tenant, session=session) ).model_dump(mode="json") @@ -357,10 +363,13 @@ class CustomConfigWorkspaceApi(Resource): @account_initialization_required @cloud_edition_billing_resource_check("workspace_custom") @with_current_tenant_id - def post(self, current_tenant_id: str): + @with_session + def post(self, session: Session, current_tenant_id: str): payload = console_ns.payload or {} args = WorkspaceCustomConfigPayload.model_validate(payload) - tenant = db.get_or_404(Tenant, current_tenant_id) + tenant = TenantService.get_tenant_by_id(current_tenant_id, session=session) + if tenant is None: + raise NotFound() custom_config_dict: TenantCustomConfigDict = { "remove_webapp_brand": args.remove_webapp_brand @@ -372,10 +381,10 @@ class CustomConfigWorkspaceApi(Resource): } tenant.custom_config_dict = custom_config_dict - db.session.commit() + session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=session) ).model_dump(mode="json") @@ -430,18 +439,21 @@ class WorkspaceInfoApi(Resource): @account_initialization_required # Change workspace name @with_current_tenant_id - def post(self, current_tenant_id: str): + @with_session + def post(self, session: Session, current_tenant_id: str): payload = console_ns.payload or {} args = WorkspaceInfoPayload.model_validate(payload) if not current_tenant_id: raise ValueError("No current tenant") - tenant = db.get_or_404(Tenant, current_tenant_id) + tenant = TenantService.get_tenant_by_id(current_tenant_id, session=session) + if tenant is None: + raise NotFound() tenant.name = args.name - db.session.commit() + session.commit() return WorkspaceTenantResultResponse( - result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=db.session()) + result="success", tenant=WorkspaceService.get_tenant_info(tenant, session=session) ).model_dump(mode="json") diff --git a/api/controllers/inner_api/app/dsl.py b/api/controllers/inner_api/app/dsl.py index 9fd111f86dc..8b206cc36be 100644 --- a/api/controllers/inner_api/app/dsl.py +++ b/api/controllers/inner_api/app/dsl.py @@ -54,9 +54,8 @@ class EnterpriseAppDSLImport(Resource): if account is None: return {"message": f"account '{args.creator_email}' not found or inactive"}, 404 - account.set_tenant_id(workspace_id) - with Session(db.engine, expire_on_commit=False) as session: + account.set_tenant_id_with_session(workspace_id, session=session) dsl_service = AppDslService(session) result = dsl_service.import_app( account=account, diff --git a/api/controllers/openapi/_input_schema.py b/api/controllers/openapi/_input_schema.py index 1b638200b8c..afaaf2e690f 100644 --- a/api/controllers/openapi/_input_schema.py +++ b/api/controllers/openapi/_input_schema.py @@ -4,9 +4,11 @@ from __future__ import annotations from typing import Any, cast +from sqlalchemy.orm import Session + from controllers.service_api.app.error import AppUnavailableError from models import App -from models.model import AppMode +from models.model import AppMode, load_annotation_reply_config JSON_SCHEMA_DRAFT = "https://json-schema.org/draft/2020-12/schema" @@ -89,13 +91,13 @@ def _form_to_jsonschema(form: list[dict[str, Any]]) -> tuple[dict[str, Any], lis return properties, required -def resolve_app_config(app: App) -> tuple[dict[str, Any], list[dict[str, Any]]]: +def resolve_app_config(app: App, *, session: Session) -> tuple[dict[str, Any], list[dict[str, Any]]]: """Resolve `(features_dict, user_input_form)` for parameters / schema derivation. Raises `AppUnavailableError` on misconfigured apps. """ if app.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app.workflow + workflow = app.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() return ( @@ -103,21 +105,22 @@ def resolve_app_config(app: App) -> tuple[dict[str, Any], list[dict[str, Any]]]: cast(list[dict[str, Any]], workflow.user_input_form(to_old_structure=True)), ) - app_model_config = app.app_model_config + app_model_config = app.app_model_config_with_session(session=session) if app_model_config is None: raise AppUnavailableError() - features_dict = cast(dict[str, Any], app_model_config.to_dict()) + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + features_dict = cast(dict[str, Any], app_model_config.to_dict(annotation_reply=annotation_reply)) return features_dict, cast(list[dict[str, Any]], features_dict.get("user_input_form", [])) -def build_input_schema(app: App) -> dict[str, Any]: +def build_input_schema(app: App, *, session: Session) -> dict[str, Any]: """Derive Draft 2020-12 JSON Schema from `user_input_form` + app mode. chat / agent-chat / advanced-chat: top-level `query` (required, minLength=1) + `inputs` object. completion / workflow: `inputs` object only. Raises `AppUnavailableError` on misconfigured apps. """ - _, user_input_form = resolve_app_config(app) + _, user_input_form = resolve_app_config(app, session=session) inputs_props, inputs_required = _form_to_jsonschema(user_input_form) properties: dict[str, Any] = {} diff --git a/api/controllers/openapi/account.py b/api/controllers/openapi/account.py index b4786f2ae25..f1b02ef115e 100644 --- a/api/controllers/openapi/account.py +++ b/api/controllers/openapi/account.py @@ -3,8 +3,10 @@ from __future__ import annotations from datetime import UTC, datetime from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound +from controllers.common.session import with_session from controllers.openapi import openapi_ns from controllers.openapi._contract import accepts, returns from controllers.openapi._models import ( @@ -18,7 +20,6 @@ from controllers.openapi._models import ( ) from controllers.openapi.auth.composition import auth_router from controllers.openapi.auth.data import AuthData -from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.oauth_bearer import ( Scope, @@ -41,14 +42,13 @@ from services.oauth_device_flow import ( class AccountApi(Resource): @auth_router.guard(scope=Scope.FULL, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, AccountResponse, description="Account info") - def get(self, *, auth_data: AuthData): + @with_session(write=False) + def get(self, session: Session, *, auth_data: AuthData): 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(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 [] - ) + account = AccountService.get_account_by_id(account_id_str, session=session) if account_id_str else None + memberships = TenantService.get_account_memberships(account_id_str, session=session) if account_id_str else [] default_ws_id = _pick_default_workspace(memberships) return AccountResponse( @@ -64,8 +64,9 @@ class AccountApi(Resource): 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(redis_client, str(auth_data.token_id), session=db.session()) + @with_session + def delete(self, session: Session, *, auth_data: AuthData): + revoke_oauth_token(redis_client, str(auth_data.token_id), session=session) return RevokeResponse(status="revoked") @@ -74,7 +75,8 @@ class AccountSessionsApi(Resource): @auth_router.guard(scope=Scope.FULL, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, SessionListResponse, description="Session list") @accepts(query=SessionListQuery) - def get(self, *, auth_data: AuthData, query: SessionListQuery): + @with_session(write=False) + def get(self, session: Session, *, auth_data: AuthData, query: SessionListQuery): # SessionListQuery enforces the advertised bounds (extra='forbid', page>=1, # 1<=limit<=MAX_PAGE_LIMIT) so the server rejects out-of-range paging rather # than silently coercing (e.g. page=0 -> empty slice). @@ -83,7 +85,7 @@ class AccountSessionsApi(Resource): page = query.page limit = query.limit - all_rows = list_active_sessions(ctx, now, session=db.session()) + all_rows = list_active_sessions(ctx, now, session=session) total = len(all_rows) sliced = all_rows[(page - 1) * limit : page * limit] @@ -114,15 +116,16 @@ class AccountSessionsApi(Resource): class AccountSessionByIdApi(Resource): @auth_router.guard(scope=Scope.FULL, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, RevokeResponse, description="Session revoked") - def delete(self, session_id: str, *, auth_data: AuthData): + @with_session + def delete(self, session: Session, session_id: str, *, auth_data: AuthData): ctx = get_auth_ctx() # 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(session_id, ctx, session=db.session()): + if not token_belongs_to_subject(session_id, ctx, session=session): raise NotFound("session not found") - revoke_oauth_token(redis_client, session_id, session=db.session()) + revoke_oauth_token(redis_client, session_id, session=session) return RevokeResponse(status="revoked") diff --git a/api/controllers/openapi/apps.py b/api/controllers/openapi/apps.py index 882b55b7041..9a3e609b70a 100644 --- a/api/controllers/openapi/apps.py +++ b/api/controllers/openapi/apps.py @@ -6,11 +6,13 @@ import uuid as _uuid from typing import Any, cast from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import Conflict, NotFound, UnprocessableEntity from configs import dify_config from controllers.common.app_access import AppAccessFilter, resolve_app_access_filter from controllers.common.fields import Parameters +from controllers.common.session import with_session from controllers.common.wraps import RBACPermission, RBACResourceScope from controllers.openapi import openapi_ns from controllers.openapi._contract import accepts, returns @@ -28,7 +30,6 @@ from controllers.openapi.auth.composition import auth_router from controllers.openapi.auth.data import AuthData, CallerKind, RBACRequirement from controllers.service_api.app.error import AppUnavailableError from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict -from extensions.ext_database import db from libs.oauth_bearer import Scope, TokenType from models import App from models.enums import AppStatus @@ -56,7 +57,7 @@ _EMPTY_PARAMETERS: dict[str, Any] = { class AppReadResource(Resource): """Base for per-app read endpoints; subclasses call `_load()` for membership/exists checks.""" - def _load(self, app_id: str, workspace_id: str | None = None) -> App: + def _load(self, session: Session, app_id: str, workspace_id: str | None = None) -> App: try: parsed_uuid = _uuid.UUID(app_id) is_uuid = True @@ -66,13 +67,13 @@ class AppReadResource(Resource): if is_uuid: # ``str(parsed_uuid)`` normalises to the canonical dashed form. - app = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) + app = AppService.get_visible_app_by_id(str(parsed_uuid), 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(name=app_id, tenant_id=workspace_id, session=db.session()) + matches = AppService.find_visible_apps_by_name(session, name=app_id, tenant_id=workspace_id) if len(matches) == 0: raise NotFound("app not found") if len(matches) > 1: @@ -86,14 +87,14 @@ class AppReadResource(Resource): return app -def parameters_payload(app: App) -> dict: +def parameters_payload(app: App, *, session: Session) -> dict: """Mirrors service_api/app/app.py::AppParameterApi response body.""" - features_dict, user_input_form = resolve_app_config(app) + features_dict, user_input_form = resolve_app_config(app, session=session) parameters = get_parameters_from_feature_dict(features_dict=features_dict, user_input_form=user_input_form) return Parameters.model_validate(parameters).model_dump(mode="json") -def build_app_describe_response(app: App, fields: set[str] | None) -> AppDescribeResponse: +def build_app_describe_response(app: App, fields: set[str] | None, *, session: Session) -> AppDescribeResponse: """Public projection of an app (name / params / input schema) — never internal config.""" want_info = fields is None or "info" in fields want_params = fields is None or "parameters" in fields @@ -117,12 +118,12 @@ def build_app_describe_response(app: App, fields: set[str] | None) -> AppDescrib input_schema: dict[str, Any] | None = None if want_params: try: - parameters = parameters_payload(app) + parameters = parameters_payload(app, session=session) except AppUnavailableError: parameters = dict(_EMPTY_PARAMETERS) if want_schema: try: - input_schema = build_input_schema(app) + input_schema = build_input_schema(app, session=session) except AppUnavailableError: input_schema = dict(EMPTY_INPUT_SCHEMA) @@ -138,10 +139,11 @@ class AppDescribeApi(AppReadResource): ) @returns(200, AppDescribeResponse, description="App description") @accepts(query=AppDescribeQuery) - def get(self, app_id: str, *, auth_data: AuthData, query: AppDescribeQuery): + @with_session(write=False) + def get(self, session: Session, app_id: str, *, auth_data: AuthData, query: AppDescribeQuery): # describe is UUID-only (workspace_id query param dropped in #37212). - app = self._load(app_id) - return build_app_describe_response(app, query.fields) + app = self._load(session, app_id) + return build_app_describe_response(app, query.fields, session=session) @openapi_ns.route("/apps") @@ -149,7 +151,8 @@ class AppListApi(Resource): @auth_router.guard_workspace(scope=Scope.APPS_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, AppListResponse, description="App list") @accepts(query=AppListQuery) - def get(self, *, auth_data: AuthData, query: AppListQuery): + @with_session(write=False) + def get(self, session: Session, *, auth_data: AuthData, query: AppListQuery): workspace_id = query.workspace_id empty = AppListResponse(page=query.page, limit=query.limit, total=0, has_more=False, data=[]) @@ -173,11 +176,15 @@ class AppListApi(Resource): ) access_filter = AppAccessFilter.unrestricted() if apply_rbac_filter: - access_filter = resolve_app_access_filter(workspace_id, str(auth_data.account_id)) + access_filter = resolve_app_access_filter( + workspace_id, + str(auth_data.account_id), + session=session, + ) tenant_name: str | None = None if parsed_uuid is not None: - app: App | None = AppService.get_visible_app_by_id(str(parsed_uuid), session=db.session()) + app: App | None = AppService.get_visible_app_by_id(str(parsed_uuid), session) if app is None or str(app.tenant_id) != workspace_id: return empty if not _is_listable(app): @@ -188,7 +195,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(workspace_id, session=db.session()) + tenant_name = TenantService.get_tenant_name(workspace_id, session=session) item = AppListRow( id=str(app.id), name=app.name, @@ -215,13 +222,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, session) if pagination is None: return empty tenant_name = None if pagination.items: - tenant_name = TenantService.get_tenant_name(workspace_id, session=db.session()) + tenant_name = TenantService.get_tenant_name(workspace_id, session=session) items = [ AppListRow( diff --git a/api/controllers/openapi/apps_permitted_external.py b/api/controllers/openapi/apps_permitted_external.py index 353a1ec1cb3..e00ec5c87f7 100644 --- a/api/controllers/openapi/apps_permitted_external.py +++ b/api/controllers/openapi/apps_permitted_external.py @@ -8,8 +8,10 @@ EE blueprint chain so this module is unreachable there. from __future__ import annotations from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound +from controllers.common.session import with_session from controllers.openapi import openapi_ns from controllers.openapi._contract import accepts, returns from controllers.openapi._models import ( @@ -22,7 +24,6 @@ from controllers.openapi._models import ( from controllers.openapi.apps import build_app_describe_response from controllers.openapi.auth.composition import auth_router from controllers.openapi.auth.data import AuthData, Edition -from extensions.ext_database import db from libs.oauth_bearer import Scope, TokenType from models import App from models.enums import AppStatus @@ -40,7 +41,8 @@ class PermittedExternalAppsListApi(Resource): ) @returns(200, PermittedExternalAppsListResponse, description="Permitted external apps list") @accepts(query=PermittedExternalAppsListQuery) - def get(self, *, auth_data: AuthData, query: PermittedExternalAppsListQuery): + @with_session(write=False) + def get(self, session: Session, *, auth_data: AuthData, query: PermittedExternalAppsListQuery): page_result = list_permitted_apps( page=query.page, limit=query.limit, @@ -55,10 +57,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(page_result.app_ids, session=db.session()) + str(a.id): a for a in AppService.find_visible_apps_by_ids(page_result.app_ids, 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(tenant_ids, session=db.session())} + tenants_by_id = {str(t.id): t for t in TenantService.get_tenants_by_ids(tenant_ids, session=session)} items: list[AppListRow] = [] for app_id in page_result.app_ids: @@ -96,9 +98,10 @@ class PermittedExternalAppDescribeApi(Resource): ) @returns(200, AppDescribeResponse, description="Permitted external app description") @accepts(query=AppDescribeQuery) - def get(self, app_id: str, *, auth_data: AuthData, query: AppDescribeQuery): + @with_session(write=False) + def get(self, session: Session, app_id: str, *, auth_data: AuthData, query: AppDescribeQuery): # App already loaded and ACL-checked by the external_sso pipeline; project it. app = auth_data.app if app is None: raise NotFound("app not found") - return build_app_describe_response(app, query.fields) + return build_app_describe_response(app, query.fields, session=session) diff --git a/api/controllers/openapi/auth/prepare.py b/api/controllers/openapi/auth/prepare.py index 96cf9a8858f..8cf63fd1c9a 100644 --- a/api/controllers/openapi/auth/prepare.py +++ b/api/controllers/openapi/auth/prepare.py @@ -6,7 +6,7 @@ from flask import request from werkzeug.exceptions import Forbidden, InternalServerError, NotFound, Unauthorized from controllers.openapi.auth.data import AuthData, CallerKind -from extensions.ext_database import db +from core.db.session_factory import session_factory from models.account import AccountStatus, TenantStatus from models.enums import AppStatus, EndUserType from services.account_service import AccountService, TenantService @@ -23,7 +23,8 @@ def load_app(data: AuthData) -> None: uuid.UUID(app_id) except ValueError: raise NotFound("app not found") - app = AppService.get_app_by_id(app_id, session=db.session()) + with session_factory.create_session() as session: + app = AppService.get_app_by_id(app_id, session) if not app or app.status != AppStatus.NORMAL: raise NotFound("app not found") data.app = app @@ -34,7 +35,8 @@ 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(str(data.app.tenant_id), session=db.session()) + with session_factory.create_session() as session: + tenant = TenantService.get_tenant_by_id(str(data.app.tenant_id), session=session) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise Forbidden("workspace unavailable") data.tenant = tenant @@ -50,7 +52,8 @@ 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(workspace_id, session=db.session()) + with session_factory.create_session() as session: + tenant = TenantService.get_tenant_by_id(workspace_id, session=session) if tenant is None or tenant.status == TenantStatus.ARCHIVE: raise NotFound("workspace not found") data.tenant = tenant @@ -59,11 +62,12 @@ 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(str(data.account_id), session=db.session()) - if account is None: - raise Unauthorized("account not found") - if data.tenant: - account.current_tenant = data.tenant + with session_factory.create_session() as session: + account = AccountService.get_account_by_id(str(data.account_id), session=session) + if account is None: + raise Unauthorized("account not found") + if data.tenant: + account.set_current_tenant_with_session(data.tenant, session=session) data.caller = account data.caller_kind = CallerKind.ACCOUNT @@ -75,7 +79,8 @@ 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(str(data.account_id), str(data.tenant.id), session=db.session()) + with session_factory.create_session() as session: + role = TenantService.get_account_role_in_tenant(str(data.account_id), str(data.tenant.id), session=session) if role is None: return data.tenant_role = role diff --git a/api/controllers/openapi/workspaces.py b/api/controllers/openapi/workspaces.py index b53776a48da..b5fec3bd4ca 100644 --- a/api/controllers/openapi/workspaces.py +++ b/api/controllers/openapi/workspaces.py @@ -15,9 +15,11 @@ from itertools import starmap from urllib import parse from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, NotFound from configs import dify_config +from controllers.common.session import with_session from controllers.openapi import openapi_ns from controllers.openapi._contract import accepts, returns from controllers.openapi._errors import MemberLicenseExceeded, MemberLimitExceeded @@ -35,7 +37,6 @@ from controllers.openapi._models import ( ) from controllers.openapi.auth.composition import auth_router from controllers.openapi.auth.data import AuthData -from extensions.ext_database import db from libs.oauth_bearer import Scope, TokenType from models import Account, Tenant, TenantAccountJoin from models.account import TenantAccountRole, TenantStatus @@ -64,15 +65,15 @@ def _member_response(account: Account) -> MemberResponse: ) -def _load_tenant(workspace_id: str) -> Tenant: - tenant = TenantService.get_tenant_by_id(workspace_id, session=db.session()) +def _load_tenant(session: Session, workspace_id: str) -> Tenant: + tenant = TenantService.get_tenant_by_id(workspace_id, session=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(str(account_id), session=db.session()) if account_id else None +def _load_account(session: Session, account_id: object) -> Account: + account = AccountService.get_account_by_id(str(account_id), session=session) if account_id else None if account is None: raise RuntimeError("authenticated account_id has no Account row") return account @@ -94,8 +95,9 @@ def _check_member_invite_quota(tenant_id: str) -> None: 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(str(auth_data.account_id), session=db.session()) + @with_session(write=False) + def get(self, session: Session, *, auth_data: AuthData): + rows = TenantService.get_workspaces_for_account(str(auth_data.account_id), session=session) return WorkspaceListResponse(workspaces=list(starmap(_workspace_summary, rows))) @@ -104,8 +106,9 @@ class WorkspacesApi(Resource): 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(str(auth_data.account_id), workspace_id, session=db.session()) + @with_session(write=False) + def get(self, session: Session, workspace_id: str, *, auth_data: AuthData): + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=session) # 404 (not 403) on non-member so workspace IDs don't leak across tenants. if row is None: raise NotFound("workspace not found") @@ -125,15 +128,16 @@ class WorkspaceSwitchApi(Resource): @auth_router.guard_workspace(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, WorkspaceDetailResponse, description="Workspace detail") - def post(self, workspace_id: str, *, auth_data: AuthData): - account = _load_account(auth_data.account_id) + @with_session + def post(self, session: Session, workspace_id: str, *, auth_data: AuthData): + account = _load_account(session, auth_data.account_id) try: - TenantService.switch_tenant(account, workspace_id, session=db.session()) + TenantService.switch_tenant(account, workspace_id, session=session) except AccountNotLinkTenantError: raise NotFound("workspace not found") - row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=db.session()) + row = TenantService.find_workspace_for_account(str(auth_data.account_id), workspace_id, session=session) if row is None: raise NotFound("workspace not found") tenant, membership = row @@ -151,9 +155,10 @@ class WorkspaceMembersApi(Resource): @auth_router.guard_workspace(scope=Scope.WORKSPACE_READ, allowed_token_types=frozenset({TokenType.OAUTH_ACCOUNT})) @returns(200, MemberListResponse, description="Member list") @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()) + @with_session(write=False) + def get(self, session: Session, workspace_id: str, *, auth_data: AuthData, query: MemberListQuery): + tenant = _load_tenant(session, workspace_id) + members = TenantService.get_tenant_members(tenant, session=session) total = len(members) start = (query.page - 1) * query.limit page_items = members[start : start + query.limit] @@ -172,9 +177,10 @@ class WorkspaceMembersApi(Resource): ) @returns(201, MemberInviteResponse, description="Member invited") @accepts(body=MemberInvitePayload) - def post(self, workspace_id: str, *, auth_data: AuthData, body: MemberInvitePayload): - inviter = _load_account(auth_data.account_id) - tenant = _load_tenant(workspace_id) + @with_session + def post(self, session: Session, workspace_id: str, *, auth_data: AuthData, body: MemberInvitePayload): + inviter = _load_account(session, auth_data.account_id) + tenant = _load_tenant(session, workspace_id) _check_member_invite_quota(str(tenant.id)) @@ -185,7 +191,7 @@ class WorkspaceMembersApi(Resource): language=None, role=body.role, inviter=inviter, - session=db.session(), + session=session, ) except AccountAlreadyInTenantError as exc: raise BadRequest(str(exc)) @@ -197,7 +203,7 @@ class WorkspaceMembersApi(Resource): raise BadRequest(str(exc)) normalized_email = body.email.lower() - member = AccountService.get_account_by_email_with_case_fallback(normalized_email, session=db.session()) + member = AccountService.get_account_by_email_with_case_fallback(normalized_email, session=session) if member is None: # invite_new_member just created or fetched this account. raise RuntimeError("invited member missing from DB after invite") @@ -229,15 +235,16 @@ class WorkspaceMemberApi(Resource): allowed_roles=frozenset({TenantAccountRole.OWNER, TenantAccountRole.ADMIN}), ) @returns(200, MemberActionResponse, description="Member removed") - 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(member_id, session=db.session()) + @with_session + def delete(self, session: Session, workspace_id: str, member_id: str, *, auth_data: AuthData): + operator = _load_account(session, auth_data.account_id) + tenant = _load_tenant(session, workspace_id) + member = AccountService.get_account_by_id(member_id, session=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=session) except CannotOperateSelfError as exc: raise BadRequest(str(exc)) except NoPermissionError as exc: @@ -254,15 +261,24 @@ class WorkspaceMemberApi(Resource): ) @returns(200, MemberActionResponse, description="Role updated") @accepts(body=MemberRoleUpdatePayload) - 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(member_id, session=db.session()) + @with_session + def patch( + self, + session: Session, + workspace_id: str, + member_id: str, + *, + auth_data: AuthData, + body: MemberRoleUpdatePayload, + ): + operator = _load_account(session, auth_data.account_id) + tenant = _load_tenant(session, workspace_id) + member = AccountService.get_account_by_id(member_id, session=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=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 520b88248eb..3d45b46f066 100644 --- a/api/controllers/service_api/app/annotation.py +++ b/api/controllers/service_api/app/annotation.py @@ -5,12 +5,13 @@ from flask import request from flask_restx import Resource from flask_restx.api import HTTPStatus from pydantic import BaseModel, Field, TypeAdapter +from sqlalchemy.orm import Session from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console.wraps import edit_permission_required from controllers.service_api import service_api_ns from controllers.service_api.wraps import validate_app_token -from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.annotation_fields import ( Annotation, @@ -205,12 +206,13 @@ class AnnotationListApi(Resource): service_api_ns.models[AnnotationList.__name__], ) @validate_app_token - def get(self, app_model: App): + @with_session(write=False) + def get(self, session: Session, app_model: App): """List annotations for the application.""" 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, session=db.session() + app_model.id, query.page, query.limit, query.keyword, session ) annotation_models = TypeAdapter(list[Annotation]).validate_python(annotation_list, from_attributes=True) return AnnotationList( @@ -247,13 +249,12 @@ class AnnotationListApi(Resource): service_api_ns.models[Annotation.__name__], ) @validate_app_token - def post(self, app_model: App): + @with_session + def post(self, session: Session, app_model: App): """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, session=db.session() - ) + annotation = AppAnnotationService.insert_app_annotation_directly(insert_args, app_model.id, session) return dump_response(Annotation, annotation), HTTPStatus.CREATED @@ -287,14 +288,15 @@ class AnnotationUpdateDeleteApi(Resource): service_api_ns.models[Annotation.__name__], ) @validate_app_token + @with_session @edit_permission_required - def put(self, app_model: App, annotation_id: UUID): + def put(self, session: Session, app_model: App, annotation_id: UUID): """Update an existing annotation.""" payload = AnnotationCreatePayload.model_validate(service_api_ns.payload or {}) 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, session) return dump_response(Annotation, annotation) @service_api_ns.doc( @@ -319,10 +321,11 @@ class AnnotationUpdateDeleteApi(Resource): } ) @validate_app_token + @with_session @edit_permission_required - def delete(self, app_model: App, annotation_id: UUID): + def delete(self, session: Session, app_model: App, annotation_id: UUID): """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, session) return "", 204 diff --git a/api/controllers/service_api/app/app.py b/api/controllers/service_api/app/app.py index 60f83d7d070..edde21b9295 100644 --- a/api/controllers/service_api/app/app.py +++ b/api/controllers/service_api/app/app.py @@ -2,6 +2,7 @@ from typing import Any, cast from flask_restx import Resource from pydantic import Field +from sqlalchemy.orm import Session from controllers.common.agent_app_parameters import get_published_agent_app_feature_dict_and_user_input_form from controllers.common.fields import Parameters @@ -13,7 +14,7 @@ from core.app.app_config.common.parameters_mapping import get_parameters_from_fe 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 models.model import App, AppMode, load_annotation_reply_config from services.app_service import AppService @@ -32,9 +33,13 @@ class AppMetaResponse(ResponseModel): register_response_schema_models(service_api_ns, Parameters, AppMetaResponse, AppInfoResponse) -def _get_agent_app_feature_dict_and_user_input_form(app_model: App) -> tuple[dict[str, Any], list[dict[str, Any]]]: +def _get_agent_app_feature_dict_and_user_input_form( + app_model: App, + *, + session: Session, +) -> tuple[dict[str, Any], list[dict[str, Any]]]: try: - return get_published_agent_app_feature_dict_and_user_input_form(app_model) + return get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session) except AgentAppNotPublishedError: raise AgentNotPublishedError() except AgentAppGeneratorError: @@ -73,23 +78,31 @@ class AppParameterApi(Resource): Returns the input form parameters and configuration for the application. """ + session = db.session() features_dict: dict[str, Any] user_input_form: list[dict[str, Any]] if app_model.mode == AppMode.AGENT: - features_dict, user_input_form = _get_agent_app_feature_dict_and_user_input_form(app_model) + features_dict, user_input_form = _get_agent_app_feature_dict_and_user_input_form( + app_model, + session=session, + ) elif app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() features_dict = workflow.features_dict user_input_form = workflow.user_input_form(to_old_structure=True) else: - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if app_model_config is None: raise AppUnavailableError() - features_dict = cast(dict[str, Any], app_model_config.to_dict()) + annotation_reply = load_annotation_reply_config(session, app_model.id) + features_dict = cast( + dict[str, Any], + app_model_config.to_dict(annotation_reply=annotation_reply), + ) user_input_form = features_dict.get("user_input_form", []) diff --git a/api/controllers/service_api/app/audio.py b/api/controllers/service_api/app/audio.py index 28ac24ee7ac..1799eb68070 100644 --- a/api/controllers/service_api/app/audio.py +++ b/api/controllers/service_api/app/audio.py @@ -104,7 +104,12 @@ class AudioApi(Resource): file = request.files["file"] try: - response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=end_user.id) + response = AudioService.transcript_asr( + app_model=app_model, + file=file, + session=db.session(), + end_user=end_user.id, + ) return dump_response(AudioTranscriptResponse, response) except services.errors.app_model_config.AppModelConfigBrokenError: diff --git a/api/controllers/service_api/app/completion.py b/api/controllers/service_api/app/completion.py index 17ca4238395..a75e4b391ec 100644 --- a/api/controllers/service_api/app/completion.py +++ b/api/controllers/service_api/app/completion.py @@ -43,7 +43,6 @@ from core.errors.error import ( ) from core.helper.trace_id_helper import get_external_trace_id, get_trace_session_id, omit_trace_session_id_from_payload from enums.cloud_plan import CloudPlan -from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError from libs import helper from libs.helper import UUIDStrOrEmpty @@ -399,7 +398,7 @@ class ChatApi(Resource): app_model=app_model, conversation_id=payload.conversation_id, user=end_user, - session=db.session(), + session=session, ) response = AppGenerateService.generate( diff --git a/api/controllers/service_api/app/conversation.py b/api/controllers/service_api/app/conversation.py index 28bd1f1ee77..d7f9071dc53 100644 --- a/api/controllers/service_api/app/conversation.py +++ b/api/controllers/service_api/app/conversation.py @@ -21,6 +21,7 @@ from fields._value_type_serializer import serialize_value_type from fields.base import ResponseModel from fields.conversation_fields import ( ConversationInfiniteScrollPagination, + ConversationResponseSource, SimpleConversation, ) from graphon.variables.types import SegmentType @@ -204,7 +205,13 @@ class ConversationApi(Resource): sort_by=query_args.sort_by, ) adapter = TypeAdapter(SimpleConversation) - conversations = [adapter.validate_python(item, from_attributes=True) for item in pagination.data] + conversations = [ + adapter.validate_python( + ConversationResponseSource(item, session=session), + from_attributes=True, + ) + for item in pagination.data + ] return ConversationInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, data=conversations ).model_dump(mode="json") @@ -294,10 +301,11 @@ class ConversationRenameApi(Resource): payload = ConversationRenamePayload.model_validate(service_api_ns.payload or {}) try: + session = db.session() conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=session ) - return dump_response(SimpleConversation, conversation) + return dump_response(SimpleConversation, ConversationResponseSource(conversation, session=session)) except services.errors.conversation.ConversationNotExistsError: raise NotFound("Conversation Not Exists.") diff --git a/api/controllers/service_api/app/message.py b/api/controllers/service_api/app/message.py index 7dd02665646..fb0c511ffa3 100644 --- a/api/controllers/service_api/app/message.py +++ b/api/controllers/service_api/app/message.py @@ -17,7 +17,7 @@ from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate 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.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import MessageInfiniteScrollPagination, MessageListItem from models.enums import FeedbackRating from models.model import App, AppMode, EndUser @@ -111,11 +111,15 @@ class MessageListApi(Resource): first_id = query_args.first_id or None try: + session = db.session() pagination = MessageService.pagination_by_first_id( - app_model, end_user, conversation_id, first_id, query_args.limit, session=db.session() + app_model, end_user, conversation_id, first_id, query_args.limit, session=session ) adapter = TypeAdapter(MessageListItem) - items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] + items = [ + adapter.validate_python(MessageResponseSource(message, session=session), from_attributes=True) + for message in pagination.data + ] return MessageInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, data=items ).model_dump(mode="json") diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index a8a47f4819f..56762ca1720 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -24,7 +24,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) -from controllers.console.app.wraps import with_session +from controllers.common.session import with_session from controllers.console.wraps import edit_permission_required from controllers.service_api import service_api_ns from controllers.service_api.dataset.error import DatasetInUseError, DatasetNameDuplicateError, InvalidActionError @@ -34,9 +34,8 @@ from controllers.service_api.wraps import ( ) from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager from core.rag.index_processor.constant.index_type import IndexTechniqueType -from extensions.ext_database import db from fields.base import ResponseModel -from fields.dataset_fields import DatasetDetailResponse +from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source from graphon.model_runtime.entities.model_entities import ModelType from libs.helper import dump_response from libs.login import current_user @@ -93,8 +92,10 @@ _SERVICE_DATASET_DETAIL_EXCLUDE = {"permission_keys"} _SERVICE_DATASET_LIST_EXCLUDE = {"data": {"__all__": _SERVICE_DATASET_DETAIL_EXCLUDE}} -def _dump_service_dataset_detail(dataset: Any) -> dict[str, Any]: - return DatasetDetailResponse.model_validate(dataset, from_attributes=True).model_dump( +def _dump_service_dataset_detail(dataset: Any, *, session: Session) -> dict[str, Any]: + return DatasetDetailResponse.model_validate( + dataset_detail_response_source(dataset, session=session), from_attributes=True + ).model_dump( mode="json", exclude=_SERVICE_DATASET_DETAIL_EXCLUDE, ) @@ -405,7 +406,8 @@ class DatasetListApi(DatasetApiResource): "Datasets retrieved successfully", service_api_ns.models[DatasetListResponse.__name__], ) - def get(self, tenant_id): + @with_session(write=False) + def get(self, session: Session, tenant_id): """Resource for getting datasets.""" query_params: dict[str, str | list[str]] = dict(request.args.to_dict()) if "tag_ids" in request.args: @@ -416,7 +418,7 @@ class DatasetListApi(DatasetApiResource): datasets, total = DatasetService.get_datasets( query.page, query.limit, - db.session(), + session, tenant_id, current_user, query.keyword, @@ -436,7 +438,7 @@ class DatasetListApi(DatasetApiResource): for embedding_model in embedding_models: model_names.append(f"{embedding_model.model}:{embedding_model.provider.provider}") - data = [_dump_service_dataset_detail(dataset) for dataset in datasets] + data = [_dump_service_dataset_detail(dataset, session=session) for dataset in datasets] for item in data: if item["indexing_technique"] == IndexTechniqueType.HIGH_QUALITY and item["embedding_model_provider"]: item["embedding_model_provider"] = str(ModelProviderID(item["embedding_model_provider"])) @@ -538,7 +540,7 @@ class DatasetListApi(DatasetApiResource): ) initialize_created_app_rbac_access_task.delay(tenant_id, current_user.id, dataset_id=dataset.id) - return _dump_service_dataset_detail(dataset), 200 + return _dump_service_dataset_detail(dataset, session=session), 200 @service_api_ns.route("/datasets/") @@ -574,16 +576,17 @@ class DatasetApi(DatasetApiResource): "Dataset retrieved successfully", service_api_ns.models[DatasetDetailWithPartialMembersResponse.__name__], ) - def get(self, _, dataset_id: UUID): + @with_session(write=False) + def get(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, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - data = _dump_service_dataset_detail(dataset) + data = _dump_service_dataset_detail(dataset, session=session) # check embedding setting assert isinstance(current_user, Account) cid = current_user.current_tenant_id @@ -612,7 +615,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, session) data.update({"partial_member_list": part_users_list}) return _dump_service_dataset_with_partial_members(data), 200 @@ -651,7 +654,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, session) if dataset is None: raise NotFound("Dataset not found.") @@ -692,7 +695,7 @@ class DatasetApi(DatasetApiResource): dataset, str(payload.permission) if payload.permission else None, payload.partial_member_list, - session=db.session(), + session=session, ) dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user, session=session) @@ -700,19 +703,19 @@ class DatasetApi(DatasetApiResource): if dataset is None: raise NotFound("Dataset not found.") - result_data = _dump_service_dataset_detail(dataset) + result_data = _dump_service_dataset_detail(dataset, session=session) assert isinstance(current_user, Account) tenant_id = current_user.current_tenant_id 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, 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, 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, session) result_data.update({"partial_member_list": partial_member_list}) return _dump_service_dataset_with_partial_members(result_data), 200 @@ -745,7 +748,8 @@ class DatasetApi(DatasetApiResource): } ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def delete(self, _, dataset_id: UUID): + @with_session + def delete(self, session: Session, _, dataset_id: UUID): """ Deletes a dataset given its ID. @@ -765,8 +769,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, session): + DatasetPermissionService.clear_partial_member_list(dataset_id_str, session) return "", 204 else: raise NotFound("Dataset not found.") @@ -812,7 +816,14 @@ class DocumentStatusApi(DatasetApiResource): } ) @service_api_ns.expect(service_api_ns.models[DocumentStatusPayload.__name__]) - def patch(self, tenant_id, dataset_id: UUID, action: Literal["enable", "disable", "archive", "un_archive"]): + @with_session + def patch( + self, + session: Session, + tenant_id, + dataset_id: UUID, + action: Literal["enable", "disable", "archive", "un_archive"], + ): """ Batch update document status. @@ -831,14 +842,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, 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, session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -850,7 +861,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, session) except services.errors.document.DocumentIndexingError as e: raise InvalidActionError(str(e)) except ValueError as e: @@ -882,12 +893,13 @@ class DatasetTagsApi(DatasetApiResource): "Tags retrieved successfully", service_api_ns.models[KnowledgeTagListResponse.__name__], ) - def get(self, _): + @with_session(write=False) + def get(self, session: Session, _): """Get all knowledge type tags.""" assert isinstance(current_user, Account) cid = current_user.current_tenant_id assert cid is not None - tags = TagService.get_tags("knowledge", cid, session=db.session()) + tags = TagService.get_tags("knowledge", cid, session=session) return dump_response(KnowledgeTagListResponse, tags), 200 @service_api_ns.doc( @@ -913,14 +925,15 @@ class DatasetTagsApi(DatasetApiResource): "Tag created successfully", service_api_ns.models[KnowledgeTagResponse.__name__], ) - def post(self, _): + @with_session + def post(self, session: Session, _): """Add a knowledge type tag.""" assert isinstance(current_user, Account) if not (current_user.has_edit_permission or current_user.is_dataset_editor): 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), session) response = KnowledgeTagResponse(id=tag.id, name=tag.name, type=tag.type, binding_count="0") return response.model_dump(mode="json"), 200 @@ -948,7 +961,8 @@ class DatasetTagsApi(DatasetApiResource): "Tag updated successfully", service_api_ns.models[KnowledgeTagResponse.__name__], ) - def patch(self, _): + @with_session + def patch(self, session: Session, _): assert isinstance(current_user, Account) if not (current_user.has_edit_permission or current_user.is_dataset_editor): raise Forbidden() @@ -956,10 +970,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, 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, session, tag_type=TagType.KNOWLEDGE) response = KnowledgeTagResponse(id=tag.id, name=tag.name, type=tag.type, binding_count=str(binding_count)) return response.model_dump(mode="json"), 200 @@ -983,10 +997,11 @@ class DatasetTagsApi(DatasetApiResource): } ) @edit_permission_required - def delete(self, _): + @with_session + def delete(self, session: Session, _): """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, session, tag_type=TagType.KNOWLEDGE) return "", 204 @@ -1011,7 +1026,8 @@ class DatasetTagBindingApi(DatasetApiResource): 403: "Forbidden - insufficient permissions", } ) - def post(self, _): + @with_session + def post(self, session: Session, _): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator assert isinstance(current_user, Account) if not (current_user.has_edit_permission or current_user.is_dataset_editor): @@ -1020,7 +1036,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(), + session, ) return "", 204 @@ -1046,7 +1062,8 @@ class DatasetTagUnbindingApi(DatasetApiResource): 403: "Forbidden - insufficient permissions", } ) - def post(self, _): + @with_session + def post(self, session: Session, _): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator assert isinstance(current_user, Account) if not (current_user.has_edit_permission or current_user.is_dataset_editor): @@ -1055,7 +1072,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(), + session, ) return "", 204 @@ -1085,14 +1102,13 @@ class DatasetTagsBindingStatusApi(DatasetApiResource): "Tags retrieved successfully", service_api_ns.models[DatasetBoundTagListResponse.__name__], ) - def get(self, _, *args, **kwargs): + @with_session(write=False) + def get(self, session: Session, _, *args, **kwargs): """Get all knowledge type tags.""" dataset_id = kwargs.get("dataset_id") 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() - ) + tags = TagService.get_tags_by_target_id("knowledge", current_user.current_tenant_id, str(dataset_id), session) response = DatasetBoundTagListResponse( data=[DatasetBoundTagResponse(id=tag.id, name=tag.name) for tag in tags], total=len(tags), diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 47e77fd5a34..44a7169fc91 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -6,6 +6,7 @@ deprecated in generated API docs so clients migrate toward the canonical paths. """ import json +from collections.abc import Mapping from contextlib import ExitStack from copy import deepcopy from typing import Annotated, Any, Literal, Self, override @@ -23,6 +24,7 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema from sqlalchemy import desc, func, select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -42,6 +44,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session from controllers.service_api import service_api_ns from controllers.service_api.app.error import ProviderNotInitializeError from controllers.service_api.dataset.error import ( @@ -65,6 +68,8 @@ from fields.document_fields import ( DocumentMetadataResponse, DocumentResponse, DocumentStatusListResponse, + document_response, + document_responses, normalize_enum, ) from libs.helper import dump_response @@ -291,6 +296,13 @@ class DocumentAndBatchResponse(ResponseModel): batch: str +def _document_and_batch_response(document: Document, batch: str, *, session: Session) -> dict[str, Any]: + return dump_response( + DocumentAndBatchResponse, + {"document": document_response(document, session=session), "batch": batch}, + ) + + # Use SkipJsonSchema to support 3 metadata modes class DocumentDetailResponse(ResponseModel): id: str @@ -357,14 +369,14 @@ register_response_schema_models( ) -def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Document, str]: +def _create_document_by_text(session: Session, tenant_id: str, dataset_id: UUID) -> tuple[Document, str]: """Create a document from text for both canonical and legacy routes.""" payload = DocumentTextCreatePayload.model_validate(service_api_ns.payload or {}) args = payload.model_dump(exclude_none=True) dataset_id_str = str(dataset_id) - tenant_id_str = tenant_id - dataset = db.session.scalar( + tenant_id_str = str(tenant_id) + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id_str, Dataset.id == dataset_id_str).limit(1) ) @@ -414,9 +426,11 @@ def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Document dataset=dataset, knowledge_config=knowledge_config, account=current_user, - dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, + dataset_process_rule=dataset.get_latest_process_rule(session=session) + if "process_rule" not in args + else None, created_from="api", - session=db.session(), + session=session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -425,10 +439,12 @@ def _create_document_by_text(tenant_id: str, dataset_id: UUID) -> tuple[Document return document, batch -def _update_document_by_text(tenant_id: str, dataset_id: UUID, document_id: UUID) -> tuple[Document, str]: +def _update_document_by_text( + session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID +) -> tuple[Document, str]: """Update a document from text for both canonical and legacy routes.""" payload = DocumentTextUpdate.model_validate(service_api_ns.payload or {}) - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == str(dataset_id)).limit(1) ) args = payload.model_dump(exclude_none=True) @@ -474,9 +490,11 @@ def _update_document_by_text(tenant_id: str, dataset_id: UUID, document_id: UUID dataset=dataset, knowledge_config=knowledge_config, account=current_user, - dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, + dataset_process_rule=dataset.get_latest_process_rule(session=session) + if "process_rule" not in args + else None, created_from="api", - session=db.session(), + session=session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -524,10 +542,11 @@ class DocumentAddByTextApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_resource_check("documents", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID): """Create document by text.""" - document, batch = _create_document_by_text(tenant_id=tenant_id, dataset_id=dataset_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _create_document_by_text(session=session, tenant_id=tenant_id, dataset_id=dataset_id) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.route("/datasets//document/create_by_text") @@ -557,10 +576,11 @@ class DeprecatedDocumentAddByTextApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_resource_check("documents", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID): """Create document by text through the deprecated underscore alias.""" - document, batch = _create_document_by_text(tenant_id=tenant_id, dataset_id=dataset_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _create_document_by_text(session=session, tenant_id=tenant_id, dataset_id=dataset_id) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.route("/datasets//documents//update-by-text") @@ -602,10 +622,13 @@ class DocumentUpdateByTextApi(DatasetApiResource): ) @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): """Update document by text.""" - document, batch = _update_document_by_text(tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _update_document_by_text( + session=session, tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id + ) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.route("/datasets//documents//update_by_text") @@ -634,10 +657,13 @@ class DeprecatedDocumentUpdateByTextApi(DatasetApiResource): ) @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): """Update document by text through the deprecated underscore alias.""" - document, batch = _update_document_by_text(tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _update_document_by_text( + session=session, tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id + ) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.route( @@ -694,9 +720,10 @@ class DocumentAddByFileApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_resource_check("documents", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id, dataset_id: UUID): + @with_session + def post(self, session: Session, tenant_id, dataset_id: UUID): """Create document by upload file.""" - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id).limit(1) ) @@ -767,7 +794,7 @@ class DocumentAddByFileApi(DatasetApiResource): knowledge_config = KnowledgeConfig.model_validate(args) DocumentService.document_create_args_validate(knowledge_config) - dataset_process_rule = dataset.latest_process_rule if "process_rule" not in args else None + dataset_process_rule = dataset.get_latest_process_rule(session=session) if "process_rule" not in args else None if not knowledge_config.original_document_id and not dataset_process_rule and not knowledge_config.process_rule: raise ValueError("process_rule is required.") @@ -775,23 +802,24 @@ class DocumentAddByFileApi(DatasetApiResource): documents, batch = DocumentService.save_document_with_dataset_id( dataset=dataset, knowledge_config=knowledge_config, - account=dataset.created_by_account, + account=dataset.get_created_by_account(session=session), dataset_process_rule=dataset_process_rule, created_from="api", - session=db.session(), + session=session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) document = documents[0] - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + return _document_and_batch_response(document, batch, session=session), 200 -def _update_document_by_file(tenant_id: str, dataset_id: UUID, document_id: UUID) -> tuple[Document, str]: +def _update_document_by_file( + session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID +) -> tuple[Document, str]: """Update a document from an uploaded file for canonical and deprecated routes.""" dataset_id_str = str(dataset_id) - tenant_id_str = tenant_id - dataset = db.session.scalar( - select(Dataset).where(Dataset.tenant_id == tenant_id_str, Dataset.id == dataset_id_str).limit(1) + dataset = session.scalar( + select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -852,10 +880,12 @@ def _update_document_by_file(tenant_id: str, dataset_id: UUID, document_id: UUID documents, _ = DocumentService.save_document_with_dataset_id( dataset=dataset, knowledge_config=knowledge_config, - account=dataset.created_by_account, - dataset_process_rule=dataset.latest_process_rule if "process_rule" not in args else None, + account=dataset.get_created_by_account(session=session), + dataset_process_rule=dataset.get_latest_process_rule(session=session) + if "process_rule" not in args + else None, created_from="api", - session=db.session(), + session=session, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -912,10 +942,13 @@ class DeprecatedDocumentUpdateByFileApi(DatasetApiResource): ) @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): """Update document by file through the deprecated file-update aliases.""" - document, batch = _update_document_by_file(tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _update_document_by_file( + session=session, tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id + ) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.route("/datasets//documents") @@ -945,11 +978,12 @@ class DocumentListApi(DatasetApiResource): @service_api_ns.response( 200, "Documents retrieved successfully", service_api_ns.models[DocumentListResponse.__name__] ) - def get(self, tenant_id, dataset_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id, dataset_id: UUID): dataset_id_str = str(dataset_id) tenant_id = str(tenant_id) query_params = query_params_from_request(DocumentListQuery) - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -967,7 +1001,7 @@ class DocumentListApi(DatasetApiResource): query = query.order_by(desc(Document.created_at), desc(Document.position)) paginated_documents = paginate_query( - query, page=query_params.page, per_page=query_params.limit, max_per_page=100 + query, session=session, page=query_params.page, per_page=query_params.limit, max_per_page=100 ) documents = paginated_documents.items @@ -975,11 +1009,11 @@ class DocumentListApi(DatasetApiResource): documents=documents, dataset=dataset, tenant_id=tenant_id, - session=db.session(), + session=session, ) response = { - "data": documents, + "data": document_responses(documents, session=session), "has_more": len(documents) == query_params.limit, "limit": query_params.limit, "total": paginated_documents.total, @@ -1020,7 +1054,8 @@ class DocumentBatchDownloadZipApi(DatasetApiResource): ) @service_api_ns.response(200, "ZIP archive generated successfully") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id, dataset_id: UUID): + @with_session(write=False) + def post(self, session: Session, tenant_id, dataset_id: UUID): payload = DocumentBatchDownloadZipPayload.model_validate(service_api_ns.payload or {}) upload_files, download_name = DocumentService.prepare_document_batch_download_zip( @@ -1028,7 +1063,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=session, ) with ExitStack() as stack: @@ -1076,23 +1111,24 @@ class DocumentIndexingStatusApi(DatasetApiResource): "Indexing status retrieved successfully", service_api_ns.models[DocumentStatusListResponse.__name__], ) - def get(self, tenant_id, dataset_id: UUID, batch: str): + @with_session(write=False) + def get(self, session: Session, tenant_id, dataset_id: UUID, batch: str): dataset_id_str = str(dataset_id) tenant_id = str(tenant_id) # get dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) 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, session) if not documents: raise NotFound("Documents not found.") documents_status = [] for document in documents: completed_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), DocumentSegment.document_id == str(document.id), @@ -1102,7 +1138,7 @@ class DocumentIndexingStatusApi(DatasetApiResource): or 0 ) total_segments = ( - db.session.scalar( + session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, @@ -1160,9 +1196,12 @@ class DocumentDownloadApi(DatasetApiResource): service_api_ns.models[UrlResponse.__name__], ) @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()) + @with_session(write=False) + def get(self, session: Session, tenant_id, dataset_id: UUID, document_id: UUID): + dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), str(tenant_id), session=session) + if not dataset: + raise NotFound("Dataset not found.") + document = DocumentService.get_document(dataset.id, str(document_id), session=session) if not document: raise NotFound("Document not found.") @@ -1170,9 +1209,7 @@ class DocumentDownloadApi(DatasetApiResource): if document.tenant_id != str(tenant_id): raise Forbidden("No permission.") - return UrlResponse(url=DocumentService.get_document_download_url(document, db.session())).model_dump( - mode="json" - ) + return UrlResponse(url=DocumentService.get_document_download_url(document, session)).model_dump(mode="json") @service_api_ns.route("/datasets//documents/") @@ -1219,13 +1256,16 @@ class DocumentApi(DatasetApiResource): "Document retrieved successfully", service_api_ns.models[DocumentDetailResponse.__name__], ) - def get(self, tenant_id, dataset_id: UUID, document_id: UUID): + @with_session(write=False) + def get(self, session: Session, tenant_id, dataset_id: UUID, document_id: UUID): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = self.get_dataset(dataset_id_str, tenant_id) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, str(tenant_id), session=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=session) if not document: raise NotFound("Document not found.") @@ -1250,17 +1290,23 @@ class DocumentApi(DatasetApiResource): document_id=document_id_str, dataset_id=dataset_id_str, tenant_id=tenant_id, - session=db.session(), + session=session, ) if metadata == "only": response_include = {"id", "doc_type", "doc_metadata"} - response = {"id": document.id, "doc_type": document.doc_type, "doc_metadata": document.doc_metadata_details} + response = { + "id": document.id, + "doc_type": document.doc_type, + "doc_metadata": document.get_doc_metadata_details(session=session), + } elif metadata == "without": + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, session) response_exclude = {"doc_type", "doc_metadata"} - 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 + document_process_rule = document.get_dataset_process_rule(session=session) + document_process_rules: Mapping[str, Any] = document_process_rule.to_dict() if document_process_rule else {} + data_source_info = document.get_data_source_detail_dict(session=session) + segment_count = document.get_segment_count(session=session) response = { "id": document.id, "position": document.position, @@ -1283,9 +1329,9 @@ class DocumentApi(DatasetApiResource): "disabled_at": int(document.disabled_at.timestamp()) if document.disabled_at else None, "disabled_by": document.disabled_by, "archived": document.archived, - "segment_count": document.segment_count, - "average_segment_length": document.average_segment_length, - "hit_count": document.hit_count, + "segment_count": segment_count, + "average_segment_length": (document.word_count or 0) // segment_count if segment_count else 0, + "hit_count": document.get_hit_count(session=session), "display_status": document.display_status, "doc_form": document.doc_form, "doc_language": document.doc_language, @@ -1293,9 +1339,11 @@ 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()) - document_process_rules = document.dataset_process_rule.to_dict() if document.dataset_process_rule else {} - data_source_info = document.data_source_detail_dict + dataset_process_rules = DatasetService.get_process_rules(dataset_id_str, session) + document_process_rule = document.get_dataset_process_rule(session=session) + document_process_rules = document_process_rule.to_dict() if document_process_rule else {} + data_source_info = document.get_data_source_detail_dict(session=session) + segment_count = document.get_segment_count(session=session) response = { "id": document.id, "position": document.position, @@ -1319,10 +1367,10 @@ class DocumentApi(DatasetApiResource): "disabled_by": document.disabled_by, "archived": document.archived, "doc_type": document.doc_type, - "doc_metadata": document.doc_metadata_details, - "segment_count": document.segment_count, - "average_segment_length": document.average_segment_length, - "hit_count": document.hit_count, + "doc_metadata": document.get_doc_metadata_details(session=session), + "segment_count": segment_count, + "average_segment_length": (document.word_count or 0) // segment_count if segment_count else 0, + "hit_count": document.get_hit_count(session=session), "display_status": document.display_status, "doc_form": document.doc_form, "doc_language": document.doc_language, @@ -1372,10 +1420,13 @@ class DocumentApi(DatasetApiResource): ) @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def patch(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def patch(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): """Update document by file on the canonical document resource.""" - document, batch = _update_document_by_file(tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id) - return dump_response(DocumentAndBatchResponse, {"document": document, "batch": batch}), 200 + document, batch = _update_document_by_file( + session=session, tenant_id=tenant_id, dataset_id=dataset_id, document_id=document_id + ) + return _document_and_batch_response(document, batch, session=session), 200 @service_api_ns.doc( summary="Delete Document", @@ -1400,21 +1451,22 @@ class DocumentApi(DatasetApiResource): } ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def delete(self, tenant_id, dataset_id: UUID, document_id: UUID): + @with_session + def delete(self, session: Session, tenant_id, dataset_id: UUID, document_id: UUID): """Delete document.""" document_id_str = str(document_id) dataset_id_str = str(dataset_id) tenant_id = str(tenant_id) # get dataset info - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) 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=session) # 404 if document not found if document is None: @@ -1426,7 +1478,7 @@ class DocumentApi(DatasetApiResource): try: # delete document - DocumentService.delete_document(document, db.session()) + DocumentService.delete_document(document, session) except services.errors.document.DocumentIndexingError: raise DocumentIndexingError("Cannot delete document during indexing.") diff --git a/api/controllers/service_api/dataset/hit_testing.py b/api/controllers/service_api/dataset/hit_testing.py index 86a64829e22..1028f174d71 100644 --- a/api/controllers/service_api/dataset/hit_testing.py +++ b/api/controllers/service_api/dataset/hit_testing.py @@ -61,7 +61,7 @@ class HitTestingApi(DatasetApiResource, DatasetsHitTestingBase): Tests retrieval performance for the specified dataset. """ dataset_id_str = str(dataset_id) - dataset = self.get_and_validate_dataset(dataset_id_str) + dataset = self.get_and_validate_dataset(session, dataset_id_str) args = self.parse_args(service_api_ns.payload) self.hit_testing_args_check(args) diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index 1d793583cc2..912071806a5 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -2,13 +2,14 @@ from typing import Literal from uuid import UUID from flask_login import current_user +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.common.controller_schemas import MetadataUpdatePayload from controllers.common.schema import register_response_schema_models, register_schema_model, register_schema_models +from controllers.common.session import with_session from controllers.service_api import service_api_ns from controllers.service_api.wraps import DatasetApiResource, cloud_edition_billing_rate_limit_check -from extensions.ext_database import db from fields.dataset_fields import ( DatasetMetadataActionResponse, DatasetMetadataBuiltInFieldsResponse, @@ -76,17 +77,18 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): 201, "Metadata created successfully", service_api_ns.models[DatasetMetadataResponse.__name__] ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id, dataset_id: UUID): + @with_session + def post(self, session: Session, tenant_id, dataset_id: UUID): """Create metadata for a dataset.""" 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, 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, session) - metadata = MetadataService.create_metadata(dataset_id_str, metadata_args, session=db.session()) + metadata = MetadataService.create_metadata(dataset_id_str, metadata_args, session=session) return dump_response(DatasetMetadataResponse, metadata), 201 @service_api_ns.doc( @@ -113,13 +115,14 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): @service_api_ns.response( 200, "Metadata retrieved successfully", service_api_ns.models[DatasetMetadataListResponse.__name__] ) - def get(self, tenant_id, dataset_id: UUID): + @with_session(write=False) + def get(self, session: Session, 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, session) if dataset is None: raise NotFound("Dataset not found.") - metadata = MetadataService.get_dataset_metadatas(dataset, session=db.session()) + metadata = MetadataService.get_dataset_metadatas(dataset, session) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -148,20 +151,19 @@ class DatasetMetadataServiceApi(DatasetApiResource): 200, "Metadata updated successfully", service_api_ns.models[DatasetMetadataResponse.__name__] ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def patch(self, tenant_id, dataset_id: UUID, metadata_id: UUID): + @with_session + def patch(self, session: Session, tenant_id, dataset_id: UUID, metadata_id: UUID): """Update metadata name.""" payload = MetadataUpdatePayload.model_validate(service_api_ns.payload or {}) 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, 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, session) - metadata = MetadataService.update_metadata_name( - dataset_id_str, metadata_id_str, payload.name, session=db.session() - ) + metadata = MetadataService.update_metadata_name(dataset_id_str, metadata_id_str, payload.name, session=session) return dump_response(DatasetMetadataResponse, metadata), 200 @service_api_ns.doc( @@ -187,16 +189,17 @@ class DatasetMetadataServiceApi(DatasetApiResource): ) @service_api_ns.response(204, "Metadata deleted successfully") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def delete(self, tenant_id, dataset_id: UUID, metadata_id: UUID): + @with_session + def delete(self, session: Session, tenant_id, dataset_id: UUID, metadata_id: UUID): """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, 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, session) - MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session=db.session()) + MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session) return "", 204 @@ -256,19 +259,20 @@ class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource): 200, "Action completed successfully", service_api_ns.models[DatasetMetadataActionResponse.__name__] ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id, dataset_id: UUID, action: Literal["enable", "disable"]): + @with_session + def post(self, session: Session, 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, 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, session) match action: case "enable": - MetadataService.enable_built_in_field(dataset, session=db.session()) + MetadataService.enable_built_in_field(dataset, session) case "disable": - MetadataService.disable_built_in_field(dataset, session=db.session()) + MetadataService.disable_built_in_field(dataset, session) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 @@ -302,16 +306,17 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): service_api_ns.models[DatasetMetadataActionResponse.__name__], ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id, dataset_id: UUID): + @with_session + def post(self, session: Session, 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, 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, session) metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {}) - MetadataService.update_documents_metadata(dataset, metadata_args, session=db.session()) + MetadataService.update_documents_metadata(dataset, metadata_args, session=session) return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 diff --git a/api/controllers/service_api/dataset/segment.py b/api/controllers/service_api/dataset/segment.py index e911c454c9e..5ed098e94b2 100644 --- a/api/controllers/service_api/dataset/segment.py +++ b/api/controllers/service_api/dataset/segment.py @@ -3,6 +3,7 @@ from uuid import UUID from pydantic import BaseModel, Field, ValidationError, field_validator from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from configs import dify_config @@ -13,6 +14,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.common.session import with_session from controllers.service_api import service_api_ns from controllers.service_api.app.error import ProviderNotInitializeError from controllers.service_api.wraps import ( @@ -24,7 +26,6 @@ from controllers.service_api.wraps import ( from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.model_manager import ModelManager from core.rag.index_processor.constant.index_type import IndexTechniqueType -from extensions.ext_database import db from fields.base import ResponseModel from fields.segment_fields import ( ChildChunkDetailResponse, @@ -129,7 +130,7 @@ register_response_schema_models( def _get_segment_for_document( - dataset: Dataset, document: Document, segment_id: str + session: Session, dataset: Dataset, document: Document, segment_id: str ) -> tuple[SegmentRef, DocumentSegment]: dataset_ref = DatasetRefService.create_dataset_ref(dataset) document_ref = DatasetRefService.create_document_ref(dataset_ref, document) @@ -137,7 +138,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, db.session()) + segment = SegmentService.get_segment_by_ref(segment_ref, session=session) if not segment: raise NotFound("Segment not found.") return segment_ref, segment @@ -179,19 +180,20 @@ class SegmentApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): _, current_tenant_id = current_account_with_tenant() """Create single segment.""" dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: 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=session) if not document: raise NotFound("Document not found.") if document.indexing_status != "completed": @@ -227,17 +229,17 @@ 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, 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, session=db.session() + segment_ids=segment_ids, dataset_id=dataset_id_str, session=session ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} response = { - "data": segment_responses_with_summaries(segments, summaries), + "data": segment_responses_with_summaries(segments, summaries, session=session), "doc_form": document.doc_form, } return dump_response(SegmentCreateListResponse, response), 200 @@ -266,7 +268,8 @@ class SegmentApi(DatasetApiResource): "Segments retrieved successfully", service_api_ns.models[SegmentListResponse.__name__], ) - def get(self, tenant_id: str, dataset_id: UUID, document_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID): _, current_tenant_id = current_account_with_tenant() """Get segments.""" # check dataset @@ -278,14 +281,14 @@ class SegmentApi(DatasetApiResource): page = args.page limit = args.limit dataset_id_str = str(dataset_id) - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: 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=session) if not document: raise NotFound("Document not found.") # check embedding model setting @@ -306,6 +309,7 @@ class SegmentApi(DatasetApiResource): raise ProviderNotInitializeError(ex.description) segments, total = SegmentService.get_segments( + session=session, document_id=document_id_str, tenant_id=current_tenant_id, status_list=args.status, @@ -317,12 +321,12 @@ 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, session=db.session() + segment_ids=segment_ids, dataset_id=dataset_id_str, session=session ) summaries = {chunk_id: record.summary_content for chunk_id, record in summary_records.items()} response = { - "data": segment_responses_with_summaries(segments, summaries), + "data": segment_responses_with_summaries(segments, summaries, session=session), "doc_form": document.doc_form, "total": total, "has_more": len(segments) == limit, @@ -354,11 +358,12 @@ class DatasetSegmentApi(DatasetApiResource): } ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def delete(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session + def delete(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): current_account_with_tenant() dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -367,12 +372,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=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()) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) + SegmentService.delete_segment(segment, document, dataset, session) return "", 204 @service_api_ns.doc( @@ -397,11 +402,12 @@ class DatasetSegmentApi(DatasetApiResource): @service_api_ns.response(200, "Segment updated successfully", service_api_ns.models[SegmentDetailResponse.__name__]) @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): _, current_tenant_id = current_account_with_tenant() dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -410,7 +416,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=session) if not document: raise NotFound("Document not found.") if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -430,16 +436,18 @@ class DatasetSegmentApi(DatasetApiResource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) segment_id_str = str(segment_id) - _, segment = _get_segment_for_document(dataset, document, segment_id_str) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) payload = SegmentUpdatePayload.model_validate(service_api_ns.payload or {}) - updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, db.session()) + updated_segment = SegmentService.update_segment(payload.segment, segment, document, dataset, session) summary = SummaryIndexService.get_segment_summary( - segment_id=updated_segment.id, dataset_id=dataset_id_str, session=db.session() + segment_id=updated_segment.id, dataset_id=dataset_id_str, session=session ) response = { - "data": segment_response_with_summary(updated_segment, summary.summary_content if summary else None), + "data": segment_response_with_summary( + updated_segment, summary.summary_content if summary else None, session=session + ), "doc_form": document.doc_form, } return dump_response(SegmentDetailResponse, response), 200 @@ -470,11 +478,12 @@ class DatasetSegmentApi(DatasetApiResource): "Segment retrieved successfully", service_api_ns.models[SegmentDetailResponse.__name__], ) - def get(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): current_account_with_tenant() dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -483,17 +492,19 @@ 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=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) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) summary = SummaryIndexService.get_segment_summary( - segment_id=segment.id, dataset_id=dataset_id_str, session=db.session() + segment_id=segment.id, dataset_id=dataset_id_str, session=session ) response = { - "data": segment_response_with_summary(segment, summary.summary_content if summary else None), + "data": segment_response_with_summary( + segment, summary.summary_content if summary else None, session=session + ), "doc_form": document.doc_form, } return dump_response(SegmentDetailResponse, response), 200 @@ -533,12 +544,13 @@ class ChildChunkApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): _, current_tenant_id = current_account_with_tenant() """Create child chunk.""" dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -546,12 +558,12 @@ 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=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) + _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # check embedding model setting if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -574,7 +586,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, session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) @@ -604,12 +616,13 @@ class ChildChunkApi(DatasetApiResource): "Child chunks retrieved successfully", service_api_ns.models[ChildChunkListResponse.__name__], ) - def get(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): + @with_session + def get(self, session: Session, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID): current_account_with_tenant() """Get child chunks.""" dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -617,12 +630,12 @@ 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=session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - _get_segment_for_document(dataset, document, segment_id_str) + _get_segment_for_document(session, dataset, document, segment_id_str) args = query_params_from_request(ChildChunkListQuery, use_defaults_for_malformed_ints=True) @@ -631,7 +644,13 @@ class ChildChunkApi(DatasetApiResource): keyword = args.keyword child_chunks = SegmentService.get_child_chunks( - segment_id_str, document_id_str, dataset_id_str, page, limit, keyword + segment_id_str, + document_id_str, + dataset_id_str, + page, + limit, + keyword, + session=session, ) response = { @@ -671,12 +690,21 @@ class DatasetChildChunkApi(DatasetApiResource): ) @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def delete(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID, child_chunk_id: UUID): + @with_session + def delete( + self, + session: Session, + tenant_id: str, + dataset_id: UUID, + document_id: UUID, + segment_id: UUID, + child_chunk_id: UUID, + ): current_account_with_tenant() """Delete child chunk.""" dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -684,21 +712,21 @@ 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=session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - segment_ref, _ = _get_segment_for_document(dataset, document, segment_id_str) + segment_ref, _ = _get_segment_for_document(session, dataset, document, segment_id_str) 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, db.session()) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, session=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, session) except ChildChunkDeleteIndexServiceError as e: raise ChildChunkDeleteIndexError(str(e)) @@ -732,12 +760,21 @@ class DatasetChildChunkApi(DatasetApiResource): @cloud_edition_billing_resource_check("vector_space", "dataset") @cloud_edition_billing_knowledge_limit_check("add_segment", "dataset") @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def patch(self, tenant_id: str, dataset_id: UUID, document_id: UUID, segment_id: UUID, child_chunk_id: UUID): + @with_session + def patch( + self, + session: Session, + tenant_id: str, + dataset_id: UUID, + document_id: UUID, + segment_id: UUID, + child_chunk_id: UUID, + ): current_account_with_tenant() """Update child chunk.""" dataset_id_str = str(dataset_id) # check dataset - dataset = db.session.scalar( + dataset = session.scalar( select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id_str).limit(1) ) if not dataset: @@ -745,16 +782,16 @@ 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=session) if not document: raise NotFound("Document not found.") segment_id_str = str(segment_id) - segment_ref, segment = _get_segment_for_document(dataset, document, segment_id_str) + segment_ref, segment = _get_segment_for_document(session, dataset, document, segment_id_str) 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, db.session()) + child_chunk = SegmentService.get_child_chunk_by_segment_ref(child_chunk_id_str, segment_ref, session=session) if not child_chunk: raise NotFound("Child chunk not found.") @@ -763,7 +800,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, session ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/service_api/wraps.py b/api/controllers/service_api/wraps.py index 8cb339a1491..c3c8e02e438 100644 --- a/api/controllers/service_api/wraps.py +++ b/api/controllers/service_api/wraps.py @@ -161,7 +161,7 @@ def validate_app_token[**P, R]( if tenant_owner_info: tenant_model, account = tenant_owner_info - account.current_tenant = tenant_model + account.set_current_tenant_with_session(tenant_model, session=db.session()) current_app.login_manager._update_request_context_with_user(account) # type: ignore user_logged_in.send(current_app._get_current_object(), user=current_user) # type: ignore else: @@ -333,7 +333,7 @@ def validate_dataset_token[R](view: Callable[..., R]) -> Callable[..., R]: account = db.session.get(Account, ta.account_id) # Login admin if account: - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=db.session()) current_app.login_manager._update_request_context_with_user(account) # type: ignore user_logged_in.send(current_app._get_current_object(), user=current_user) # type: ignore else: diff --git a/api/controllers/web/app.py b/api/controllers/web/app.py index 6804d072ef0..80cf720581a 100644 --- a/api/controllers/web/app.py +++ b/api/controllers/web/app.py @@ -15,7 +15,7 @@ from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPu 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 +from models.model import App, AppMode, EndUser, load_annotation_reply_config from services.app_service import AppService from services.enterprise.enterprise_service import EnterpriseService from services.feature_service import FeatureService @@ -77,28 +77,36 @@ class AppParameterApi(WebApiResource): @web_ns.response(200, "Success", web_ns.models[fields.Parameters.__name__]) def get(self, app_model: App, end_user: EndUser): """Retrieve app parameters.""" + session = db.session() features_dict: dict[str, Any] user_input_form: list[dict[str, Any]] if app_model.mode == AppMode.AGENT: try: - features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(app_model) + features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form( + app_model, + session=session, + ) except AgentAppNotPublishedError: raise AgentNotPublishedError() except AgentAppGeneratorError: raise AppUnavailableError() elif app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise AppUnavailableError() features_dict = workflow.features_dict user_input_form = workflow.user_input_form(to_old_structure=True) else: - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if app_model_config is None: raise AppUnavailableError() - features_dict = cast(dict[str, Any], app_model_config.to_dict()) + annotation_reply = load_annotation_reply_config(session, app_model.id) + features_dict = cast( + dict[str, Any], + app_model_config.to_dict(annotation_reply=annotation_reply), + ) user_input_form = features_dict.get("user_input_form", []) diff --git a/api/controllers/web/audio.py b/api/controllers/web/audio.py index defda5ab0d9..3e5c939968c 100644 --- a/api/controllers/web/audio.py +++ b/api/controllers/web/audio.py @@ -79,7 +79,12 @@ class AudioApi(WebApiResource): file = request.files["file"] try: - response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=end_user.external_user_id) + response = AudioService.transcript_asr( + app_model=app_model, + file=file, + session=db.session(), + end_user=end_user.external_user_id, + ) return dump_response(AudioToTextResponse, response) except services.errors.app_model_config.AppModelConfigBrokenError: diff --git a/api/controllers/web/completion.py b/api/controllers/web/completion.py index baae7e0b181..58083ed48a6 100644 --- a/api/controllers/web/completion.py +++ b/api/controllers/web/completion.py @@ -30,7 +30,6 @@ 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 @@ -224,7 +223,7 @@ class ChatApi(WebApiResource): app_model=app_model, conversation_id=payload.conversation_id, user=end_user, - session=db.session(), + session=session, ) response = AppGenerateService.generate( diff --git a/api/controllers/web/conversation.py b/api/controllers/web/conversation.py index 09a3a508824..75aae01a576 100644 --- a/api/controllers/web/conversation.py +++ b/api/controllers/web/conversation.py @@ -15,6 +15,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom from extensions.ext_database import db from fields.conversation_fields import ( ConversationInfiniteScrollPagination, + ConversationResponseSource, ResultResponse, SimpleConversation, ) @@ -80,7 +81,13 @@ class ConversationListApi(WebApiResource): sort_by=query.sort_by, ) adapter = TypeAdapter(SimpleConversation) - conversations = [adapter.validate_python(item, from_attributes=True) for item in pagination.data] + conversations = [ + adapter.validate_python( + ConversationResponseSource(item, session=session), + from_attributes=True, + ) + for item in pagination.data + ] return ConversationInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, @@ -156,12 +163,13 @@ class ConversationRenameApi(WebApiResource): payload = ConversationRenamePayload.model_validate(web_ns.payload or {}) try: + session = db.session() conversation = ConversationService.rename( - app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=db.session() + app_model, conversation_id, end_user, payload.name, payload.auto_generate, session=session ) return ( TypeAdapter(SimpleConversation) - .validate_python(conversation, from_attributes=True) + .validate_python(ConversationResponseSource(conversation, session=session), from_attributes=True) .model_dump(mode="json") ) except ConversationNotExistsError: diff --git a/api/controllers/web/message.py b/api/controllers/web/message.py index 3edb8436691..658cb0f14df 100644 --- a/api/controllers/web/message.py +++ b/api/controllers/web/message.py @@ -26,7 +26,7 @@ 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.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import SuggestedQuestionsResponse, WebMessageInfiniteScrollPagination, WebMessageListItem from graphon.model_runtime.errors.invoke import InvokeError from libs import helper @@ -86,11 +86,15 @@ class MessageListApi(WebApiResource): query = MessageListQuery.model_validate(raw_args) try: + session = db.session() pagination = MessageService.pagination_by_first_id( - app_model, end_user, query.conversation_id, query.first_id, query.limit, session=db.session() + app_model, end_user, query.conversation_id, query.first_id, query.limit, session=session ) adapter = TypeAdapter(WebMessageListItem) - items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] + items = [ + adapter.validate_python(MessageResponseSource(message, session=session), from_attributes=True) + for message in pagination.data + ] return WebMessageInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, diff --git a/api/controllers/web/saved_message.py b/api/controllers/web/saved_message.py index cc453961c43..99ea48bf7c9 100644 --- a/api/controllers/web/saved_message.py +++ b/api/controllers/web/saved_message.py @@ -10,7 +10,7 @@ from controllers.web import web_ns from controllers.web.error import NotCompletionAppError from controllers.web.wraps import WebApiResource from extensions.ext_database import db -from fields.conversation_fields import ResultResponse +from fields.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import SavedMessageInfiniteScrollPagination, SavedMessageItem from models.model import App, EndUser from services.errors.message import MessageNotExistsError @@ -43,11 +43,15 @@ class SavedMessageListApi(WebApiResource): raw_args = request.args.to_dict() query = SavedMessageListQuery.model_validate(raw_args) + session = db.session() pagination = SavedMessageService.pagination_by_last_id( - app_model, end_user, query.last_id, query.limit, session=db.session() + app_model, end_user, query.last_id, query.limit, session=session ) adapter = TypeAdapter(SavedMessageItem) - items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data] + items = [ + adapter.validate_python(MessageResponseSource(message, session=session), from_attributes=True) + for message in pagination.data + ] return SavedMessageInfiniteScrollPagination( limit=pagination.limit, has_more=pagination.has_more, data=items ).model_dump(mode="json") diff --git a/api/core/agent/base_agent_runner.py b/api/core/agent/base_agent_runner.py index 850858bfcef..2d20ca75a5e 100644 --- a/api/core/agent/base_agent_runner.py +++ b/api/core/agent/base_agent_runner.py @@ -42,7 +42,7 @@ from graphon.model_runtime.entities.message_entities import ImagePromptMessageCo from graphon.model_runtime.entities.model_entities import ModelFeature from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel from models.enums import CreatorUserRole -from models.model import Conversation, Message, MessageAgentThought, MessageFile +from models.model import Conversation, Message, MessageAgentThought, MessageFile, load_annotation_reply_config logger = logging.getLogger(__name__) _file_access_controller = DatabaseFileAccessController() @@ -76,7 +76,9 @@ class BaseAgentRunner(AppRunner): self.message = message self.user_id = user_id self.memory = memory - self.history_prompt_messages = self.organize_agent_history(prompt_messages=prompt_messages or []) + self.history_prompt_messages = self.organize_agent_history( + session=session, prompt_messages=prompt_messages or [] + ) self.model_instance = model_instance # init callback @@ -104,7 +106,7 @@ class BaseAgentRunner(AppRunner): ) # get how many agent thoughts have been created self.agent_thought_count = ( - db.session.scalar( + session.scalar( select(func.count()) .select_from(MessageAgentThought) .where( @@ -113,7 +115,7 @@ class BaseAgentRunner(AppRunner): ) or 0 ) - db.session.close() + session.close() # check if model supports stream tool call llm_model = cast(LargeLanguageModel, model_instance.model_type_instance) @@ -350,7 +352,7 @@ class BaseAgentRunner(AppRunner): db.session.commit() db.session.close() - def organize_agent_history(self, prompt_messages: list[PromptMessage]) -> list[PromptMessage]: + def organize_agent_history(self, prompt_messages: list[PromptMessage], *, session: Session) -> list[PromptMessage]: """ Organize agent history """ @@ -362,7 +364,7 @@ class BaseAgentRunner(AppRunner): messages = ( ( - db.session.execute( + session.execute( select(Message) .where(Message.conversation_id == self.message.conversation_id) .order_by(Message.created_at.desc()) @@ -378,8 +380,8 @@ class BaseAgentRunner(AppRunner): if message.id == self.message.id: continue - result.append(self.organize_agent_user_prompt(message)) - agent_thoughts = message.agent_thoughts + result.append(self.organize_agent_user_prompt(message, session=session)) + agent_thoughts = message.agent_thoughts_with_session(session=session) if agent_thoughts: for agent_thought in agent_thoughts: tool_names_raw = agent_thought.tool @@ -441,17 +443,21 @@ class BaseAgentRunner(AppRunner): if message.answer: result.append(AssistantPromptMessage(content=message.answer)) - db.session.close() + session.close() return result - def organize_agent_user_prompt(self, message: Message) -> UserPromptMessage: + def organize_agent_user_prompt(self, message: Message, *, session: Session) -> UserPromptMessage: stmt = select(MessageFile).where(MessageFile.message_id == message.id) - files = db.session.scalars(stmt).all() + files = session.scalars(stmt).all() if not files: return UserPromptMessage(content=message.query) - if message.app_model_config: - file_extra_config = FileUploadConfigManager.convert(message.app_model_config.to_dict()) + app_model_config = message.app_model_config_with_session(session=session) + if app_model_config: + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + file_extra_config = FileUploadConfigManager.convert( + app_model_config.to_dict(annotation_reply=annotation_reply) + ) else: file_extra_config = None diff --git a/api/core/agent/cot_agent_runner.py b/api/core/agent/cot_agent_runner.py index 4d823ca79c0..a1ccf75386b 100644 --- a/api/core/agent/cot_agent_runner.py +++ b/api/core/agent/cot_agent_runner.py @@ -114,7 +114,11 @@ class CotAgentRunner(BaseAgentRunner, ABC): message_file_ids: list[str] = [] agent_thought_id = self.create_agent_thought( - message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids + message_id=message.id, + message="", + tool_name="", + tool_input="", + messages_ids=message_file_ids, ) if iteration_step > 1: @@ -125,6 +129,11 @@ class CotAgentRunner(BaseAgentRunner, ABC): # recalc llm max tokens prompt_messages = self._organize_prompt_messages() self.recalc_llm_max_tokens(self.model_config, prompt_messages) + + # Release any setup/tool transaction before waiting on the provider stream. + session.commit() + session.close() + # invoke model chunks = model_instance.invoke_llm( prompt_messages=prompt_messages, @@ -333,6 +342,8 @@ class CotAgentRunner(BaseAgentRunner, ABC): agent_tool_callback=self.agent_callback, trace_manager=trace_manager, ) + session.commit() + session.close() # publish files for message_file_id in message_files: diff --git a/api/core/agent/fc_agent_runner.py b/api/core/agent/fc_agent_runner.py index 0b92daf93a4..5bffa0002bf 100644 --- a/api/core/agent/fc_agent_runner.py +++ b/api/core/agent/fc_agent_runner.py @@ -87,12 +87,21 @@ class FunctionCallAgentRunner(BaseAgentRunner): message_file_ids: list[str] = [] agent_thought_id = self.create_agent_thought( - message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids + message_id=message.id, + message="", + tool_name="", + tool_input="", + messages_ids=message_file_ids, ) # recalc llm max tokens prompt_messages = self._organize_prompt_messages() self.recalc_llm_max_tokens(self.model_config, prompt_messages) + + # Release any setup/tool transaction before waiting on the provider stream. + session.commit() + session.close() + # invoke model chunks: Union[Generator[LLMResultChunk, None, None], LLMResult] = model_instance.invoke_llm( prompt_messages=prompt_messages, @@ -256,6 +265,8 @@ class FunctionCallAgentRunner(BaseAgentRunner): message_id=self.message.id, conversation_id=self.conversation.id, ) + session.commit() + session.close() # publish files for message_file_id in message_files: # publish message file 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 0108e7d7c72..497f1d44cd5 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 @@ -1,6 +1,8 @@ import uuid from typing import Any, Literal, cast +from sqlalchemy.orm import Session + from core.app.app_config.entities import ( DatasetEntity, DatasetRetrieveConfigEntity, @@ -9,7 +11,6 @@ from core.app.app_config.entities import ( ) from core.entities.agent_entities import PlanningStrategy from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict -from extensions.ext_database import db from models.model import AppMode, AppModelConfigDict from services.dataset_service import DatasetService @@ -140,7 +141,7 @@ class DatasetConfigManager: @classmethod def validate_and_set_defaults( - cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any] + cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any], session: Session ) -> tuple[dict[str, Any], list[str]]: """ Validate and set defaults for dataset feature @@ -150,7 +151,7 @@ class DatasetConfigManager: :param config: app model config args """ # Extract dataset config for legacy compatibility - config = cls.extract_dataset_config_for_legacy_compatibility(tenant_id, app_mode, config) + config = cls.extract_dataset_config_for_legacy_compatibility(tenant_id, app_mode, config, session) # dataset_configs if "dataset_configs" not in config or not config.get("dataset_configs"): @@ -175,7 +176,9 @@ class DatasetConfigManager: return config, ["agent_mode", "dataset_configs", "dataset_query_variable"] @classmethod - def extract_dataset_config_for_legacy_compatibility(cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any]): + def extract_dataset_config_for_legacy_compatibility( + cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any], session: Session + ): """ Extract dataset config for legacy compatibility @@ -238,7 +241,7 @@ class DatasetConfigManager: except ValueError: raise ValueError("id in dataset must be of UUID type") - if not cls.is_dataset_exists(tenant_id, tool_item["id"]): + if not cls.is_dataset_exists(tenant_id, tool_item["id"], session): raise ValueError("Dataset ID does not exist, please check your permission.") has_datasets = True @@ -255,9 +258,9 @@ class DatasetConfigManager: return config @classmethod - def is_dataset_exists(cls, tenant_id: str, dataset_id: str) -> bool: + def is_dataset_exists(cls, tenant_id: str, dataset_id: str, session: Session) -> bool: # verify if the dataset ID exists - dataset = DatasetService.get_dataset(dataset_id, db.session()) + dataset = DatasetService.get_dataset(dataset_id, 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 23acfb7ea62..f88034e2b43 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -86,6 +86,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_run_id: str, streaming: Literal[False], pause_state_config: PauseStateLayerConfig | None = None, + *, + session: Session, ) -> Mapping[str, Any]: ... @overload @@ -99,6 +101,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_run_id: str, streaming: Literal[True], pause_state_config: PauseStateLayerConfig | None = None, + *, + session: Session, ) -> Generator[Mapping | str, None, None]: ... @overload @@ -112,6 +116,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_run_id: str, streaming: bool, pause_state_config: PauseStateLayerConfig | None = None, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping, None, None]: ... def generate( @@ -124,6 +130,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_run_id: str, streaming: bool = True, pause_state_config: PauseStateLayerConfig | None = None, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping, None, None]: """ Generate App response. @@ -134,6 +142,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): :param args: request args :param invoke_from: invoke from source :param streaming: is stream + :param session: database session supplied by the caller """ if not args.get("query"): raise ValueError("query is required") @@ -157,7 +166,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): if conversation_id: try: conversation = ConversationService.get_conversation( - app_model=app_model, conversation_id=conversation_id, user=user, session=db.session() + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) except ConversationNotExistsError: if invoke_from == InvokeFrom.SERVICE_API: @@ -255,6 +264,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): conversation=conversation, stream=streaming, pause_state_config=pause_state_config, + session=session, ) def resume( @@ -265,6 +275,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, conversation: Conversation, message: Message, + session: Session, application_generate_entity: AdvancedChatAppGenerateEntity, workflow_execution_repository: WorkflowExecutionRepository, workflow_node_execution_repository: WorkflowNodeExecutionRepository, @@ -301,6 +312,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): pause_state_config=pause_state_config, graph_runtime_state=graph_runtime_state, response_stream_filter=response_stream_filter, + session=session, ) def single_iteration_generate( @@ -311,6 +323,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, args: Mapping[str, Any], streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -321,6 +335,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): :param user: account or end user :param args: request args :param streaming: is streamed + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -377,7 +392,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): tenant_id=application_generate_entity.app_config.tenant_id, user_id=user.id, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) return self._generate( @@ -390,6 +405,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): conversation=None, stream=streaming, variable_loader=var_loader, + session=session, ) def single_loop_generate( @@ -400,6 +416,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, args: LoopNodeRunPayload, streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -410,6 +428,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): :param user: account or end user :param args: request args :param streaming: is stream + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -464,7 +483,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): tenant_id=application_generate_entity.app_config.tenant_id, user_id=user.id, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) return self._generate( @@ -477,6 +496,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): conversation=None, stream=streaming, variable_loader=var_loader, + session=session, ) def _generate( @@ -486,6 +506,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, invoke_from: InvokeFrom, application_generate_entity: AdvancedChatAppGenerateEntity, + session: Session, workflow_execution_repository: WorkflowExecutionRepository, workflow_node_execution_repository: WorkflowNodeExecutionRepository, conversation: Conversation | None = None, @@ -504,6 +525,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): :param user: account or end user :param invoke_from: invoke from source :param application_generate_entity: application generate entity + :param session: database session supplied by the caller :param workflow_execution_repository: repository for workflow execution :param workflow_node_execution_repository: repository for workflow node execution :param conversation: conversation @@ -519,18 +541,22 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): if conversation is not None and message is not None: pass else: - conversation, message = self._init_generate_records(application_generate_entity, conversation) + conversation, message = self._init_generate_records( + application_generate_entity, + conversation, + session=session, + ) if is_first_conversation: # update conversation features conversation.override_model_configs = workflow.features - db.session.commit() - db.session.refresh(conversation) + session.commit() + session.refresh(conversation) # get conversation dialogue count # NOTE: dialogue_count should not start from 0, # because during the first conversation, dialogue_count should be 1. - self._dialogue_count = get_thread_messages_length(conversation.id) + 1 + self._dialogue_count = get_thread_messages_length(conversation.id, session=session) + 1 # init queue manager queue_manager = MessageBasedAppQueueManager( @@ -582,7 +608,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): workflow_snapshot = WorkflowSnapshot.from_workflow(workflow) conversation_snapshot = ConversationSnapshot.from_conversation(conversation) message_snapshot = MessageSnapshot.from_message(message) - db.session.close() + session.close() # return response or stream generator response = self._handle_advanced_chat_response( diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index 249cb33a98c..31a65578d49 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -39,7 +39,6 @@ from core.workflow.system_variables import ( ) from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool from core.workflow.workflow_entry import WorkflowEntry -from extensions.ext_database import db from extensions.ext_redis import redis_client from extensions.otel import WorkflowAppRunnerHandler, trace_span from extensions.workflow_warm_shutdown import WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, celery_warm_shutdown_started @@ -173,12 +172,21 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): ) # annotation reply - if self.handle_annotation_reply( - app_record=self._app, - message=self.message, - query=new_query, - app_generate_entity=self.application_generate_entity, - ): + with create_session() as session: + annotation_reply = self.handle_annotation_reply( + app_record=self._app, + message=self.message, + query=new_query, + app_generate_entity=self.application_generate_entity, + session=session, + ) + session.commit() + if annotation_reply: + self._publish_event(QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id)) + self._complete_with_stream_output( + text=annotation_reply.content, + stopped_by=QueueStopEvent.StopBy.ANNOTATION_REPLY, + ) return # Initialize conversation variables @@ -212,10 +220,6 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): trace_session_id=self.application_generate_entity.extras.get("trace_session_id"), ) - # Release the Flask scoped session before workflow execution so a checked-out DB connection - # is not held for the lifetime of the graph run. - db.session.close() - # RUN WORKFLOW # Create Redis command channel for this workflow execution task_id = self.application_generate_entity.task_id @@ -300,26 +304,22 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): return False, new_inputs, new_query def handle_annotation_reply( - self, app_record: App, message: Message, query: str, app_generate_entity: AdvancedChatAppGenerateEntity - ) -> bool: - annotation_reply = self.query_app_annotations_to_reply( + self, + app_record: App, + message: Message, + query: str, + app_generate_entity: AdvancedChatAppGenerateEntity, + session: Session, + ) -> MessageAnnotation | None: + return self.query_app_annotations_to_reply( app_record=app_record, message=message, query=query, user_id=app_generate_entity.user_id, invoke_from=app_generate_entity.invoke_from, + session=session, ) - if annotation_reply: - self._publish_event(QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id)) - - self._complete_with_stream_output( - text=annotation_reply.content, stopped_by=QueueStopEvent.StopBy.ANNOTATION_REPLY - ) - return True - - return False - def _complete_with_stream_output(self, text: str, stopped_by: QueueStopEvent.StopBy): """ Direct output @@ -329,7 +329,13 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): self._publish_event(QueueStopEvent(stopped_by=stopped_by)) def query_app_annotations_to_reply( - self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom + self, + app_record: App, + message: Message, + query: str, + user_id: str, + invoke_from: InvokeFrom, + session: Session, ) -> MessageAnnotation | None: """ Query app annotations to reply @@ -342,7 +348,12 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): """ annotation_reply_feature = AnnotationReplyFeature() return annotation_reply_feature.query( - app_record=app_record, message=message, query=query, user_id=user_id, invoke_from=invoke_from + app_record=app_record, + message=message, + query=query, + user_id=user_id, + invoke_from=invoke_from, + session=session, ) def moderation_for_inputs( @@ -395,7 +406,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): existing_variables = self._create_all_conversation_variables(session) else: # Check and add any missing variables from the workflow - existing_variables = self._sync_missing_conversation_variables(session, existing_variables) + existing_variables = self._sync_missing_conversation_variables(existing_variables, session) # Convert to Variable objects for use in the workflow conversation_variables = [var.to_variable() for var in existing_variables] @@ -435,7 +446,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): return new_variables def _sync_missing_conversation_variables( - self, session: Session, existing_variables: list[ConversationVariable] + self, existing_variables: list[ConversationVariable], session: Session ) -> list[ConversationVariable]: """ Sync missing conversation variables from the workflow definition. diff --git a/api/core/app/apps/advanced_chat/generate_task_pipeline.py b/api/core/app/apps/advanced_chat/generate_task_pipeline.py index 2c9350ca774..97e54155219 100644 --- a/api/core/app/apps/advanced_chat/generate_task_pipeline.py +++ b/api/core/app/apps/advanced_chat/generate_task_pipeline.py @@ -9,7 +9,7 @@ from threading import Thread from typing import Any, Union from sqlalchemy import select, update -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.orm import Session from constants.tts_auto_play_timeout import TTS_AUTO_PLAY_TIMEOUT, TTS_AUTO_PLAY_YIELD_CPU_TIME from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom @@ -71,12 +71,12 @@ from core.app.entities.task_entities import ( from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline from core.app.task_pipeline.message_cycle_manager import MessageCycleManager from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk +from core.db.session_factory import session_factory from core.ops.ops_trace_manager import TraceQueueManager from core.repositories.human_input_repository import HumanInputFormRepositoryImpl from core.workflow.file_reference import resolve_file_record_id from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from core.workflow.system_variables import build_system_variables -from extensions.ext_database import db from graphon.enums import WorkflowExecutionStatus from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.model_runtime.utils.encoders import jsonable_encoder @@ -399,8 +399,13 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport): @contextmanager def _database_session(self): """Context manager for database sessions.""" - with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: - yield session + with session_factory.create_session() as session: + try: + yield session + session.commit() + except Exception: + session.rollback() + raise def _ensure_workflow_initialized(self): """Fluent validation for workflow state.""" @@ -825,7 +830,8 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport): self, event: QueueAnnotationReplyEvent, **kwargs ) -> Generator[StreamResponse, None, None]: """Handle annotation reply events.""" - self._message_cycle_manager.handle_annotation_reply(event) + with self._database_session() as session: + self._message_cycle_manager.handle_annotation_reply(event, session) yield from () def _handle_message_replace_event( diff --git a/api/core/app/apps/agent_app/app_config_manager.py b/api/core/app/apps/agent_app/app_config_manager.py index 71224534133..b993e689d26 100644 --- a/api/core/app/apps/agent_app/app_config_manager.py +++ b/api/core/app/apps/agent_app/app_config_manager.py @@ -24,7 +24,7 @@ from core.app.app_config.entities import ( from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features from core.app.apps.agent_app.app_variable_projection import agent_app_variables_to_user_input_form from models.agent_config_entities import AgentSoulConfig -from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation +from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation class AgentAppConfig(EasyUIBasedAppConfig): @@ -43,11 +43,16 @@ class AgentAppConfigManager(BaseAppConfigManager): *, app_model: App, agent_soul: AgentSoulConfig, + annotation_reply: AnnotationReplyConfig | None, app_model_config: AppModelConfig | None = None, conversation: Conversation | None = None, ) -> AgentAppConfig: """Build the Agent App config from the Agent Soul (+ optional feature flags).""" - config_dict = cls._synthesize_config_dict(agent_soul, app_model_config) + config_dict = cls._synthesize_config_dict( + agent_soul, + app_model_config, + annotation_reply=annotation_reply, + ) # The synthesized dict is shaped like an app_model_config; the EasyUI # sub-managers type their param as AppModelConfigDict (a TypedDict). typed_config = cast(AppModelConfigDict, config_dict) @@ -77,6 +82,8 @@ class AgentAppConfigManager(BaseAppConfigManager): def _synthesize_config_dict( agent_soul: AgentSoulConfig, app_model_config: AppModelConfig | None, + *, + annotation_reply: AnnotationReplyConfig | None, ) -> dict[str, Any]: """Shape a Soul + feature flags into an ``app_model_config``-style dict. @@ -84,7 +91,11 @@ class AgentAppConfigManager(BaseAppConfigManager): ``app_model_config`` when one exists; model + prompt always come from the Agent Soul (the single source of truth for those). """ - base = merge_agent_app_features(agent_soul=agent_soul, app_model_config=app_model_config) + base = merge_agent_app_features( + agent_soul=agent_soul, + app_model_config=app_model_config, + annotation_reply=annotation_reply, + ) model = agent_soul.model if model is not None: diff --git a/api/core/app/apps/agent_app/app_feature_projection.py b/api/core/app/apps/agent_app/app_feature_projection.py index cb8efd9f290..517dcb42b96 100644 --- a/api/core/app/apps/agent_app/app_feature_projection.py +++ b/api/core/app/apps/agent_app/app_feature_projection.py @@ -1,12 +1,14 @@ from typing import Any from models.agent_config_entities import AgentSoulConfig +from models.model import AnnotationReplyConfig def merge_agent_app_features( *, agent_soul: AgentSoulConfig, app_model_config: Any | None, + annotation_reply: AnnotationReplyConfig | None, ) -> dict[str, Any]: """Project public Agent App features from legacy config plus Agent Soul. @@ -14,7 +16,12 @@ def merge_agent_app_features( opening statements. Agent Soul is the source of truth for Agent-owned features like file upload, so Soul fields override same-named legacy keys. """ - features: dict[str, Any] = dict(app_model_config.to_dict()) if app_model_config else {} + if app_model_config is None: + features: dict[str, Any] = {} + else: + if annotation_reply is None: + raise ValueError("Annotation reply config is required") + features = dict(app_model_config.to_dict(annotation_reply=annotation_reply)) soul_features = agent_soul.app_features.model_dump(mode="json", exclude_none=True) features.update(soul_features) return features diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index da802680c08..68f6edbd550 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -21,6 +21,7 @@ from typing import Any, Literal from flask import Flask, current_app from pydantic import JsonValue from sqlalchemy import and_, or_, select +from sqlalchemy.orm import Session from clients.agent_backend import AgentBackendRunEventAdapter from clients.agent_backend.factory import create_agent_backend_run_client @@ -46,10 +47,11 @@ from core.app.entities.app_invoke_entities import ( UserFrom, ) from core.app.llm.model_access import build_dify_model_access +from core.db.session_factory import session_factory from core.ops.ops_trace_manager import TraceQueueManager from core.workflow.file_reference import build_file_reference, is_canonical_file_reference from extensions.ext_database import db -from models import Account, App, EndUser, Message +from models import Account, App, AppModelConfig, EndUser, Message, MessageAnnotation from models.agent import ( APP_BACKED_AGENT_SOURCES, Agent, @@ -61,6 +63,7 @@ from models.agent import ( AgentStatus, ) from models.agent_config_entities import AgentSoulConfig +from models.model import load_annotation_reply_config from services.conversation_service import ConversationService logger = logging.getLogger(__name__) @@ -137,6 +140,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, + session: Session, streaming: bool = True, ) -> Mapping[str, Any] | Generator[Mapping | str, None, None]: if not streaming: @@ -152,6 +156,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): invoke_from=invoke_from, draft_type=args.get("draft_type"), user=user, + session=session, ) runtime_session_snapshot_id = self._runtime_session_snapshot_id( invoke_from=invoke_from, @@ -162,15 +167,19 @@ 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, session=db.session() + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) # Build the EasyUI-shaped config from the Agent Soul so the chat pipeline # can persist usage; the answer itself comes from the agent backend. - app_model_config = app_model.app_model_config + app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) + annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None app_config = AgentAppConfigManager.get_app_config( app_model=app_model, agent_soul=agent_soul, + annotation_reply=annotation_reply, app_model_config=app_model_config, conversation=conversation, ) @@ -210,7 +219,11 @@ class AgentAppGenerator(MessageBasedAppGenerator): agent_runtime_exit_intent=agent_runtime_exit_intent, ) - conversation, message = self._init_generate_records(application_generate_entity, conversation) + conversation, message = self._init_generate_records( + application_generate_entity, + conversation, + session=session, + ) queue_manager = MessageBasedAppQueueManager( task_id=application_generate_entity.task_id, @@ -253,6 +266,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): user: Account | EndUser, conversation_id: str, invoke_from: InvokeFrom, + session: Session, ) -> None: """Resume an Agent App conversation after a submitted ask_human HITL form. @@ -263,19 +277,28 @@ 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, session=db.session() + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( app_model, invoke_from=invoke_from, - draft_type=self._resume_draft_type(app_model=app_model, conversation=conversation, user=user), + draft_type=self._resume_draft_type( + app_model=app_model, conversation=conversation, user=user, session=session + ), user=user, + session=session, ) + app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) + annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None + app_config = AgentAppConfigManager.get_app_config( app_model=app_model, agent_soul=agent_soul, - app_model_config=app_model.app_model_config, + annotation_reply=annotation_reply, + app_model_config=app_model_config, conversation=conversation, ) model_conf = ModelConfigConverter.convert(app_config) @@ -287,7 +310,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): # turn's query); the continuation is driven by deferred_tool_results and # the restored snapshot, not by re-processing this prompt. A blank prompt # would drop the user-prompt layer and fail the snapshot match. - paused_message = db.session.scalar( + paused_message = session.scalar( select(Message) .where(Message.conversation_id == conversation.id, Message.query != "") .order_by(Message.created_at.desc()) @@ -318,7 +341,11 @@ class AgentAppGenerator(MessageBasedAppGenerator): agent_config_version_kind=agent_config_version_kind, ) - conversation, message = self._init_generate_records(application_generate_entity, conversation) + conversation, message = self._init_generate_records( + application_generate_entity, + conversation, + session=session, + ) queue_manager = MessageBasedAppQueueManager( task_id=application_generate_entity.task_id, @@ -357,7 +384,9 @@ class AgentAppGenerator(MessageBasedAppGenerator): ) @staticmethod - def _resume_draft_type(*, app_model: App, conversation: Any, user: Account | EndUser) -> str | None: + def _resume_draft_type( + *, app_model: App, conversation: Any, user: Account | EndUser, session: Session + ) -> str | None: if conversation.invoke_from != InvokeFrom.DEBUGGER: return None active_session = AgentAppRuntimeSessionStore().load_active_session_for_conversation( @@ -367,7 +396,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): ) snapshot_id = active_session.scope.agent_config_snapshot_id if active_session is not None else None if snapshot_id and isinstance(user, Account): - draft = db.session.scalar( + draft = session.scalar( select(AgentConfigDraft).where( AgentConfigDraft.tenant_id == app_model.tenant_id, AgentConfigDraft.id == snapshot_id, @@ -413,15 +442,31 @@ class AgentAppGenerator(MessageBasedAppGenerator): # Apply app-level input guards (content moderation + annotation # reply) before reaching the Agent backend, mirroring the EasyUI # chat / agent-chat runners. These can short-circuit the turn. - app_model = db.session.get(App, app_config.app_id) - if app_model is None: - raise AgentAppGeneratorError("App not found") - handled, query = self._run_input_guards( - application_generate_entity=application_generate_entity, - app_model=app_model, - message=message, - queue_manager=queue_manager, - ) + with session_factory.get_session_maker().begin() as session: + app_model = session.get(App, app_config.app_id) + if app_model is None: + raise AgentAppGeneratorError("App not found") + handled, query, annotation_reply = self._run_input_guards( + session=session, + application_generate_entity=application_generate_entity, + app_model=app_model, + message=message, + queue_manager=queue_manager, + ) + if annotation_reply: + from core.app.apps.agent_app.app_runner import publish_text_answer + from core.app.entities.queue_entities import QueueAnnotationReplyEvent + + queue_manager.publish( + QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id), + PublishFrom.APPLICATION_MANAGER, + ) + publish_text_answer( + queue_manager=queue_manager, + model_name=application_generate_entity.model_conf.model, + answer=annotation_reply.content, + user_query=query, + ) if handled: return query = _append_prompt_file_mappings( @@ -436,11 +481,13 @@ class AgentAppGenerator(MessageBasedAppGenerator): user_from=user_from, invoke_from=application_generate_entity.invoke_from, ) - _, _, agent_soul = self._resolve_agent_by_id( - tenant_id=app_config.tenant_id, - agent_id=application_generate_entity.agent_id, - snapshot_id=application_generate_entity.agent_config_snapshot_id, - ) + with session_factory.create_session() as session: + _, _, agent_soul = self._resolve_agent_by_id( + tenant_id=app_config.tenant_id, + agent_id=application_generate_entity.agent_id, + snapshot_id=application_generate_entity.agent_config_snapshot_id, + session=session, + ) runner = self._build_runner(dify_context) runner.run( @@ -502,20 +549,18 @@ class AgentAppGenerator(MessageBasedAppGenerator): def _run_input_guards( self, *, + session: Session, application_generate_entity: AgentAppGenerateEntity, app_model: App, message: Message, queue_manager: AppQueueManager, - ) -> tuple[bool, str]: + ) -> tuple[bool, str, MessageAnnotation | None]: """Apply input moderation + annotation reply before the backend call. - Returns ``(handled, query)``: when ``handled`` is True a direct answer - has already been published (a blocked/preset moderation response or a - matched annotation) and the backend turn must be skipped. Otherwise - ``query`` is the possibly moderation-overridden query to send onward. + Returns ``(handled, query, annotation_reply)``. Annotation output is + published by the caller only after this transaction commits. """ from core.app.apps.agent_app.app_runner import publish_text_answer - from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature from core.moderation.base import ModerationError from core.moderation.input_moderation import InputModeration @@ -538,7 +583,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): ) except ModerationError as e: publish_text_answer(queue_manager=queue_manager, model_name=model_name, answer=str(e), user_query=query) - return True, query + return True, query, None # annotation reply: a matching annotation answers the turn deterministically. if query: @@ -548,21 +593,12 @@ class AgentAppGenerator(MessageBasedAppGenerator): query=query, user_id=application_generate_entity.user_id, invoke_from=application_generate_entity.invoke_from, + session=session, ) if annotation_reply: - queue_manager.publish( - QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id), - PublishFrom.APPLICATION_MANAGER, - ) - publish_text_answer( - queue_manager=queue_manager, - model_name=model_name, - answer=annotation_reply.content, - user_query=query, - ) - return True, query + return True, query, annotation_reply - return False, query + return False, query, None def _resolve_agent( self, @@ -571,8 +607,9 @@ class AgentAppGenerator(MessageBasedAppGenerator): invoke_from: InvokeFrom, draft_type: Any, user: Account | EndUser, + session: Session, ) -> tuple[Agent, str, Literal["snapshot", "draft", "build_draft"], AgentSoulConfig]: - agent = db.session.scalar( + agent = session.scalar( select(Agent) .where( Agent.tenant_id == app_model.tenant_id, @@ -603,6 +640,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): agent=agent, draft_type=draft_type, account_id=user.id if isinstance(user, Account) else None, + session=session, ) agent_soul = AgentSoulConfig.model_validate(draft.config_snapshot_dict) config_version_kind: Literal["snapshot", "draft", "build_draft"] = ( @@ -617,6 +655,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): tenant_id=app_model.tenant_id, agent_id=agent.id, snapshot_id=agent.active_config_snapshot_id, + session=session, ) return agent, snapshot.id, "snapshot", agent_soul @@ -633,7 +672,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): @staticmethod def _resolve_debug_draft( - *, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None + *, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None, session: Session ) -> AgentConfigDraft: effective_draft_type = ( AgentConfigDraftType.DEBUG_BUILD @@ -651,7 +690,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): stmt = stmt.where(AgentConfigDraft.account_id == account_id) else: stmt = stmt.where(AgentConfigDraft.account_id.is_(None)) - draft = db.session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) + draft = session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) if draft is not None: return draft if effective_draft_type == AgentConfigDraftType.DEBUG_BUILD: @@ -660,6 +699,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): tenant_id=tenant_id, agent_id=agent.id, snapshot_id=agent.active_config_snapshot_id, + session=session, ) draft = AgentConfigDraft( tenant_id=tenant_id, @@ -672,20 +712,20 @@ class AgentAppGenerator(MessageBasedAppGenerator): created_by=agent.created_by, updated_by=agent.updated_by, ) - db.session.add(draft) - db.session.flush() + session.add(draft) + session.flush() return draft @staticmethod def _resolve_agent_by_id( - *, tenant_id: str, agent_id: str, snapshot_id: str | None + *, tenant_id: str, agent_id: str, snapshot_id: str | None, session: Session ) -> tuple[Agent, AgentConfigSnapshot | AgentConfigDraft, AgentSoulConfig]: - agent = db.session.scalar(select(Agent).where(Agent.id == agent_id, Agent.tenant_id == tenant_id)) + agent = session.scalar(select(Agent).where(Agent.id == agent_id, Agent.tenant_id == tenant_id)) if agent is None: raise AgentAppGeneratorError("Agent not found") if not snapshot_id: raise AgentAppGeneratorError("Agent has no published version") - snapshot = db.session.scalar( + snapshot = session.scalar( select(AgentConfigSnapshot).where( AgentConfigSnapshot.tenant_id == tenant_id, AgentConfigSnapshot.agent_id == agent_id, @@ -695,7 +735,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): if snapshot is not None: agent_soul = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict) return agent, snapshot, agent_soul - draft = db.session.scalar( + draft = session.scalar( select(AgentConfigDraft).where( AgentConfigDraft.tenant_id == tenant_id, AgentConfigDraft.agent_id == agent_id, diff --git a/api/core/app/apps/agent_chat/app_config_manager.py b/api/core/app/apps/agent_chat/app_config_manager.py index e269b98bad9..80e6e5a322b 100644 --- a/api/core/app/apps/agent_chat/app_config_manager.py +++ b/api/core/app/apps/agent_chat/app_config_manager.py @@ -2,6 +2,8 @@ import uuid from collections.abc import Mapping from typing import Any, cast +from sqlalchemy.orm import Session + from core.agent.entities import AgentEntity from core.app.app_config.base_app_config_manager import BaseAppConfigManager from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager @@ -20,7 +22,7 @@ from core.app.app_config.features.suggested_questions_after_answer.manager impor ) from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager from core.entities.agent_entities import PlanningStrategy -from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation +from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation OLD_TOOLS = ["dataset", "google_search", "web_reader", "wikipedia", "current_datetime"] @@ -41,6 +43,8 @@ class AgentChatAppConfigManager(BaseAppConfigManager): app_model_config: AppModelConfig, conversation: Conversation | None = None, override_config_dict: AppModelConfigDict | None = None, + *, + annotation_reply: AnnotationReplyConfig | None, ) -> AgentChatAppConfig: """ Convert app model config to agent chat app config @@ -58,7 +62,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG if config_from != EasyUIBasedAppModelConfigFrom.ARGS: - app_model_config_dict = app_model_config.to_dict() + app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply) config_dict = app_model_config_dict.copy() else: if not override_config_dict: @@ -88,7 +92,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): return app_config @classmethod - def config_validate(cls, tenant_id: str, config: Mapping[str, Any]) -> AppModelConfigDict: + def config_validate(cls, tenant_id: str, config: Mapping[str, Any], session: Session) -> AppModelConfigDict: """ Validate for agent chat app model config @@ -116,7 +120,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): related_config_keys.extend(current_related_config_keys) # agent_mode - config, current_related_config_keys = cls.validate_agent_mode_and_set_defaults(tenant_id, config) + config, current_related_config_keys = cls.validate_agent_mode_and_set_defaults(tenant_id, config, session) related_config_keys.extend(current_related_config_keys) # opening_statement @@ -144,7 +148,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): # dataset configs # dataset_query_variable config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults( - tenant_id, app_mode, config + tenant_id, app_mode, config, session ) related_config_keys.extend(current_related_config_keys) @@ -163,7 +167,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): @classmethod def validate_agent_mode_and_set_defaults( - cls, tenant_id: str, config: dict[str, Any] + cls, tenant_id: str, config: dict[str, Any], session: Session ) -> tuple[dict[str, Any], list[str]]: """ Validate agent_mode and set defaults for agent feature @@ -220,7 +224,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager): except ValueError: raise ValueError("id in dataset must be of UUID type") - if not DatasetConfigManager.is_dataset_exists(tenant_id, tool_item["id"]): + if not DatasetConfigManager.is_dataset_exists(tenant_id, tool_item["id"], session): raise ValueError("Dataset ID does not exist, please check your permission.") else: # latest style, use key-value pair diff --git a/api/core/app/apps/agent_chat/app_generator.py b/api/core/app/apps/agent_chat/app_generator.py index 1fbef2a4092..1c639758519 100644 --- a/api/core/app/apps/agent_chat/app_generator.py +++ b/api/core/app/apps/agent_chat/app_generator.py @@ -21,13 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity, InvokeFrom +from core.db.session_factory import session_factory from core.helper.trace_id_helper import extract_trace_session_id_from_args from core.ops.ops_trace_manager import TraceQueueManager -from extensions.ext_database import db from factories import file_factory from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from libs.flask_utils import preserve_flask_contexts from models import Account, App, EndUser +from models.model import load_annotation_reply_config from services.conversation_service import ConversationService logger = logging.getLogger(__name__) @@ -43,6 +44,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[False], + session: Session, ) -> Mapping[str, Any]: ... @overload @@ -54,6 +56,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[True], + session: Session, ) -> Generator[Mapping | str, None, None]: ... @overload @@ -65,6 +68,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool, + session: Session, ) -> Mapping | Generator[Mapping | str, None, None]: ... def generate( @@ -75,6 +79,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = True, + session: Session, ) -> Mapping | Generator[Mapping | str, None, None]: """ Generate App response. @@ -108,10 +113,14 @@ 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, session=db.session() + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) # get app model config - app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) + app_model_config = self._get_app_model_config( + app_model=app_model, + conversation=conversation, + session=session, + ) # validate override model config override_model_config_dict = None @@ -123,11 +132,16 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): override_model_config_dict = AgentChatAppConfigManager.config_validate( tenant_id=app_model.tenant_id, config=args["model_config"], + session=session, ) # always enable retriever resource in debugger mode override_model_config_dict["retriever_resource"] = {"enabled": True} + annotation_reply = ( + None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id) + ) + # parse files # TODO(QuantumGhost): Move file parsing logic to the API controller layer # for better separation of concerns. @@ -137,7 +151,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from): files = args.get("files") or [] file_extra_config = FileUploadConfigManager.convert( - override_model_config_dict or app_model_config.to_dict() + override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply) ) if file_extra_config: file_objs = file_factory.build_from_mappings( @@ -155,6 +169,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): app_model_config=app_model_config, conversation=conversation, override_config_dict=override_model_config_dict, + annotation_reply=annotation_reply, ) # get tracing instance @@ -186,7 +201,11 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): ) # init generate records - (conversation, message) = self._init_generate_records(application_generate_entity, conversation) + (conversation, message) = self._init_generate_records( + application_generate_entity, + conversation, + session=session, + ) # init queue manager queue_manager = MessageBasedAppQueueManager( @@ -205,7 +224,6 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): target=self._generate_worker, kwargs={ "flask_app": current_app._get_current_object(), # type: ignore - "session": db.session(), "context": context, "application_generate_entity": application_generate_entity, "queue_manager": queue_manager, @@ -230,7 +248,6 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, - session: Session, context: contextvars.Context, application_generate_entity: AgentChatAppGenerateEntity, queue_manager: AppQueueManager, @@ -255,13 +272,14 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): # chatbot app runner = AgentChatAppRunner() - runner.run( - session=session, - application_generate_entity=application_generate_entity, - queue_manager=queue_manager, - conversation=conversation, - message=message, - ) + with session_factory.create_session() as session: + runner.run( + application_generate_entity=application_generate_entity, + queue_manager=queue_manager, + conversation=conversation, + message=message, + session=session, + ) except GenerateTaskStoppedError: pass except InvokeAuthorizationError: @@ -278,5 +296,3 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): except Exception as e: logger.exception("Unknown Error when generating") queue_manager.publish_error(e, PublishFrom.APPLICATION_MANAGER) - finally: - db.session.close() diff --git a/api/core/app/apps/agent_chat/app_runner.py b/api/core/app/apps/agent_chat/app_runner.py index 6bbc20388dd..ddfc2dac96a 100644 --- a/api/core/app/apps/agent_chat/app_runner.py +++ b/api/core/app/apps/agent_chat/app_runner.py @@ -32,14 +32,17 @@ class AgentChatAppRunner(AppRunner): def run( self, - session: Session, application_generate_entity: AgentChatAppGenerateEntity, queue_manager: AppQueueManager, conversation: Conversation, message: Message, + session: Session, ): - """ - Run assistant application + """Run the assistant application with bounded explicit transactions. + + The setup session is committed and released before the multi-step agent + runner begins model or tool I/O. + :param application_generate_entity: application generate entity :param queue_manager: application queue manager :param conversation: conversation @@ -49,10 +52,10 @@ class AgentChatAppRunner(AppRunner): app_config = application_generate_entity.app_config app_config = cast(AgentChatAppConfig, app_config) app_stmt = select(App).where(App.id == app_config.app_id) - with create_session() as session: - app_record = session.scalar(app_stmt) + with create_session() as read_session: + app_record = read_session.scalar(app_stmt) if app_record: - session.expunge(app_record) + read_session.expunge(app_record) if not app_record: raise ValueError("App not found") @@ -112,7 +115,10 @@ class AgentChatAppRunner(AppRunner): query=query, user_id=application_generate_entity.user_id, invoke_from=application_generate_entity.invoke_from, + session=session, ) + session.commit() + session.close() if annotation_reply: queue_manager.publish( @@ -191,15 +197,15 @@ class AgentChatAppRunner(AppRunner): agent_entity.strategy = AgentEntity.Strategy.FUNCTION_CALLING conversation_stmt = select(Conversation).where(Conversation.id == conversation.id) msg_stmt = select(Message).where(Message.id == message.id) - with create_session() as session: - conversation_result = session.scalar(conversation_stmt) + with create_session() as read_session: + conversation_result = read_session.scalar(conversation_stmt) if conversation_result is None: raise ValueError("Conversation not found") - message_result = session.scalar(msg_stmt) + message_result = read_session.scalar(msg_stmt) if message_result is not None: - session.expunge(message_result) - session.expunge(conversation_result) + read_session.expunge(message_result) + read_session.expunge(conversation_result) if message_result is None: raise ValueError("Message not found") @@ -234,6 +240,9 @@ class AgentChatAppRunner(AppRunner): model_instance=model_instance, ) + session.commit() + session.close() + invoke_result = runner.run( session=session, message=message, diff --git a/api/core/app/apps/base_app_runner.py b/api/core/app/apps/base_app_runner.py index 941ae6b330b..3a32db0cc1b 100644 --- a/api/core/app/apps/base_app_runner.py +++ b/api/core/app/apps/base_app_runner.py @@ -5,7 +5,7 @@ from collections.abc import Generator, Mapping, Sequence from mimetypes import guess_extension from typing import TYPE_CHECKING, Any, Union -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session from core.app.app_config.entities import ExternalDataVariableEntity, PromptTemplateEntity from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom @@ -24,6 +24,7 @@ from core.app.entities.queue_entities import ( ) from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature from core.app.features.hosting_moderation.hosting_moderation import HostingModerationFeature +from core.db.session_factory import session_factory from core.external_data_tool.external_data_fetch import ExternalDataFetch from core.memory.token_buffer_memory import TokenBufferMemory from core.model_manager import ModelInstance @@ -32,7 +33,6 @@ from core.prompt.advanced_prompt_transform import AdvancedPromptTransform from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate, MemoryConfig from core.prompt.simple_prompt_transform import ModelMode, SimplePromptTransform from core.tools.tool_file_manager import ToolFileManager -from extensions.ext_database import db from graphon.file import FileTransferMethod, FileType from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.message_entities import ( @@ -310,20 +310,30 @@ class AppRunner: case list(): for content in message.content: match content: - case str(): - text += content case TextPromptMessageContent(): text += content.data case ImagePromptMessageContent(): if message_id and user_id and tenant_id: try: - self._handle_multimodal_image_content( - content=content, - message_id=message_id, - user_id=user_id, - tenant_id=tenant_id, - queue_manager=queue_manager, - ) + with session_factory.create_session() as session: + message_file_id = self._handle_multimodal_image_content( + session=session, + content=content, + message_id=message_id, + user_id=user_id, + tenant_id=tenant_id, + queue_manager=queue_manager, + ) + session.commit() + if message_file_id: + queue_manager.publish( + QueueMessageFileEvent(message_file_id=message_file_id), + PublishFrom.APPLICATION_MANAGER, + ) + _logger.info( + "QueueMessageFileEvent published for message_file_id: %s", + message_file_id, + ) except Exception: _logger.exception("Failed to handle multimodal image output") else: @@ -365,7 +375,8 @@ class AppRunner: user_id: str, tenant_id: str, queue_manager: AppQueueManager, - ): + session: Session, + ) -> str | None: """ Handle multimodal image content from LLM response. Save the image and create a MessageFile record. @@ -386,7 +397,7 @@ class AppRunner: if not image_url and not base64_data: _logger.warning("Image content has neither URL nor base64 data") - return + return None tool_file_manager = ToolFileManager() @@ -420,10 +431,10 @@ class AppRunner: ) _logger.info("Image saved successfully, tool_file_id: %s", tool_file.id) else: - return + return None except Exception: _logger.exception("Failed to save image file") - return + return None # Create MessageFile record. # Use an independent session so this side-effect write does not @@ -441,16 +452,9 @@ class AppRunner: created_by=user_id, ) - with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: - session.add(message_file) - - # Publish QueueMessageFileEvent - queue_manager.publish( - QueueMessageFileEvent(message_file_id=message_file.id), - PublishFrom.APPLICATION_MANAGER, - ) - - _logger.info("QueueMessageFileEvent published for message_file_id: %s", message_file.id) + session.add(message_file) + session.flush() + return message_file.id def moderation_for_inputs( self, @@ -536,7 +540,7 @@ class AppRunner: ) def query_app_annotations_to_reply( - self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom + self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom, session: Session ) -> MessageAnnotation | None: """ Query app annotations to reply @@ -549,5 +553,10 @@ class AppRunner: """ annotation_reply_feature = AnnotationReplyFeature() return annotation_reply_feature.query( - app_record=app_record, message=message, query=query, user_id=user_id, invoke_from=invoke_from + app_record=app_record, + message=message, + query=query, + user_id=user_id, + invoke_from=invoke_from, + session=session, ) diff --git a/api/core/app/apps/chat/app_config_manager.py b/api/core/app/apps/chat/app_config_manager.py index 35507d65ab9..56c17b4b591 100644 --- a/api/core/app/apps/chat/app_config_manager.py +++ b/api/core/app/apps/chat/app_config_manager.py @@ -1,5 +1,7 @@ from typing import Any, cast +from sqlalchemy.orm import Session + from core.app.app_config.base_app_config_manager import BaseAppConfigManager from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager from core.app.app_config.easy_ui_based_app.dataset.manager import DatasetConfigManager @@ -15,7 +17,7 @@ from core.app.app_config.features.suggested_questions_after_answer.manager impor SuggestedQuestionsAfterAnswerConfigManager, ) from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager -from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation +from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation class ChatAppConfig(EasyUIBasedAppConfig): @@ -34,6 +36,8 @@ class ChatAppConfigManager(BaseAppConfigManager): app_model_config: AppModelConfig, conversation: Conversation | None = None, override_config_dict: AppModelConfigDict | None = None, + *, + annotation_reply: AnnotationReplyConfig | None, ) -> ChatAppConfig: """ Convert app model config to chat app config @@ -51,7 +55,7 @@ class ChatAppConfigManager(BaseAppConfigManager): config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG if config_from != EasyUIBasedAppModelConfigFrom.ARGS: - app_model_config_dict = app_model_config.to_dict() + app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply) config_dict = app_model_config_dict.copy() else: if not override_config_dict: @@ -81,7 +85,7 @@ class ChatAppConfigManager(BaseAppConfigManager): return app_config @classmethod - def config_validate(cls, tenant_id: str, config: dict[str, Any]) -> AppModelConfigDict: + def config_validate(cls, tenant_id: str, config: dict[str, Any], session: Session) -> AppModelConfigDict: """ Validate for chat app model config @@ -110,7 +114,7 @@ class ChatAppConfigManager(BaseAppConfigManager): # dataset_query_variable config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults( - tenant_id, app_mode, config + tenant_id, app_mode, config, session=session ) related_config_keys.extend(current_related_config_keys) diff --git a/api/core/app/apps/chat/app_generator.py b/api/core/app/apps/chat/app_generator.py index 678525e0f77..a67b246233f 100644 --- a/api/core/app/apps/chat/app_generator.py +++ b/api/core/app/apps/chat/app_generator.py @@ -21,13 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.entities.app_invoke_entities import ChatAppGenerateEntity, InvokeFrom +from core.db.session_factory import session_factory from core.helper.trace_id_helper import extract_trace_session_id_from_args from core.ops.ops_trace_manager import TraceQueueManager from extensions.ext_database import db from factories import file_factory from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from models import Account -from models.model import App, EndUser +from models.model import App, EndUser, load_annotation_reply_config from services.conversation_service import ConversationService logger = logging.getLogger(__name__) @@ -37,44 +38,48 @@ class ChatAppGenerator(MessageBasedAppGenerator): @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[True], + *, + session: Session, ) -> Generator[Mapping | str, None, None]: ... @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[False], + *, + session: Session, ) -> Mapping[str, Any]: ... @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool, + *, + session: Session, ) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]: ... def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]: """ Generate App response. @@ -105,10 +110,14 @@ 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, session=db.session() + app_model=app_model, conversation_id=conversation_id, user=user, session=session ) # get app model config - app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) + app_model_config = self._get_app_model_config( + app_model=app_model, + conversation=conversation, + session=session, + ) # validate override model config override_model_config_dict = None @@ -118,12 +127,16 @@ class ChatAppGenerator(MessageBasedAppGenerator): # validate config override_model_config_dict = ChatAppConfigManager.config_validate( - tenant_id=app_model.tenant_id, config=args.get("model_config", {}) + tenant_id=app_model.tenant_id, config=args.get("model_config", {}), session=session ) # always enable retriever resource in debugger mode override_model_config_dict["retriever_resource"] = {"enabled": True} + annotation_reply = ( + None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id) + ) + # parse files # TODO(QuantumGhost): Move file parsing logic to the API controller layer # for better separation of concerns. @@ -133,7 +146,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from): files = args["files"] if args.get("files") else [] file_extra_config = FileUploadConfigManager.convert( - override_model_config_dict or app_model_config.to_dict() + override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply) ) if file_extra_config: file_objs = file_factory.build_from_mappings( @@ -151,6 +164,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): app_model_config=app_model_config, conversation=conversation, override_config_dict=override_model_config_dict, + annotation_reply=annotation_reply, ) # get tracing instance @@ -183,7 +197,11 @@ class ChatAppGenerator(MessageBasedAppGenerator): ) # init generate records - (conversation, message) = self._init_generate_records(application_generate_entity, conversation) + (conversation, message) = self._init_generate_records( + application_generate_entity, + conversation, + session=session, + ) # init queue manager queue_manager = MessageBasedAppQueueManager( @@ -202,7 +220,6 @@ class ChatAppGenerator(MessageBasedAppGenerator): def worker_with_context(): return context.run( self._generate_worker, - session=session, flask_app=current_app._get_current_object(), # type: ignore application_generate_entity=application_generate_entity, queue_manager=queue_manager, @@ -229,7 +246,6 @@ class ChatAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, - session: Session, application_generate_entity: ChatAppGenerateEntity, queue_manager: AppQueueManager, conversation_id: str, @@ -252,13 +268,14 @@ class ChatAppGenerator(MessageBasedAppGenerator): # chatbot app runner = ChatAppRunner() - runner.run( - session=session, - application_generate_entity=application_generate_entity, - queue_manager=queue_manager, - conversation=conversation, - message=message, - ) + with session_factory.create_session() as session: + runner.run( + application_generate_entity=application_generate_entity, + queue_manager=queue_manager, + conversation=conversation, + message=message, + session=session, + ) except GenerateTaskStoppedError: pass except InvokeAuthorizationError: diff --git a/api/core/app/apps/chat/app_runner.py b/api/core/app/apps/chat/app_runner.py index 3ee037c82c9..25ade27e76f 100644 --- a/api/core/app/apps/chat/app_runner.py +++ b/api/core/app/apps/chat/app_runner.py @@ -17,7 +17,6 @@ from core.memory.token_buffer_memory import TokenBufferMemory from core.model_manager import ModelInstance from core.moderation.base import ModerationError from core.rag.retrieval.dataset_retrieval import DatasetRetrieval -from extensions.ext_database import db from graphon.file import File from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent from models.model import App, Conversation, Message @@ -32,14 +31,17 @@ class ChatAppRunner(AppRunner): def run( self, - session: Session, application_generate_entity: ChatAppGenerateEntity, queue_manager: AppQueueManager, conversation: Conversation, message: Message, + session: Session, ): - """ - Run application + """Run the application without retaining ``session`` during model I/O. + + Database preparation is committed and the connection is released before + the provider response is requested or consumed. + :param application_generate_entity: application generate entity :param queue_manager: application queue manager :param conversation: conversation @@ -49,10 +51,10 @@ class ChatAppRunner(AppRunner): app_config = application_generate_entity.app_config app_config = cast(ChatAppConfig, app_config) stmt = select(App).where(App.id == app_config.app_id) - with create_session() as session: - app_record = session.scalar(stmt) + with create_session() as read_session: + app_record = read_session.scalar(stmt) if app_record: - session.expunge(app_record) + read_session.expunge(app_record) if not app_record: raise ValueError("App not found") @@ -123,7 +125,10 @@ class ChatAppRunner(AppRunner): query=query, user_id=application_generate_entity.user_id, invoke_from=application_generate_entity.invoke_from, + session=session, ) + session.commit() + session.close() if annotation_reply: queue_manager.publish( @@ -188,6 +193,9 @@ class ChatAppRunner(AppRunner): ) context_files = retrieved_files or [] + session.commit() + session.close() + # reorganize all inputs and template to prompt messages # Include: prompt template, inputs, query(optional), files(optional) # memory(optional), external data, dataset context(optional) @@ -223,10 +231,6 @@ class ChatAppRunner(AppRunner): model=application_generate_entity.model_conf.model, ) - # Release the Flask scoped session before LLM streaming so a checked-out DB connection - # is not held for the lifetime of the provider response. - db.session.close() - invoke_result = model_instance.invoke_llm( prompt_messages=prompt_messages, model_parameters=application_generate_entity.model_conf.parameters, diff --git a/api/core/app/apps/completion/app_config_manager.py b/api/core/app/apps/completion/app_config_manager.py index fcfb38e8c80..9f9b8702c4f 100644 --- a/api/core/app/apps/completion/app_config_manager.py +++ b/api/core/app/apps/completion/app_config_manager.py @@ -1,5 +1,7 @@ from typing import Any, cast +from sqlalchemy.orm import Session + from core.app.app_config.base_app_config_manager import BaseAppConfigManager from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager from core.app.app_config.easy_ui_based_app.dataset.manager import DatasetConfigManager @@ -10,7 +12,7 @@ from core.app.app_config.entities import EasyUIBasedAppConfig, EasyUIBasedAppMod from core.app.app_config.features.file_upload.manager import FileUploadConfigManager from core.app.app_config.features.more_like_this.manager import MoreLikeThisConfigManager from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager -from models.model import App, AppMode, AppModelConfig, AppModelConfigDict +from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict class CompletionAppConfig(EasyUIBasedAppConfig): @@ -24,7 +26,12 @@ class CompletionAppConfig(EasyUIBasedAppConfig): class CompletionAppConfigManager(BaseAppConfigManager): @classmethod def get_app_config( - cls, app_model: App, app_model_config: AppModelConfig, override_config_dict: AppModelConfigDict | None = None + cls, + app_model: App, + app_model_config: AppModelConfig, + override_config_dict: AppModelConfigDict | None = None, + *, + annotation_reply: AnnotationReplyConfig | None, ) -> CompletionAppConfig: """ Convert app model config to completion app config @@ -39,7 +46,7 @@ class CompletionAppConfigManager(BaseAppConfigManager): config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG if config_from != EasyUIBasedAppModelConfigFrom.ARGS: - app_model_config_dict = app_model_config.to_dict() + app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply) config_dict = app_model_config_dict.copy() else: if not override_config_dict: @@ -68,7 +75,7 @@ class CompletionAppConfigManager(BaseAppConfigManager): return app_config @classmethod - def config_validate(cls, tenant_id: str, config: dict[str, Any]) -> AppModelConfigDict: + def config_validate(cls, tenant_id: str, config: dict[str, Any], session: Session) -> AppModelConfigDict: """ Validate for completion app model config @@ -97,7 +104,7 @@ class CompletionAppConfigManager(BaseAppConfigManager): # dataset_query_variable config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults( - tenant_id, app_mode, config + tenant_id, app_mode, config, session ) related_config_keys.extend(current_related_config_keys) diff --git a/api/core/app/apps/completion/app_generator.py b/api/core/app/apps/completion/app_generator.py index 5096f323354..54634fe2664 100644 --- a/api/core/app/apps/completion/app_generator.py +++ b/api/core/app/apps/completion/app_generator.py @@ -21,12 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.entities.app_invoke_entities import CompletionAppGenerateEntity, InvokeFrom +from core.db.session_factory import session_factory from core.helper.trace_id_helper import extract_trace_session_id_from_args from core.ops.ops_trace_manager import TraceQueueManager from extensions.ext_database import db from factories import file_factory from graphon.model_runtime.errors.invoke import InvokeAuthorizationError -from models import Account, App, EndUser, Message +from models import Account, App, AppModelConfig, Conversation, EndUser, Message +from models.model import load_annotation_reply_config from services.errors.app import MoreLikeThisDisabledError from services.errors.message import MessageNotExistsError @@ -37,44 +39,48 @@ class CompletionAppGenerator(MessageBasedAppGenerator): @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[True], + *, + session: Session, ) -> Generator[str | Mapping[str, Any], None, None]: ... @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: Literal[False], + *, + session: Session, ) -> Mapping[str, Any]: ... @overload def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = False, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: ... def generate( self, - session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], invoke_from: InvokeFrom, streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -96,7 +102,11 @@ class CompletionAppGenerator(MessageBasedAppGenerator): conversation = None # get app model config - app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation) + app_model_config = self._get_app_model_config( + app_model=app_model, + conversation=conversation, + session=session, + ) # validate override model config override_model_config_dict = None @@ -106,9 +116,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator): # validate config override_model_config_dict = CompletionAppConfigManager.config_validate( - tenant_id=app_model.tenant_id, config=args.get("model_config", {}) + tenant_id=app_model.tenant_id, config=args.get("model_config", {}), session=session ) + annotation_reply = ( + None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id) + ) + # parse files # TODO(QuantumGhost): Move file parsing logic to the API controller layer # for better separation of concerns. @@ -118,7 +132,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from): files = args["files"] if args.get("files") else [] file_extra_config = FileUploadConfigManager.convert( - override_model_config_dict or app_model_config.to_dict() + override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply) ) if file_extra_config: file_objs = file_factory.build_from_mappings( @@ -132,7 +146,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator): # convert to app config app_config = CompletionAppConfigManager.get_app_config( - app_model=app_model, app_model_config=app_model_config, override_config_dict=override_model_config_dict + app_model=app_model, + app_model_config=app_model_config, + override_config_dict=override_model_config_dict, + annotation_reply=annotation_reply, ) # get tracing instance @@ -161,7 +178,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator): ) # init generate records - (conversation, message) = self._init_generate_records(application_generate_entity) + (conversation, message) = self._init_generate_records( + application_generate_entity, + session=session, + ) # init queue manager queue_manager = MessageBasedAppQueueManager( @@ -180,7 +200,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def worker_with_context(): return context.run( self._generate_worker, - session=session, flask_app=current_app._get_current_object(), # type: ignore application_generate_entity=application_generate_entity, queue_manager=queue_manager, @@ -206,7 +225,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, - session: Session, application_generate_entity: CompletionAppGenerateEntity, queue_manager: AppQueueManager, message_id: str, @@ -226,12 +244,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator): # chatbot app runner = CompletionAppRunner() - runner.run( - session=session, - application_generate_entity=application_generate_entity, - queue_manager=queue_manager, - message=message, - ) + with session_factory.create_session() as session: + runner.run( + application_generate_entity=application_generate_entity, + queue_manager=queue_manager, + message=message, + session=session, + ) except GenerateTaskStoppedError: pass except InvokeAuthorizationError: @@ -253,12 +272,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def generate_more_like_this( self, - session: Session, app_model: App, message_id: str, user: Account | EndUser, invoke_from: InvokeFrom, stream: bool = True, + *, + session: Session, ) -> Mapping | Generator[Mapping | str, None, None]: """ Generate App response. @@ -276,12 +296,14 @@ class CompletionAppGenerator(MessageBasedAppGenerator): Message.from_end_user_id == (user.id if isinstance(user, EndUser) else None), Message.from_account_id == (user.id if isinstance(user, Account) else None), ) - message = db.session.scalar(stmt) + message = session.scalar(stmt) if not message: raise MessageNotExistsError() - current_app_model_config = app_model.app_model_config + current_app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) if not current_app_model_config: raise MoreLikeThisDisabledError() @@ -290,10 +312,16 @@ class CompletionAppGenerator(MessageBasedAppGenerator): if not current_app_model_config.more_like_this or more_like_this.get("enabled", False) is False: raise MoreLikeThisDisabledError() - app_model_config = message.app_model_config + conversation = session.get(Conversation, message.conversation_id) if message.conversation_id else None + app_model_config = ( + session.get(AppModelConfig, conversation.app_model_config_id) + if conversation and conversation.app_model_config_id + else None + ) if not app_model_config: raise ValueError("Message app_model_config is None") - override_model_config_dict = app_model_config.to_dict() + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + override_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply) model_dict = override_model_config_dict["model"] completion_params = model_dict.get("completion_params", {}) completion_params["temperature"] = 0.9 @@ -305,7 +333,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): file_extra_config = FileUploadConfigManager.convert(override_model_config_dict) if file_extra_config: file_objs = file_factory.build_from_mappings( - mappings=message.message_files, + mappings=message.message_files_with_session(session=session), tenant_id=app_model.tenant_id, config=file_extra_config, access_controller=self._file_access_controller, @@ -315,7 +343,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator): # convert to app config app_config = CompletionAppConfigManager.get_app_config( - app_model=app_model, app_model_config=app_model_config, override_config_dict=override_model_config_dict + app_model=app_model, + app_model_config=app_model_config, + override_config_dict=override_model_config_dict, + annotation_reply=annotation_reply, ) # init application generate entity @@ -323,7 +354,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): task_id=str(uuid.uuid4()), app_config=app_config, model_conf=ModelConfigConverter.convert(app_config), - inputs=message.inputs, + inputs=message.inputs_with_session(session=session), query=message.query, files=list(file_objs), user_id=user.id, @@ -333,7 +364,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator): ) # init generate records - (conversation, message) = self._init_generate_records(application_generate_entity) + (conversation, message) = self._init_generate_records( + application_generate_entity, + session=session, + ) # init queue manager queue_manager = MessageBasedAppQueueManager( @@ -352,7 +386,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def worker_with_context(): return context.run( self._generate_worker, - session=session, flask_app=current_app._get_current_object(), # type: ignore application_generate_entity=application_generate_entity, queue_manager=queue_manager, diff --git a/api/core/app/apps/completion/app_runner.py b/api/core/app/apps/completion/app_runner.py index b9c76569ba8..572468fba3d 100644 --- a/api/core/app/apps/completion/app_runner.py +++ b/api/core/app/apps/completion/app_runner.py @@ -15,7 +15,6 @@ from core.db.session_factory import create_session from core.model_manager import ModelInstance from core.moderation.base import ModerationError from core.rag.retrieval.dataset_retrieval import DatasetRetrieval -from extensions.ext_database import db from graphon.file import File from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent from models.model import App, Message @@ -30,13 +29,16 @@ class CompletionAppRunner(AppRunner): def run( self, - session: Session, application_generate_entity: CompletionAppGenerateEntity, queue_manager: AppQueueManager, message: Message, + session: Session, ): - """ - Run application + """Run the application without retaining ``session`` during model I/O. + + Database preparation is committed and the connection is released before + the provider response is requested or consumed. + :param application_generate_entity: application generate entity :param queue_manager: application queue manager :param message: message @@ -45,10 +47,10 @@ class CompletionAppRunner(AppRunner): app_config = application_generate_entity.app_config app_config = cast(CompletionAppConfig, app_config) stmt = select(App).where(App.id == app_config.app_id) - with create_session() as session: - app_record = session.scalar(stmt) + with create_session() as read_session: + app_record = read_session.scalar(stmt) if app_record: - session.expunge(app_record) + read_session.expunge(app_record) if not app_record: raise ValueError("App not found") @@ -150,6 +152,9 @@ class CompletionAppRunner(AppRunner): ) context_files = retrieved_files or [] + session.commit() + session.close() + # reorganize all inputs and template to prompt messages # Include: prompt template, inputs, query(optional), files(optional) # memory(optional), external data, dataset context(optional) @@ -184,10 +189,6 @@ class CompletionAppRunner(AppRunner): model=application_generate_entity.model_conf.model, ) - # Release the Flask scoped session before LLM streaming so a checked-out DB connection - # is not held for the lifetime of the provider response. - db.session.close() - invoke_result = model_instance.invoke_llm( prompt_messages=prompt_messages, model_parameters=application_generate_entity.model_conf.parameters, diff --git a/api/core/app/apps/message_based_app_generator.py b/api/core/app/apps/message_based_app_generator.py index fe61224ada5..6988ef124ea 100644 --- a/api/core/app/apps/message_based_app_generator.py +++ b/api/core/app/apps/message_based_app_generator.py @@ -89,12 +89,18 @@ class MessageBasedAppGenerator(BaseAppGenerator): logger.exception("Failed to handle response, conversation_id: %s", conversation.id) raise e - def _get_app_model_config(self, app_model: App, conversation: Conversation | None = None) -> AppModelConfig: + def _get_app_model_config( + self, + app_model: App, + conversation: Conversation | None = None, + *, + session: Session, + ) -> AppModelConfig: if conversation: stmt = select(AppModelConfig).where( AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id ) - app_model_config = db.session.scalar(stmt) + app_model_config = session.scalar(stmt) if not app_model_config: raise AppModelConfigBrokenError() @@ -102,7 +108,7 @@ class MessageBasedAppGenerator(BaseAppGenerator): if app_model.app_model_config_id is None: raise AppModelConfigBrokenError() - app_model_config = app_model.app_model_config + app_model_config = session.get(AppModelConfig, app_model.app_model_config_id) if not app_model_config: raise AppModelConfigBrokenError() @@ -118,6 +124,8 @@ class MessageBasedAppGenerator(BaseAppGenerator): AdvancedChatAppGenerateEntity, ], conversation: Conversation | None = None, + *, + session: Session, ) -> tuple[Conversation, Message]: """ Initialize generate records @@ -183,9 +191,9 @@ class MessageBasedAppGenerator(BaseAppGenerator): from_account_id=account_id, ) - db.session.add(conversation) - db.session.flush() - db.session.refresh(conversation) + session.add(conversation) + session.flush() + session.refresh(conversation) else: conversation.updated_at = naive_utc_now() @@ -216,9 +224,9 @@ class MessageBasedAppGenerator(BaseAppGenerator): app_mode=app_config.app_mode, ) - db.session.add(message) - db.session.flush() - db.session.refresh(message) + session.add(message) + session.flush() + session.refresh(message) message_files = [] for file in application_generate_entity.files: @@ -235,16 +243,16 @@ class MessageBasedAppGenerator(BaseAppGenerator): message_files.append(message_file) if message_files: - db.session.add_all(message_files) + session.add_all(message_files) - db.session.commit() + session.commit() if isinstance(application_generate_entity, ConversationAppGenerateEntity): application_generate_entity.conversation_id = conversation.id application_generate_entity.is_new_conversation = created_new_conversation return conversation, message except Exception: - db.session.rollback() + session.rollback() raise def _get_conversation_introduction(self, application_generate_entity: AppGenerateEntity) -> str: diff --git a/api/core/app/apps/pipeline/pipeline_generator.py b/api/core/app/apps/pipeline/pipeline_generator.py index cafc95d035a..7cf7949a4d5 100644 --- a/api/core/app/apps/pipeline/pipeline_generator.py +++ b/api/core/app/apps/pipeline/pipeline_generator.py @@ -64,6 +64,7 @@ class PipelineGenerator(BaseAppGenerator): def generate( self, *, + session: Session, pipeline: Pipeline, workflow: Workflow, user: Account | EndUser, @@ -79,6 +80,7 @@ class PipelineGenerator(BaseAppGenerator): def generate( self, *, + session: Session, pipeline: Pipeline, workflow: Workflow, user: Account | EndUser, @@ -94,6 +96,7 @@ class PipelineGenerator(BaseAppGenerator): def generate( self, *, + session: Session, pipeline: Pipeline, workflow: Workflow, user: Account | EndUser, @@ -108,6 +111,7 @@ class PipelineGenerator(BaseAppGenerator): def generate( self, *, + session: Session, pipeline: Pipeline, workflow: Workflow, user: Account | EndUser, @@ -120,10 +124,9 @@ class PipelineGenerator(BaseAppGenerator): ) -> Mapping[str, Any] | Generator[Mapping | str, None, None] | None: # Add null check for dataset - with Session(db.engine, expire_on_commit=False) as session: - dataset = pipeline.retrieve_dataset(session) - if not dataset: - raise ValueError("Pipeline dataset is required") + dataset = pipeline.retrieve_dataset(session) + if not dataset: + raise ValueError("Pipeline dataset is required") inputs: Mapping[str, Any] = args["inputs"] start_node_id: str = args["start_node_id"] datasource_type = DatasourceProviderType(args["datasource_type"]) @@ -157,9 +160,9 @@ class PipelineGenerator(BaseAppGenerator): batch=batch, document_form=dataset.chunk_structure, ) - db.session.add(document) + session.add(document) documents.append(document) - db.session.commit() + session.flush() # run in child thread rag_pipeline_invoke_entities = [] @@ -177,8 +180,7 @@ class PipelineGenerator(BaseAppGenerator): pipeline_id=pipeline.id, created_by=user.id, ) - db.session.add(document_pipeline_execution_log) - db.session.commit() + session.add(document_pipeline_execution_log) application_generate_entity = RagPipelineGenerateEntity( task_id=str(uuid.uuid4()), app_config=pipeline_config, @@ -227,6 +229,7 @@ class PipelineGenerator(BaseAppGenerator): ) if invoke_from == InvokeFrom.DEBUGGER or is_retry: return self._generate( + session=session, flask_app=current_app._get_current_object(), # type: ignore context=contextvars.copy_context(), pipeline=pipeline, @@ -253,6 +256,8 @@ class PipelineGenerator(BaseAppGenerator): ) ) + if invoke_from == InvokeFrom.PUBLISHED_PIPELINE and not is_retry: + session.commit() if rag_pipeline_invoke_entities: RagPipelineTaskProxy(dataset.tenant_id, user.id, rag_pipeline_invoke_entities).delay() # return batch, dataset, documents @@ -282,6 +287,7 @@ class PipelineGenerator(BaseAppGenerator): def _generate( self, *, + session: Session, flask_app: Flask, context: contextvars.Context, pipeline: Pipeline, @@ -310,7 +316,7 @@ class PipelineGenerator(BaseAppGenerator): """ with preserve_flask_contexts(flask_app, context_vars=context): # init queue manager - workflow = db.session.get(Workflow, workflow_id) + workflow = session.get(Workflow, workflow_id) if not workflow: raise ValueError(f"Workflow not found: {workflow_id}") queue_manager = PipelineQueueManager( @@ -362,6 +368,8 @@ class PipelineGenerator(BaseAppGenerator): user: Account | EndUser, args: Mapping[str, Any], streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -372,6 +380,7 @@ class PipelineGenerator(BaseAppGenerator): :param user: account or end user :param args: request args :param streaming: is streamed + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -384,10 +393,9 @@ class PipelineGenerator(BaseAppGenerator): pipeline=pipeline, workflow=workflow, start_node_id=args.get("start_node_id", "shared") ) - with Session(db.engine) as session: - dataset = pipeline.retrieve_dataset(session) - if not dataset: - raise ValueError("Pipeline dataset is required") + dataset = pipeline.retrieve_dataset(session) + if not dataset: + raise ValueError("Pipeline dataset is required") # init application generate entity - use RagPipelineGenerateEntity instead application_generate_entity = RagPipelineGenerateEntity( @@ -428,7 +436,7 @@ class PipelineGenerator(BaseAppGenerator): app_id=application_generate_entity.app_config.app_id, triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) var_loader = DraftVarLoader( engine=db.engine, @@ -438,6 +446,7 @@ class PipelineGenerator(BaseAppGenerator): ) return self._generate( + session=session, flask_app=current_app._get_current_object(), # type: ignore pipeline=pipeline, workflow_id=workflow.id, @@ -459,6 +468,8 @@ class PipelineGenerator(BaseAppGenerator): user: Account | EndUser, args: Mapping[str, Any], streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -469,6 +480,7 @@ class PipelineGenerator(BaseAppGenerator): :param user: account or end user :param args: request args :param streaming: is streamed + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -476,10 +488,9 @@ class PipelineGenerator(BaseAppGenerator): if args.get("inputs") is None: raise ValueError("inputs is required") - with Session(db.engine) as session: - dataset = pipeline.retrieve_dataset(session) - if not dataset: - raise ValueError("Pipeline dataset is required") + dataset = pipeline.retrieve_dataset(session) + if not dataset: + raise ValueError("Pipeline dataset is required") # convert to app config pipeline_config = PipelineConfigManager.get_pipeline_config( @@ -524,7 +535,7 @@ class PipelineGenerator(BaseAppGenerator): app_id=application_generate_entity.app_config.app_id, triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) var_loader = DraftVarLoader( engine=db.engine, @@ -534,6 +545,7 @@ class PipelineGenerator(BaseAppGenerator): ) return self._generate( + session=session, flask_app=current_app._get_current_object(), # type: ignore pipeline=pipeline, workflow_id=workflow.id, diff --git a/api/core/app/apps/workflow/app_generator.py b/api/core/app/apps/workflow/app_generator.py index 168b0e525d6..e8eca44cf88 100644 --- a/api/core/app/apps/workflow/app_generator.py +++ b/api/core/app/apps/workflow/app_generator.py @@ -419,6 +419,8 @@ class WorkflowAppGenerator(BaseAppGenerator): user: Account | EndUser, args: Mapping[str, Any], streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -429,6 +431,7 @@ class WorkflowAppGenerator(BaseAppGenerator): :param user: account or end user :param args: request args :param streaming: is streamed + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -478,7 +481,7 @@ class WorkflowAppGenerator(BaseAppGenerator): app_id=application_generate_entity.app_config.app_id, triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) var_loader = DraftVarLoader( engine=db.engine, @@ -508,6 +511,8 @@ class WorkflowAppGenerator(BaseAppGenerator): user: Account | EndUser, args: LoopNodeRunPayload, streaming: bool = True, + *, + session: Session, ) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: """ Generate App response. @@ -518,6 +523,7 @@ class WorkflowAppGenerator(BaseAppGenerator): :param user: account or end user :param args: request args :param streaming: is streamed + :param session: database session supplied by the caller """ if not node_id: raise ValueError("node_id is required") @@ -565,7 +571,7 @@ class WorkflowAppGenerator(BaseAppGenerator): app_id=application_generate_entity.app_config.app_id, triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, ) - draft_var_srv = WorkflowDraftVariableService(db.session()) + draft_var_srv = WorkflowDraftVariableService(session) draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id) var_loader = DraftVarLoader( engine=db.engine, diff --git a/api/core/app/features/annotation_reply/annotation_reply.py b/api/core/app/features/annotation_reply/annotation_reply.py index 3af2211188a..ca95b267015 100644 --- a/api/core/app/features/annotation_reply/annotation_reply.py +++ b/api/core/app/features/annotation_reply/annotation_reply.py @@ -1,15 +1,14 @@ import logging +from typing import cast -from sqlalchemy import select from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.rag.datasource.vdb.vector_factory import Vector from core.rag.index_processor.constant.index_type import IndexTechniqueType -from extensions.ext_database import db -from models.dataset import Dataset, DatasetCollectionBinding +from models.dataset import Dataset from models.enums import CollectionBindingType, ConversationFromSource -from models.model import App, AppAnnotationSetting, Message, MessageAnnotation +from models.model import AnnotationReplyEnabledConfig, App, Message, MessageAnnotation, load_annotation_reply_config from services.annotation_service import AppAnnotationService from services.dataset_service import DatasetCollectionBindingService @@ -25,34 +24,27 @@ class AnnotationReplyFeature: user_id: str, invoke_from: InvokeFrom, *, - session: Session | None = None, + session: Session, ) -> MessageAnnotation | None: """Return the closest annotation reply and record a hit in ``session``. - The caller may provide its transaction so the setting lookup, annotation - lookup, and hit-history write share one session. Runtime callers that do - not provide one continue to use Flask-SQLAlchemy's scoped session. - Vector-search failures are logged and return ``None``; transaction - cleanup remains the caller's responsibility. + The setting lookup, vector access, annotation lookup, and hit-history + write share the caller-owned session. Vector-search failures are logged + and return ``None``; transaction cleanup remains the caller's responsibility. """ - if session is None: - session = db.session() - - stmt = select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_record.id) - annotation_setting = session.scalar(stmt) - - if not annotation_setting: + try: + annotation_reply_config = load_annotation_reply_config(session, app_record.id) + except ValueError: return None - collection_binding_detail = session.get(DatasetCollectionBinding, annotation_setting.collection_binding_id) - - if not collection_binding_detail: + if not annotation_reply_config["enabled"]: return None + enabled_config = cast(AnnotationReplyEnabledConfig, annotation_reply_config) try: - score_threshold = annotation_setting.score_threshold or 1 - embedding_provider_name = collection_binding_detail.provider_name - embedding_model_name = collection_binding_detail.model_name + score_threshold = enabled_config["score_threshold"] or 1 + embedding_provider_name = enabled_config["embedding_model"]["embedding_provider_name"] + embedding_model_name = enabled_config["embedding_model"]["embedding_model_name"] dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding( embedding_provider_name, embedding_model_name, session, CollectionBindingType.ANNOTATION @@ -67,7 +59,7 @@ class AnnotationReplyFeature: collection_binding_id=dataset_collection_binding.id, ) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) documents = vector.search_by_vector( query=query, top_k=1, score_threshold=score_threshold, filter={"group_id": [dataset.id]} diff --git a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py index 3af48947f78..7b38d973943 100644 --- a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py +++ b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py @@ -5,7 +5,7 @@ from threading import Thread from typing import Any, cast from sqlalchemy import select -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.orm import Session from constants.tts_auto_play_timeout import TTS_AUTO_PLAY_TIMEOUT, TTS_AUTO_PLAY_YIELD_CPU_TIME from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom @@ -44,15 +44,15 @@ from core.app.entities.task_entities import ( ) from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline from core.app.task_pipeline.message_cycle_manager import MessageCycleManager -from core.app.task_pipeline.message_file_utils import MessageFileInfoDict, prepare_file_dict +from core.app.task_pipeline.message_file_utils import prepare_file_dict from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk +from core.db.session_factory import session_factory from core.model_manager import ModelInstance from core.ops.entities.trace_entity import TraceTaskName from core.ops.ops_trace_manager import TraceQueueManager, TraceTask from core.prompt.utils.prompt_message_util import PromptMessageUtil from core.prompt.utils.prompt_template_parser import PromptTemplateParser from events.message_event import message_was_created -from extensions.ext_database import db from graphon.file import FileTransferMethod from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.message_entities import ( @@ -269,8 +269,9 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat match event: case QueueErrorEvent(): - with sessionmaker(bind=db.engine).begin() as session: + with session_factory.create_session() as session: err = self.handle_error(event=event, session=session, message_id=self._message_id) + session.commit() yield self.error_to_stream_response(err) break case QueueStopEvent() | QueueMessageEndEvent(): @@ -290,17 +291,22 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat answer=output_moderation_answer ) - with sessionmaker(bind=db.engine).begin() as session: + with session_factory.create_session() as session: # Save message self._save_message(session=session, trace_manager=trace_manager) + session.commit() message_end_resp = self._message_end_to_stream_response() yield message_end_resp case QueueRetrieverResourcesEvent(): self._message_cycle_manager.handle_retriever_resources(event) case QueueAnnotationReplyEvent(): - annotation = self._message_cycle_manager.handle_annotation_reply(event) - if annotation: - self._task_state.llm_result.message.content = annotation.content + annotation_content = None + with session_factory.create_session() as session: + annotation = self._message_cycle_manager.handle_annotation_reply(event, session) + if annotation: + annotation_content = annotation.content + if annotation_content: + self._task_state.llm_result.message.content = annotation_content case QueueAgentThoughtEvent(): agent_thought_response = self._agent_thought_to_stream_response(event) if agent_thought_response is not None: @@ -477,8 +483,8 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat metadata_dict = self._task_state.metadata.model_dump(exclude_none=True) # Fetch files associated with this message - files: list[MessageFileInfoDict] = [] - with Session(db.engine, expire_on_commit=False) as session: + files: Sequence[Mapping[str, Any]] = [] + with session_factory.create_session() as session: message_files = session.scalars(select(MessageFile).where(MessageFile.message_id == self._message_id)).all() if message_files: @@ -500,13 +506,13 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat file_dict = prepare_file_dict(message_file, upload_files_map) files_list.append(file_dict) - files = files_list + files = cast(Sequence[Mapping[str, Any]], files_list) return MessageEndStreamResponse( task_id=self._application_generate_entity.task_id, id=self._message_id, metadata=metadata_dict, - files=cast(Sequence[Mapping[str, Any]], files), + files=files, ) def _agent_message_to_stream_response(self, answer: str, message_id: str) -> AgentMessageStreamResponse: @@ -526,7 +532,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat :param event: agent thought event :return: """ - with Session(db.engine, expire_on_commit=False) as session: + with session_factory.create_session() as session: agent_thought: MessageAgentThought | None = session.scalar( select(MessageAgentThought).where(MessageAgentThought.id == event.agent_thought_id).limit(1) ) diff --git a/api/core/app/task_pipeline/message_cycle_manager.py b/api/core/app/task_pipeline/message_cycle_manager.py index 5ada7d0ba2d..3bffe83ee28 100644 --- a/api/core/app/task_pipeline/message_cycle_manager.py +++ b/api/core/app/task_pipeline/message_cycle_manager.py @@ -35,7 +35,8 @@ from core.tools.signature import sign_tool_file from extensions.ext_database import db from extensions.ext_redis import redis_client from models.enums import MessageFileBelongsTo -from models.model import AppMode, Conversation, MessageAnnotation, MessageFile +from models.model import App, AppMode, Conversation, MessageAnnotation, MessageFile +from services.account_service import AccountService from services.annotation_service import AppAnnotationService logger = logging.getLogger(__name__) @@ -115,48 +116,50 @@ class MessageCycleManager: def _generate_conversation_name_worker(self, flask_app: Flask, conversation_id: str, query: str): with flask_app.app_context(): - # get conversation and message - stmt = select(Conversation).where(Conversation.id == conversation_id) - conversation = db.session.scalar(stmt) + with session_factory.create_session() as session: + # get conversation and message + stmt = select(Conversation).where(Conversation.id == conversation_id) + conversation = session.scalar(stmt) - if not conversation: - return - - if conversation.mode != AppMode.COMPLETION: - app_model = conversation.app - if not app_model: + if not conversation: return - # generate conversation name - query_hash = hashlib.md5(query.encode()).hexdigest()[:16] - cache_key = f"conv_name:{conversation_id}:{query_hash}" + if conversation.mode != AppMode.COMPLETION: + app_model = session.get(App, conversation.app_id) + if not app_model: + return - cached_name = redis_client.get(cache_key) - if cached_name: - name = cached_name.decode("utf-8") - else: - try: - name = LLMGenerator.generate_conversation_name( - app_model.tenant_id, query, conversation_id, conversation.app_id - ) - redis_client.setex(cache_key, 3600, name) - except Exception: - if dify_config.DEBUG: - logger.exception("generate conversation name failed, conversation_id: %s", conversation_id) - name = query[:47] + "..." if len(query) > 50 else query - conversation.name = name - db.session.commit() - db.session.close() + # generate conversation name + query_hash = hashlib.md5(query.encode()).hexdigest()[:16] + cache_key = f"conv_name:{conversation_id}:{query_hash}" - def handle_annotation_reply(self, event: QueueAnnotationReplyEvent) -> MessageAnnotation | None: + cached_name = redis_client.get(cache_key) + if cached_name: + name = cached_name.decode("utf-8") + else: + try: + name = LLMGenerator.generate_conversation_name( + app_model.tenant_id, query, conversation_id, conversation.app_id + ) + redis_client.setex(cache_key, 3600, name) + except Exception: + if dify_config.DEBUG: + logger.exception( + "generate conversation name failed, conversation_id: %s", conversation_id + ) + name = query[:47] + "..." if len(query) > 50 else query + conversation.name = name + session.commit() + + def handle_annotation_reply(self, event: QueueAnnotationReplyEvent, session: Session) -> MessageAnnotation | None: """ Handle annotation reply. :param event: event :return: """ - annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session=db.session()) + annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session) if annotation: - account = annotation.account + account = AccountService.get_account_by_id(annotation.account_id, session=session) self._task_state.metadata.annotation_reply = AnnotationReply( id=annotation.id, account=AnnotationReplyAccount( diff --git a/api/core/indexing_runner.py b/api/core/indexing_runner.py index 0ed91e77913..65f0b9a9523 100644 --- a/api/core/indexing_runner.py +++ b/api/core/indexing_runner.py @@ -10,9 +10,11 @@ from typing import Any from flask import Flask, current_app from sqlalchemy import delete, func, select, update +from sqlalchemy.orm import Session from sqlalchemy.orm.exc import ObjectDeletedError from configs import dify_config +from core.db.session_factory import session_factory from core.entities.knowledge_entities import IndexingEstimate, PreviewDetail, QAPreviewDetail from core.errors.error import ProviderTokenNotInitError from core.model_manager import ModelInstance, ModelManager @@ -31,7 +33,6 @@ from core.rag.splitter.fixed_text_splitter import ( ) from core.rag.splitter.text_splitter import TextSplitter from core.tools.utils.web_reader_tool import get_image_upload_file_ids -from extensions.ext_database import db from extensions.ext_redis import redis_client from extensions.ext_storage import storage from graphon.model_runtime.entities.model_entities import ModelType @@ -54,30 +55,34 @@ class IndexingRunner: def _get_model_manager(tenant_id: str) -> ModelManager: return ModelManager.for_tenant(tenant_id=tenant_id) - def _handle_indexing_error(self, document_id: str, error: Exception) -> None: + def _handle_indexing_error(self, document_id: str, error: Exception, session: Session) -> None: """Handle indexing errors by updating document status.""" logger.exception("consume document failed") - document = db.session.get(DatasetDocument, document_id) + document = session.get(DatasetDocument, document_id) if document: document.indexing_status = IndexingStatus.ERROR error_message = getattr(error, "description", str(error)) document.error = str(error_message) document.stopped_at = naive_utc_now() - db.session.commit() + session.flush() - def run(self, dataset_documents: list[DatasetDocument]): - """Run the indexing process.""" + def run(self, dataset_documents: list[DatasetDocument], session: Session): + """Run indexing with commits before slow transforms and parallel index workers. + + The phase commits keep document locks short and make newly created segments + visible to the worker sessions used for keyword and vector indexing. + """ for dataset_document in dataset_documents: document_id = dataset_document.id try: # Re-query the document to ensure it's bound to the current session - requeried_document = db.session.get(DatasetDocument, document_id) + requeried_document = session.get(DatasetDocument, document_id) if not requeried_document: logger.warning("Document not found, skipping document id: %s", document_id) continue # get dataset - dataset = db.session.get(Dataset, requeried_document.dataset_id) + dataset = session.get(Dataset, requeried_document.dataset_id) if not dataset: raise ValueError("no dataset found") @@ -85,19 +90,20 @@ class IndexingRunner: stmt = select(DatasetProcessRule).where( DatasetProcessRule.id == requeried_document.dataset_process_rule_id ) - processing_rule = db.session.scalar(stmt) + processing_rule = session.scalar(stmt) if not processing_rule: raise ValueError("no process rule found") index_type = requeried_document.doc_form index_processor = IndexProcessorFactory(index_type).init_index_processor() # extract - text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict()) + text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict(), session) + session.commit() # transform - current_user = db.session.get(Account, requeried_document.created_by) + current_user = session.get(Account, requeried_document.created_by) if not current_user: raise ValueError("no current user found") - current_user.set_tenant_id(dataset.tenant_id) + current_user.set_tenant_id_with_session(dataset.tenant_id, session=session) documents = self._transform( index_processor, dataset, @@ -105,9 +111,11 @@ class IndexingRunner: requeried_document.doc_language, processing_rule.to_dict(), current_user=current_user, + session=session, ) # save segment - self._load_segments(dataset, requeried_document, documents) + self._load_segments(dataset, requeried_document, documents, session) + session.commit() # load self._load( @@ -115,34 +123,35 @@ class IndexingRunner: dataset=dataset, dataset_document=requeried_document, documents=documents, + session=session, ) except DocumentIsPausedError: raise DocumentIsPausedError(f"Document paused, document id: {document_id}") except ProviderTokenNotInitError as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) except ObjectDeletedError: logger.warning("Document deleted, document id: %s", document_id) except Exception as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) - def run_in_splitting_status(self, dataset_document: DatasetDocument): + def run_in_splitting_status(self, dataset_document: DatasetDocument, session: Session): """Run the indexing process when the index_status is splitting.""" document_id = dataset_document.id try: # Re-query the document to ensure it's bound to the current session - requeried_document = db.session.get(DatasetDocument, document_id) + requeried_document = session.get(DatasetDocument, document_id) if not requeried_document: logger.warning("Document not found: %s", document_id) return # get dataset - dataset = db.session.get(Dataset, requeried_document.dataset_id) + dataset = session.get(Dataset, requeried_document.dataset_id) if not dataset: raise ValueError("no dataset found") # get exist document_segment list and delete - document_segments = db.session.scalars( + document_segments = session.scalars( select(DocumentSegment).where( DocumentSegment.dataset_id == dataset.id, DocumentSegment.document_id == requeried_document.id, @@ -150,27 +159,28 @@ class IndexingRunner: ).all() for document_segment in document_segments: - db.session.delete(document_segment) + session.delete(document_segment) if requeried_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: # delete child chunks - db.session.execute(delete(ChildChunk).where(ChildChunk.segment_id == document_segment.id)) - db.session.commit() + session.execute(delete(ChildChunk).where(ChildChunk.segment_id == document_segment.id)) + session.commit() # get the process rule stmt = select(DatasetProcessRule).where(DatasetProcessRule.id == requeried_document.dataset_process_rule_id) - processing_rule = db.session.scalar(stmt) + processing_rule = session.scalar(stmt) if not processing_rule: raise ValueError("no process rule found") index_type = requeried_document.doc_form index_processor = IndexProcessorFactory(index_type).init_index_processor() # extract - text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict()) + text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict(), session) + session.commit() # transform - current_user = db.session.get(Account, requeried_document.created_by) + current_user = session.get(Account, requeried_document.created_by) if not current_user: raise ValueError("no current user found") - current_user.set_tenant_id(dataset.tenant_id) + current_user.set_tenant_id_with_session(dataset.tenant_id, session=session) documents = self._transform( index_processor, dataset, @@ -178,9 +188,11 @@ class IndexingRunner: requeried_document.doc_language, processing_rule.to_dict(), current_user=current_user, + session=session, ) # save segment - self._load_segments(dataset, requeried_document, documents) + self._load_segments(dataset, requeried_document, documents, session) + session.commit() # load self._load( @@ -188,32 +200,33 @@ class IndexingRunner: dataset=dataset, dataset_document=requeried_document, documents=documents, + session=session, ) except DocumentIsPausedError: raise DocumentIsPausedError(f"Document paused, document id: {document_id}") except ProviderTokenNotInitError as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) except Exception as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) - def run_in_indexing_status(self, dataset_document: DatasetDocument): + def run_in_indexing_status(self, dataset_document: DatasetDocument, session: Session): """Run the indexing process when the index_status is indexing.""" document_id = dataset_document.id try: # Re-query the document to ensure it's bound to the current session - requeried_document = db.session.get(DatasetDocument, document_id) + requeried_document = session.get(DatasetDocument, document_id) if not requeried_document: logger.warning("Document not found: %s", document_id) return # get dataset - dataset = db.session.get(Dataset, requeried_document.dataset_id) + dataset = session.get(Dataset, requeried_document.dataset_id) if not dataset: raise ValueError("no dataset found") # get exist document_segment list and delete - document_segments = db.session.scalars( + document_segments = session.scalars( select(DocumentSegment).where( DocumentSegment.dataset_id == dataset.id, DocumentSegment.document_id == requeried_document.id, @@ -235,7 +248,7 @@ class IndexingRunner: }, ) if requeried_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = document_segment.get_child_chunks() + child_chunks = document_segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -259,13 +272,14 @@ class IndexingRunner: dataset=dataset, dataset_document=requeried_document, documents=documents, + session=session, ) except DocumentIsPausedError: raise DocumentIsPausedError(f"Document paused, document id: {document_id}") except ProviderTokenNotInitError as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) except Exception as e: - self._handle_indexing_error(document_id, e) + self._handle_indexing_error(document_id, e, session) def indexing_estimate( self, @@ -276,6 +290,8 @@ class IndexingRunner: doc_language: str = "English", dataset_id: str | None = None, indexing_technique: str = IndexTechniqueType.ECONOMY, + *, + session: Session, ) -> IndexingEstimate: """ Estimate the indexing for the document. @@ -289,7 +305,7 @@ class IndexingRunner: embedding_model_instance = None if dataset_id: - dataset = db.session.get(Dataset, dataset_id) + dataset = session.get(Dataset, dataset_id) if not dataset: raise ValueError("Dataset not found.") if IndexTechniqueType.HIGH_QUALITY in {dataset.indexing_technique, indexing_technique}: @@ -316,7 +332,6 @@ class IndexingRunner: qa_preview_texts: list[QAPreviewDetail] = [] total_segments = 0 - deleted_preview_images = False # doc_form represents the segmentation method (general, parent-child, QA) index_type = doc_form index_processor = IndexProcessorFactory(index_type).init_index_processor() @@ -328,7 +343,9 @@ class IndexingRunner: "rules": tmp_processing_rule.get("rules"), } # Extract document content - text_docs = index_processor.extract(extract_setting, process_rule_mode=tmp_processing_rule["mode"]) + text_docs = index_processor.extract( + extract_setting, process_rule_mode=tmp_processing_rule["mode"], session=session + ) # Cleaning and segmentation documents = index_processor.transform( text_docs, @@ -338,6 +355,7 @@ class IndexingRunner: tenant_id=tenant_id, doc_language=doc_language, preview=True, + session=session, ) total_segments += len(documents) for document in documents: @@ -357,7 +375,7 @@ class IndexingRunner: image_upload_file_ids = get_image_upload_file_ids(document.page_content) for upload_file_id in image_upload_file_ids: stmt = select(UploadFile).where(UploadFile.id == upload_file_id) - image_file = db.session.scalar(stmt) + image_file = session.scalar(stmt) if image_file is None: continue try: @@ -368,11 +386,11 @@ class IndexingRunner: image_upload_file_is: %s", upload_file_id, ) - db.session.delete(image_file) - deleted_preview_images = True + session.delete(image_file) - if deleted_preview_images: - db.session.commit() + # Persist preview cleanup and release the caller transaction before + # summary workers query through their own sessions. + session.commit() if doc_form and doc_form == "qa_model": return IndexingEstimate(total_segments=total_segments * 20, qa_preview=qa_preview_texts, preview=[]) @@ -381,13 +399,17 @@ class IndexingRunner: summary_index_setting = tmp_processing_rule.get("summary_index_setting") if summary_index_setting and summary_index_setting.get("enable") and preview_texts: preview_texts = index_processor.generate_summary_preview( - tenant_id, preview_texts, summary_index_setting, doc_language + tenant_id, preview_texts, summary_index_setting, doc_language, session=session ) return IndexingEstimate(total_segments=total_segments, preview=preview_texts) def _extract( - self, index_processor: BaseIndexProcessor, dataset_document: DatasetDocument, process_rule: Mapping[str, Any] + self, + index_processor: BaseIndexProcessor, + dataset_document: DatasetDocument, + process_rule: Mapping[str, Any], + session: Session, ) -> list[Document]: data_source_info = dataset_document.data_source_info_dict text_docs = [] @@ -396,7 +418,7 @@ class IndexingRunner: if not data_source_info or "upload_file_id" not in data_source_info: raise ValueError("no upload file found") stmt = select(UploadFile).where(UploadFile.id == data_source_info["upload_file_id"]) - file_detail = db.session.scalars(stmt).one_or_none() + file_detail = session.scalars(stmt).one_or_none() if file_detail: extract_setting = ExtractSetting( @@ -404,7 +426,9 @@ class IndexingRunner: upload_file=file_detail, document_model=dataset_document.doc_form, ) - text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"]) + text_docs = index_processor.extract( + extract_setting, process_rule_mode=process_rule["mode"], session=session + ) case DataSourceType.NOTION_IMPORT: if ( not data_source_info @@ -426,7 +450,9 @@ class IndexingRunner: ), document_model=dataset_document.doc_form, ) - text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"]) + text_docs = index_processor.extract( + extract_setting, process_rule_mode=process_rule["mode"], session=session + ) case DataSourceType.WEBSITE_CRAWL: if ( not data_source_info @@ -449,11 +475,14 @@ class IndexingRunner: ), document_model=dataset_document.doc_form, ) - text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"]) + text_docs = index_processor.extract( + extract_setting, process_rule_mode=process_rule["mode"], session=session + ) case _: return [] # update document status to splitting self._update_document_index_status( + session=session, document_id=dataset_document.id, after_indexing_status=IndexingStatus.SPLITTING, extra_update_params={ @@ -576,6 +605,7 @@ class IndexingRunner: dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document], + session: Session, ): """ insert index and update document/segment status to completed @@ -625,10 +655,10 @@ class IndexingRunner: executor.submit( self._process_chunk, current_app._get_current_object(), # type: ignore - index_processor, + dataset_document.doc_form, chunk_documents, - dataset, - dataset_document, + dataset.id, + dataset_document.id, embedding_model_instance, ) ) @@ -645,6 +675,7 @@ class IndexingRunner: # update document status to completed self._update_document_index_status( + session=session, document_id=dataset_document.id, after_indexing_status=IndexingStatus.COMPLETED, extra_update_params={ @@ -656,20 +687,80 @@ class IndexingRunner: ) @staticmethod - def _process_keyword_index(flask_app, dataset_id, document_id, documents): + def _process_keyword_index(flask_app: Flask, dataset_id: str, document_id: str, documents: list[Document]): with flask_app.app_context(): - dataset = db.session.get(Dataset, dataset_id) - if not dataset: - raise ValueError("no dataset found") - keyword = Keyword(dataset) - keyword.create(documents) - if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY: - document_ids = [document.metadata["doc_id"] for document in documents] - db.session.execute( + with session_factory.create_session() as session: + dataset = session.get(Dataset, dataset_id) + if not dataset: + raise ValueError("no dataset found") + keyword = Keyword(dataset) + keyword.create(documents, session) + if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY: + document_ids = [document.metadata["doc_id"] for document in documents] + session.execute( + update(DocumentSegment) + .where( + DocumentSegment.document_id == document_id, + DocumentSegment.dataset_id == dataset_id, + DocumentSegment.index_node_id.in_(document_ids), + DocumentSegment.status == SegmentStatus.INDEXING, + ) + .values( + status=SegmentStatus.COMPLETED, + enabled=True, + completed_at=naive_utc_now(), + ) + ) + session.commit() + + def _process_chunk( + self, + flask_app: Flask, + index_type: str, + chunk_documents: list[Document], + dataset_id: str, + dataset_document_id: str, + embedding_model_instance: ModelInstance | None, + ): + with flask_app.app_context(): + with session_factory.create_session() as session: + dataset = session.get(Dataset, dataset_id) + if not dataset: + raise ValueError("no dataset found") + + dataset_document = session.get(DatasetDocument, dataset_document_id) + if not dataset_document: + raise ValueError("no document found") + + # check document is paused + self._check_document_paused_status(dataset_document.id) + + tokens = 0 + if embedding_model_instance: + page_content_list = [document.page_content for document in chunk_documents] + tokens += sum(embedding_model_instance.get_text_embedding_num_tokens(page_content_list)) + + multimodal_documents = [] + for document in chunk_documents: + if document.attachments and dataset.is_multimodal: + multimodal_documents.extend(document.attachments) + + # load index + index_processor = IndexProcessorFactory(index_type).init_index_processor() + index_processor.load( + dataset, + chunk_documents, + multimodal_documents=multimodal_documents, + with_keywords=False, + session=session, + ) + + document_ids = [document.metadata["doc_id"] for document in chunk_documents] + session.execute( update(DocumentSegment) .where( - DocumentSegment.document_id == document_id, - DocumentSegment.dataset_id == dataset_id, + DocumentSegment.document_id == dataset_document.id, + DocumentSegment.dataset_id == dataset.id, DocumentSegment.index_node_id.in_(document_ids), DocumentSegment.status == SegmentStatus.INDEXING, ) @@ -680,55 +771,9 @@ class IndexingRunner: ) ) - db.session.commit() + session.commit() - def _process_chunk( - self, - flask_app: Flask, - index_processor: BaseIndexProcessor, - chunk_documents: list[Document], - dataset: Dataset, - dataset_document: DatasetDocument, - embedding_model_instance: ModelInstance | None, - ): - with flask_app.app_context(): - # check document is paused - self._check_document_paused_status(dataset_document.id) - - tokens = 0 - if embedding_model_instance: - page_content_list = [document.page_content for document in chunk_documents] - tokens += sum(embedding_model_instance.get_text_embedding_num_tokens(page_content_list)) - - multimodal_documents = [] - for document in chunk_documents: - if document.attachments and dataset.is_multimodal: - multimodal_documents.extend(document.attachments) - - # load index - index_processor.load( - dataset, chunk_documents, multimodal_documents=multimodal_documents, with_keywords=False - ) - - document_ids = [document.metadata["doc_id"] for document in chunk_documents] - db.session.execute( - update(DocumentSegment) - .where( - DocumentSegment.document_id == dataset_document.id, - DocumentSegment.dataset_id == dataset.id, - DocumentSegment.index_node_id.in_(document_ids), - DocumentSegment.status == SegmentStatus.INDEXING, - ) - .values( - status=SegmentStatus.COMPLETED, - enabled=True, - completed_at=naive_utc_now(), - ) - ) - - db.session.commit() - - return tokens + return tokens @staticmethod def _check_document_paused_status(document_id: str): @@ -742,12 +787,14 @@ class IndexingRunner: document_id: str, after_indexing_status: IndexingStatus, extra_update_params: Mapping[Any, Any] | None = None, + *, + session: Session, ): """ Update the document indexing status. """ count = ( - db.session.scalar( + session.scalar( select(func.count()) .select_from(DatasetDocument) .where(DatasetDocument.id == document_id, DatasetDocument.is_paused == True) @@ -756,7 +803,7 @@ class IndexingRunner: ) if count > 0: raise DocumentIsPausedError() - document = db.session.get(DatasetDocument, document_id) + document = session.get(DatasetDocument, document_id) if not document: raise DocumentIsDeletedPausedError() @@ -764,18 +811,18 @@ class IndexingRunner: if extra_update_params: update_params.update(extra_update_params) - db.session.execute(update(DatasetDocument).where(DatasetDocument.id == document_id).values(update_params)) # type: ignore - db.session.commit() + session.execute(update(DatasetDocument).where(DatasetDocument.id == document_id).values(update_params)) # type: ignore + session.flush() @staticmethod - def _update_segments_by_document(dataset_document_id: str, update_params: Mapping[Any, Any]): + def _update_segments_by_document(dataset_document_id: str, update_params: Mapping[Any, Any], session: Session): """ Update the document segment by document id. """ - db.session.execute( + session.execute( update(DocumentSegment).where(DocumentSegment.document_id == dataset_document_id).values(update_params) ) - db.session.commit() + session.flush() def _transform( self, @@ -785,6 +832,8 @@ class IndexingRunner: doc_language: str, process_rule: Mapping[str, Any], current_user: Account | None = None, + *, + session: Session, ) -> list[Document]: # get embedding model instance embedding_model_instance = None @@ -809,11 +858,14 @@ class IndexingRunner: process_rule=process_rule, tenant_id=dataset.tenant_id, doc_language=doc_language, + session=session, ) return documents - def _load_segments(self, dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document]): + def _load_segments( + self, dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document], session: Session + ): # save node to document segment doc_store = DatasetDocumentStore( dataset=dataset, user_id=dataset_document.created_by, document_id=dataset_document.id @@ -821,12 +873,15 @@ class IndexingRunner: # add document segments doc_store.add_documents( - docs=documents, save_child=dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX + docs=documents, + save_child=dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX, + session=session, ) # update document status to indexing cur_time = naive_utc_now() self._update_document_index_status( + session=session, document_id=dataset_document.id, after_indexing_status=IndexingStatus.INDEXING, extra_update_params={ @@ -838,6 +893,7 @@ class IndexingRunner: # update segment status to indexing self._update_segments_by_document( + session=session, dataset_document_id=dataset_document.id, update_params={ DocumentSegment.status: SegmentStatus.INDEXING, diff --git a/api/core/ops/base_trace_instance.py b/api/core/ops/base_trace_instance.py index a1f96b9edf4..c0095858d46 100644 --- a/api/core/ops/base_trace_instance.py +++ b/api/core/ops/base_trace_instance.py @@ -63,6 +63,6 @@ class BaseTraceInstance(ABC): ) if not current_tenant: raise ValueError(f"Current tenant not found for account {service_account.id}") - service_account.set_tenant_id(current_tenant.tenant_id) + service_account.set_tenant_id_with_session(current_tenant.tenant_id, session=session) return service_account diff --git a/api/core/plugin/backwards_invocation/app.py b/api/core/plugin/backwards_invocation/app.py index a74be9be2d8..8e9eeb64fe0 100644 --- a/api/core/plugin/backwards_invocation/app.py +++ b/api/core/plugin/backwards_invocation/app.py @@ -61,7 +61,6 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): @classmethod def invoke_app( cls, - session: Session, app_id: str, user_id: str, tenant_id: str, @@ -70,6 +69,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): stream: bool, inputs: Mapping, files: list[dict], + session: Session, ) -> Generator[Mapping | str, None, None] | Mapping: """ invoke app @@ -91,21 +91,20 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): if not query: raise ValueError("missing query") - return cls.invoke_chat_app(session, app, user, conversation_id, query, stream, inputs, files) + return cls.invoke_chat_app(app, user, conversation_id, query, stream, inputs, files, session) case AppMode.WORKFLOW: workflow = cls._get_workflow(app) if not workflow: raise ValueError("unexpected app type") return cls.invoke_workflow_app(app, workflow, user, stream, inputs, files) case AppMode.COMPLETION: - return cls.invoke_completion_app(session, app, user, stream, inputs, files) + return cls.invoke_completion_app(app, user, stream, inputs, files, session) case _: raise ValueError("unexpected app type") @classmethod def invoke_chat_app( cls, - session: Session, app: App, user: Account | EndUser, conversation_id: str, @@ -113,6 +112,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): stream: bool, inputs: Mapping, files: list[dict], + session: Session, ) -> Generator[Mapping | str, None, None] | Mapping: """ invoke chat app @@ -142,6 +142,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): workflow_run_id=str(uuid.uuid4()), streaming=stream, pause_state_config=pause_config, + session=session, ) case AppMode.AGENT_CHAT: return AgentChatAppGenerator().generate( @@ -155,10 +156,10 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): }, invoke_from=InvokeFrom.SERVICE_API, streaming=stream, + session=session, ) case AppMode.CHAT: return ChatAppGenerator().generate( - session=session, app_model=app, user=user, args={ @@ -169,6 +170,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): }, invoke_from=InvokeFrom.SERVICE_API, streaming=stream, + session=session, ) case _: raise ValueError("unexpected app type") @@ -205,23 +207,23 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): @classmethod def invoke_completion_app( cls, - session: Session, app: App, user: EndUser | Account, stream: bool, inputs: Mapping, files: list[dict], + session: Session, ) -> Generator[Mapping | str, None, None] | Mapping: """ invoke completion app """ return CompletionAppGenerator().generate( - session=session, app_model=app, user=user, args={"inputs": inputs, "files": files}, invoke_from=InvokeFrom.SERVICE_API, streaming=stream, + session=session, ) @classmethod diff --git a/api/core/prompt/utils/get_thread_messages_length.py b/api/core/prompt/utils/get_thread_messages_length.py index de64c27a734..6e4b3b1e550 100644 --- a/api/core/prompt/utils/get_thread_messages_length.py +++ b/api/core/prompt/utils/get_thread_messages_length.py @@ -1,18 +1,18 @@ from sqlalchemy import select +from sqlalchemy.orm import Session from core.prompt.utils.extract_thread_messages import extract_thread_messages -from extensions.ext_database import db from models.model import Message -def get_thread_messages_length(conversation_id: str) -> int: +def get_thread_messages_length(conversation_id: str, *, session: Session) -> int: """ Get the number of thread messages based on the parent message id. """ # Fetch all messages related to the conversation stmt = select(Message).where(Message.conversation_id == conversation_id).order_by(Message.created_at.desc()) - messages = db.session.scalars(stmt).all() + messages = session.scalars(stmt).all() # Extract thread messages thread_messages = extract_thread_messages(messages) diff --git a/api/core/rag/data_post_processor/data_post_processor.py b/api/core/rag/data_post_processor/data_post_processor.py index ca530748edf..ace81bd79e1 100644 --- a/api/core/rag/data_post_processor/data_post_processor.py +++ b/api/core/rag/data_post_processor/data_post_processor.py @@ -1,5 +1,7 @@ from typing import TypedDict +from sqlalchemy.orm import Session + from core.model_manager import ModelInstance, ModelManager from core.rag.data_post_processor.reorder import ReorderRunner from core.rag.index_processor.constant.query_type import QueryType @@ -42,8 +44,12 @@ class DataPostProcessor: reranking_model: RerankingModelDict | None = None, weights: WeightsDict | None = None, reorder_enabled: bool = False, + *, + session: Session, ): - self.rerank_runner = self._get_rerank_runner(reranking_mode, tenant_id, reranking_model, weights) + self.rerank_runner = self._get_rerank_runner( + reranking_mode, tenant_id, reranking_model, weights, session=session + ) self.reorder_runner = self._get_reorder_runner(reorder_enabled) def invoke( @@ -68,6 +74,8 @@ class DataPostProcessor: tenant_id: str, reranking_model: RerankingModelDict | None = None, weights: WeightsDict | None = None, + *, + session: Session, ) -> BaseRerankRunner | None: if reranking_mode == RerankMode.WEIGHTED_SCORE and weights: runner = RerankRunnerFactory.create_rerank_runner( @@ -90,7 +98,9 @@ class DataPostProcessor: if rerank_model_instance is None: return None runner = RerankRunnerFactory.create_rerank_runner( - runner_type=reranking_mode, rerank_model_instance=rerank_model_instance + runner_type=reranking_mode, + rerank_model_instance=rerank_model_instance, + session=session, ) return runner return None diff --git a/api/core/rag/datasource/keyword/jieba/jieba.py b/api/core/rag/datasource/keyword/jieba/jieba.py index be83f69c48b..62e79d8c84d 100644 --- a/api/core/rag/datasource/keyword/jieba/jieba.py +++ b/api/core/rag/datasource/keyword/jieba/jieba.py @@ -4,12 +4,12 @@ from typing import Any, TypedDict, override import orjson from pydantic import BaseModel from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.rag.datasource.keyword.jieba.jieba_keyword_table_handler import JiebaKeywordTableHandler from core.rag.datasource.keyword.keyword_base import BaseKeyword from core.rag.models.document import Document -from extensions.ext_database import db from extensions.ext_redis import redis_client from extensions.ext_storage import storage from models.dataset import Dataset, DatasetKeywordTable, DocumentSegment @@ -30,32 +30,32 @@ class Jieba(BaseKeyword): self._config = KeywordTableConfig() @override - def create(self, texts: list[Document], **kwargs) -> BaseKeyword: + def create(self, texts: list[Document], session: Session, **kwargs: Any) -> BaseKeyword: lock_name = f"keyword_indexing_lock_{self.dataset.id}" with redis_client.lock(lock_name, timeout=600): keyword_table_handler = JiebaKeywordTableHandler() - keyword_table = self._get_dataset_keyword_table() + keyword_table = self._get_dataset_keyword_table(session=session) keyword_number = self.dataset.keyword_number or self._config.max_keywords_per_chunk for text in texts: keywords = keyword_table_handler.extract_keywords(text.page_content, keyword_number) if text.metadata is not None: - self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords)) + self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords), session) keyword_table = self._add_text_to_keyword_table( keyword_table or {}, text.metadata["doc_id"], list(keywords) ) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) return self @override - def add_texts(self, texts: list[Document], **kwargs): + def add_texts(self, texts: list[Document], session: Session, **kwargs: Any): lock_name = f"keyword_indexing_lock_{self.dataset.id}" with redis_client.lock(lock_name, timeout=600): keyword_table_handler = JiebaKeywordTableHandler() - keyword_table = self._get_dataset_keyword_table() + keyword_table = self._get_dataset_keyword_table(session=session) keywords_list = kwargs.get("keywords_list") keyword_number = self.dataset.keyword_number or self._config.max_keywords_per_chunk for i in range(len(texts)): @@ -67,33 +67,47 @@ class Jieba(BaseKeyword): else: keywords = keyword_table_handler.extract_keywords(text.page_content, keyword_number) if text.metadata is not None: - self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords)) + self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords), session) keyword_table = self._add_text_to_keyword_table( keyword_table or {}, text.metadata["doc_id"], list(keywords) ) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) @override - def text_exists(self, id: str) -> bool: - keyword_table = self._get_dataset_keyword_table() + def text_exists(self, id: str, *, session: Session) -> bool: + dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session) + keyword_table = None + keyword_table_dict = ( + dataset_keyword_table.get_keyword_table_dict(session=session) if dataset_keyword_table else None + ) + if keyword_table_dict: + data: Any = keyword_table_dict["__data__"] + keyword_table = dict(data["table"]) if keyword_table is None: return False return id in set.union(*keyword_table.values()) @override - def delete_by_ids(self, ids: list[str]): + def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any): lock_name = f"keyword_indexing_lock_{self.dataset.id}" with redis_client.lock(lock_name, timeout=600): - keyword_table = self._get_dataset_keyword_table() + keyword_table = self._get_dataset_keyword_table(session) if keyword_table is not None: keyword_table = self._delete_ids_from_keyword_table(keyword_table, ids) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) @override - def search(self, query: str, **kwargs: Any) -> list[Document]: - keyword_table = self._get_dataset_keyword_table() + def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]: + dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session) + keyword_table = None + keyword_table_dict = ( + dataset_keyword_table.get_keyword_table_dict(session=session) if dataset_keyword_table else None + ) + if keyword_table_dict: + data: Any = keyword_table_dict["__data__"] + keyword_table = dict(data["table"]) k = kwargs.get("top_k", 4) document_ids_filter = kwargs.get("document_ids_filter") @@ -107,7 +121,7 @@ class Jieba(BaseKeyword): if document_ids_filter: segment_query_stmt = segment_query_stmt.where(DocumentSegment.document_id.in_(document_ids_filter)) - segments = db.session.scalars(segment_query_stmt).all() + segments = session.scalars(segment_query_stmt).all() segment_map = {segment.index_node_id: segment for segment in segments} for chunk_index in sorted_chunk_indices: segment = segment_map.get(chunk_index) @@ -128,39 +142,43 @@ class Jieba(BaseKeyword): return documents @override - def delete(self): + def delete(self, *, session: Session): lock_name = f"keyword_indexing_lock_{self.dataset.id}" with redis_client.lock(lock_name, timeout=600): - dataset_keyword_table = self.dataset.dataset_keyword_table + dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session) if dataset_keyword_table: - db.session.delete(dataset_keyword_table) - db.session.commit() + session.delete(dataset_keyword_table) + session.commit() if dataset_keyword_table.data_source_type != "database": file_key = "keyword_files/" + self.dataset.tenant_id + "/" + self.dataset.id + ".txt" storage.delete(file_key) - def _save_dataset_keyword_table(self, keyword_table: dict[str, set[str]] | None): + def _save_dataset_keyword_table(self, keyword_table: dict[str, set[str]] | None, session: Session): keyword_table_dict = { "__type__": "keyword_table", "__data__": {"index_id": self.dataset.id, "summary": None, "table": keyword_table}, } - dataset_keyword_table = self.dataset.dataset_keyword_table + dataset_keyword_table = session.scalar( + select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.dataset.id) + ) keyword_data_source_type = dataset_keyword_table.data_source_type if dataset_keyword_table else "file" if keyword_data_source_type == "database": if dataset_keyword_table is None: return dataset_keyword_table.keyword_table = dumps_with_sets(keyword_table_dict) - db.session.commit() + session.flush() else: file_key = "keyword_files/" + self.dataset.tenant_id + "/" + self.dataset.id + ".txt" if storage.exists(file_key): storage.delete(file_key) storage.save(file_key, dumps_with_sets(keyword_table_dict).encode("utf-8")) - def _get_dataset_keyword_table(self) -> dict[str, set[str]] | None: - dataset_keyword_table = self.dataset.dataset_keyword_table + def _get_dataset_keyword_table(self, session: Session) -> dict[str, set[str]] | None: + dataset_keyword_table = session.scalar( + select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.dataset.id) + ) if dataset_keyword_table: - keyword_table_dict = dataset_keyword_table.keyword_table_dict + keyword_table_dict = dataset_keyword_table.get_keyword_table_dict(session=session) if keyword_table_dict: data: Any = keyword_table_dict["__data__"] return dict(data["table"]) @@ -178,8 +196,8 @@ class Jieba(BaseKeyword): "__data__": {"index_id": self.dataset.id, "summary": None, "table": {}}, } ) - db.session.add(dataset_keyword_table) - db.session.commit() + session.add(dataset_keyword_table) + session.flush() return {} @@ -228,25 +246,25 @@ class Jieba(BaseKeyword): return sorted_chunk_indices[:k] - def _update_segment_keywords(self, dataset_id: str, node_id: str, keywords: list[str]): + def _update_segment_keywords(self, dataset_id: str, node_id: str, keywords: list[str], session: Session): stmt = select(DocumentSegment).where( DocumentSegment.dataset_id == dataset_id, DocumentSegment.index_node_id == node_id ) - document_segment = db.session.scalar(stmt) + document_segment = session.scalar(stmt) if document_segment: document_segment.keywords = keywords - db.session.add(document_segment) - db.session.commit() + session.add(document_segment) + session.flush() - def create_segment_keywords(self, node_id: str, keywords: list[str]): - keyword_table = self._get_dataset_keyword_table() - self._update_segment_keywords(self.dataset.id, node_id, keywords) + def create_segment_keywords(self, node_id: str, keywords: list[str], session: Session): + keyword_table = self._get_dataset_keyword_table(session) + self._update_segment_keywords(self.dataset.id, node_id, keywords, session) keyword_table = self._add_text_to_keyword_table(keyword_table or {}, node_id, keywords) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) - def multi_create_segment_keywords(self, pre_segment_data_list: list[PreSegmentData]): + def multi_create_segment_keywords(self, pre_segment_data_list: list[PreSegmentData], session: Session): keyword_table_handler = JiebaKeywordTableHandler() - keyword_table = self._get_dataset_keyword_table() + keyword_table = self._get_dataset_keyword_table(session) for pre_segment_data in pre_segment_data_list: segment = pre_segment_data["segment"] if pre_segment_data["keywords"]: @@ -264,12 +282,12 @@ class Jieba(BaseKeyword): keyword_table = self._add_text_to_keyword_table( keyword_table or {}, segment.index_node_id, list(keywords) ) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) - def update_segment_keywords_index(self, node_id: str, keywords: list[str]): - keyword_table = self._get_dataset_keyword_table() + def update_segment_keywords_index(self, node_id: str, keywords: list[str], session: Session): + keyword_table = self._get_dataset_keyword_table(session) keyword_table = self._add_text_to_keyword_table(keyword_table or {}, node_id, keywords) - self._save_dataset_keyword_table(keyword_table) + self._save_dataset_keyword_table(keyword_table, session) def set_orjson_default(obj: Any): diff --git a/api/core/rag/datasource/keyword/keyword_base.py b/api/core/rag/datasource/keyword/keyword_base.py index 0a59855306e..f9bddeb55d5 100644 --- a/api/core/rag/datasource/keyword/keyword_base.py +++ b/api/core/rag/datasource/keyword/keyword_base.py @@ -3,6 +3,8 @@ from __future__ import annotations from abc import ABC, abstractmethod from typing import Any +from sqlalchemy.orm import Session + from core.rag.models.document import Document from models.dataset import Dataset @@ -12,35 +14,35 @@ class BaseKeyword(ABC): self.dataset = dataset @abstractmethod - def create(self, texts: list[Document], **kwargs) -> BaseKeyword: + def create(self, texts: list[Document], session: Session, **kwargs: Any) -> BaseKeyword: raise NotImplementedError @abstractmethod - def add_texts(self, texts: list[Document], **kwargs): + def add_texts(self, texts: list[Document], session: Session, **kwargs: Any): raise NotImplementedError @abstractmethod - def text_exists(self, id: str) -> bool: + def text_exists(self, id: str, *, session: Session) -> bool: raise NotImplementedError @abstractmethod - def delete_by_ids(self, ids: list[str]): + def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any): raise NotImplementedError @abstractmethod - def delete(self): + def delete(self, *, session: Session): raise NotImplementedError @abstractmethod - def search(self, query: str, **kwargs: Any) -> list[Document]: + def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]: raise NotImplementedError - def _filter_duplicate_texts(self, texts: list[Document]) -> list[Document]: + def _filter_duplicate_texts(self, texts: list[Document], *, session: Session) -> list[Document]: for text in texts.copy(): if text.metadata is None: continue doc_id = text.metadata["doc_id"] - exists_duplicate_node = self.text_exists(doc_id) + exists_duplicate_node = self.text_exists(doc_id, session=session) if exists_duplicate_node: texts.remove(text) diff --git a/api/core/rag/datasource/keyword/keyword_factory.py b/api/core/rag/datasource/keyword/keyword_factory.py index b2e1a55eecc..463ebebe166 100644 --- a/api/core/rag/datasource/keyword/keyword_factory.py +++ b/api/core/rag/datasource/keyword/keyword_factory.py @@ -1,5 +1,7 @@ from typing import Any +from sqlalchemy.orm import Session + from configs import dify_config from core.rag.datasource.keyword.keyword_base import BaseKeyword from core.rag.datasource.keyword.keyword_type import KeyWordType @@ -27,23 +29,23 @@ class Keyword: case _: raise ValueError(f"Keyword store {keyword_type} is not supported.") - def create(self, texts: list[Document], **kwargs): - self._keyword_processor.create(texts, **kwargs) + def create(self, texts: list[Document], session: Session, **kwargs: Any): + self._keyword_processor.create(texts, session, **kwargs) - def add_texts(self, texts: list[Document], **kwargs): - self._keyword_processor.add_texts(texts, **kwargs) + def add_texts(self, texts: list[Document], session: Session, **kwargs: Any): + self._keyword_processor.add_texts(texts, session, **kwargs) - def text_exists(self, id: str) -> bool: - return self._keyword_processor.text_exists(id) + def text_exists(self, id: str, *, session: Session) -> bool: + return self._keyword_processor.text_exists(id, session=session) - def delete_by_ids(self, ids: list[str]): - self._keyword_processor.delete_by_ids(ids) + def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any): + self._keyword_processor.delete_by_ids(ids, session, **kwargs) - def delete(self): - self._keyword_processor.delete() + def delete(self, *, session: Session): + self._keyword_processor.delete(session=session) - def search(self, query: str, **kwargs: Any) -> list[Document]: - return self._keyword_processor.search(query, **kwargs) + def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]: + return self._keyword_processor.search(query, session=session, **kwargs) def __getattr__(self, name): if self._keyword_processor is not None: diff --git a/api/core/rag/datasource/retrieval_service.py b/api/core/rag/datasource/retrieval_service.py index 3b20f8bc530..ea6eee6f68a 100644 --- a/api/core/rag/datasource/retrieval_service.py +++ b/api/core/rag/datasource/retrieval_service.py @@ -12,7 +12,6 @@ from sqlalchemy.orm import Session, load_only from configs import dify_config from core.app.file_access import grant_upload_file_access -from core.db.session_factory import session_factory from core.model_manager import ModelManager from core.rag.data_post_processor.data_post_processor import DataPostProcessor, RerankingModelDict, WeightsDict from core.rag.datasource.keyword.keyword_factory import Keyword @@ -303,9 +302,13 @@ class RetrievalService: keyword = Keyword(dataset=dataset) - documents = keyword.search( - cls.escape_query_for_search(query), top_k=top_k, document_ids_filter=document_ids_filter - ) + with Session(db.engine) as session: + documents = keyword.search( + cls.escape_query_for_search(query), + session=session, + top_k=top_k, + document_ids_filter=document_ids_filter, + ) all_documents.extend(documents) except Exception as e: logger.error(e, exc_info=True) @@ -333,7 +336,6 @@ class RetrievalService: if not dataset: raise ValueError("dataset not found") - vector = Vector(dataset=dataset) documents = [] # Hybrid search merges keyword / full-text / vector hits and then reranks # (weighted fusion or reranking model). Applying the user score threshold at @@ -342,29 +344,31 @@ class RetrievalService: embedding_score_threshold = ( 0.0 if retrieval_method == RetrievalMethod.HYBRID_SEARCH else score_threshold ) - if query_type == QueryType.TEXT_QUERY: - documents.extend( - vector.search_by_vector( - query, - search_type="similarity_score_threshold", - top_k=top_k, - score_threshold=embedding_score_threshold, - filter={"group_id": [dataset.id]}, - document_ids_filter=document_ids_filter, + with Session(db.engine) as session: + vector = Vector(dataset=dataset, session=session) + if query_type == QueryType.TEXT_QUERY: + documents.extend( + vector.search_by_vector( + query, + search_type="similarity_score_threshold", + top_k=top_k, + score_threshold=embedding_score_threshold, + filter={"group_id": [dataset.id]}, + document_ids_filter=document_ids_filter, + ) ) - ) - if query_type == QueryType.IMAGE_QUERY: - if not dataset.is_multimodal: - return - documents.extend( - vector.search_by_file( - file_id=query, - top_k=top_k, - score_threshold=embedding_score_threshold, - filter={"group_id": [dataset.id]}, - document_ids_filter=document_ids_filter, + if query_type == QueryType.IMAGE_QUERY: + if not dataset.is_multimodal: + return + documents.extend( + vector.search_by_file( + file_id=query, + top_k=top_k, + score_threshold=embedding_score_threshold, + filter={"group_id": [dataset.id]}, + document_ids_filter=document_ids_filter, + ) ) - ) if documents: if ( @@ -373,18 +377,37 @@ class RetrievalService: and reranking_model["reranking_provider_name"] and retrieval_method == RetrievalMethod.SEMANTIC_SEARCH ): - data_post_processor = DataPostProcessor( - str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False - ) - if dataset.is_multimodal: - model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id) - is_support_vision = model_manager.check_model_support_vision( - tenant_id=dataset.tenant_id, - provider=reranking_model["reranking_provider_name"], - model=reranking_model["reranking_model_name"], - model_type=ModelType.RERANK, + with Session(db.engine) as rerank_session: + data_post_processor = DataPostProcessor( + str(dataset.tenant_id), + str(RerankMode.RERANKING_MODEL), + reranking_model, + None, + False, + session=rerank_session, ) - if is_support_vision: + if dataset.is_multimodal: + model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id) + is_support_vision = model_manager.check_model_support_vision( + tenant_id=dataset.tenant_id, + provider=reranking_model["reranking_provider_name"], + model=reranking_model["reranking_model_name"], + model_type=ModelType.RERANK, + ) + if is_support_vision: + all_documents.extend( + data_post_processor.invoke( + query=query, + documents=documents, + score_threshold=score_threshold, + top_n=len(documents), + query_type=query_type, + ) + ) + else: + # not effective, return original documents + all_documents.extend(documents) + else: all_documents.extend( data_post_processor.invoke( query=query, @@ -394,19 +417,6 @@ class RetrievalService: query_type=query_type, ) ) - else: - # not effective, return original documents - all_documents.extend(documents) - else: - all_documents.extend( - data_post_processor.invoke( - query=query, - documents=documents, - score_threshold=score_threshold, - top_n=len(documents), - query_type=query_type, - ) - ) else: all_documents.extend(documents) except Exception as e: @@ -434,7 +444,8 @@ class RetrievalService: if not dataset: raise ValueError("dataset not found") - vector_processor = Vector(dataset=dataset) + with Session(db.engine) as session: + vector_processor = Vector(dataset=dataset, session=session) documents = vector_processor.search_by_full_text( cls.escape_query_for_search(query), top_k=top_k, document_ids_filter=document_ids_filter @@ -446,17 +457,23 @@ class RetrievalService: and reranking_model["reranking_provider_name"] and retrieval_method == RetrievalMethod.FULL_TEXT_SEARCH ): - data_post_processor = DataPostProcessor( - str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False - ) - all_documents.extend( - data_post_processor.invoke( - query=query, - documents=documents, - score_threshold=score_threshold, - top_n=len(documents), + with Session(db.engine) as rerank_session: + data_post_processor = DataPostProcessor( + str(dataset.tenant_id), + str(RerankMode.RERANKING_MODEL), + reranking_model, + None, + False, + session=rerank_session, + ) + all_documents.extend( + data_post_processor.invoke( + query=query, + documents=documents, + score_threshold=score_threshold, + top_n=len(documents), + ) ) - ) else: all_documents.extend(documents) except Exception as e: @@ -468,7 +485,7 @@ class RetrievalService: return query.replace('"', '\\"') @classmethod - def format_retrieval_documents(cls, documents: list[Document]) -> list[RetrievalSegments]: + def format_retrieval_documents(cls, session: Session, documents: list[Document]) -> list[RetrievalSegments]: """Format retrieval documents with optimized batch processing""" if not documents: return [] @@ -482,7 +499,7 @@ class RetrievalService: # Batch query dataset documents dataset_documents = { doc.id: doc - for doc in db.session.scalars( + for doc in session.scalars( select(DatasetDocument) .where(DatasetDocument.id.in_(document_ids)) .options(load_only(DatasetDocument.id, DatasetDocument.doc_form, DatasetDocument.dataset_id)) @@ -558,84 +575,83 @@ class RetrievalService: doc_segment_map: dict[str, list[str]] = {} segment_summary_map: dict[str, str] = {} # Map segment_id to summary content - with session_factory.create_session() as session: - attachments = cls.get_segment_attachment_infos(image_doc_ids, session) + attachments = cls.get_segment_attachment_infos(image_doc_ids, session) - for attachment in attachments: - segment_ids.append(attachment["segment_id"]) - if attachment["segment_id"] in attachment_map: - attachment_map[attachment["segment_id"]].append(attachment["attachment_info"]) - else: - attachment_map[attachment["segment_id"]] = [attachment["attachment_info"]] - if attachment["segment_id"] in doc_segment_map: - doc_segment_map[attachment["segment_id"]].append(attachment["attachment_id"]) - else: - doc_segment_map[attachment["segment_id"]] = [attachment["attachment_id"]] + for attachment in attachments: + segment_ids.append(attachment["segment_id"]) + if attachment["segment_id"] in attachment_map: + attachment_map[attachment["segment_id"]].append(attachment["attachment_info"]) + else: + attachment_map[attachment["segment_id"]] = [attachment["attachment_info"]] + if attachment["segment_id"] in doc_segment_map: + doc_segment_map[attachment["segment_id"]].append(attachment["attachment_id"]) + else: + doc_segment_map[attachment["segment_id"]] = [attachment["attachment_id"]] - child_chunk_stmt = select(ChildChunk).where(ChildChunk.index_node_id.in_(child_index_node_ids)) - child_index_nodes = session.execute(child_chunk_stmt).scalars().all() + child_chunk_stmt = select(ChildChunk).where(ChildChunk.index_node_id.in_(child_index_node_ids)) + child_index_nodes = session.execute(child_chunk_stmt).scalars().all() - for i in child_index_nodes: - assert i.index_node_id - segment_ids.append(i.segment_id) - if i.segment_id in child_chunk_map: - child_chunk_map[i.segment_id].append(i) - else: - child_chunk_map[i.segment_id] = [i] - if i.segment_id in doc_segment_map: - doc_segment_map[i.segment_id].append(i.index_node_id) - else: - doc_segment_map[i.segment_id] = [i.index_node_id] + for i in child_index_nodes: + assert i.index_node_id + segment_ids.append(i.segment_id) + if i.segment_id in child_chunk_map: + child_chunk_map[i.segment_id].append(i) + else: + child_chunk_map[i.segment_id] = [i] + if i.segment_id in doc_segment_map: + doc_segment_map[i.segment_id].append(i.index_node_id) + else: + doc_segment_map[i.segment_id] = [i.index_node_id] - if index_node_ids: - document_segment_stmt = select(DocumentSegment).where( - DocumentSegment.enabled == True, - DocumentSegment.status == "completed", - DocumentSegment.index_node_id.in_(index_node_ids), + if index_node_ids: + document_segment_stmt = select(DocumentSegment).where( + DocumentSegment.enabled == True, + DocumentSegment.status == "completed", + DocumentSegment.index_node_id.in_(index_node_ids), + ) + index_node_segments = session.execute(document_segment_stmt).scalars().all() + for index_node_segment in index_node_segments: + assert index_node_segment.index_node_id + doc_segment_map[index_node_segment.id] = [index_node_segment.index_node_id] + + if segment_ids: + document_segment_stmt = select(DocumentSegment).where( + DocumentSegment.enabled == True, + DocumentSegment.status == "completed", + DocumentSegment.id.in_(segment_ids), + ) + segments = session.execute(document_segment_stmt).scalars().all() # type: ignore + + if index_node_segments: + segments.extend(index_node_segments) + + # Handle summary documents: query segments by original_chunk_id + if summary_segment_ids: + summary_segment_ids_list = list(summary_segment_ids) + summary_segment_stmt = select(DocumentSegment).where( + DocumentSegment.enabled == True, + DocumentSegment.status == "completed", + DocumentSegment.id.in_(summary_segment_ids_list), + ) + summary_segments = session.execute(summary_segment_stmt).scalars().all() # type: ignore + segments.extend(summary_segments) + # Add summary segment IDs to segment_ids for summary query + for seg in summary_segments: + if seg.id not in segment_ids: + segment_ids.append(seg.id) + + # Batch query summaries for segments retrieved via summary (only enabled summaries) + if summary_segment_ids: + summaries = session.scalars( + select(DocumentSegmentSummary).where( + DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)), + DocumentSegmentSummary.status == "completed", + DocumentSegmentSummary.enabled.is_(True), # Only retrieve enabled summaries ) - index_node_segments = session.execute(document_segment_stmt).scalars().all() - for index_node_segment in index_node_segments: - assert index_node_segment.index_node_id - doc_segment_map[index_node_segment.id] = [index_node_segment.index_node_id] - - if segment_ids: - document_segment_stmt = select(DocumentSegment).where( - DocumentSegment.enabled == True, - DocumentSegment.status == "completed", - DocumentSegment.id.in_(segment_ids), - ) - segments = session.execute(document_segment_stmt).scalars().all() # type: ignore - - if index_node_segments: - segments.extend(index_node_segments) - - # Handle summary documents: query segments by original_chunk_id - if summary_segment_ids: - summary_segment_ids_list = list(summary_segment_ids) - summary_segment_stmt = select(DocumentSegment).where( - DocumentSegment.enabled == True, - DocumentSegment.status == "completed", - DocumentSegment.id.in_(summary_segment_ids_list), - ) - summary_segments = session.execute(summary_segment_stmt).scalars().all() # type: ignore - segments.extend(summary_segments) - # Add summary segment IDs to segment_ids for summary query - for seg in summary_segments: - if seg.id not in segment_ids: - segment_ids.append(seg.id) - - # Batch query summaries for segments retrieved via summary (only enabled summaries) - if summary_segment_ids: - summaries = session.scalars( - select(DocumentSegmentSummary).where( - DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)), - DocumentSegmentSummary.status == "completed", - DocumentSegmentSummary.enabled.is_(True), # Only retrieve enabled summaries - ) - ).all() - for summary in summaries: - if summary.summary_content: - segment_summary_map[summary.chunk_id] = summary.summary_content + ).all() + for summary in summaries: + if summary.summary_content: + segment_summary_map[summary.chunk_id] = summary.summary_content include_segment_ids = set() segment_child_map: dict[str, SegmentChildMapDetail] = {} @@ -774,7 +790,7 @@ class RetrievalService: return sorted(result, key=lambda x: x.score if x.score is not None else 0.0, reverse=True) except Exception as e: - db.session.rollback() + session.rollback() raise e @trace_span() @@ -882,9 +898,6 @@ class RetrievalService: if attachment_id and reranking_mode == RerankMode.WEIGHTED_SCORE: all_documents.extend(all_documents_item) all_documents_item = self._deduplicate_documents(all_documents_item) - data_post_processor = DataPostProcessor( - str(dataset.tenant_id), reranking_mode, reranking_model, weights, False - ) if query: rerank_query = query @@ -894,17 +907,26 @@ class RetrievalService: query_type = QueryType.IMAGE_QUERY else: return - all_documents_item = data_post_processor.invoke( - query=rerank_query, - documents=all_documents_item, - score_threshold=score_threshold, - top_n=top_k, - query_type=query_type, - ) - if not data_post_processor.rerank_runner and score_threshold: - all_documents_item = self._filter_documents_by_vector_score_threshold( - all_documents_item, score_threshold + with Session(db.engine) as rerank_session: + data_post_processor = DataPostProcessor( + str(dataset.tenant_id), + reranking_mode, + reranking_model, + weights, + False, + session=rerank_session, ) + all_documents_item = data_post_processor.invoke( + query=rerank_query, + documents=all_documents_item, + score_threshold=score_threshold, + top_n=top_k, + query_type=query_type, + ) + if not data_post_processor.rerank_runner and score_threshold: + all_documents_item = self._filter_documents_by_vector_score_threshold( + all_documents_item, score_threshold + ) all_documents.extend(all_documents_item) diff --git a/api/core/rag/datasource/vdb/vector_factory.py b/api/core/rag/datasource/vdb/vector_factory.py index 4d65951d9a9..1bd0d8cafbb 100644 --- a/api/core/rag/datasource/vdb/vector_factory.py +++ b/api/core/rag/datasource/vdb/vector_factory.py @@ -5,6 +5,7 @@ from abc import ABC, abstractmethod from typing import Any, override from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.model_manager import ModelManager @@ -15,7 +16,6 @@ from core.rag.embedding.cached_embedding import CacheEmbedding from core.rag.embedding.embedding_base import Embeddings from core.rag.index_processor.constant.doc_type import DocType from core.rag.models.document import Document -from extensions.ext_database import db from extensions.ext_redis import redis_client from extensions.ext_storage import storage from extensions.otel import trace_span @@ -99,7 +99,7 @@ class _LazyEmbeddings(Embeddings): class Vector: - def __init__(self, dataset: Dataset, attributes: list | None = None): + def __init__(self, dataset: Dataset, attributes: list | None = None, *, session: Session): if attributes is None: # `is_summary` and `original_chunk_id` are stored on summary vectors # by `SummaryIndexService` and read back by `RetrievalService` to @@ -120,14 +120,15 @@ class Vector: ] self._dataset = dataset # Use a lazy proxy so cleanup paths (delete_by_ids / delete / text_exists) - # never transitively trigger billing API calls during ``Vector(dataset)`` + # never transitively trigger billing API calls during ``Vector(dataset, session=...)`` # construction. The real embedding model is materialized only when an # ``embed_*`` method is actually invoked (i.e. create / search paths). self._embeddings: Embeddings = _LazyEmbeddings(dataset) self._attributes = attributes - self._vector_processor = self._init_vector() + self._session = session + self._vector_processor = self._init_vector(session=session) - def _init_vector(self) -> BaseVector: + def _init_vector(self, *, session: Session) -> BaseVector: vector_type = dify_config.VECTOR_STORE if self._dataset.index_struct_dict: @@ -137,7 +138,7 @@ class Vector: stmt = select(Whitelist).where( Whitelist.tenant_id == self._dataset.tenant_id, Whitelist.category == "vector_db" ) - whitelist = db.session.scalars(stmt).one_or_none() + whitelist = session.scalars(stmt).one_or_none() if whitelist: vector_type = VectorType.TIDB_ON_QDRANT @@ -194,7 +195,7 @@ class Vector: # Batch query all upload files to avoid N+1 queries attachment_ids = [doc.metadata["doc_id"] for doc in batch] stmt = select(UploadFile).where(UploadFile.id.in_(attachment_ids)) - upload_files = db.session.scalars(stmt).all() + upload_files = self._session.scalars(stmt).all() upload_file_map = {str(f.id): f for f in upload_files} file_base64_list = [] @@ -252,7 +253,7 @@ class Vector: return self._vector_processor.search_by_vector(query_vector, **kwargs) def search_by_file(self, file_id: str, **kwargs: Any) -> list[Document]: - upload_file: UploadFile | None = db.session.get(UploadFile, file_id) + upload_file: UploadFile | None = self._session.get(UploadFile, file_id) if not upload_file: return [] diff --git a/api/core/rag/docstore/dataset_docstore.py b/api/core/rag/docstore/dataset_docstore.py index c7d52d74cb1..9b269e782e7 100644 --- a/api/core/rag/docstore/dataset_docstore.py +++ b/api/core/rag/docstore/dataset_docstore.py @@ -4,11 +4,11 @@ from collections.abc import Sequence from typing import Any from sqlalchemy import delete, func, select +from sqlalchemy.orm import Session from core.model_manager import ModelManager from core.rag.index_processor.constant.index_type import IndexTechniqueType 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.dataset import ChildChunk, Dataset, DocumentSegment, SegmentAttachmentBinding from models.enums import SegmentType @@ -45,8 +45,11 @@ class DatasetDocumentStore: @property def docs(self) -> dict[str, Document]: + raise ValueError("session is required; use get_docs(session)") + + def get_docs(self, session: Session) -> dict[str, Document]: stmt = select(DocumentSegment).where(DocumentSegment.dataset_id == self._dataset.id) - document_segments = db.session.scalars(stmt).all() + document_segments = session.scalars(stmt).all() output = {} for document_segment in document_segments: @@ -64,8 +67,14 @@ class DatasetDocumentStore: return output - def add_documents(self, docs: Sequence[Document], allow_update: bool = True, save_child: bool = False): - max_position = db.session.scalar( + def add_documents( + self, + docs: Sequence[Document], + session: Session, + allow_update: bool = True, + save_child: bool = False, + ): + max_position = session.scalar( select(func.max(DocumentSegment.position)).where(DocumentSegment.document_id == self._document_id) ) @@ -94,7 +103,7 @@ class DatasetDocumentStore: if doc.metadata is None: raise ValueError("doc.metadata must be a dict") - segment_document = self.get_document_segment(doc_id=doc.metadata["doc_id"]) + segment_document = self.get_document_segment(doc_id=doc.metadata["doc_id"], session=session) # NOTE: doc could already exist in the store, but we overwrite it if not allow_update and segment_document: @@ -121,10 +130,10 @@ class DatasetDocumentStore: if doc.metadata.get("answer"): segment_document.answer = doc.metadata.pop("answer", "") - db.session.add(segment_document) - db.session.flush() + session.add(segment_document) + session.flush() self.add_multimodel_documents_binding( - segment_id=segment_document.id, multimodel_documents=doc.attachments + segment_id=segment_document.id, multimodel_documents=doc.attachments, session=session ) if save_child: if doc.children: @@ -143,7 +152,7 @@ class DatasetDocumentStore: type=SegmentType.AUTOMATIC, created_by=self._user_id, ) - db.session.add(child_segment) + session.add(child_segment) else: segment_document.content = doc.page_content if doc.metadata.get("answer"): @@ -152,11 +161,11 @@ class DatasetDocumentStore: segment_document.word_count = len(doc.page_content) segment_document.tokens = tokens self.add_multimodel_documents_binding( - segment_id=segment_document.id, multimodel_documents=doc.attachments + segment_id=segment_document.id, multimodel_documents=doc.attachments, session=session ) if save_child and doc.children: # delete the existing child chunks - db.session.execute( + session.execute( delete(ChildChunk).where( ChildChunk.tenant_id == self._dataset.tenant_id, ChildChunk.dataset_id == self._dataset.id, @@ -180,17 +189,17 @@ class DatasetDocumentStore: type=SegmentType.AUTOMATIC, created_by=self._user_id, ) - db.session.add(child_segment) + session.add(child_segment) - db.session.commit() + session.flush() - def document_exists(self, doc_id: str) -> bool: + def document_exists(self, doc_id: str, session: Session) -> bool: """Check if document exists.""" - result = self.get_document_segment(doc_id) + result = self.get_document_segment(doc_id, session=session) return result is not None - def get_document(self, doc_id: str, raise_error: bool = True) -> Document | None: - document_segment = self.get_document_segment(doc_id) + def get_document(self, doc_id: str, session: Session, raise_error: bool = True) -> Document | None: + document_segment = self.get_document_segment(doc_id, session=session) if document_segment is None: if raise_error: @@ -208,8 +217,8 @@ class DatasetDocumentStore: }, ) - def delete_document(self, doc_id: str, raise_error: bool = True): - document_segment = self.get_document_segment(doc_id) + def delete_document(self, doc_id: str, session: Session, raise_error: bool = True): + document_segment = self.get_document_segment(doc_id, session=session) if document_segment is None: if raise_error: @@ -217,37 +226,39 @@ class DatasetDocumentStore: else: return None - db.session.delete(document_segment) - db.session.commit() + session.delete(document_segment) + session.flush() - def set_document_hash(self, doc_id: str, doc_hash: str): + def set_document_hash(self, doc_id: str, doc_hash: str, session: Session): """Set the hash for a given doc_id.""" - document_segment = self.get_document_segment(doc_id) + document_segment = self.get_document_segment(doc_id, session=session) if document_segment is None: return None document_segment.index_node_hash = doc_hash - db.session.commit() + session.flush() - def get_document_hash(self, doc_id: str) -> str | None: + def get_document_hash(self, doc_id: str, session: Session) -> str | None: """Get the stored hash for a document, if it exists.""" - document_segment = self.get_document_segment(doc_id) + document_segment = self.get_document_segment(doc_id, session=session) if document_segment is None: return None data: str | None = document_segment.index_node_hash return data - def get_document_segment(self, doc_id: str) -> DocumentSegment | None: + def get_document_segment(self, doc_id: str, session: Session) -> DocumentSegment | None: stmt = select(DocumentSegment).where( DocumentSegment.dataset_id == self._dataset.id, DocumentSegment.index_node_id == doc_id ) - document_segment = db.session.scalar(stmt) + document_segment = session.scalar(stmt) return document_segment - def add_multimodel_documents_binding(self, segment_id: str, multimodel_documents: list[AttachmentDocument] | None): + def add_multimodel_documents_binding( + self, segment_id: str, multimodel_documents: list[AttachmentDocument] | None, session: Session + ): if multimodel_documents and self._document_id is not None: for multimodel_document in multimodel_documents: binding = SegmentAttachmentBinding( @@ -257,4 +268,4 @@ class DatasetDocumentStore: segment_id=segment_id, attachment_id=multimodel_document.metadata["doc_id"], ) - db.session.add(binding) + session.add(binding) diff --git a/api/core/rag/extractor/extract_processor.py b/api/core/rag/extractor/extract_processor.py index 36d879427a5..7eed135bbf0 100644 --- a/api/core/rag/extractor/extract_processor.py +++ b/api/core/rag/extractor/extract_processor.py @@ -4,6 +4,8 @@ from pathlib import Path from typing import Literal, overload from urllib.parse import unquote +from sqlalchemy.orm import Session + from configs import dify_config from core.file import remote_fetcher from core.rag.extractor.csv_extractor import CSVExtractor @@ -111,7 +113,12 @@ class ExtractProcessor: @classmethod def extract( - cls, extract_setting: ExtractSetting, is_automatic: bool = False, file_path: str | None = None + cls, + extract_setting: ExtractSetting, + is_automatic: bool = False, + file_path: str | None = None, + *, + session: Session | None = None, ) -> list[Document]: if extract_setting.datasource_type == DatasourceType.FILE: upload_file = extract_setting.upload_file @@ -141,7 +148,9 @@ class ExtractProcessor: ) elif file_extension == ".pdf": assert upload_file is not None - extractor = PdfExtractor(file_path, upload_file.tenant_id, upload_file.created_by) + extractor = PdfExtractor( + file_path, upload_file.tenant_id, upload_file.created_by, session=session + ) elif file_extension in {".md", ".markdown", ".mdx"}: extractor = ( UnstructuredMarkdownExtractor(file_path, unstructured_api_url, unstructured_api_key) @@ -152,7 +161,9 @@ class ExtractProcessor: extractor = HtmlExtractor(file_path) elif file_extension == ".docx": assert upload_file is not None - extractor = WordExtractor(file_path, upload_file.tenant_id, upload_file.created_by) + extractor = WordExtractor( + file_path, upload_file.tenant_id, upload_file.created_by, session=session + ) elif file_extension == ".doc": extractor = UnstructuredWordExtractor(file_path, unstructured_api_url, unstructured_api_key) elif file_extension == ".csv": @@ -184,14 +195,18 @@ class ExtractProcessor: ) elif file_extension == ".pdf": assert upload_file is not None - extractor = PdfExtractor(file_path, upload_file.tenant_id, upload_file.created_by) + extractor = PdfExtractor( + file_path, upload_file.tenant_id, upload_file.created_by, session=session + ) elif file_extension in {".md", ".markdown", ".mdx"}: extractor = MarkdownExtractor(file_path, autodetect_encoding=True) elif file_extension in {".htm", ".html"}: extractor = HtmlExtractor(file_path) elif file_extension == ".docx": assert upload_file is not None - extractor = WordExtractor(file_path, upload_file.tenant_id, upload_file.created_by) + extractor = WordExtractor( + file_path, upload_file.tenant_id, upload_file.created_by, session=session + ) elif file_extension == ".csv": extractor = CSVExtractor(file_path, autodetect_encoding=True) elif file_extension == ".epub": diff --git a/api/core/rag/extractor/pdf_extractor.py b/api/core/rag/extractor/pdf_extractor.py index a79854a7353..587d1220696 100644 --- a/api/core/rag/extractor/pdf_extractor.py +++ b/api/core/rag/extractor/pdf_extractor.py @@ -9,6 +9,7 @@ from typing import override import pypdfium2 import pypdfium2.raw as pdfium_c +from sqlalchemy.orm import Session from configs import dify_config from core.rag.extractor.blob.blob import Blob @@ -33,6 +34,7 @@ class PdfExtractor(BaseExtractor): tenant_id: Workspace ID. user_id: ID of the user performing the extraction. file_cache_key: Optional cache key for the extracted text. + session: Session used to persist extracted images. """ # Magic bytes for image format detection: (magic_bytes, extension, mime_type) @@ -48,13 +50,23 @@ class PdfExtractor(BaseExtractor): (b"MM\x00+", "tiff", "image/tiff"), ) MAX_MAGIC_LEN = max(len(m) for m, _, _ in IMAGE_FORMATS) + _session: Session | None - def __init__(self, file_path: str, tenant_id: str, user_id: str, file_cache_key: str | None = None): + def __init__( + self, + file_path: str, + tenant_id: str, + user_id: str, + file_cache_key: str | None = None, + *, + session: Session | None = None, + ): """Initialize PdfExtractor.""" self._file_path = file_path self._tenant_id = tenant_id self._user_id = user_id self._file_cache_key = file_cache_key + self._session = session @override def extract(self) -> list[Document]: @@ -174,6 +186,8 @@ class PdfExtractor(BaseExtractor): except Exception as e: logger.warning("Failed to get objects from PDF page: %s", e) if upload_files: - db.session.add_all(upload_files) - db.session.commit() + session = self._session or db.session + session.add_all(upload_files) + if self._session is None: + session.commit() return "\n".join(image_content) diff --git a/api/core/rag/extractor/word_extractor.py b/api/core/rag/extractor/word_extractor.py index c2edc7c4a7b..e57e52f8975 100644 --- a/api/core/rag/extractor/word_extractor.py +++ b/api/core/rag/extractor/word_extractor.py @@ -18,6 +18,7 @@ from docx.oxml.ns import qn from docx.table import Table from docx.text.paragraph import Paragraph from docx.text.run import Run +from sqlalchemy.orm import Session from configs import dify_config from core.file import remote_fetcher @@ -38,16 +39,19 @@ class WordExtractor(BaseExtractor): Args: file_path: Path to the file to load. + session: Session used to persist extracted images. """ _closed: bool + _session: Session | None - def __init__(self, file_path: str, tenant_id: str, user_id: str): + def __init__(self, file_path: str, tenant_id: str, user_id: str, *, session: Session | None = None): """Initialize with file path.""" self._closed = False self.file_path = file_path self.tenant_id = tenant_id self.user_id = user_id + self._session = session if "~" in self.file_path: self.file_path = os.path.expanduser(self.file_path) @@ -112,8 +116,10 @@ class WordExtractor(BaseExtractor): return bool(parsed.netloc) and bool(parsed.scheme) def _extract_images_from_docx(self, doc): + session = self._session or db.session image_count = 0 image_map = {} + upload_files: list[UploadFile] = [] base_url = dify_config.FILES_URL for r_id, rel in doc.part.rels.items(): @@ -152,7 +158,7 @@ class WordExtractor(BaseExtractor): used_by=self.user_id, used_at=naive_utc_now(), ) - db.session.add(upload_file) + upload_files.append(upload_file) image_map[r_id] = f"![image]({base_url}/files/{upload_file.id}/file-preview)" else: image_ext = rel.target_ref.split(".")[-1] @@ -180,9 +186,12 @@ class WordExtractor(BaseExtractor): used_by=self.user_id, used_at=naive_utc_now(), ) - db.session.add(upload_file) + upload_files.append(upload_file) image_map[rel.target_part] = f"![image]({base_url}/files/{upload_file.id}/file-preview)" - db.session.commit() + if upload_files: + session.add_all(upload_files) + if self._session is None: + session.commit() return image_map def _table_to_markdown(self, table, image_map): diff --git a/api/core/rag/index_processor/index_processor.py b/api/core/rag/index_processor/index_processor.py index 757134e7349..39254bb0ac7 100644 --- a/api/core/rag/index_processor/index_processor.py +++ b/api/core/rag/index_processor/index_processor.py @@ -7,6 +7,7 @@ from typing import Any from flask import current_app from sqlalchemy import delete, func, select, update +from sqlalchemy.orm import Session from core.db.session_factory import session_factory from core.rag.index_processor.constant.index_type import IndexTechniqueType @@ -61,75 +62,77 @@ class IndexProcessor: chunks: Mapping[str, Any], batch: Any, summary_index_setting: SummaryIndexSettingDict | None = None, + *, + session: Session, ) -> IndexingResultDict: - with session_factory.create_session() as session: - document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) - if not document: - raise KnowledgeIndexNodeError(f"Document {document_id} not found.") + document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) + if not document: + raise KnowledgeIndexNodeError(f"Document {document_id} not found.") - dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) - if not dataset: - raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.") + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if not dataset: + raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.") - dataset_name_value = dataset.name - document_name_value = document.name - created_at_value = document.created_at - if summary_index_setting is None: - summary_index_setting = dataset.summary_index_setting - index_node_ids = [] + dataset_name_value = dataset.name + document_name_value = document.name + created_at_value = document.created_at + if summary_index_setting is None: + summary_index_setting = dataset.summary_index_setting + index_node_ids = [] - index_processor = IndexProcessorFactory(dataset.chunk_structure).init_index_processor() - if original_document_id: - segments = session.scalars( - select(DocumentSegment).where(DocumentSegment.document_id == original_document_id) - ).all() - if segments: - index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] + index_processor = IndexProcessorFactory(dataset.chunk_structure).init_index_processor() + if original_document_id: + segments = session.scalars( + select(DocumentSegment).where(DocumentSegment.document_id == original_document_id) + ).all() + if segments: + index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] indexing_start_at = time.perf_counter() + # The metadata reads above must not keep a transaction open across vector I/O. + session.commit() # delete from vector index if index_node_ids: - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) + index_processor.clean( + dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session + ) + session.commit() + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == original_document_id) + session.execute(segment_delete_stmt) + session.commit() - with session_factory.create_session() as session, session.begin(): - if index_node_ids: - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == original_document_id) - session.execute(segment_delete_stmt) - - index_processor.index(dataset, document, chunks) + index_processor.index(dataset, document, chunks, session) + session.commit() indexing_end_at = time.perf_counter() - with session_factory.create_session() as session, session.begin(): - document.indexing_latency = indexing_end_at - indexing_start_at - document.indexing_status = "completed" - document.completed_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None) - document.word_count = ( - session.scalar( - select(func.sum(DocumentSegment.word_count)).where( - DocumentSegment.document_id == document_id, - DocumentSegment.dataset_id == dataset_id, - ) - ) - ) or 0 - # Update need_summary based on dataset's summary_index_setting - if summary_index_setting and summary_index_setting.get("enable") is True: - document.need_summary = True - else: - document.need_summary = False - session.add(document) - # update document segment status - session.execute( - update(DocumentSegment) - .where( + document.indexing_latency = indexing_end_at - indexing_start_at + document.indexing_status = "completed" + document.completed_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None) + document.word_count = ( + session.scalar( + select(func.sum(DocumentSegment.word_count)).where( DocumentSegment.document_id == document_id, DocumentSegment.dataset_id == dataset_id, ) - .values( - status="completed", - enabled=True, - completed_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None), - ) ) + ) or 0 + # Update need_summary based on dataset's summary_index_setting + document.need_summary = bool(summary_index_setting and summary_index_setting.get("enable") is True) + session.add(document) + # update document segment status + session.execute( + update(DocumentSegment) + .where( + DocumentSegment.document_id == document_id, + DocumentSegment.dataset_id == dataset_id, + ) + .values( + status="completed", + enabled=True, + completed_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None), + ) + ) + session.flush() result: IndexingResultDict = { "dataset_id": dataset_id, @@ -149,25 +152,27 @@ class IndexProcessor: document_id: str, chunk_structure: str, summary_index_setting: SummaryIndexSettingDict | None, + *, + session: Session, ) -> Preview: doc_language = None - with session_factory.create_session() as session: - if document_id: - document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) - else: - document = None + if document_id: + document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) + else: + document = None - dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) - if not dataset: - raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.") + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if not dataset: + raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.") - if summary_index_setting is None: - summary_index_setting = dataset.summary_index_setting + if summary_index_setting is None: + summary_index_setting = dataset.summary_index_setting - if document: - doc_language = document.doc_language - indexing_technique = dataset.indexing_technique - tenant_id = dataset.tenant_id + if document: + doc_language = document.doc_language + indexing_technique = dataset.indexing_technique + tenant_id = dataset.tenant_id + session.commit() preview_output = self.format_preview(chunk_structure, chunks) if indexing_technique != IndexTechniqueType.HIGH_QUALITY: @@ -194,23 +199,13 @@ class IndexProcessor: """Generate summary for a single chunk.""" if flask_app: with flask_app.app_context(): - if preview_item.content is not None: - # Set Flask application context in worker thread - summary, _ = ParagraphIndexProcessor.generate_summary( - tenant_id=tenant_id, - text=preview_item.content, - summary_index_setting=summary_index_setting, - document_language=doc_language, - ) - if summary: - preview_item.summary = summary - - else: + with session_factory.create_session() as worker_session: summary, _ = ParagraphIndexProcessor.generate_summary( tenant_id=tenant_id, text=preview_item.content if preview_item.content is not None else "", summary_index_setting=summary_index_setting, document_language=doc_language, + session=worker_session, ) if summary: preview_item.summary = summary diff --git a/api/core/rag/index_processor/index_processor_base.py b/api/core/rag/index_processor/index_processor_base.py index 8da401b226a..7af2c517b84 100644 --- a/api/core/rag/index_processor/index_processor_base.py +++ b/api/core/rag/index_processor/index_processor_base.py @@ -12,6 +12,7 @@ from urllib.parse import unquote, urlparse import httpx from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.entities.knowledge_entities import PreviewDetail @@ -46,11 +47,13 @@ class BaseIndexProcessor(ABC): """Interface for extract files.""" @abstractmethod - def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]: + def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]: raise NotImplementedError @abstractmethod - def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]: + def transform( + self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs + ) -> list[Document]: raise NotImplementedError @abstractmethod @@ -60,6 +63,8 @@ class BaseIndexProcessor(ABC): preview_texts: list[PreviewDetail], summary_index_setting: SummaryIndexSettingDict, doc_language: str | None = None, + *, + session: Session, ) -> list[PreviewDetail]: """ For each segment in preview_texts, generate a summary using LLM and attach it to the segment. @@ -71,6 +76,7 @@ class BaseIndexProcessor(ABC): preview_texts: List of preview details to generate summaries for summary_index_setting: Summary index configuration doc_language: Optional document language to ensure summary is generated in the correct language + session: SQLAlchemy session used for summary image lookups """ raise NotImplementedError @@ -81,16 +87,20 @@ class BaseIndexProcessor(ABC): documents: list[Document], multimodal_documents: list[AttachmentDocument] | None = None, with_keywords: bool = True, + *, + session: Session, **kwargs, ) -> None: raise NotImplementedError @abstractmethod - def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None: + def clean( + self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs + ) -> None: raise NotImplementedError @abstractmethod - def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None: + def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None: raise NotImplementedError @abstractmethod @@ -136,7 +146,9 @@ class BaseIndexProcessor(ABC): return character_splitter - def _get_content_files(self, document: Document, current_user: Account | None = None) -> list[AttachmentDocument]: + def _get_content_files( + self, document: Document, current_user: Account | None = None, *, session: Session + ) -> list[AttachmentDocument]: """ Get the content files from the document. """ @@ -173,7 +185,7 @@ class BaseIndexProcessor(ABC): if match: if current_user: tool_file_id = match.group(1) - upload_file_id = self._download_tool_file(tool_file_id, current_user) + upload_file_id = self._download_tool_file(tool_file_id, current_user, session=session) if upload_file_id: upload_file_id_list.append(upload_file_id) continue @@ -187,7 +199,7 @@ class BaseIndexProcessor(ABC): # Get unique IDs for database query unique_upload_file_ids = list(set(upload_file_id_list)) - upload_files = db.session.scalars(select(UploadFile).where(UploadFile.id.in_(unique_upload_file_ids))).all() + upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(unique_upload_file_ids))).all() # Create a mapping from ID to UploadFile for quick lookup upload_file_map = {upload_file.id: upload_file for upload_file in upload_files} @@ -293,13 +305,13 @@ class BaseIndexProcessor(ABC): logging.warning("Unexpected error downloading image from %s", image_url, exc_info=True) return None - def _download_tool_file(self, tool_file_id: str, current_user: Account) -> str | None: + def _download_tool_file(self, tool_file_id: str, current_user: Account, *, session: Session) -> str | None: """ Download the tool file from the ID. """ from services.file_service import FileService - tool_file = db.session.get(ToolFile, tool_file_id) + tool_file = session.get(ToolFile, tool_file_id) if not tool_file: return None blob = storage.load_once(tool_file.file_key) 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 b31c1bb634b..9f2ab8e8d73 100644 --- a/api/core/rag/index_processor/processor/paragraph_index_processor.py +++ b/api/core/rag/index_processor/processor/paragraph_index_processor.py @@ -5,14 +5,12 @@ import re import uuid from typing import Any, TypedDict, cast, override -from sqlalchemy.orm import Session - -logger = logging.getLogger(__name__) - from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.file_access import DatabaseFileAccessController from core.app.llm import deduct_llm_quota +from core.db.session_factory import session_factory from core.entities.knowledge_entities import PreviewDetail from core.llm_generator.prompts import DEFAULT_GENERATOR_SUMMARY_PROMPT from core.model_manager import ModelInstance @@ -30,7 +28,6 @@ from core.rag.index_processor.index_processor_base import BaseIndexProcessor, Su from core.rag.models.document import AttachmentDocument, Document, MultimodalGeneralStructureChunk from core.tools.utils.text_processing_utils import remove_leading_symbols from core.workflow.file_reference import build_file_reference -from extensions.ext_database import db from factories.file_factory import build_from_mapping from graphon.file import File, FileTransferMethod, FileType, file_manager from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage @@ -50,6 +47,9 @@ from models.dataset import Document as DatasetDocument from services.account_service import AccountService from services.summary_index_service import SummaryIndexService +logger = logging.getLogger(__name__) + + _file_access_controller = DatabaseFileAccessController() @@ -61,18 +61,21 @@ class ParagraphFormatPreviewDict(TypedDict): class ParagraphIndexProcessor(BaseIndexProcessor): @override - def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]: + def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]: text_docs = ExtractProcessor.extract( extract_setting=extract_setting, is_automatic=( kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical" ), + session=session, ) return text_docs @override - def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]: + def transform( + self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs + ) -> list[Document]: process_rule = kwargs.get("process_rule") if not process_rule: raise ValueError("No process rule found.") @@ -109,7 +112,9 @@ class ParagraphIndexProcessor(BaseIndexProcessor): document_node.metadata["doc_id"] = doc_id document_node.metadata["doc_hash"] = hash multimodal_documents = ( - self._get_content_files(document_node, current_user) if document_node.metadata else None + self._get_content_files(document_node, current_user, session=session) + if document_node.metadata + else None ) if multimodal_documents: document_node.attachments = multimodal_documents @@ -128,10 +133,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor): documents: list[Document], multimodal_documents: list[AttachmentDocument] | None = None, with_keywords: bool = True, + *, + session: Session, **kwargs, ) -> None: if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.create(documents) if multimodal_documents and dataset.is_multimodal: vector.create_multimodal(multimodal_documents) @@ -140,12 +147,14 @@ class ParagraphIndexProcessor(BaseIndexProcessor): keywords_list = kwargs.get("keywords_list") keyword = Keyword(dataset) if keywords_list and len(keywords_list) > 0: - keyword.add_texts(documents, keywords_list=keywords_list) + keyword.add_texts(documents, session, keywords_list=keywords_list) else: - keyword.add_texts(documents) + keyword.add_texts(documents, session) @override - def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None: + def clean( + self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs + ) -> None: # Note: Summary indexes are now disabled (not deleted) when segments are disabled. # This method is called for actual deletion scenarios (e.g., when segment is deleted). # For disable operations, disable_summaries_for_segments is called directly in the task. @@ -154,7 +163,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): if delete_summaries: if node_ids: # Find segments by index_node_id - segments = db.session.scalars( + segments = session.scalars( select(DocumentSegment).where( DocumentSegment.dataset_id == dataset.id, DocumentSegment.index_node_id.in_(node_ids), @@ -162,13 +171,13 @@ class ParagraphIndexProcessor(BaseIndexProcessor): ).all() segment_ids = [segment.id for segment in segments] if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) + SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) + SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) if node_ids: vector.delete_by_ids(node_ids) else: @@ -177,12 +186,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor): if with_keywords: keyword = Keyword(dataset) if node_ids: - keyword.delete_by_ids(node_ids) + keyword.delete_by_ids(node_ids, session) else: - keyword.delete() + keyword.delete(session=session) @override - def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None: + def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None: documents: list[Any] = [] all_multimodal_documents: list[Any] = [] if isinstance(chunks, list): @@ -194,7 +203,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): "doc_hash": helper.generate_text_hash(content), } doc = Document(page_content=content, metadata=metadata) - attachments = self._get_content_files(doc) + attachments = self._get_content_files(doc, session=session) if attachments: doc.attachments = attachments all_multimodal_documents.extend(attachments) @@ -226,10 +235,11 @@ class ParagraphIndexProcessor(BaseIndexProcessor): all_multimodal_documents.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session()) + with session_factory.create_session() as account_session: + account = AccountService.load_user(document.created_by, account_session) if not account: raise ValueError("Invalid account") - doc.attachments = self._get_content_files(doc, current_user=account) + doc.attachments = self._get_content_files(doc, current_user=account, session=session) if doc.attachments: all_multimodal_documents.extend(doc.attachments) documents.append(doc) @@ -237,15 +247,16 @@ class ParagraphIndexProcessor(BaseIndexProcessor): # save node to document segment doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id) # add document segments - doc_store.add_documents(docs=documents, save_child=False) + doc_store.add_documents(docs=documents, save_child=False, session=session) + session.commit() if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.create(documents) if all_multimodal_documents and dataset.is_multimodal: vector.create_multimodal(all_multimodal_documents) elif dataset.indexing_technique == IndexTechniqueType.ECONOMY: keyword = Keyword(dataset) - keyword.add_texts(documents) + keyword.add_texts(documents, session) @override def format_preview(self, chunks: Any) -> ParagraphFormatPreviewDict: @@ -269,6 +280,8 @@ class ParagraphIndexProcessor(BaseIndexProcessor): preview_texts: list[PreviewDetail], summary_index_setting: SummaryIndexSettingDict, doc_language: str | None = None, + *, + session: Session, ) -> list[PreviewDetail]: """ For each segment, concurrently call generate_summary to generate a summary @@ -291,15 +304,25 @@ class ParagraphIndexProcessor(BaseIndexProcessor): if flask_app: # Ensure Flask app context in worker thread with flask_app.app_context(): - summary, _ = self.generate_summary( - tenant_id, preview.content, summary_index_setting, document_language=doc_language - ) + with session_factory.create_session() as worker_session: + summary, _ = self.generate_summary( + tenant_id, + preview.content, + summary_index_setting, + document_language=doc_language, + session=worker_session, + ) preview.summary = summary else: # Fallback: try without app context (may fail) - summary, _ = self.generate_summary( - tenant_id, preview.content, summary_index_setting, document_language=doc_language - ) + with session_factory.create_session() as worker_session: + summary, _ = self.generate_summary( + tenant_id, + preview.content, + summary_index_setting, + document_language=doc_language, + session=worker_session, + ) preview.summary = summary # Generate summaries concurrently using ThreadPoolExecutor @@ -354,6 +377,8 @@ class ParagraphIndexProcessor(BaseIndexProcessor): summary_index_setting: SummaryIndexSettingDict | None = None, segment_id: str | None = None, document_language: str | None = None, + *, + session: Session, ) -> tuple[str, LLMUsage]: """ Generate summary for the given text using ModelInstance.invoke_llm and the default or custom summary prompt, @@ -366,6 +391,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor): segment_id: Optional segment ID to fetch attachments from SegmentAttachmentBinding table document_language: Optional document language (e.g., "Chinese", "English") to ensure summary is generated in the correct language + session: SQLAlchemy session used for summary image lookups Returns: Tuple of (summary_content, llm_usage) where llm_usage is LLMUsage object @@ -414,12 +440,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, 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, session) # Build prompt messages prompt_messages = [] 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 aecb4154d6f..fa2826f7a57 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 @@ -6,6 +6,7 @@ import uuid from typing import Any, TypedDict, override from sqlalchemy import delete, select +from sqlalchemy.orm import Session from configs import dify_config from core.db.session_factory import session_factory @@ -21,7 +22,6 @@ from core.rag.index_processor.constant.doc_type import DocType from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.index_processor.index_processor_base import BaseIndexProcessor, SummaryIndexSettingDict from core.rag.models.document import AttachmentDocument, ChildDocument, Document, ParentChildStructureChunk -from extensions.ext_database import db from libs import helper from models import Account from models.dataset import ChildChunk, Dataset, DatasetProcessRule, DocumentSegment @@ -42,18 +42,21 @@ class ParentChildFormatPreviewDict(TypedDict): class ParentChildIndexProcessor(BaseIndexProcessor): @override - def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]: + def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]: text_docs = ExtractProcessor.extract( extract_setting=extract_setting, is_automatic=( kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical" ), + session=session, ) return text_docs @override - def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]: + def transform( + self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs + ) -> list[Document]: process_rule = kwargs.get("process_rule") if not process_rule: raise ValueError("No process rule found.") @@ -95,7 +98,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): page_content = page_content if len(page_content) > 0: document_node.page_content = page_content - multimodel_documents = self._get_content_files(document_node, current_user) + multimodel_documents = self._get_content_files(document_node, current_user, session=session) if multimodel_documents: document_node.attachments = multimodel_documents # parse document to child nodes @@ -108,7 +111,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): elif rules.parent_mode == ParentMode.FULL_DOC: page_content = "\n".join([document.page_content for document in documents]) document = Document(page_content=page_content, metadata=documents[0].metadata) - multimodel_documents = self._get_content_files(document) + multimodel_documents = self._get_content_files(document, session=session) if multimodel_documents: document.attachments = multimodel_documents # parse document to child nodes @@ -135,10 +138,12 @@ class ParentChildIndexProcessor(BaseIndexProcessor): documents: list[Document], multimodal_documents: list[AttachmentDocument] | None = None, with_keywords: bool = True, + *, + session: Session, **kwargs, ) -> None: if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) for document in documents: child_documents = document.children if child_documents: @@ -150,7 +155,9 @@ class ParentChildIndexProcessor(BaseIndexProcessor): vector.create_multimodal(multimodal_documents) @override - def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None: + def clean( + self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs + ) -> None: # node_ids is segment's node_ids # Note: Summary indexes are now disabled (not deleted) when segments are disabled. # This method is called for actual deletion scenarios (e.g., when segment is deleted). @@ -160,24 +167,23 @@ class ParentChildIndexProcessor(BaseIndexProcessor): if delete_summaries: if node_ids: # Find segments by index_node_id - with session_factory.create_session() as session: - segments = session.scalars( - select(DocumentSegment).where( - DocumentSegment.dataset_id == dataset.id, - DocumentSegment.index_node_id.in_(node_ids), - ) - ).all() - segment_ids = [segment.id for segment in segments] - if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) + segments = session.scalars( + select(DocumentSegment).where( + DocumentSegment.dataset_id == dataset.id, + DocumentSegment.index_node_id.in_(node_ids), + ) + ).all() + segment_ids = [segment.id for segment in segments] + if segment_ids: + SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) + SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: delete_child_chunks = kwargs.get("delete_child_chunks") or False precomputed_child_node_ids = kwargs.get("precomputed_child_node_ids") - vector = Vector(dataset) + vector = Vector(dataset, session=session) if node_ids: # Use precomputed child_node_ids if available (to avoid race conditions) @@ -185,7 +191,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): child_node_ids = precomputed_child_node_ids else: # Fallback to original query (may fail if segments are already deleted) - rows = db.session.execute( + rows = session.execute( select(ChildChunk.index_node_id) .join(DocumentSegment, ChildChunk.segment_id == DocumentSegment.id) .where( @@ -202,23 +208,23 @@ class ParentChildIndexProcessor(BaseIndexProcessor): # Delete from database if delete_child_chunks and child_node_ids: - db.session.execute( + session.execute( delete(ChildChunk).where( ChildChunk.dataset_id == dataset.id, ChildChunk.index_node_id.in_(child_node_ids) ) ) - db.session.commit() + session.flush() else: vector.delete() if delete_child_chunks: # Use existing compound index: (tenant_id, dataset_id, ...) - db.session.execute( + session.execute( delete(ChildChunk).where( ChildChunk.tenant_id == dataset.tenant_id, ChildChunk.dataset_id == dataset.id ) ) - db.session.commit() + session.flush() def _split_child_nodes( self, @@ -257,7 +263,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): return child_nodes @override - def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None: + def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None: parent_childs = ParentChildStructureChunk.model_validate(chunks) documents = [] for parent_child in parent_childs.parent_child_chunks: @@ -291,10 +297,11 @@ class ParentChildIndexProcessor(BaseIndexProcessor): attachments.append(file_document) doc.attachments = attachments else: - account = AccountService.load_user(document.created_by, db.session()) + with session_factory.create_session() as account_session: + account = AccountService.load_user(document.created_by, account_session) if not account: raise ValueError("Invalid account") - doc.attachments = self._get_content_files(doc, current_user=account) + doc.attachments = self._get_content_files(doc, current_user=account, session=session) documents.append(doc) if documents: # update document parent mode @@ -308,14 +315,14 @@ class ParentChildIndexProcessor(BaseIndexProcessor): ), created_by=document.created_by, ) - db.session.add(dataset_process_rule) - db.session.flush() + session.add(dataset_process_rule) + session.flush() document.dataset_process_rule_id = dataset_process_rule.id - db.session.commit() # save node to document segment doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id) # add document segments - doc_store.add_documents(docs=documents, save_child=True) + doc_store.add_documents(docs=documents, save_child=True, session=session) + session.commit() if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: all_child_documents = [] all_multimodal_documents = [] @@ -324,7 +331,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor): all_child_documents.extend(doc.children) if doc.attachments: all_multimodal_documents.extend(doc.attachments) - vector = Vector(dataset) + vector = Vector(dataset, session=session) if all_child_documents: vector.create(all_child_documents) if all_multimodal_documents and dataset.is_multimodal: @@ -351,6 +358,8 @@ class ParentChildIndexProcessor(BaseIndexProcessor): preview_texts: list[PreviewDetail], summary_index_setting: SummaryIndexSettingDict, doc_language: str | None = None, + *, + session: Session, ) -> list[PreviewDetail]: """ For each parent chunk in preview_texts, concurrently call generate_summary to generate a summary @@ -377,21 +386,25 @@ class ParentChildIndexProcessor(BaseIndexProcessor): if flask_app: # Ensure Flask app context in worker thread with flask_app.app_context(): + with session_factory.create_session() as worker_session: + summary, _ = ParagraphIndexProcessor.generate_summary( + tenant_id=tenant_id, + text=preview.content, + summary_index_setting=summary_index_setting, + document_language=doc_language, + session=worker_session, + ) + preview.summary = summary + else: + # Fallback: try without app context (may fail) + with session_factory.create_session() as worker_session: summary, _ = ParagraphIndexProcessor.generate_summary( tenant_id=tenant_id, text=preview.content, summary_index_setting=summary_index_setting, document_language=doc_language, + session=worker_session, ) - preview.summary = summary - else: - # Fallback: try without app context (may fail) - summary, _ = ParagraphIndexProcessor.generate_summary( - tenant_id=tenant_id, - text=preview.content, - summary_index_setting=summary_index_setting, - document_language=doc_language, - ) preview.summary = summary # Generate summaries concurrently using ThreadPoolExecutor 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 7b7443a621f..8a70ec63773 100644 --- a/api/core/rag/index_processor/processor/qa_index_processor.py +++ b/api/core/rag/index_processor/processor/qa_index_processor.py @@ -9,9 +9,9 @@ from typing import Any, TypedDict, override import pandas as pd from flask import Flask, current_app from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage -from core.db.session_factory import session_factory from core.entities.knowledge_entities import PreviewDetail from core.llm_generator.llm_generator import LLMGenerator from core.rag.cleaner.clean_processor import CleanProcessor @@ -41,17 +41,20 @@ class QAFormatPreviewDict(TypedDict): class QAIndexProcessor(BaseIndexProcessor): @override - def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]: + def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]: text_docs = ExtractProcessor.extract( extract_setting=extract_setting, is_automatic=( kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical" ), + session=session, ) return text_docs @override - def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]: + def transform( + self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs + ) -> list[Document]: preview = kwargs.get("preview") process_rule = kwargs.get("process_rule") if not process_rule: @@ -145,16 +148,20 @@ class QAIndexProcessor(BaseIndexProcessor): documents: list[Document], multimodal_documents: list[AttachmentDocument] | None = None, with_keywords: bool = True, + *, + session: Session, **kwargs, ) -> None: if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.create(documents) if multimodal_documents and dataset.is_multimodal: vector.create_multimodal(multimodal_documents) @override - def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None: + def clean( + self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs + ) -> None: # Note: Summary indexes are now disabled (not deleted) when segments are disabled. # This method is called for actual deletion scenarios (e.g., when segment is deleted). # For disable operations, disable_summaries_for_segments is called directly in the task. @@ -164,28 +171,27 @@ class QAIndexProcessor(BaseIndexProcessor): if delete_summaries: if node_ids: # Find segments by index_node_id - with session_factory.create_session() as session: - segments = session.scalars( - select(DocumentSegment).where( - DocumentSegment.dataset_id == dataset.id, - DocumentSegment.index_node_id.in_(node_ids), - ) - ).all() - segment_ids = [segment.id for segment in segments] - if segment_ids: - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids) + segments = session.scalars( + select(DocumentSegment).where( + DocumentSegment.dataset_id == dataset.id, + DocumentSegment.index_node_id.in_(node_ids), + ) + ).all() + segment_ids = [segment.id for segment in segments] + if segment_ids: + SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session) else: # Delete all summaries for the dataset - SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None) + SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session) - vector = Vector(dataset) + vector = Vector(dataset, session=session) if node_ids: vector.delete_by_ids(node_ids) else: vector.delete() @override - def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None: + def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None: qa_chunks = QAStructureChunk.model_validate(chunks) documents = [] for qa_chunk in qa_chunks.qa_chunks: @@ -201,9 +207,10 @@ class QAIndexProcessor(BaseIndexProcessor): if documents: # save node to document segment doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id) - doc_store.add_documents(docs=documents, save_child=False) + doc_store.add_documents(docs=documents, save_child=False, session=session) + session.commit() if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.create(documents) else: raise ValueError("Indexing technique must be high quality.") @@ -228,6 +235,8 @@ class QAIndexProcessor(BaseIndexProcessor): preview_texts: list[PreviewDetail], summary_index_setting: SummaryIndexSettingDict, doc_language: str | None = None, + *, + session: Session, ) -> list[PreviewDetail]: """ QA model doesn't generate summaries, so this method returns preview_texts unchanged. diff --git a/api/core/rag/rerank/rerank_model.py b/api/core/rag/rerank/rerank_model.py index 8552e7f65dd..ae7ffaad248 100644 --- a/api/core/rag/rerank/rerank_model.py +++ b/api/core/rag/rerank/rerank_model.py @@ -1,12 +1,13 @@ import base64 from typing import override +from sqlalchemy.orm import Session + from core.model_manager import ModelInstance, ModelManager from core.rag.index_processor.constant.doc_type import DocType from core.rag.index_processor.constant.query_type import QueryType from core.rag.models.document import Document from core.rag.rerank.rerank_base import BaseRerankRunner -from extensions.ext_database import db from extensions.ext_storage import storage from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult @@ -14,8 +15,11 @@ from models.model import UploadFile class RerankModelRunner(BaseRerankRunner): - def __init__(self, rerank_model_instance: ModelInstance): + _session: Session + + def __init__(self, rerank_model_instance: ModelInstance, *, session: Session): self.rerank_model_instance = rerank_model_instance + self._session = session @override def run( @@ -134,8 +138,7 @@ class RerankModelRunner(BaseRerankRunner): and document.metadata["doc_id"] not in doc_ids ): if document.metadata.get("doc_type") == DocType.IMAGE: - # Query file info within db.session context to ensure thread-safe access - upload_file = db.session.get(UploadFile, document.metadata["doc_id"]) + upload_file = self._session.get(UploadFile, document.metadata["doc_id"]) if upload_file: blob = storage.load_once(upload_file.key) document_file_base64 = base64.b64encode(blob).decode() @@ -169,8 +172,7 @@ class RerankModelRunner(BaseRerankRunner): rerank_result, unique_documents = self.fetch_text_rerank(query, documents, score_threshold, top_n) return rerank_result, unique_documents elif query_type == QueryType.IMAGE_QUERY: - # Query file info within db.session context to ensure thread-safe access - upload_file = db.session.get(UploadFile, query) + upload_file = self._session.get(UploadFile, query) if upload_file: blob = storage.load_once(upload_file.key) file_query = base64.b64encode(blob).decode() diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index c1ba9964a3a..c758bee219d 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -274,7 +274,8 @@ class DatasetRetrieval: retrieval_resource_list.append(source) # deal with dify documents if dify_documents: - records = RetrievalService.format_retrieval_documents(dify_documents) + with Session(bind=session.get_bind()) as format_session: + records = RetrievalService.format_retrieval_documents(format_session, dify_documents) dataset_ids = [i.segment.dataset_id for i in records] document_ids = [i.segment.document_id for i in records] @@ -491,7 +492,8 @@ class DatasetRetrieval: retrieval_resource_list.append(source) # deal with dify documents if dify_documents: - records = RetrievalService.format_retrieval_documents(dify_documents) + with Session(bind=session.get_bind()) as format_session: + records = RetrievalService.format_retrieval_documents(format_session, dify_documents) if records: for record in records: segment = record.segment @@ -1225,7 +1227,11 @@ class DatasetRetrieval: continue # pass if dataset is not available - if dataset and dataset.provider != "external" and dataset.available_document_count == 0: + if ( + dataset + and dataset.provider != "external" + and dataset.get_total_available_documents(session=session) == 0 + ): continue available_datasets.append(dataset) @@ -1859,23 +1865,31 @@ class DatasetRetrieval: # Skip second reranking when there is only one dataset if reranking_enable and dataset_count > 1: # do rerank for searched documents - data_post_processor = DataPostProcessor(tenant_id, reranking_mode, reranking_model, weights, False) - if query: - all_documents_item = data_post_processor.invoke( - query=query, - documents=all_documents_item, - score_threshold=score_threshold, - top_n=top_k, - query_type=QueryType.TEXT_QUERY, - ) - if attachment_id: - all_documents_item = data_post_processor.invoke( - documents=all_documents_item, - score_threshold=score_threshold, - top_n=top_k, - query_type=QueryType.IMAGE_QUERY, - query=attachment_id, + with session_factory.create_session() as session: + data_post_processor = DataPostProcessor( + tenant_id, + reranking_mode, + reranking_model, + weights, + False, + session=session, ) + if query: + all_documents_item = data_post_processor.invoke( + query=query, + documents=all_documents_item, + score_threshold=score_threshold, + top_n=top_k, + query_type=QueryType.TEXT_QUERY, + ) + if attachment_id: + all_documents_item = data_post_processor.invoke( + documents=all_documents_item, + score_threshold=score_threshold, + top_n=top_k, + query_type=QueryType.IMAGE_QUERY, + query=attachment_id, + ) else: if index_type == IndexTechniqueType.ECONOMY: if not query: 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 c26523b9be5..da7c10d5ac4 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 @@ -76,11 +76,11 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): model=self.reranking_model_name, ) - rerank_runner = RerankModelRunner(rerank_model_instance) + rerank_runner = RerankModelRunner(rerank_model_instance, session=session) 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, session) document_score_list = {} for item in all_documents: @@ -96,7 +96,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): DocumentSegment.enabled == True, DocumentSegment.index_node_id.in_(index_node_ids), ) - segments = db.session.scalars(document_segment_stmt).all() + segments = session.scalars(document_segment_stmt).all() if segments: index_node_id_to_position = {id: position for position, id in enumerate(index_node_ids)} @@ -112,13 +112,13 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): context_list: list[RetrievalSourceMetadata] = [] resource_number = 1 for segment in sorted_segments: - dataset = db.session.get(Dataset, segment.dataset_id) + dataset = session.get(Dataset, segment.dataset_id) document_stmt = select(Document).where( Document.id == segment.document_id, Document.enabled == True, Document.archived == False, ) - document = db.session.scalar(document_stmt) + document = session.scalar(document_stmt) if dataset and document: source = RetrievalSourceMetadata( position=resource_number, 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 d7e390ca877..8f07310f3c7 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py @@ -12,7 +12,6 @@ from core.rag.models.document import Document as RetrievalDocument from core.rag.retrieval.dataset_retrieval import DatasetRetrieval from core.rag.retrieval.retrieval_methods import RetrievalMethod from core.tools.utils.dataset_retriever.dataset_retriever_base_tool import DatasetRetrieverBaseTool -from extensions.ext_database import db from models.dataset import Dataset from models.dataset import Document as DatasetDocument from services.external_knowledge_service import ExternalDatasetService @@ -60,12 +59,12 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): @override def _run(self, session: Session, query: str) -> str: dataset_stmt = select(Dataset).where(Dataset.tenant_id == self.tenant_id, Dataset.id == self.dataset_id) - dataset = db.session.scalar(dataset_stmt) + dataset = session.scalar(dataset_stmt) 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, session) dataset_retrieval = DatasetRetrieval() metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition( session, @@ -162,14 +161,15 @@ 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, session) document_score_list = {} if dataset.indexing_technique != IndexTechniqueType.ECONOMY: for item in documents: if item.metadata is not None and item.metadata.get("score"): document_score_list[item.metadata["doc_id"]] = item.metadata["score"] document_context_list: list[DocumentContext] = [] - records = RetrievalService.format_retrieval_documents(documents) + with Session(bind=session.get_bind()) as format_session: + records = RetrievalService.format_retrieval_documents(format_session, documents) if records: for record in records: segment = record.segment @@ -195,13 +195,13 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): if self.return_resource: for record in records: segment = record.segment - dataset = db.session.get(Dataset, segment.dataset_id) + dataset = session.get(Dataset, segment.dataset_id) dataset_document_stmt = select(DatasetDocument).where( DatasetDocument.id == segment.document_id, DatasetDocument.enabled == True, DatasetDocument.archived == False, ) - document = db.session.scalar(dataset_document_stmt) + document = session.scalar(dataset_document_stmt) if dataset and document: source = RetrievalSourceMetadata( dataset_id=dataset.id, diff --git a/api/core/tools/workflow_as_tool/tool.py b/api/core/tools/workflow_as_tool/tool.py index 3aded386594..b5a70148b8b 100644 --- a/api/core/tools/workflow_as_tool/tool.py +++ b/api/core/tools/workflow_as_tool/tool.py @@ -281,7 +281,7 @@ class WorkflowTool(Tool): user_stmt = select(Account).where(Account.id == user_id) user = session.scalar(user_stmt) if user: - user.current_tenant = tenant + user.set_current_tenant_with_session(tenant, session=session) session.expunge(user) return user diff --git a/api/core/workflow/nodes/agent_v2/validators.py b/api/core/workflow/nodes/agent_v2/validators.py index 7b915fe02be..83dce82a58d 100644 --- a/api/core/workflow/nodes/agent_v2/validators.py +++ b/api/core/workflow/nodes/agent_v2/validators.py @@ -147,7 +147,7 @@ class WorkflowAgentNodeValidator: ) cls._validate_agent_soul_env(binding=binding, agent_soul=agent_soul) cls._validate_agent_soul_tools(binding=binding, agent_soul=agent_soul) - cls._validate_agent_soul_knowledge(binding=binding, agent_soul=agent_soul) + cls._validate_agent_soul_knowledge(session=session, binding=binding, agent_soul=agent_soul) node_job = WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict) cls.validate_node_job(session=session, binding=binding, node_job=node_job, topology=topology) @@ -370,11 +370,13 @@ class WorkflowAgentNodeValidator: def _validate_agent_soul_knowledge( cls, *, + session: Session, binding: WorkflowAgentNodeBinding, agent_soul: AgentSoulConfig, ) -> None: """Validate knowledge set dataset rows against the publishing tenant.""" missing_ids = list_missing_tenant_knowledge_dataset_ids( + session=session, tenant_id=binding.tenant_id, agent_soul=agent_soul, ) diff --git a/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py b/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py index 86854c01827..e0721a15ee5 100644 --- a/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py +++ b/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py @@ -2,6 +2,9 @@ import logging from collections.abc import Mapping from typing import TYPE_CHECKING, Any, override +from sqlalchemy.orm import Session + +from core.db.session_factory import session_factory from core.rag.index_processor.index_processor import IndexProcessor from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict from core.rag.summary_index.summary_index import SummaryIndex @@ -83,9 +86,15 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]): # Get indexing_technique and summary_index_setting from node_data (workflow graph config) # or fallback to dataset if not available in node_data - outputs = self.index_processor.get_preview_output( - chunks, dataset_id, document_id, node_data.chunk_structure, summary_index_setting - ) + with session_factory.create_session() as session: + outputs = self.index_processor.get_preview_output( + chunks, + dataset_id, + document_id, + node_data.chunk_structure, + summary_index_setting, + session=session, + ) return NodeRunResult( status=WorkflowNodeExecutionStatus.SUCCEEDED, inputs=variables, @@ -97,15 +106,17 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]): if not batch: raise KnowledgeIndexNodeError("Batch is required.") - results = self._invoke_knowledge_index( - dataset_id=dataset_id, - document_id=document_id, - original_document_id=original_document_id_segment.value if original_document_id_segment else "", - is_preview=is_preview, - batch=batch.value, - chunks=chunks, - summary_index_setting=summary_index_setting, - ) + with session_factory.create_session() as session: + results = self._invoke_knowledge_index( + session=session, + dataset_id=dataset_id, + document_id=document_id, + original_document_id=original_document_id_segment.value if original_document_id_segment else "", + is_preview=is_preview, + batch=batch.value, + chunks=chunks, + summary_index_setting=summary_index_setting, + ) return NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED, inputs=variables, outputs=results) except KnowledgeIndexNodeError as e: @@ -134,12 +145,16 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]): batch: Any, chunks: Mapping[str, Any], summary_index_setting: SummaryIndexSettingDict | None = None, + *, + session: Session, ): if not document_id: raise KnowledgeIndexNodeError("document_id is required.") rst = self.index_processor.index_and_clean( - dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting + dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session ) + # Summary generation opens independent sessions and must see the indexed rows. + session.commit() self.summary_index_service.generate_and_vectorize_summary( dataset_id, document_id, is_preview, summary_index_setting ) diff --git a/api/core/workflow/nodes/knowledge_index/protocols.py b/api/core/workflow/nodes/knowledge_index/protocols.py index d04e79c2a8c..7bac1283b87 100644 --- a/api/core/workflow/nodes/knowledge_index/protocols.py +++ b/api/core/workflow/nodes/knowledge_index/protocols.py @@ -2,6 +2,7 @@ from collections.abc import Mapping from typing import Any, Protocol, TypedDict from pydantic import BaseModel, Field +from sqlalchemy.orm import Session class IndexingResultDict(TypedDict): @@ -44,6 +45,8 @@ class IndexProcessorProtocol(Protocol): chunks: Mapping[str, Any], batch: Any, summary_index_setting: dict[str, Any] | None = None, + *, + session: Session, ) -> IndexingResultDict: ... def get_preview_output( @@ -53,6 +56,8 @@ class IndexProcessorProtocol(Protocol): document_id: str, chunk_structure: str, summary_index_setting: dict[str, Any] | None, + *, + session: Session, ) -> Preview: ... diff --git a/api/events/event_handlers/create_document_index.py b/api/events/event_handlers/create_document_index.py index a9a5bc20c35..8bc4240251f 100644 --- a/api/events/event_handlers/create_document_index.py +++ b/api/events/event_handlers/create_document_index.py @@ -19,31 +19,33 @@ logger = logging.getLogger(__name__) def handle(sender, **kwargs): dataset_id = sender document_ids = kwargs.get("document_ids", []) - documents = [] start_at = time.perf_counter() - with session_factory.create_session() as session: - for document_id in document_ids: - logger.info(click.style(f"Start process document: {document_id}", fg="green")) - - document = session.scalar( - select(Document).where( - Document.id == document_id, - Document.dataset_id == dataset_id, - ) - ) - - if not document: - raise NotFound("Document not found") - - document.indexing_status = IndexingStatus.PARSING - document.processing_started_at = naive_utc_now() - documents.append(document) - session.add(document) - session.commit() - try: indexing_runner = IndexingRunner() - indexing_runner.run(documents) + with session_factory.create_session() as session: + documents = [] + for document_id in document_ids: + logger.info(click.style(f"Start process document: {document_id}", fg="green")) + + document = session.scalar( + select(Document).where( + Document.id == document_id, + Document.dataset_id == dataset_id, + ) + ) + + if not document: + raise NotFound("Document not found") + + document.indexing_status = IndexingStatus.PARSING + document.processing_started_at = naive_utc_now() + documents.append(document) + session.add(document) + # Persist the status transition before extraction and indexing begin. + session.commit() + + indexing_runner.run(documents, session) + session.commit() end_at = time.perf_counter() logger.info(click.style(f"Processed dataset: {dataset_id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: diff --git a/api/events/event_handlers/create_installed_app_when_app_created.py b/api/events/event_handlers/create_installed_app_when_app_created.py index 38e102d5fd2..2e49c81e415 100644 --- a/api/events/event_handlers/create_installed_app_when_app_created.py +++ b/api/events/event_handlers/create_installed_app_when_app_created.py @@ -1,17 +1,24 @@ -from core.db.session_factory import session_factory +from sqlalchemy import select +from sqlalchemy.orm import Session + from events.app_event import app_was_created from models.model import InstalledApp @app_was_created.connect -def handle(sender, **kwargs): +def handle(sender, *, session: Session, **_kwargs) -> None: """Create an installed app when an app is created.""" app = sender + installed_app_id = session.scalar( + select(InstalledApp.id).where(InstalledApp.tenant_id == app.tenant_id, InstalledApp.app_id == app.id).limit(1) + ) + if installed_app_id: + return + installed_app = InstalledApp( tenant_id=app.tenant_id, app_id=app.id, app_owner_tenant_id=app.tenant_id, ) - with session_factory.create_session() as session: - session.add(installed_app) - session.commit() + session.add(installed_app) + session.flush() diff --git a/api/events/event_handlers/create_site_record_when_app_created.py b/api/events/event_handlers/create_site_record_when_app_created.py index 5e2a456dce3..69149543fd7 100644 --- a/api/events/event_handlers/create_site_record_when_app_created.py +++ b/api/events/event_handlers/create_site_record_when_app_created.py @@ -1,11 +1,12 @@ -from core.db.session_factory import session_factory +from sqlalchemy.orm import Session + from events.app_event import app_was_created from models.enums import CustomizeTokenStrategy from models.model import Site @app_was_created.connect -def handle(sender, **kwargs): +def handle(sender, *, session: Session, **kwargs) -> None: """Create site record when an app is created.""" app = sender account = kwargs.get("account") @@ -18,10 +19,9 @@ def handle(sender, **kwargs): icon_background=app.icon_background, default_language=account.interface_language, customize_token_strategy=CustomizeTokenStrategy.NOT_ALLOW, - code=Site.generate_code(16), + code=Site.generate_code(16, session=session), created_by=app.created_by, updated_by=app.updated_by, ) - with session_factory.create_session() as session: - session.add(site) - session.commit() + session.add(site) + session.flush() diff --git a/api/events/event_handlers/update_app_dataset_join_when_app_model_config_updated.py b/api/events/event_handlers/update_app_dataset_join_when_app_model_config_updated.py index 4709534ae62..6124bf06992 100644 --- a/api/events/event_handlers/update_app_dataset_join_when_app_model_config_updated.py +++ b/api/events/event_handlers/update_app_dataset_join_when_app_model_config_updated.py @@ -1,15 +1,16 @@ from typing import Any, cast from sqlalchemy import delete, select +from sqlalchemy.orm import Session from events.app_event import app_model_config_was_updated -from extensions.ext_database import db from models.dataset import AppDatasetJoin from models.model import AppModelConfig @app_model_config_was_updated.connect -def handle(sender, **kwargs): +def handle(sender, *, session: Session, **kwargs) -> None: + """Update dataset joins with the caller-provided session.""" app = sender app_model_config = kwargs.get("app_model_config") if app_model_config is None: @@ -17,7 +18,7 @@ def handle(sender, **kwargs): dataset_ids = get_dataset_ids_from_model_config(app_model_config) - app_dataset_joins = db.session.scalars(select(AppDatasetJoin).where(AppDatasetJoin.app_id == app.id)).all() + app_dataset_joins = session.scalars(select(AppDatasetJoin).where(AppDatasetJoin.app_id == app.id)).all() removed_dataset_ids: set[str] = set() if not app_dataset_joins: @@ -31,16 +32,14 @@ def handle(sender, **kwargs): if removed_dataset_ids: for dataset_id in removed_dataset_ids: - db.session.execute( + session.execute( delete(AppDatasetJoin).where(AppDatasetJoin.app_id == app.id, AppDatasetJoin.dataset_id == dataset_id) ) if added_dataset_ids: for dataset_id in added_dataset_ids: app_dataset_join = AppDatasetJoin(app_id=app.id, dataset_id=dataset_id) - db.session.add(app_dataset_join) - - db.session.commit() + session.add(app_dataset_join) def get_dataset_ids_from_model_config(app_model_config: AppModelConfig) -> set[str]: diff --git a/api/extensions/ext_login.py b/api/extensions/ext_login.py index 6515b22eb36..3bc8b7baf17 100644 --- a/api/extensions/ext_login.py +++ b/api/extensions/ext_login.py @@ -5,12 +5,13 @@ import flask_login from flask import Request, Response, request from flask_login import user_loaded_from_request, user_logged_in from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound, Unauthorized from configs import dify_config from constants import HEADER_NAME_APP_CODE +from core.db.session_factory import session_factory from dify_app import DifyApp -from extensions.ext_database import db from libs.passport import PassportService from libs.token import extract_access_token, extract_console_cookie_token, extract_webapp_passport from models import Account, Tenant, TenantAccountJoin @@ -46,6 +47,12 @@ login_manager = DifyLoginManager() @login_manager.request_loader def load_user_from_request(request_from_flask_login: Request) -> LoginUser | None: """Load user based on the request.""" + with session_factory.create_session() as session: + return _load_user_from_request(request_from_flask_login, session) + + +def _load_user_from_request(request_from_flask_login: Request, session: Session) -> LoginUser | None: + """Load user based on the request using an explicit database session.""" del request_from_flask_login # Skip authentication for documentation endpoints @@ -60,7 +67,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non if admin_api_key and admin_api_key == auth_token: workspace_id = request.headers.get("X-WORKSPACE-ID") if workspace_id: - tenant_account_join = db.session.execute( + tenant_account_join = session.execute( select(Tenant, TenantAccountJoin) .where(Tenant.id == workspace_id) .where(TenantAccountJoin.tenant_id == Tenant.id) @@ -68,9 +75,9 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non ).one_or_none() if tenant_account_join: tenant, ta = tenant_account_join - account = db.session.scalar(select(Account).where(Account.id == ta.account_id)) + account = session.scalar(select(Account).where(Account.id == ta.account_id)) if account: - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) return account if request.blueprint in {"console", "inner_api"}: @@ -84,7 +91,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=session) return logged_in_account elif request.blueprint == "openapi": # Account-branch device-flow approval routes (approve / deny / @@ -103,7 +110,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=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 @@ -113,7 +120,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non end_user_id = decoded.get("end_user_id") if not end_user_id: raise Unauthorized("Invalid Authorization token.") - end_user = db.session.scalar(select(EndUser).where(EndUser.id == end_user_id)) + end_user = session.scalar(select(EndUser).where(EndUser.id == end_user_id)) if not end_user: raise NotFound("End user not found.") return end_user @@ -123,7 +130,7 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non decoded = PassportService().verify(auth_token) end_user_id = decoded.get("end_user_id") if end_user_id: - end_user = db.session.scalar(select(EndUser).where(EndUser.id == end_user_id)) + end_user = session.scalar(select(EndUser).where(EndUser.id == end_user_id)) if not end_user: raise NotFound("End user not found.") return end_user @@ -133,10 +140,10 @@ def load_user_from_request(request_from_flask_login: Request) -> LoginUser | Non server_code = request.view_args.get("server_code") if request.view_args else None if not server_code: raise Unauthorized("Invalid Authorization token.") - app_mcp_server = db.session.scalar(select(AppMCPServer).where(AppMCPServer.server_code == server_code).limit(1)) + app_mcp_server = session.scalar(select(AppMCPServer).where(AppMCPServer.server_code == server_code).limit(1)) if not app_mcp_server: raise NotFound("App MCP server not found.") - end_user = db.session.scalar( + end_user = session.scalar( select(EndUser).where(EndUser.session_id == app_mcp_server.id, EndUser.type == EndUserType.MCP).limit(1) ) if not end_user: @@ -152,7 +159,7 @@ def on_user_logged_in(_sender: object, user: LoginUser) -> None: """Called when a user logged in. Note: AccountService.load_logged_in_account will populate user.current_tenant_id - through the load_user method, which calls account.set_tenant_id(). + through the load_user method, which calls account.set_tenant_id_with_session(). """ # tenant_id context variable removed - using current_user.current_tenant_id directly pass diff --git a/api/fields/conversation_fields.py b/api/fields/conversation_fields.py index 72eec3f1c7d..073305d2dd9 100644 --- a/api/fields/conversation_fields.py +++ b/api/fields/conversation_fields.py @@ -1,17 +1,145 @@ from __future__ import annotations +from collections.abc import Sequence from datetime import datetime from typing import Any from pydantic import Field, field_validator, model_validator +from sqlalchemy.orm import Session from fields.base import ResponseModel from graphon.file import File from libs.helper import to_timestamp +from models.account import Account +from models.model import ( + AppModelConfigDict, + MessageAgentThought, + MessageAnnotation, + MessageFeedback, + MessageFileInfo, +) +from models.model import ( + Conversation as ConversationModel, +) +from models.model import ( + Message as MessageModel, +) type JSONValue = Any +class _SessionResponseSource[SourceT]: + def __init__(self, source: SourceT, *, session: Session) -> None: + self._source = source + self._session = session + + def __getattr__(self, name: str) -> object: + return getattr(self._source, name) # noqa: no-new-getattr response adapter delegates model fields + + +class _FeedbackResponseSource(_SessionResponseSource[MessageFeedback]): + @property + def from_account(self) -> Account | None: + return self._source.from_account_with_session(session=self._session) + + +class _AnnotationResponseSource(_SessionResponseSource[MessageAnnotation]): + @property + def account(self) -> Account | None: + return self._source.account_with_session(session=self._session) + + @property + def annotation_create_account(self) -> Account | None: + return self._source.annotation_create_account_with_session(session=self._session) + + +class MessageResponseSource(_SessionResponseSource[MessageModel]): + @property + def inputs(self) -> dict[str, Any]: + return self._source.inputs_with_session(session=self._session) + + @property + def feedbacks(self) -> list[_FeedbackResponseSource]: + return [ + _FeedbackResponseSource(feedback, session=self._session) + for feedback in self._source.feedbacks_with_session(session=self._session) + ] + + @property + def user_feedback(self) -> MessageFeedback | None: + return self._source.user_feedback_with_session(session=self._session) + + @property + def annotation(self) -> _AnnotationResponseSource | None: + annotation = self._source.annotation_with_session(session=self._session) + return _AnnotationResponseSource(annotation, session=self._session) if annotation else None + + @property + def annotation_hit_history(self) -> _AnnotationResponseSource | None: + annotation = self._source.annotation_hit_history_with_session(session=self._session) + return _AnnotationResponseSource(annotation, session=self._session) if annotation else None + + @property + def agent_thoughts(self) -> Sequence[MessageAgentThought]: + return self._source.agent_thoughts_with_session(session=self._session) + + @property + def message_files(self) -> list[MessageFileInfo]: + return self._source.message_files_with_session(session=self._session) + + +class ConversationResponseSource(_SessionResponseSource[ConversationModel]): + @property + def inputs(self) -> dict[str, Any]: + return self._source.inputs_with_session(session=self._session) + + @property + def model_config(self) -> AppModelConfigDict: + return self._source.model_config_with_session(session=self._session) + + @property + def summary_or_query(self) -> str: + return self._source.summary_or_query_with_session(session=self._session) + + @property + def annotated(self) -> bool: + return self._source.annotated_with_session(session=self._session) + + @property + def annotation(self) -> _AnnotationResponseSource | None: + annotation = self._source.annotation_with_session(session=self._session) + return _AnnotationResponseSource(annotation, session=self._session) if annotation else None + + @property + def message_count(self) -> int: + return self._source.message_count_with_session(session=self._session) + + @property + def user_feedback_stats(self) -> dict[str, int]: + return self._source.user_feedback_stats_with_session(session=self._session) + + @property + def admin_feedback_stats(self) -> dict[str, int]: + return self._source.admin_feedback_stats_with_session(session=self._session) + + @property + def status_count(self) -> dict[str, int] | None: + return self._source.status_count_with_session(session=self._session) + + @property + def first_message(self) -> MessageResponseSource | None: + message = self._source.first_message_with_session(session=self._session) + return MessageResponseSource(message, session=self._session) if message else None + + @property + def from_end_user_session_id(self) -> str | None: + return self._source.from_end_user_session_id_with_session(session=self._session) + + @property + def from_account_name(self) -> str | None: + return self._source.from_account_name_with_session(session=self._session) + + class MessageFile(ResponseModel): id: str filename: str diff --git a/api/fields/dataset_fields.py b/api/fields/dataset_fields.py index f97f5b79460..4846aa9689c 100644 --- a/api/fields/dataset_fields.py +++ b/api/fields/dataset_fields.py @@ -1,6 +1,9 @@ +from dataclasses import dataclass from datetime import datetime +from typing import Any from pydantic import Field, field_validator +from sqlalchemy.orm import Session from fields.base import ResponseModel from libs.helper import to_timestamp @@ -170,3 +173,62 @@ class DatasetDetailResponse(ResponseModel): @classmethod def _expand_null_nested(cls, value: object) -> object: return {} if value is None else value + + +@dataclass(frozen=True) +class DatasetDetailResponseSource: + """Expose session-backed dataset fields during response validation.""" + + dataset: Any + session: Session + + @property + def app_count(self) -> int: + return self.dataset.get_app_count(session=self.session) + + @property + def document_count(self) -> int: + return self.dataset.get_document_count(session=self.session) + + @property + def word_count(self) -> int: + return self.dataset.get_word_count(session=self.session) + + @property + def author_name(self) -> str | None: + return self.dataset.get_author_name(session=self.session) + + @property + def tags(self) -> Any: + return self.dataset.get_tags(session=self.session) + + @property + def doc_form(self) -> str | None: + return self.dataset.get_doc_form(session=self.session) + + @property + def external_knowledge_info(self) -> Any: + return self.dataset.get_external_knowledge_info(session=self.session) + + @property + def doc_metadata(self) -> Any: + return self.dataset.get_doc_metadata(session=self.session) + + @property + def is_published(self) -> bool: + return self.dataset.get_is_published(session=self.session) + + @property + def total_documents(self) -> int: + return self.dataset.get_total_documents(session=self.session) + + @property + def total_available_documents(self) -> int: + return self.dataset.get_total_available_documents(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.dataset, name) # noqa: no-new-getattr response adapter delegates model fields + + +def dataset_detail_response_source(dataset: Any, *, session: Session) -> DatasetDetailResponseSource: + return DatasetDetailResponseSource(dataset=dataset, session=session) diff --git a/api/fields/document_fields.py b/api/fields/document_fields.py index a565d19ae62..aa3b4135ec6 100644 --- a/api/fields/document_fields.py +++ b/api/fields/document_fields.py @@ -1,12 +1,16 @@ """Response schemas for dataset document endpoints.""" +from collections.abc import Iterable +from dataclasses import dataclass from datetime import datetime from typing import Any from pydantic import Field, field_validator +from sqlalchemy.orm import Session from fields.base import ResponseModel from libs.helper import to_timestamp +from models.dataset import DocMetadataDetailItem, Document def normalize_enum(value: Any) -> Any: @@ -66,6 +70,39 @@ class DocumentResponse(ResponseModel): return to_timestamp(value) +@dataclass(frozen=True) +class DocumentWithSession: + """Expose session-backed document fields during response validation.""" + + document: Document + session: Session + + @property + def data_source_detail_dict(self) -> dict[str, Any]: + return self.document.get_data_source_detail_dict(session=self.session) + + @property + def hit_count(self) -> int: + return self.document.get_hit_count(session=self.session) + + @property + def doc_metadata_details(self) -> list[DocMetadataDetailItem] | None: + return self.document.get_doc_metadata_details(session=self.session) + + def __getattr__(self, name: str) -> Any: + return getattr(self.document, name) # noqa: no-new-getattr response adapter delegates model fields + + +def document_response(document: Document, *, session: Session) -> DocumentResponse: + return DocumentResponse.model_validate( + DocumentWithSession(document=document, session=session), from_attributes=True + ) + + +def document_responses(documents: Iterable[Document], *, session: Session) -> list[DocumentResponse]: + return [document_response(document, session=session) for document in documents] + + class DocumentListResponse(ResponseModel): data: list[DocumentResponse] has_more: bool diff --git a/api/fields/segment_fields.py b/api/fields/segment_fields.py index b5c99754008..ecee5335e53 100644 --- a/api/fields/segment_fields.py +++ b/api/fields/segment_fields.py @@ -4,6 +4,7 @@ from datetime import datetime from typing import Any from pydantic import field_serializer +from sqlalchemy.orm import Session from fields.base import ResponseModel from libs.helper import to_timestamp @@ -73,23 +74,36 @@ class SegmentResponse(ResponseModel): @dataclass(frozen=True) class SegmentWithSummary: + """Expose session-backed segment relations during response validation.""" + segment: Any summary: str | None + session: Session + + @property + def child_chunks(self) -> Any: + return self.segment.get_child_chunks(session=self.session, include_full_doc=False) + + @property + def attachments(self) -> Any: + return self.segment.get_attachments(session=self.session) def __getattr__(self, name: str) -> Any: return getattr(self.segment, name) -def segment_response_with_summary(segment: Any, summary: str | None) -> SegmentResponse: - response_source = SegmentWithSummary(segment=segment, summary=summary) +def segment_response_with_summary(segment: Any, summary: str | None, *, session: Session) -> SegmentResponse: + response_source = SegmentWithSummary(segment=segment, summary=summary, session=session) return SegmentResponse.model_validate(response_source, from_attributes=True) def segment_responses_with_summaries( segments: Iterable[Any], summaries: Mapping[str, str | None], + *, + session: Session, ) -> list[SegmentResponse]: - return [segment_response_with_summary(segment, summaries.get(segment.id)) for segment in segments] + return [segment_response_with_summary(segment, summaries.get(segment.id), session=session) for segment in segments] class SegmentDetailResponse(ResponseModel): diff --git a/api/libs/pagination.py b/api/libs/pagination.py index c38297efc2b..a30ddc68d70 100644 --- a/api/libs/pagination.py +++ b/api/libs/pagination.py @@ -4,7 +4,7 @@ import math from dataclasses import dataclass from sqlalchemy import Select, func, select -from sqlalchemy.orm import Session, scoped_session +from sqlalchemy.orm import Session @dataclass @@ -38,10 +38,10 @@ class PaginatedResult[T]: def paginate_query( stmt: Select, *, + session: Session, page: int = 1, per_page: int = 20, max_per_page: int | None = None, - session: Session | scoped_session | None = None, ) -> PaginatedResult: """Execute *stmt* as a paginated query using plain SQLAlchemy. @@ -56,13 +56,8 @@ def paginate_query( max_per_page: Hard ceiling for *per_page*; ``None`` means no cap. session: - The session to use. Falls back to ``db.session`` when omitted. + SQLAlchemy session used to execute the count and page queries. """ - if session is None: - from extensions.ext_database import db - - session = db.session - if max_per_page is not None: per_page = min(per_page, max_per_page) diff --git a/api/models/account.py b/api/models/account.py index 350b811bd46..919ee7da820 100644 --- a/api/models/account.py +++ b/api/models/account.py @@ -129,13 +129,8 @@ class Account(UserMixin, TypeBase): return self._current_tenant @current_tenant.setter - def current_tenant(self, tenant: "Tenant"): + def current_tenant(self, tenant: "Tenant") -> None: with Session(db.engine, expire_on_commit=False) as session: - tenant_join_query = select(TenantAccountJoin).where( - TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == self.id - ) - tenant_join = session.scalar(tenant_join_query) - tenant_query = select(Tenant).where(Tenant.id == tenant.id) # TODO: A workaround to reload the tenant with `expire_on_commit=False`, allowing # access to it after the session has been closed. # This prevents `DetachedInstanceError` when accessing the tenant outside @@ -143,7 +138,16 @@ class Account(UserMixin, TypeBase): # (The `tenant` argument is typically loaded by `db.session` without the # `expire_on_commit=False` flag, meaning its lifetime is tied to the web # request's lifecycle.) - tenant_reloaded = session.scalars(tenant_query).one() + self.set_current_tenant_with_session(tenant, session=session) + + def set_current_tenant_with_session(self, tenant: "Tenant", *, session: Session) -> None: + """Set the current tenant and role using the caller-owned session.""" + tenant_join_query = select(TenantAccountJoin).where( + TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == self.id + ) + tenant_join = session.scalar(tenant_join_query) + tenant_query = select(Tenant).where(Tenant.id == tenant.id) + tenant_reloaded = session.scalars(tenant_query).one() if tenant_join: self.role = TenantAccountRole(tenant_join.role) @@ -155,20 +159,24 @@ class Account(UserMixin, TypeBase): def current_tenant_id(self) -> str | None: return self._current_tenant.id if self._current_tenant else None - def set_tenant_id(self, tenant_id: str): + def set_tenant_id(self, tenant_id: str) -> None: + with Session(db.engine, expire_on_commit=False) as session: + self.set_tenant_id_with_session(tenant_id, session=session) + + def set_tenant_id_with_session(self, tenant_id: str, *, session: Session) -> None: + """Set the current tenant by id using the caller-owned session.""" query = ( select(Tenant, TenantAccountJoin) .where(Tenant.id == tenant_id) .where(TenantAccountJoin.tenant_id == Tenant.id) .where(TenantAccountJoin.account_id == self.id) ) - with Session(db.engine, expire_on_commit=False) as session: - tenant_account_join = session.execute(query).first() - if not tenant_account_join: - return - tenant, join = tenant_account_join - self.role = TenantAccountRole(join.role) - self._current_tenant = tenant + tenant_account_join = session.execute(query).first() + if not tenant_account_join: + return + tenant, join = tenant_account_join + self.role = TenantAccountRole(join.role) + self._current_tenant = tenant @property def current_role(self): @@ -270,9 +278,9 @@ class Tenant(TypeBase): DateTime, server_default=func.current_timestamp(), init=False, onupdate=func.current_timestamp() ) - def get_accounts(self) -> list[Account]: + def get_accounts(self, *, session: Session) -> list[Account]: return list( - db.session.scalars( + session.scalars( select(Account).where( Account.id == TenantAccountJoin.account_id, TenantAccountJoin.tenant_id == self.id ) diff --git a/api/models/dataset.py b/api/models/dataset.py index c72f1aadd38..1432d819486 100644 --- a/api/models/dataset.py +++ b/api/models/dataset.py @@ -7,6 +7,7 @@ import os import pickle import re import time +from collections.abc import Sequence from datetime import datetime from json import JSONDecodeError from typing import Any, ClassVar, TypedDict, cast, override @@ -212,13 +213,19 @@ class Dataset(Base): is_multimodal = mapped_column(sa.Boolean, default=False, nullable=False, server_default=sa.text("false")) @property - def total_documents(self): - return db.session.scalar(select(func.count(Document.id)).where(Document.dataset_id == self.id)) or 0 + def total_documents(self) -> int: + return self.get_total_documents(session=db.session()) + + def get_total_documents(self, *, session: Session) -> int: + return self.get_document_count(session=session) @property - def total_available_documents(self): + def total_available_documents(self) -> int: + return self.get_total_available_documents(session=db.session()) + + def get_total_available_documents(self, *, session: Session) -> int: return ( - db.session.scalar( + session.scalar( select(func.count(Document.id)).where( Document.dataset_id == self.id, Document.indexing_status == "completed", @@ -229,9 +236,8 @@ class Dataset(Base): or 0 ) - @property - def dataset_keyword_table(self): - return db.session.scalar(select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.id)) + def get_dataset_keyword_table(self, *, session: Session) -> "DatasetKeywordTable | None": + return session.scalar(select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.id)) @property def index_struct_dict(self): @@ -247,18 +253,27 @@ class Dataset(Base): @property def created_by_account(self): - return db.session.get(Account, self.created_by) + return self.get_created_by_account(session=db.session()) + + def get_created_by_account(self, *, session: Session) -> Account | None: + return session.get(Account, self.created_by) @property def author_name(self) -> str | None: - account = db.session.get(Account, self.created_by) + return self.get_author_name(session=db.session()) + + def get_author_name(self, *, session: Session) -> str | None: + account = self.get_created_by_account(session=session) if account: return account.name return None @property def latest_process_rule(self): - return db.session.scalar( + return self.get_latest_process_rule(session=db.session()) + + def get_latest_process_rule(self, *, session: Session) -> "DatasetProcessRule | None": + return session.scalar( select(DatasetProcessRule) .where(DatasetProcessRule.dataset_id == self.id) .order_by(DatasetProcessRule.created_at.desc()) @@ -266,9 +281,12 @@ class Dataset(Base): ) @property - def app_count(self): + def app_count(self) -> int: + return self.get_app_count(session=db.session()) + + def get_app_count(self, *, session: Session) -> int: return ( - db.session.scalar( + session.scalar( select(func.count(AppDatasetJoin.id)).where( AppDatasetJoin.dataset_id == self.id, App.id == AppDatasetJoin.app_id ) @@ -277,8 +295,11 @@ class Dataset(Base): ) @property - def document_count(self): - return db.session.scalar(select(func.count(Document.id)).where(Document.dataset_id == self.id)) or 0 + def document_count(self) -> int: + return self.get_document_count(session=db.session()) + + def get_document_count(self, *, session: Session) -> int: + return session.scalar(select(func.count(Document.id)).where(Document.dataset_id == self.id)) or 0 @property def available_document_count(self): @@ -308,19 +329,25 @@ class Dataset(Base): ) @property - def word_count(self): - return db.session.scalar( - select(func.coalesce(func.sum(Document.word_count), 0)).where(Document.dataset_id == self.id) + def word_count(self) -> int: + return self.get_word_count(session=db.session()) + + def get_word_count(self, *, session: Session) -> int: + return ( + session.scalar( + select(func.coalesce(func.sum(Document.word_count), 0)).where(Document.dataset_id == self.id) + ) + or 0 ) @property def doc_form(self) -> str | None: + return self.get_doc_form(session=db.session()) + + def get_doc_form(self, *, session: Session) -> str | None: if self.chunk_structure: return self.chunk_structure - document = db.session.scalar(select(Document).where(Document.dataset_id == self.id).limit(1)) - if document: - return document.doc_form - return None + return session.scalar(select(Document.doc_form).where(Document.dataset_id == self.id).limit(1)) @property def retrieval_model_dict(self): @@ -343,8 +370,11 @@ class Dataset(Base): return {**default_retrieval_model, **self.retrieval_model} @property - def tags(self): - tags = db.session.scalars( + def tags(self) -> Sequence[Tag]: + return self.get_tags(session=db.session()) + + def get_tags(self, *, session: Session) -> Sequence[Tag]: + tags = session.scalars( select(Tag) .join(TagBinding, Tag.id == TagBinding.tag_id) .where( @@ -358,10 +388,13 @@ class Dataset(Base): return tags or [] @property - def external_knowledge_info(self): + def external_knowledge_info(self) -> dict[str, Any] | None: + return self.get_external_knowledge_info(session=db.session()) + + def get_external_knowledge_info(self, *, session: Session) -> dict[str, Any] | None: if self.provider != "external": return None - external_knowledge_binding = db.session.scalar( + external_knowledge_binding = session.scalar( select(ExternalKnowledgeBindings).where( ExternalKnowledgeBindings.dataset_id == self.id, ExternalKnowledgeBindings.tenant_id == self.tenant_id, @@ -369,7 +402,7 @@ class Dataset(Base): ) if not external_knowledge_binding: return None - external_knowledge_api = db.session.scalar( + external_knowledge_api = session.scalar( select(ExternalKnowledgeApis).where( ExternalKnowledgeApis.id == external_knowledge_binding.external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == self.tenant_id, @@ -385,18 +418,22 @@ class Dataset(Base): } @property - def is_published(self): + def is_published(self) -> bool: + return self.get_is_published(session=db.session()) + + def get_is_published(self, *, session: Session) -> bool: if self.pipeline_id: - pipeline = db.session.scalar(select(Pipeline).where(Pipeline.id == self.pipeline_id)) + pipeline = session.scalar(select(Pipeline).where(Pipeline.id == self.pipeline_id)) if pipeline: return pipeline.is_published return False @property - def doc_metadata(self): - dataset_metadatas = db.session.scalars( - select(DatasetMetadata).where(DatasetMetadata.dataset_id == self.id) - ).all() + def doc_metadata(self) -> list[dict[str, str]]: + return self.get_doc_metadata(session=db.session()) + + def get_doc_metadata(self, *, session: Session) -> list[dict[str, str]]: + dataset_metadatas = session.scalars(select(DatasetMetadata).where(DatasetMetadata.dataset_id == self.id)).all() doc_metadata = [ { @@ -603,10 +640,13 @@ class Document(Base): @property def data_source_detail_dict(self) -> dict[str, Any]: + return self.get_data_source_detail_dict(session=db.session()) + + def get_data_source_detail_dict(self, *, session: Session) -> dict[str, Any]: if self.data_source_info: if self.data_source_type == "upload_file": data_source_info_dict: dict[str, Any] = json.loads(self.data_source_info) - file_detail = db.session.scalar( + file_detail = session.scalar( select(UploadFile).where(UploadFile.id == data_source_info_dict["upload_file_id"]) ) if file_detail: @@ -634,29 +674,48 @@ class Document(Base): @property def dataset_process_rule(self): + return self.get_dataset_process_rule(session=db.session()) + + def get_dataset_process_rule(self, *, session: Session) -> "DatasetProcessRule | None": if self.dataset_process_rule_id: - return db.session.get(DatasetProcessRule, self.dataset_process_rule_id) + return session.get(DatasetProcessRule, self.dataset_process_rule_id) return None @property - def dataset(self): - return db.session.scalar(select(Dataset).where(Dataset.id == self.dataset_id)) + def dataset(self) -> Dataset | None: + return self.get_dataset(session=db.session()) + + def get_dataset(self, *, session: Session) -> Dataset | None: + """Load the owning dataset with the caller-owned database session.""" + return session.get(Dataset, self.dataset_id) @property def segment_count(self): - return ( - db.session.scalar(select(func.count(DocumentSegment.id)).where(DocumentSegment.document_id == self.id)) or 0 - ) + return self.get_segment_count(session=db.session()) + + def get_segment_count(self, *, session: Session) -> int: + return session.scalar(select(func.count(DocumentSegment.id)).where(DocumentSegment.document_id == self.id)) or 0 @property def hit_count(self): - return db.session.scalar( - select(func.coalesce(func.sum(DocumentSegment.hit_count), 0)).where(DocumentSegment.document_id == self.id) + return self.get_hit_count(session=db.session()) + + def get_hit_count(self, *, session: Session) -> int: + return ( + session.scalar( + select(func.coalesce(func.sum(DocumentSegment.hit_count), 0)).where( + DocumentSegment.document_id == self.id + ) + ) + or 0 ) @property def uploader(self): - user = db.session.scalar(select(Account).where(Account.id == self.created_by)) + return self.get_uploader(session=db.session()) + + def get_uploader(self, *, session: Session) -> str | None: + user = session.scalar(select(Account).where(Account.id == self.created_by)) return user.name if user else None @property @@ -669,8 +728,11 @@ class Document(Base): @property def doc_metadata_details(self) -> list[DocMetadataDetailItem] | None: + return self.get_doc_metadata_details(session=db.session()) + + def get_doc_metadata_details(self, *, session: Session) -> list[DocMetadataDetailItem] | None: if self.doc_metadata: - document_metadatas = db.session.scalars( + document_metadatas = session.scalars( select(DatasetMetadata) .join(DatasetMetadataBinding, DatasetMetadataBinding.metadata_id == DatasetMetadata.id) .where( @@ -687,7 +749,7 @@ class Document(Base): } metadata_list.append(metadata_dict) # deal built-in fields - metadata_list.extend(self.get_built_in_fields()) + metadata_list.extend(self.get_built_in_fields(session=session)) return metadata_list return None @@ -698,7 +760,7 @@ class Document(Base): return self.dataset_process_rule.to_dict() return None - def get_built_in_fields(self) -> list[DocMetadataDetailItem]: + def get_built_in_fields(self, *, session: Session) -> list[DocMetadataDetailItem]: built_in_fields: list[DocMetadataDetailItem] = [] built_in_fields.append( { @@ -713,7 +775,7 @@ class Document(Base): "id": "built-in", "name": BuiltInField.uploader, "type": "string", - "value": self.uploader, + "value": self.get_uploader(session=session), } ) built_in_fields.append( @@ -889,12 +951,20 @@ class DocumentSegment(TypeBase): hit_count: Mapped[int] = mapped_column(sa.Integer, nullable=False, default=0) @property - def dataset(self): - return db.session.scalar(select(Dataset).where(Dataset.id == self.dataset_id)) + def dataset(self) -> Dataset | None: + return self.get_dataset(session=db.session()) + + def get_dataset(self, *, session: Session) -> Dataset | None: + """Load the owning dataset with the caller-owned database session.""" + return session.get(Dataset, self.dataset_id) @property - def document(self): - return db.session.scalar(select(Document).where(Document.id == self.document_id)) + def document(self) -> Document | None: + return self.get_document(session=db.session()) + + def get_document(self, *, session: Session) -> Document | None: + """Load the owning document with the caller-owned database session.""" + return session.get(Document, self.document_id) @property def previous_segment(self): @@ -914,30 +984,20 @@ class DocumentSegment(TypeBase): @property def child_chunks(self): - if not self.document: - return [] - process_rule = self.document.dataset_process_rule - if process_rule and process_rule.mode == "hierarchical": - rules_dict = process_rule.rules_dict - if rules_dict: - rules = Rule.model_validate(rules_dict) - if rules.parent_mode and rules.parent_mode != ParentMode.FULL_DOC: - child_chunks = db.session.scalars( - select(ChildChunk).where(ChildChunk.segment_id == self.id).order_by(ChildChunk.position.asc()) - ).all() - return child_chunks or [] - return [] + return self.get_child_chunks(session=db.session(), include_full_doc=False) - def get_child_chunks(self): - if not self.document: + def get_child_chunks(self, *, session: Session, include_full_doc: bool = True) -> Sequence["ChildChunk"]: + """Load hierarchical child chunks with the caller-owned database session.""" + document = session.get(Document, self.document_id) + if not document: return [] - process_rule = self.document.dataset_process_rule + process_rule = document.get_dataset_process_rule(session=session) if process_rule and process_rule.mode == "hierarchical": rules_dict = process_rule.rules_dict if rules_dict: rules = Rule.model_validate(rules_dict) - if rules.parent_mode: - child_chunks = db.session.scalars( + if rules.parent_mode and (include_full_doc or rules.parent_mode != ParentMode.FULL_DOC): + child_chunks = session.scalars( select(ChildChunk).where(ChildChunk.segment_id == self.id).order_by(ChildChunk.position.asc()) ).all() return child_chunks or [] @@ -1014,8 +1074,12 @@ class DocumentSegment(TypeBase): @property def attachments(self) -> list[AttachmentItem]: + return self.get_attachments(session=db.session()) + + def get_attachments(self, *, session: Session) -> list[AttachmentItem]: + """Load attachment metadata with the caller-owned database session.""" # Use JOIN to fetch attachments in a single query instead of two separate queries - attachments_with_bindings = db.session.execute( + attachments_with_bindings = session.execute( select(SegmentAttachmentBinding, UploadFile) .join(UploadFile, UploadFile.id == SegmentAttachmentBinding.attachment_id) .where( @@ -1167,12 +1231,15 @@ class DatasetQuery(TypeBase): @property def queries(self) -> list[dict[str, Any]]: + return self.get_queries(session=db.session()) + + def get_queries(self, *, session: Session) -> list[dict[str, Any]]: try: queries = json.loads(self.content) if isinstance(queries, list): for query in queries: if query["content_type"] == QueryType.IMAGE_QUERY: - file_info = db.session.scalar(select(UploadFile).where(UploadFile.id == query["content"])) + file_info = session.scalar(select(UploadFile).where(UploadFile.id == query["content"])) if file_info: query["file_info"] = { "id": file_info.id, @@ -1218,8 +1285,7 @@ class DatasetKeywordTable(TypeBase): String(255), nullable=False, server_default=sa.text("'database'"), default="database" ) - @property - def keyword_table_dict(self) -> dict[str, set[Any]] | None: + def get_keyword_table_dict(self, *, session: Session) -> dict[str, set[Any]] | None: class SetDecoder(json.JSONDecoder): def __init__(self, *args: Any, **kwargs: Any) -> None: def object_hook(dct: Any) -> Any: @@ -1237,7 +1303,7 @@ class DatasetKeywordTable(TypeBase): super().__init__(object_hook=object_hook, *args, **kwargs) # get dataset - dataset = db.session.scalar(select(Dataset).where(Dataset.id == self.dataset_id)) + dataset = session.scalar(select(Dataset).where(Dataset.id == self.dataset_id)) if not dataset: return None if self.data_source_type == "database": @@ -1438,11 +1504,14 @@ class ExternalKnowledgeApis(TypeBase): @property def dataset_bindings(self) -> list[DatasetBindingItem]: - external_knowledge_bindings = db.session.scalars( + return self.get_dataset_bindings(session=db.session()) + + def get_dataset_bindings(self, *, session: Session) -> list[DatasetBindingItem]: + external_knowledge_bindings = session.scalars( select(ExternalKnowledgeBindings).where(ExternalKnowledgeBindings.external_knowledge_api_id == self.id) ).all() dataset_ids = [binding.dataset_id for binding in external_knowledge_bindings] - datasets = db.session.scalars(select(Dataset).where(Dataset.id.in_(dataset_ids))).all() + datasets = session.scalars(select(Dataset).where(Dataset.id.in_(dataset_ids))).all() dataset_bindings: list[DatasetBindingItem] = [] for dataset in datasets: dataset_bindings.append({"id": dataset.id, "name": dataset.name}) diff --git a/api/models/model.py b/api/models/model.py index 2d7c6e0c053..fa74e9b5ead 100644 --- a/api/models/model.py +++ b/api/models/model.py @@ -15,7 +15,7 @@ import sqlalchemy as sa from flask import request from flask_login import UserMixin # type: ignore[import-untyped] from sqlalchemy import BigInteger, Float, Index, PrimaryKeyConstraint, String, exists, func, select, text -from sqlalchemy.orm import Mapped, Session, mapped_column, sessionmaker +from sqlalchemy.orm import Mapped, Session, mapped_column from configs import dify_config from constants import DEFAULT_FILE_NUMBER_LIMITS @@ -69,20 +69,25 @@ def _get_file_access_controller(): return DatabaseFileAccessController() -def _resolve_app_tenant_id(app_id: str) -> str: - resolved_tenant_id = db.session.scalar(select(App.tenant_id).where(App.id == app_id)) +def _resolve_app_tenant_id(app_id: str, *, session: Session) -> str: + resolved_tenant_id = session.scalar(select(App.tenant_id).where(App.id == app_id)) if not resolved_tenant_id: raise ValueError(f"Unable to resolve tenant_id for app {app_id}") return resolved_tenant_id -def _build_app_tenant_resolver(app_id: str, owner_tenant_id: str | None = None) -> Callable[[], str]: +def _build_app_tenant_resolver( + app_id: str, + *, + session: Session, + owner_tenant_id: str | None = None, +) -> Callable[[], str]: resolved_tenant_id = owner_tenant_id def resolve_owner_tenant_id() -> str: nonlocal resolved_tenant_id if resolved_tenant_id is None: - resolved_tenant_id = _resolve_app_tenant_id(app_id) + resolved_tenant_id = _resolve_app_tenant_id(app_id, session=session) return resolved_tenant_id return resolve_owner_tenant_id @@ -441,10 +446,13 @@ class App(Base): @property def desc_or_prompt(self) -> str: + return self.desc_or_prompt_with_session(session=db.session()) + + def desc_or_prompt_with_session(self, *, session: Session) -> str: if self.description: return self.description else: - app_model_config = self.app_model_config + app_model_config = self.app_model_config_with_session(session=session) if app_model_config: pre_prompt = app_model_config.pre_prompt or "" # Truncate to 200 characters with ellipsis if using prompt as description @@ -456,26 +464,38 @@ class App(Base): @property def site(self) -> Site | None: - return db.session.scalar(select(Site).where(Site.app_id == self.id)) + return self.site_with_session(session=db.session()) + + def site_with_session(self, *, session: Session) -> Site | None: + return session.scalar(select(Site).where(Site.app_id == self.id)) @property def app_model_config(self) -> AppModelConfig | None: + return self.app_model_config_with_session(session=db.session()) + + def app_model_config_with_session(self, *, session: Session) -> AppModelConfig | None: if self.app_model_config_id: - return db.session.scalar(select(AppModelConfig).where(AppModelConfig.id == self.app_model_config_id)) + return session.scalar(select(AppModelConfig).where(AppModelConfig.id == self.app_model_config_id)) return None @property def workflow(self) -> Workflow | None: + return self.workflow_with_session(session=db.session()) + + def workflow_with_session(self, *, session: Session) -> Workflow | None: if self.workflow_id: from .workflow import Workflow - return db.session.scalar(select(Workflow).where(Workflow.id == self.workflow_id)) + return session.scalar(select(Workflow).where(Workflow.id == self.workflow_id)) return None @property def bound_agent_id(self) -> str | None: + return self.bound_agent_id_with_session(session=db.session()) + + def bound_agent_id_with_session(self, *, session: Session) -> str | None: """For an Agent App (mode=agent), the roster Agent it is backed by. Resolved via ``Agent.app_id`` so the console can open the Composer in @@ -485,7 +505,7 @@ class App(Base): return None from .agent import APP_BACKED_AGENT_SOURCES, Agent, AgentScope, AgentStatus - agent = db.session.scalar( + agent = session.scalar( select(Agent).where( Agent.tenant_id == self.tenant_id, sa.or_( @@ -512,7 +532,11 @@ class App(Base): @property def is_agent(self) -> bool: - app_model_config = self.app_model_config + return self.is_agent_with_session(session=db.session()) + + def is_agent_with_session(self, *, session: Session) -> bool: + """Detect legacy agent mode, committing the compatible app mode through the supplied session.""" + app_model_config = session.get(AppModelConfig, self.app_model_config_id) if self.app_model_config_id else None if not app_model_config: return False if not app_model_config.agent_mode: @@ -521,25 +545,32 @@ class App(Base): if app_model_config.agent_mode_dict.get("enabled", False) and app_model_config.agent_mode_dict.get( "strategy", "" ) in {"function_call", "react"}: + session.execute(sa.update(App).where(App.id == self.id).values(mode=AppMode.AGENT_CHAT)) + session.commit() self.mode = AppMode.AGENT_CHAT - db.session.commit() return True return False @property def mode_compatible_with_agent(self) -> str: - if self.mode == AppMode.CHAT and self.is_agent: + return self.mode_compatible_with_agent_with_session(session=db.session()) + + def mode_compatible_with_agent_with_session(self, *, session: Session) -> str: + if self.mode == AppMode.CHAT and self.is_agent_with_session(session=session): return AppMode.AGENT_CHAT return str(self.mode) @property def deleted_tools(self) -> list[DeletedToolInfo]: + return self.deleted_tools_with_session(session=db.session()) + + def deleted_tools_with_session(self, *, session: Session) -> list[DeletedToolInfo]: from core.plugin.plugin_service import PluginService from core.tools.tool_manager import ToolManager, ToolProviderType # get agent mode tools - app_model_config = self.app_model_config + app_model_config = self.app_model_config_with_session(session=session) if not app_model_config: return [] @@ -582,17 +613,16 @@ class App(Base): if not api_provider_ids and not builtin_provider_ids: return [] - with sessionmaker(db.engine).begin() as session: - if api_provider_ids: - existing_api_providers = [ - str(api_provider.id) - for api_provider in session.execute( - text("SELECT id FROM tool_api_providers WHERE id IN :provider_ids"), - {"provider_ids": tuple(api_provider_ids)}, - ).fetchall() - ] - else: - existing_api_providers = [] + if api_provider_ids: + existing_api_providers = [ + str(api_provider.id) + for api_provider in session.execute( + text("SELECT id FROM tool_api_providers WHERE id IN :provider_ids"), + {"provider_ids": tuple(api_provider_ids)}, + ).fetchall() + ] + else: + existing_api_providers = [] if builtin_provider_ids: # get the non-hardcoded builtin providers @@ -649,7 +679,10 @@ class App(Base): @property def tags(self) -> Sequence[Tag]: - tags = db.session.scalars( + return self.tags_with_session(session=db.session()) + + def tags_with_session(self, *, session: Session) -> Sequence[Tag]: + tags = session.scalars( select(Tag) .join(TagBinding, Tag.id == TagBinding.tag_id) .where( @@ -664,8 +697,11 @@ class App(Base): @property def author_name(self) -> str | None: + return self.author_name_with_session(session=db.session()) + + def author_name_with_session(self, *, session: Session) -> str | None: if self.created_by: - account = db.session.scalar(select(Account).where(Account.id == self.created_by)) + account = session.scalar(select(Account).where(Account.id == self.created_by)) if account: return account.name @@ -744,7 +780,10 @@ class AppModelConfig(TypeBase): @property def app(self) -> App | None: - return db.session.scalar(select(App).where(App.id == self.app_id)) + return self.app_with_session(session=db.session()) + + def app_with_session(self, *, session: Session) -> App | None: + return session.scalar(select(App).where(App.id == self.app_id)) @property def model_dict(self) -> ModelConfig: @@ -981,7 +1020,10 @@ class InstalledApp(TypeBase): @property def app(self) -> App | None: - return db.session.scalar(select(App).where(App.id == self.app_id)) + return self.app_with_session(session=db.session()) + + def app_with_session(self, *, session: Session) -> App | None: + return session.scalar(select(App).where(App.id == self.app_id)) @property def tenant(self) -> Tenant | None: @@ -1009,7 +1051,10 @@ class TrialApp(TypeBase): @property def app(self) -> App | None: - return db.session.scalar(select(App).where(App.id == self.app_id)) + return self.app_with_session(session=db.session()) + + def app_with_session(self, *, session: Session) -> App | None: + return session.scalar(select(App).where(App.id == self.app_id)) class AccountTrialAppRecord(TypeBase): @@ -1158,6 +1203,21 @@ class Conversation(Base): @property def inputs(self) -> dict[str, Any]: + return self.inputs_with_session(session=db.session()) + + @inputs.setter + def inputs(self, value: Mapping[str, Any]): + inputs = dict(value) + for k, v in inputs.items(): + match v: + case File(): + inputs[k] = v.model_dump() + case list(): + if all(isinstance(item, File) for item in v): + inputs[k] = [item.model_dump() for item in v if isinstance(item, File)] + self._inputs = inputs + + def inputs_with_session(self, *, session: Session) -> dict[str, Any]: inputs = self._inputs.copy() # Compatibility bridge: stored input payloads may come from before or after the # graph-layer file refactor. Newer rows may omit `tenant_id`, so keep tenant @@ -1165,6 +1225,7 @@ class Conversation(Base): # into `graphon.file.File`. tenant_resolver = _build_app_tenant_resolver( app_id=self.app_id, + session=session, owner_tenant_id=cast(str | None, getattr(self, "_owner_tenant_id", None)), ) @@ -1199,20 +1260,11 @@ class Conversation(Base): return inputs - @inputs.setter - def inputs(self, value: Mapping[str, Any]): - inputs = dict(value) - for k, v in inputs.items(): - match v: - case File(): - inputs[k] = v.model_dump() - case list(): - if all(isinstance(item, File) for item in v): - inputs[k] = [item.model_dump() for item in v if isinstance(item, File)] - self._inputs = inputs - @property def model_config(self) -> AppModelConfigDict: + return self.model_config_with_session(session=db.session()) + + def model_config_with_session(self, *, session: Session) -> AppModelConfigDict: model_config = cast(AppModelConfigDict, {}) app_model_config: AppModelConfig | None = None @@ -1229,15 +1281,17 @@ class Conversation(Base): app_model_config = AppModelConfig(app_id=self.app_id).from_model_config_dict( cast(AppModelConfigDict, override_model_configs) ) - model_config = app_model_config.to_dict() + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + model_config = app_model_config.to_dict(annotation_reply=annotation_reply) else: model_config["configs"] = override_model_configs # type: ignore[typeddict-unknown-key] else: - app_model_config = db.session.scalar( + app_model_config = session.scalar( select(AppModelConfig).where(AppModelConfig.id == self.app_model_config_id) ) if app_model_config: - model_config = app_model_config.to_dict() + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + model_config = app_model_config.to_dict(annotation_reply=annotation_reply) model_config["model_id"] = self.model_id model_config["provider"] = self.model_provider @@ -1246,39 +1300,62 @@ class Conversation(Base): @property def summary_or_query(self): + return self.summary_or_query_with_session(session=db.session()) + + def summary_or_query_with_session(self, *, session: Session) -> str: if self.summary: return self.summary else: - first_message = self.first_message + first_message = self.first_message_with_session(session=session) if first_message: return first_message.query else: return "" @property - def annotated(self): + def annotated(self) -> bool: + return self.annotated_with_session(session=db.session()) + + def annotated_with_session(self, *, session: Session) -> bool: return ( - db.session.scalar( - select(func.count(MessageAnnotation.id)).where(MessageAnnotation.conversation_id == self.id) - ) + session.scalar(select(func.count(MessageAnnotation.id)).where(MessageAnnotation.conversation_id == self.id)) or 0 ) > 0 @property - def annotation(self): - return db.session.scalar(select(MessageAnnotation).where(MessageAnnotation.conversation_id == self.id).limit(1)) + def annotation(self) -> MessageAnnotation | None: + return self.annotation_with_session(session=db.session()) + + def annotation_with_session(self, *, session: Session) -> MessageAnnotation | None: + return session.scalar(select(MessageAnnotation).where(MessageAnnotation.conversation_id == self.id).limit(1)) @property - def message_count(self): - return db.session.scalar(select(func.count(Message.id)).where(Message.conversation_id == self.id)) or 0 + def message_count(self) -> int: + return self.message_count_with_session(session=db.session()) + + def message_count_with_session(self, *, session: Session) -> int: + return session.scalar(select(func.count(Message.id)).where(Message.conversation_id == self.id)) or 0 @property - def user_feedback_stats(self): + def user_feedback_stats(self) -> dict[str, int]: + return self.user_feedback_stats_with_session(session=db.session()) + + def user_feedback_stats_with_session(self, *, session: Session) -> dict[str, int]: + return self._feedback_stats_with_session(session=session, from_source=FeedbackFromSource.USER) + + @property + def admin_feedback_stats(self) -> dict[str, int]: + return self.admin_feedback_stats_with_session(session=db.session()) + + def admin_feedback_stats_with_session(self, *, session: Session) -> dict[str, int]: + return self._feedback_stats_with_session(session=session, from_source=FeedbackFromSource.ADMIN) + + def _feedback_stats_with_session(self, *, session: Session, from_source: FeedbackFromSource) -> dict[str, int]: like = ( - db.session.scalar( + session.scalar( select(func.count(MessageFeedback.id)).where( MessageFeedback.conversation_id == self.id, - MessageFeedback.from_source == "user", + MessageFeedback.from_source == from_source, MessageFeedback.rating == FeedbackRating.LIKE, ) ) @@ -1286,36 +1363,10 @@ class Conversation(Base): ) dislike = ( - db.session.scalar( + session.scalar( select(func.count(MessageFeedback.id)).where( MessageFeedback.conversation_id == self.id, - MessageFeedback.from_source == "user", - MessageFeedback.rating == FeedbackRating.DISLIKE, - ) - ) - or 0 - ) - - return {"like": like, "dislike": dislike} - - @property - def admin_feedback_stats(self): - like = ( - db.session.scalar( - select(func.count(MessageFeedback.id)).where( - MessageFeedback.conversation_id == self.id, - MessageFeedback.from_source == "admin", - MessageFeedback.rating == FeedbackRating.LIKE, - ) - ) - or 0 - ) - - dislike = ( - db.session.scalar( - select(func.count(MessageFeedback.id)).where( - MessageFeedback.conversation_id == self.id, - MessageFeedback.from_source == "admin", + MessageFeedback.from_source == from_source, MessageFeedback.rating == FeedbackRating.DISLIKE, ) ) @@ -1326,10 +1377,13 @@ class Conversation(Base): @property def status_count(self): + return self.status_count_with_session(session=db.session()) + + def status_count_with_session(self, *, session: Session) -> dict[str, int] | None: from models.workflow import WorkflowRun # Get all messages with workflow_run_id for this conversation - messages = db.session.scalars( + messages = session.scalars( select(Message).where(Message.conversation_id == self.id, Message.workflow_run_id.isnot(None)) ).all() @@ -1341,7 +1395,7 @@ class Conversation(Base): workflow_runs = {} if workflow_run_ids: - workflow_runs_query = db.session.scalars( + workflow_runs_query = session.scalars( select(WorkflowRun).where( WorkflowRun.id.in_(workflow_run_ids), WorkflowRun.app_id == self.app_id, # Filter by this conversation's app_id @@ -1380,8 +1434,11 @@ class Conversation(Base): } @property - def first_message(self): - return db.session.scalar( + def first_message(self) -> Message | None: + return self.first_message_with_session(session=db.session()) + + def first_message_with_session(self, *, session: Session) -> Message | None: + return session.scalar( select(Message).where(Message.conversation_id == self.id).order_by(Message.created_at.asc()) ) @@ -1391,9 +1448,12 @@ class Conversation(Base): return session.scalar(select(App).where(App.id == self.app_id)) @property - def from_end_user_session_id(self): + def from_end_user_session_id(self) -> str | None: + return self.from_end_user_session_id_with_session(session=db.session()) + + def from_end_user_session_id_with_session(self, *, session: Session) -> str | None: if self.from_end_user_id: - end_user = db.session.scalar(select(EndUser).where(EndUser.id == self.from_end_user_id)) + end_user = session.scalar(select(EndUser).where(EndUser.id == self.from_end_user_id)) if end_user: return end_user.session_id @@ -1401,8 +1461,11 @@ class Conversation(Base): @property def from_account_name(self) -> str | None: + return self.from_account_name_with_session(session=db.session()) + + def from_account_name_with_session(self, *, session: Session) -> str | None: if self.from_account_id: - account = db.session.scalar(select(Account).where(Account.id == self.from_account_id)) + account = session.scalar(select(Account).where(Account.id == self.from_account_id)) if account: return account.name @@ -1501,12 +1564,29 @@ class Message(Base): @property def inputs(self) -> dict[str, Any]: + return self.inputs_with_session(session=db.session()) + + @inputs.setter + def inputs(self, value: Mapping[str, Any]): + inputs = dict(value) + for k, v in inputs.items(): + match v: + case File(): + inputs[k] = v.model_dump() + case list(): + v_list = v + if all(isinstance(item, File) for item in v_list): + inputs[k] = [item.model_dump() for item in v_list if isinstance(item, File)] + self._inputs = inputs + + def inputs_with_session(self, *, session: Session) -> dict[str, Any]: inputs = self._inputs.copy() # Compatibility bridge: message inputs are persisted as JSON and must remain # readable across file payload shape changes. Do not assume `tenant_id` # is serialized into each file mapping going forward. tenant_resolver = _build_app_tenant_resolver( app_id=self.app_id, + session=session, owner_tenant_id=cast(str | None, getattr(self, "_owner_tenant_id", None)), ) for key, value in inputs.items(): @@ -1538,19 +1618,6 @@ class Message(Base): inputs[key] = file_list return inputs - @inputs.setter - def inputs(self, value: Mapping[str, Any]): - inputs = dict(value) - for k, v in inputs.items(): - match v: - case File(): - inputs[k] = v.model_dump() - case list(): - v_list = v - if all(isinstance(item, File) for item in v_list): - inputs[k] = [item.model_dump() for item in v_list if isinstance(item, File)] - self._inputs = inputs - @property def re_sign_file_url_answer(self) -> str: if not self.answer: @@ -1626,45 +1693,59 @@ class Message(Base): return re_sign_file_url_answer @property - def user_feedback(self): - return db.session.scalar( + def user_feedback(self) -> MessageFeedback | None: + return self.user_feedback_with_session(session=db.session()) + + def user_feedback_with_session(self, *, session: Session) -> MessageFeedback | None: + return session.scalar( select(MessageFeedback).where(MessageFeedback.message_id == self.id, MessageFeedback.from_source == "user") ) @property - def admin_feedback(self): - return db.session.scalar( + def admin_feedback(self) -> MessageFeedback | None: + return self.admin_feedback_with_session(session=db.session()) + + def admin_feedback_with_session(self, *, session: Session) -> MessageFeedback | None: + return session.scalar( select(MessageFeedback).where(MessageFeedback.message_id == self.id, MessageFeedback.from_source == "admin") ) @property - def feedbacks(self): - feedbacks = db.session.scalars(select(MessageFeedback).where(MessageFeedback.message_id == self.id)).all() - return feedbacks + def feedbacks(self) -> Sequence[MessageFeedback]: + return self.feedbacks_with_session(session=db.session()) + + def feedbacks_with_session(self, *, session: Session) -> Sequence[MessageFeedback]: + return session.scalars(select(MessageFeedback).where(MessageFeedback.message_id == self.id)).all() @property - def annotation(self): - annotation = db.session.scalar(select(MessageAnnotation).where(MessageAnnotation.message_id == self.id)) - return annotation + def annotation(self) -> MessageAnnotation | None: + return self.annotation_with_session(session=db.session()) + + def annotation_with_session(self, *, session: Session) -> MessageAnnotation | None: + return session.scalar(select(MessageAnnotation).where(MessageAnnotation.message_id == self.id)) @property - def annotation_hit_history(self): - annotation_history = db.session.scalar( + def annotation_hit_history(self) -> MessageAnnotation | None: + return self.annotation_hit_history_with_session(session=db.session()) + + def annotation_hit_history_with_session(self, *, session: Session) -> MessageAnnotation | None: + annotation_history = session.scalar( select(AppAnnotationHitHistory).where(AppAnnotationHitHistory.message_id == self.id) ) if annotation_history: - return db.session.scalar( + return session.scalar( select(MessageAnnotation).where(MessageAnnotation.id == annotation_history.annotation_id) ) return None @property - def app_model_config(self): - conversation = db.session.scalar(select(Conversation).where(Conversation.id == self.conversation_id)) + def app_model_config(self) -> AppModelConfig | None: + return self.app_model_config_with_session(session=db.session()) + + def app_model_config_with_session(self, *, session: Session) -> AppModelConfig | None: + conversation = session.scalar(select(Conversation).where(Conversation.id == self.conversation_id)) if conversation: - return db.session.scalar( - select(AppModelConfig).where(AppModelConfig.id == conversation.app_model_config_id) - ) + return session.scalar(select(AppModelConfig).where(AppModelConfig.id == conversation.app_model_config_id)) return None @@ -1678,7 +1759,10 @@ class Message(Base): @property def agent_thoughts(self) -> Sequence[MessageAgentThought]: - return db.session.scalars( + return self.agent_thoughts_with_session(session=db.session()) + + def agent_thoughts_with_session(self, *, session: Session) -> Sequence[MessageAgentThought]: + return session.scalars( select(MessageAgentThought) .where(MessageAgentThought.message_id == self.id) .order_by(MessageAgentThought.position.asc()) @@ -1690,10 +1774,13 @@ class Message(Base): @property def message_files(self) -> list[MessageFileInfo]: + return self.message_files_with_session(session=db.session()) + + def message_files_with_session(self, *, session: Session) -> list[MessageFileInfo]: from factories import file_factory - message_files = db.session.scalars(select(MessageFile).where(MessageFile.message_id == self.id)).all() - current_app = db.session.scalar(select(App).where(App.id == self.app_id)) + message_files = session.scalars(select(MessageFile).where(MessageFile.message_id == self.id)).all() + current_app = session.scalar(select(App).where(App.id == self.app_id)) if not current_app: raise ValueError(f"App {self.app_id} not found") @@ -1756,7 +1843,7 @@ class Message(Base): ], ) - db.session.commit() + session.commit() return result # TODO(QuantumGhost): dirty hacks, fix this later. @@ -1861,7 +1948,10 @@ class MessageFeedback(TypeBase): @property def from_account(self) -> Account | None: - return db.session.scalar(select(Account).where(Account.id == self.from_account_id)) + return self.from_account_with_session(session=db.session()) + + def from_account_with_session(self, *, session: Session) -> Account | None: + return session.scalar(select(Account).where(Account.id == self.from_account_id)) def to_dict(self) -> MessageFeedbackDict: return { @@ -1946,12 +2036,18 @@ class MessageAnnotation(TypeBase): return self.question or self.content @property - def account(self): - return db.session.scalar(select(Account).where(Account.id == self.account_id)) + def account(self) -> Account | None: + return self.account_with_session(session=db.session()) + + def account_with_session(self, *, session: Session) -> Account | None: + return session.scalar(select(Account).where(Account.id == self.account_id)) @property - def annotation_create_account(self): - return db.session.scalar(select(Account).where(Account.id == self.account_id)) + def annotation_create_account(self) -> Account | None: + return self.annotation_create_account_with_session(session=db.session()) + + def annotation_create_account_with_session(self, *, session: Session) -> Account | None: + return session.scalar(select(Account).where(Account.id == self.account_id)) class AppAnnotationHitHistory(TypeBase): @@ -2220,10 +2316,10 @@ class Site(Base): self._custom_disclaimer = value @staticmethod - def generate_code(n: int) -> str: + def generate_code(n: int, *, session: Session) -> str: while True: result = generate_string(n) - while (db.session.scalar(select(func.count(Site.id)).where(Site.code == result)) or 0) > 0: + while (session.scalar(select(func.count(Site.id)).where(Site.code == result)) or 0) > 0: result = generate_string(n) return result @@ -2251,10 +2347,10 @@ class ApiToken(Base): # bug: this uses setattr so idk the field. created_at = mapped_column(sa.DateTime, nullable=False, server_default=func.current_timestamp()) @staticmethod - def generate_api_key(prefix: str, n: int) -> str: + def generate_api_key(prefix: str, n: int, *, session: Session) -> str: while True: result = prefix + generate_string(n) - if db.session.scalar(select(exists().where(ApiToken.token == result))): + if session.scalar(select(exists().where(ApiToken.token == result))): continue return result diff --git a/api/models/workflow.py b/api/models/workflow.py index 2923dfdb022..d002d5eec7f 100644 --- a/api/models/workflow.py +++ b/api/models/workflow.py @@ -285,12 +285,18 @@ class Workflow(Base): # bug return workflow @property - def created_by_account(self): - return db.session.get(Account, self.created_by) + def created_by_account(self) -> Account | None: + return self.get_created_by_account(session=db.session()) + + def get_created_by_account(self, *, session: orm.Session) -> Account | None: + return session.get(Account, self.created_by) @property - def updated_by_account(self): - return db.session.get(Account, self.updated_by) if self.updated_by else None + def updated_by_account(self) -> Account | None: + return self.get_updated_by_account(session=db.session()) + + def get_updated_by_account(self, *, session: orm.Session) -> Account | None: + return session.get(Account, self.updated_by) if self.updated_by else None @property def kind_or_standard(self) -> str: @@ -552,6 +558,9 @@ class Workflow(Base): # bug "not if this specific workflow version is the one being used by the tool." ) def tool_published(self) -> bool: + return self.get_tool_published(session=db.session()) + + def get_tool_published(self, *, session: orm.Session) -> bool: """ DEPRECATED: This property is not accurate for determining if a workflow is published as a tool. It only checks if there's a WorkflowToolProvider for the app, not if this specific workflow version @@ -567,7 +576,7 @@ class Workflow(Base): # bug WorkflowToolProvider.app_id == self.app_id, ) ) - return db.session.execute(stmt).scalar_one() + return session.execute(stmt).scalar_one() @property def environment_variables( diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index 2a4f2a404d6..19ed007db00 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -3058,7 +3058,7 @@ Get message details by ID | 404 | Message not found | | ### [POST] /apps/{app_id}/model-config -**Modify app model config** +**Modify the app model config and dataset joins in one request transaction** Update application model configuration diff --git a/api/providers/trace/trace-tencent/src/dify_trace_tencent/tencent_trace.py b/api/providers/trace/trace-tencent/src/dify_trace_tencent/tencent_trace.py index 305e7851c8b..9221ee9d7e9 100644 --- a/api/providers/trace/trace-tencent/src/dify_trace_tencent/tencent_trace.py +++ b/api/providers/trace/trace-tencent/src/dify_trace_tencent/tencent_trace.py @@ -260,7 +260,7 @@ class TencentDataTrace(BaseTraceInstance): if not current_tenant: raise ValueError(f"Current tenant not found for account {service_account.id}") - service_account.set_tenant_id(current_tenant.tenant_id) + service_account.set_tenant_id_with_session(current_tenant.tenant_id, session=session) repository = SQLAlchemyWorkflowNodeExecutionRepository( session_factory=session_maker, diff --git a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py index 4d4898b1173..4aae27f6b4c 100644 --- a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py +++ b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py @@ -441,7 +441,7 @@ class TestTencentDataTrace: results = tencent_data_trace._get_workflow_node_executions(trace_info) assert results == mock_executions - account.set_tenant_id.assert_called_once_with("tenant-1") + account.set_tenant_id_with_session.assert_called_once_with("tenant-1", session=session) def test_get_workflow_node_executions_no_app_id(self, tencent_data_trace, caplog: pytest.LogCaptureFixture): trace_info = MagicMock(spec=WorkflowTraceInfo) diff --git a/api/providers/vdb/vdb-weaviate/tests/unit_tests/test_weaviate_vector.py b/api/providers/vdb/vdb-weaviate/tests/unit_tests/test_weaviate_vector.py index 79c06ea6028..84dbddbeb3a 100644 --- a/api/providers/vdb/vdb-weaviate/tests/unit_tests/test_weaviate_vector.py +++ b/api/providers/vdb/vdb-weaviate/tests/unit_tests/test_weaviate_vector.py @@ -866,7 +866,7 @@ class TestVectorDefaultAttributes(unittest.TestCase): mock_dataset = MagicMock() mock_dataset.index_struct_dict = None - vector = Vector(dataset=mock_dataset) + vector = Vector(dataset=mock_dataset, session=MagicMock()) assert "doc_type" in vector._attributes, f"doc_type should be in default attributes, got: {vector._attributes}" diff --git a/api/schedule/clean_unused_datasets_task.py b/api/schedule/clean_unused_datasets_task.py index 03417647724..9bdb074647c 100644 --- a/api/schedule/clean_unused_datasets_task.py +++ b/api/schedule/clean_unused_datasets_task.py @@ -8,9 +8,9 @@ from sqlalchemy.exc import SQLAlchemyError import app from configs import dify_config +from core.db.session_factory import session_factory from core.rag.index_processor.index_processor_factory import IndexProcessorFactory from enums.cloud_plan import CloudPlan -from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.pagination import paginate_query from models.dataset import Dataset, DatasetAutoDisableLog, DatasetQuery, Document @@ -89,68 +89,78 @@ def clean_unused_datasets_task(): .order_by(Dataset.created_at.desc()) ) - datasets = paginate_query(stmt, page=page, per_page=50) + with session_factory.create_session() as session: + datasets = paginate_query(stmt, page=page, per_page=50, session=session) + + if datasets is None or datasets.items is None or len(datasets.items) == 0: + break + + for dataset in datasets: + dataset_query = session.scalars( + select(DatasetQuery).where( + DatasetQuery.created_at > clean_day, DatasetQuery.dataset_id == dataset.id + ) + ).all() + + if not dataset_query or len(dataset_query) == 0: + try: + should_clean = True + + # Check plan filter if specified + if plan_filter: + features_cache_key = f"features:{dataset.tenant_id}" + plan_cache = redis_client.get(features_cache_key) + if plan_cache is None: + features = FeatureService.get_features( + dataset.tenant_id, exclude_vector_space=True + ) + redis_client.setex(features_cache_key, 600, features.billing.subscription.plan) + plan = features.billing.subscription.plan + else: + plan = plan_cache.decode() + should_clean = plan == plan_filter + + if should_clean: + # Add auto disable log if required + if add_logs: + documents = session.scalars( + select(Document).where( + Document.dataset_id == dataset.id, + Document.enabled == True, + Document.archived == False, + ) + ).all() + for document in documents: + dataset_auto_disable_log = DatasetAutoDisableLog( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id=document.id, + ) + session.add(dataset_auto_disable_log) + + # Remove index + index_processor = IndexProcessorFactory( + dataset.get_doc_form(session=session) + ).init_index_processor() + index_processor.clean(dataset, None, session=session) + + # Update document + session.execute( + update(Document).where(Document.dataset_id == dataset.id).values(enabled=False) + ) + session.commit() + click.echo( + click.style(f"Cleaned unused dataset {dataset.id} from db success!", fg="green") + ) + except Exception as e: + session.rollback() + click.echo( + click.style(f"clean dataset index error: {e.__class__.__name__} {str(e)}", fg="red") + ) except SQLAlchemyError: raise - if datasets is None or datasets.items is None or len(datasets.items) == 0: - break - - for dataset in datasets: - dataset_query = db.session.scalars( - select(DatasetQuery).where( - DatasetQuery.created_at > clean_day, DatasetQuery.dataset_id == dataset.id - ) - ).all() - - if not dataset_query or len(dataset_query) == 0: - try: - should_clean = True - - # Check plan filter if specified - if plan_filter: - features_cache_key = f"features:{dataset.tenant_id}" - plan_cache = redis_client.get(features_cache_key) - if plan_cache is None: - features = FeatureService.get_features(dataset.tenant_id, exclude_vector_space=True) - redis_client.setex(features_cache_key, 600, features.billing.subscription.plan) - plan = features.billing.subscription.plan - else: - plan = plan_cache.decode() - should_clean = plan == plan_filter - - if should_clean: - # Add auto disable log if required - if add_logs: - documents = db.session.scalars( - select(Document).where( - Document.dataset_id == dataset.id, - Document.enabled == True, - Document.archived == False, - ) - ).all() - for document in documents: - dataset_auto_disable_log = DatasetAutoDisableLog( - tenant_id=dataset.tenant_id, - dataset_id=dataset.id, - document_id=document.id, - ) - db.session.add(dataset_auto_disable_log) - - # Remove index - index_processor = IndexProcessorFactory(dataset.doc_form).init_index_processor() - index_processor.clean(dataset, None) - - # Update document - db.session.execute( - update(Document).where(Document.dataset_id == dataset.id).values(enabled=False) - ) - db.session.commit() - click.echo(click.style(f"Cleaned unused dataset {dataset.id} from db success!", fg="green")) - except Exception as e: - click.echo(click.style(f"clean dataset index error: {e.__class__.__name__} {str(e)}", fg="red")) - page += 1 end_at = time.perf_counter() diff --git a/api/services/account_service.py b/api/services/account_service.py index ee7d1feabfd..83716915a98 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -330,7 +330,7 @@ class AccountService: .limit(1) ) if current_tenant: - account.set_tenant_id(current_tenant.tenant_id) + account.set_tenant_id_with_session(current_tenant.tenant_id, session=session) else: available_ta = session.scalar( select(TenantAccountJoin) @@ -341,7 +341,7 @@ class AccountService: if not available_ta: return None - account.set_tenant_id(available_ta.tenant_id) + account.set_tenant_id_with_session(available_ta.tenant_id, session=session) available_ta.current = True available_ta.last_opened_at = naive_utc_now() session.commit() @@ -350,6 +350,8 @@ class AccountService: # NOTE: make sure account is accessible outside of a db session # This ensures that it will work correctly after upgrading to Flask version 3.1.2 session.refresh(account) + if session.expire_on_commit and account.current_tenant is not None: + session.refresh(account.current_tenant) session.close() return account @@ -1328,7 +1330,7 @@ class TenantService: role_ids=[owner_role_id], session=session, ) - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) session.commit() tenant_was_created.send(tenant) @@ -1550,7 +1552,7 @@ class TenantService: tenant_account_join.current = True tenant_account_join.last_opened_at = naive_utc_now() # Set the current tenant for the account - account.set_tenant_id(tenant_account_join.tenant_id) + account.set_tenant_id_with_session(tenant_account_join.tenant_id, session=session) session.commit() @staticmethod @@ -1984,7 +1986,7 @@ class RegisterService: try: tenant = TenantService.create_tenant(f"{account.name}'s Workspace", session=session) TenantService.create_tenant_member(tenant, account, session, role="owner") - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) tenant_was_created.send(tenant) except Exception: _try_join_enterprise_default_workspace(str(account.id)) diff --git a/api/services/agent/composer_service.py b/api/services/agent/composer_service.py index bd73fca171e..b9cd1db148c 100644 --- a/api/services/agent/composer_service.py +++ b/api/services/agent/composer_service.py @@ -7,7 +7,6 @@ from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from sqlalchemy.sql.elements import ColumnElement -from extensions.ext_database import db from libs.helper import to_timestamp from models import Account from models.agent import ( @@ -108,43 +107,43 @@ class AgentComposerService: def load_workflow_composer( cls, *, + session: Session, 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, session=session) + workflow = cls._get_draft_workflow(session=session, 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, session=session + session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id ) 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, session=session) + agent = cls._get_agent_if_present(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) version = cls._workflow_composer_version( + session=session, 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 + session=session, binding=binding, agent=agent, version=version, account_id=account_id ) @classmethod def _workflow_composer_version( cls, *, + session: Session, tenant_id: str, binding: WorkflowAgentNodeBinding, agent: Agent | None, snapshot_id: str | None, - session: Session, ) -> AgentConfigSnapshot | None: if snapshot_id: if agent is None: @@ -162,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, session=session) + return cls._require_version(session=session, tenant_id=tenant_id, agent_id=agent.id, version_id=snapshot_id) version_id = ( agent.active_config_snapshot_id @@ -170,22 +169,22 @@ class AgentComposerService: else binding.current_snapshot_id ) return cls._get_version_if_present( + session=session, tenant_id=tenant_id, agent_id=agent.id if agent else None, version_id=version_id, - session=session, ) @classmethod def save_workflow_composer( cls, *, + session: Session, 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") @@ -193,15 +192,16 @@ class AgentComposerService: _backfill_cli_tool_ids(payload.agent_soul) _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, session=session) + cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul) + workflow = cls._get_draft_workflow(session=session, 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, session=session + session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id ) match payload.save_strategy: case ComposerSaveStrategy.NODE_JOB_ONLY: binding = cls._save_node_job_only( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow.id, @@ -209,18 +209,18 @@ 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, session=session + session=session, tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload ) case ComposerSaveStrategy.SAVE_AS_NEW_VERSION: binding = cls._save_as_new_version( - tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload, session=session + session=session, tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload ) case ComposerSaveStrategy.SAVE_AS_NEW_AGENT: binding = cls._save_as_new_agent( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow.id, @@ -228,34 +228,33 @@ 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, session=session + session=session, tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload ) - session.commit() - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) + session.flush() + agent = cls._get_agent_if_present(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) version_id = ( agent.active_config_snapshot_id if agent and binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT else binding.current_snapshot_id ) version = cls._get_version_if_present( + session=session, 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 + session=session, binding=binding, agent=agent, version=version, account_id=account_id ) state["validation"] = cls.collect_validation_findings( + session=session, tenant_id=tenant_id, payload=payload, agent_id=binding.agent_id, - session=session, ) return state @@ -263,6 +262,7 @@ class AgentComposerService: def copy_workflow_composer_from_roster( cls, *, + session: Session, tenant_id: str, app_id: str, node_id: str, @@ -270,23 +270,22 @@ 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, session=session) + workflow = cls._get_draft_workflow(session=session, tenant_id=tenant_id, app_id=app_id) binding = cls._require_binding( - cls._get_workflow_binding(tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session) + cls._get_workflow_binding(session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id) ) 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, session=session) + agent = cls._get_agent_if_present(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) version = cls._get_version_if_present( + session=session, 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 + session=session, binding=binding, agent=agent, version=version, account_id=account_id ) if binding.binding_type != WorkflowAgentBindingType.ROSTER_AGENT: @@ -294,20 +293,21 @@ class AgentComposerService: 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, session=session) + source_agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=source_agent_id) 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( + session=session, 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() agent_soul = AgentSoulConfig.model_validate(source_version.config_snapshot_dict) inline_agent = cls._create_workflow_only_agent( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow.id, @@ -320,16 +320,15 @@ 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( + session=session, tenant_id=tenant_id, source_agent_id=source_agent.id, target_agent_id=inline_agent.id, 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 @@ -337,37 +336,36 @@ class AgentComposerService: binding.current_snapshot_id = inline_agent.active_config_snapshot_id binding.updated_by = account_id session.flush() - session.commit() version = cls._require_version( + session=session, 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, session=session + session=session, binding=binding, agent=inline_agent, version=version, account_id=account_id ) @classmethod - 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) + def load_agent_app_composer(cls, *, session: Session, tenant_id: str, app_id: str) -> dict[str, Any]: + agent = cls._require_agent_app_agent(session=session, tenant_id=tenant_id, app_id=app_id) + return cls._load_agent_composer_for_agent(session=session, tenant_id=tenant_id, agent=agent) @classmethod - 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) + def load_agent_composer(cls, *, session: Session, tenant_id: str, agent_id: str) -> dict[str, Any]: + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) + return cls._load_agent_composer_for_agent(session=session, tenant_id=tenant_id, agent=agent) @classmethod def load_agent_soul_for_debug( cls, *, + session: Session, tenant_id: str, agent_id: str, account_id: str, draft_type: AgentConfigDraftType, - session: Session, ) -> AgentSoulConfig: """Load the same normal or account-owned build draft used by Agent debug chat.""" if draft_type == AgentConfigDraftType.DEBUG_BUILD: @@ -386,17 +384,17 @@ class AgentComposerService: return AgentSoulConfig.model_validate(state["agent_soul"]) @classmethod - def _load_agent_composer_for_agent(cls, *, tenant_id: str, agent: Agent, session: Session) -> dict[str, Any]: + def _load_agent_composer_for_agent(cls, *, session: Session, tenant_id: str, agent: Agent) -> dict[str, Any]: draft = cls._get_or_create_agent_draft( + session=session, 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, session=session + session=session, tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id ) return { "variant": ComposerVariant.AGENT_APP.value, @@ -413,13 +411,7 @@ class AgentComposerService: @classmethod def save_agent_app_composer( - cls, - *, - tenant_id: str, - app_id: str, - account_id: str, - payload: ComposerSavePayload, - session: Session, + cls, *, session: Session, tenant_id: str, app_id: str, account_id: str, payload: ComposerSavePayload ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent App composer endpoint only accepts agent_app variant") @@ -432,7 +424,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, session=session) + agent = cls._get_agent_app_agent(session=session, tenant_id=tenant_id, app_id=app_id) if not agent: agent = Agent( tenant_id=tenant_id, @@ -454,22 +446,16 @@ class AgentComposerService: session.rollback() raise AgentNameConflictError() from exc return cls._save_agent_composer_for_agent( + session=session, 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, - session: Session, + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent composer endpoint only accepts agent_app variant") @@ -481,51 +467,45 @@ 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, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) return cls._save_agent_composer_for_agent( + session=session, 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, - session: Session, + cls, *, session: Session, tenant_id: str, agent: Agent, account_id: str, payload: ComposerSavePayload ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") cls._save_agent_draft( + session=session, tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, 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( + session=session, tenant_id=tenant_id, agent=agent, agent_soul=payload.agent_soul, - session=session, ) - session.commit() - state = cls.load_agent_composer(tenant_id=tenant_id, agent_id=agent.id, session=session) + session.flush() + state = cls.load_agent_composer(session=session, tenant_id=tenant_id, agent_id=agent.id) state["validation"] = cls.collect_validation_findings( + session=session, tenant_id=tenant_id, payload=payload, agent_id=agent.id, - session=session, ) return state @@ -533,23 +513,24 @@ class AgentComposerService: def _agent_soul_matches_active_config( cls, *, + session: Session, tenant_id: str, agent: Agent, agent_soul: AgentSoulConfig, - session: Session, ) -> bool: if not agent.active_config_snapshot_id: return False active_version = cls._get_version_if_present( + session=session, tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, - session=session, ) if not active_version: return False if agent.source in APP_BACKED_AGENT_SOURCES and not cls._has_publish_visible_revision( + session=session, tenant_id=tenant_id, agent_id=agent.id, snapshot_id=agent.active_config_snapshot_id, @@ -559,8 +540,10 @@ class AgentComposerService: return _agent_soul_config_json(agent_soul) == _agent_soul_config_json(active_version.config_snapshot_dict) @classmethod - def _has_publish_visible_revision(cls, *, tenant_id: str, agent_id: str, snapshot_id: str) -> bool: - revisions = db.session.scalars( + def _has_publish_visible_revision( + cls, *, session: Session, tenant_id: str, agent_id: str, snapshot_id: str + ) -> bool: + revisions = session.scalars( select(AgentConfigRevision.operation).where( AgentConfigRevision.tenant_id == tenant_id, AgentConfigRevision.agent_id == agent_id, @@ -581,24 +564,18 @@ class AgentComposerService: @classmethod def publish_agent_app_draft( - cls, - *, - tenant_id: str, - agent_id: str, - account_id: str, - version_note: str | None = None, - session: Session, + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, version_note: str | None = None ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) if agent.scope != AgentScope.ROSTER or agent.source not in APP_BACKED_AGENT_SOURCES: raise AgentNotFoundError() draft = cls._get_or_create_agent_draft( + session=session, tenant_id=tenant_id, agent=agent, 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( @@ -611,8 +588,9 @@ class AgentComposerService: ) if not agent_soul_has_model(agent_soul): raise AgentModelNotConfiguredError() - cls.validate_knowledge_datasets(tenant_id=tenant_id, agent_soul=agent_soul) + cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=agent_soul) version = cls._create_config_version( + session=session, tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, @@ -620,7 +598,6 @@ 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) @@ -628,7 +605,7 @@ class AgentComposerService: agent.updated_by = account_id draft.base_snapshot_id = version.id draft.updated_by = account_id - session.commit() + session.flush() return { "result": "success", "active_config_snapshot_id": version.id, @@ -638,29 +615,23 @@ class AgentComposerService: @classmethod def checkout_agent_app_build_draft( - cls, - *, - tenant_id: str, - agent_id: str, - account_id: str, - force: bool = False, - session: Session, + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, force: bool = False ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) normal_draft = cls._get_or_create_agent_draft( + session=session, 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( + session=session, 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) @@ -677,19 +648,19 @@ class AgentComposerService: 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 - session.commit() + session.flush() 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, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str ) -> dict[str, Any]: build_draft = cls._get_agent_draft( + session=session, 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() @@ -697,47 +668,42 @@ class AgentComposerService: @classmethod def save_agent_app_build_draft( - cls, - *, - tenant_id: str, - agent_id: str, - account_id: str, - payload: ComposerSavePayload, - session: Session, + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, payload: ComposerSavePayload ) -> 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, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) build_draft = cls._save_agent_draft( + session=session, tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, agent_soul=payload.agent_soul, account_id_for_audit=account_id, - session=session, ) - session.commit() + session.flush() 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, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str ) -> dict[str, Any]: - agent = cls._require_agent(tenant_id=tenant_id, agent_id=agent_id, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) build_draft = cls._get_agent_draft( + session=session, 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() applied_agent_soul = AgentSoulConfig.model_validate(build_draft.config_snapshot_dict) normal_draft = cls._save_agent_draft( + session=session, tenant_id=tenant_id, agent=agent, draft_type=AgentConfigDraftType.DRAFT, @@ -745,43 +711,42 @@ 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( + session=session, tenant_id=tenant_id, agent=agent, agent_soul=applied_agent_soul, - session=session, ) agent.updated_by = account_id session.delete(build_draft) - session.commit() + session.flush() 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, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str ) -> dict[str, Any]: build_draft = cls._get_agent_draft( + session=session, 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: session.delete(build_draft) - session.commit() + session.flush() return {"result": "success"} @classmethod def collect_validation_findings( cls, *, + session: Session, 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 = ( @@ -796,16 +761,18 @@ class AgentComposerService: if agent_id and payload.agent_soul is not None: findings["warnings"].extend( cls._drive_mention_findings( + session=session, tenant_id=tenant_id, agent_id=agent_id, prompt=payload.agent_soul.prompt.system_prompt, - session=session, ) ) return findings @classmethod - def validate_knowledge_datasets(cls, *, tenant_id: str, agent_soul: AgentSoulConfig | None) -> None: + def validate_knowledge_datasets( + cls, *, session: Session, tenant_id: str, agent_soul: AgentSoulConfig | None + ) -> None: """Hard-validate tenant-scoped knowledge set datasets before saving. DTO validators own set shape, duplicate set ids/names, and duplicate @@ -815,7 +782,9 @@ class AgentComposerService: """ if agent_soul is None: return - missing_ids = list_missing_tenant_knowledge_dataset_ids(tenant_id=tenant_id, agent_soul=agent_soul) + missing_ids = list_missing_tenant_knowledge_dataset_ids( + session=session, tenant_id=tenant_id, agent_soul=agent_soul + ) if missing_ids: raise InvalidComposerConfigError( "knowledge_dataset_not_found: knowledge sets reference missing or out-of-scope datasets: " @@ -823,7 +792,7 @@ class AgentComposerService: ) @classmethod - def resolve_bound_agent_id(cls, *, tenant_id: str, app_id: str, session: Session) -> str | None: + def resolve_bound_agent_id(cls, *, session: Session, tenant_id: str, app_id: str) -> str | None: """The Agent App's bound roster agent id, if any (validate-endpoint context).""" return session.scalar( select(Agent.id) @@ -839,15 +808,15 @@ class AgentComposerService: @classmethod def resolve_workflow_node_agent_id( - cls, *, tenant_id: str, app_id: str, node_id: str, session: Session + cls, *, session: Session, tenant_id: str, app_id: str, node_id: str ) -> 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, session=session) + workflow = cls._get_draft_workflow(session=session, tenant_id=tenant_id, app_id=app_id) except ValueError: return None binding = cls._get_workflow_binding( - tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id, session=session + session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id ) return binding.agent_id if binding else None @@ -855,10 +824,10 @@ class AgentComposerService: def _drive_mention_findings( cls, *, + session: Session, 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 @@ -901,19 +870,13 @@ class AgentComposerService: @classmethod def get_workflow_candidates( - cls, - *, - tenant_id: str, - app_id: str, - node_id: str, - user_id: str, - session: Session, + cls, *, session: Session, tenant_id: str, app_id: str, node_id: str, user_id: str ) -> 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, session=session) + workflow = cls._get_draft_workflow(session=session, tenant_id=tenant_id, app_id=app_id) except ValueError: workflow = None @@ -921,37 +884,35 @@ class AgentComposerService: 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, session=session + session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id ) if binding is not None: node_job = cls._parse_node_job(binding) - agent_soul = cls._load_binding_soul(tenant_id=tenant_id, binding=binding, session=session) + agent_soul = cls._load_binding_soul(session=session, tenant_id=tenant_id, binding=binding) truncated = False previous_outputs: list[dict[str, Any]] = [] if workflow is not None: - draft_variable_session = cls._draft_variable_session() - try: - previous_outputs, outputs_truncated = previous_node_output_candidates( - 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, 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 - ), - system_variables_loader=lambda: cls._draft_system_variables( - session=draft_variable_session, app_id=app_id, user_id=user_id - ), - ) - finally: - draft_variable_session.close() + previous_outputs, outputs_truncated = previous_node_output_candidates( + graph=workflow.graph_dict, + node_id=node_id, + declared_outputs_loader=lambda nid: cls._binding_declared_outputs( + session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=nid + ), + draft_variables_loader=lambda nid: cls._draft_node_variables( + session=session, app_id=app_id, node_id=nid, user_id=user_id + ), + system_variables_loader=lambda: cls._draft_system_variables( + session=session, app_id=app_id, user_id=user_id + ), + ) truncated = truncated or outputs_truncated soul_lists, soul_truncated = soul_candidates( agent_soul=agent_soul, - dataset_lookup=lambda ids: get_tenant_knowledge_dataset_rows(tenant_id=tenant_id, dataset_ids=ids), + dataset_lookup=lambda ids: get_tenant_knowledge_dataset_rows( + session=session, tenant_id=tenant_id, dataset_ids=ids + ), workspace_tools_loader=lambda: cls._workspace_dify_tools(tenant_id=tenant_id, user_id=user_id), ) truncated = truncated or soul_truncated @@ -972,15 +933,17 @@ class AgentComposerService: @classmethod def get_agent_app_candidates( - cls, *, tenant_id: str, agent_id: str, user_id: str, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str, user_id: str ) -> 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, session=session) + agent_soul = cls._load_agent_soul(session=session, tenant_id=tenant_id, agent_id=agent_id) 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), + dataset_lookup=lambda ids: get_tenant_knowledge_dataset_rows( + session=session, tenant_id=tenant_id, dataset_ids=ids + ), workspace_tools_loader=lambda: cls._workspace_dify_tools(tenant_id=tenant_id, user_id=user_id), ) response = ComposerCandidatesResponse( @@ -1003,29 +966,29 @@ class AgentComposerService: @classmethod def _load_binding_soul( - cls, *, tenant_id: str, binding: WorkflowAgentNodeBinding, session: Session + cls, *, session: Session, tenant_id: str, binding: WorkflowAgentNodeBinding ) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=binding.agent_id, session=session) + agent = cls._get_agent_if_present(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) version = cls._get_version_if_present( + session=session, 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, session: Session) -> AgentSoulConfig | None: - agent = cls._get_agent_if_present(tenant_id=tenant_id, agent_id=agent_id, session=session) + def _load_agent_soul(cls, *, session: Session, tenant_id: str, agent_id: str) -> AgentSoulConfig | None: + agent = cls._get_agent_if_present(session=session, tenant_id=tenant_id, agent_id=agent_id) if agent is None: return None draft = cls._get_or_create_agent_draft( + session=session, tenant_id=tenant_id, agent=agent, 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) @@ -1041,10 +1004,10 @@ class AgentComposerService: @classmethod def _binding_declared_outputs( - cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session + cls, *, session: Session, tenant_id: str, workflow_id: str, node_id: str ) -> list[DeclaredOutputConfig] | None: binding = cls._get_workflow_binding( - tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id, session=session + session=session, tenant_id=tenant_id, workflow_id=workflow_id, node_id=node_id ) if binding is None: return None @@ -1053,12 +1016,6 @@ class AgentComposerService: return None return list(_effective_declared_outputs(node_job.declared_outputs)) - @staticmethod - def _draft_variable_session(): - from sqlalchemy.orm import sessionmaker - - return sessionmaker(bind=db.engine, expire_on_commit=False)() - @staticmethod def _draft_node_variables(*, session: Any, app_id: str, node_id: str, user_id: str) -> list[tuple[str, str | None]]: from services.workflow_draft_variable_service import WorkflowDraftVariableService @@ -1120,7 +1077,7 @@ class AgentComposerService: return tools @classmethod - def calculate_impact(cls, *, tenant_id: str, current_snapshot_id: str, session: Session) -> dict[str, Any]: + def calculate_impact(cls, *, session: Session, tenant_id: str, current_snapshot_id: str) -> dict[str, Any]: snapshot = session.scalar( select(AgentConfigSnapshot) .where( @@ -1161,6 +1118,7 @@ class AgentComposerService: def _save_node_job_only( cls, *, + session: Session, tenant_id: str, app_id: str, workflow_id: str, @@ -1168,12 +1126,12 @@ class AgentComposerService: account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, - session: Session, ) -> WorkflowAgentNodeBinding: node_job = payload.node_job or WorkflowNodeJobConfig() if binding: if cls._is_start_from_scratch_request(binding=binding, payload=payload): return cls._switch_roster_binding_to_inline_agent( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow_id, @@ -1181,25 +1139,24 @@ 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: current_snapshot = cls._require_version( + session=session, tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, - session=session, ) version = cls._update_current_version( + session=session, current_snapshot=current_snapshot, account_id=account_id, 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, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) 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 @@ -1212,13 +1169,13 @@ class AgentComposerService: agent_soul = payload.agent_soul or AgentSoulConfig() agent = cls._create_workflow_only_agent( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow_id, node_id=node_id, account_id=account_id, agent_soul=agent_soul, - session=session, ) binding = WorkflowAgentNodeBinding( tenant_id=tenant_id, @@ -1249,6 +1206,7 @@ class AgentComposerService: def _switch_roster_binding_to_inline_agent( cls, *, + session: Session, tenant_id: str, app_id: str, workflow_id: str, @@ -1256,20 +1214,19 @@ 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.") agent_soul = payload.agent_soul or AgentSoulConfig() agent = cls._create_workflow_only_agent( + session=session, tenant_id=tenant_id, app_id=app_id, workflow_id=workflow_id, 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 @@ -1283,30 +1240,30 @@ class AgentComposerService: def _save_to_current_version( cls, *, + session: Session, tenant_id: str, account_id: str, binding: WorkflowAgentNodeBinding | None, payload: ComposerSavePayload, - session: Session, ) -> WorkflowAgentNodeBinding: binding = cls._require_binding(binding) if payload.agent_soul is None: raise ValueError("agent_soul is required") current_snapshot = cls._require_version( + session=session, tenant_id=tenant_id, agent_id=binding.agent_id, version_id=binding.current_snapshot_id, - session=session, ) version = cls._update_current_version( + session=session, current_snapshot=current_snapshot, account_id=account_id, 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, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) 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 @@ -1321,25 +1278,25 @@ class AgentComposerService: def _save_as_new_version( cls, *, + session: Session, tenant_id: str, 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: raise ValueError("agent_id and agent_soul are required") version = cls._create_config_version( + session=session, tenant_id=tenant_id, agent_id=binding.agent_id, account_id=account_id, 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, session=session) + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) 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 @@ -1354,6 +1311,7 @@ class AgentComposerService: def _save_as_new_agent( cls, *, + session: Session, tenant_id: str, app_id: str, workflow_id: str, @@ -1361,12 +1319,12 @@ 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") agent_name = payload.new_agent_name or "Untitled Agent" agent = cls._create_roster_agent_for_composer( + session=session, tenant_id=tenant_id, account_id=account_id, name=agent_name, @@ -1378,7 +1336,6 @@ 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: @@ -1403,23 +1360,24 @@ class AgentComposerService: def _save_to_roster( cls, *, + session: Session, tenant_id: str, 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, session=session) + source_agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) source_version = cls._require_version( + session=session, 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 roster_agent = cls._create_roster_agent_for_composer( + session=session, tenant_id=tenant_id, account_id=account_id, name=agent_name, @@ -1433,16 +1391,15 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_TO_ROSTER, version_note=payload.version_note, - session=session, ) cls._copy_agent_drive_rows( + session=session, tenant_id=tenant_id, source_agent_id=source_agent.id, target_agent_id=roster_agent.id, 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 @@ -1456,6 +1413,7 @@ class AgentComposerService: def _create_workflow_only_agent( cls, *, + session: Session, tenant_id: str, app_id: str, workflow_id: str, @@ -1468,7 +1426,6 @@ class AgentComposerService: icon_type: Any | None = None, icon: str | None = None, icon_background: str | None = None, - session: Session, ) -> Agent: backing_app = AgentRosterService(session).create_hidden_backing_app_for_workflow_agent( tenant_id=tenant_id, @@ -1501,13 +1458,13 @@ class AgentComposerService: session.add(agent) session.flush() version = cls._create_config_version( + session=session, tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, 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) @@ -1518,13 +1475,13 @@ class AgentComposerService: def _copy_agent_drive_rows( cls, *, + session: Session, tenant_id: str, source_agent_id: str, target_agent_id: str, 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]] = [] @@ -1611,6 +1568,7 @@ class AgentComposerService: def _create_roster_agent_for_composer( cls, *, + session: Session, tenant_id: str, account_id: str, name: str, @@ -1622,9 +1580,8 @@ 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, session=session) + account = cls._require_account(session=session, account_id=account_id) try: app = AppService().create_app( tenant_id, @@ -1649,18 +1606,18 @@ class AgentComposerService: raise AgentNotFoundError() current_snapshot = cls._require_version( + session=session, tenant_id=tenant_id, agent_id=agent.id, version_id=agent.active_config_snapshot_id, - session=session, ) version = cls._update_current_version( + session=session, current_snapshot=current_snapshot, account_id=account_id, 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) @@ -1672,6 +1629,7 @@ class AgentComposerService: def _create_config_version( cls, *, + session: Session, tenant_id: str, agent_id: str, account_id: str, @@ -1679,7 +1637,6 @@ class AgentComposerService: operation: AgentConfigRevisionOperation, version_note: str | None, previous_snapshot_id: str | None = None, - session: Session, ) -> AgentConfigSnapshot: next_version = ( session.scalar( @@ -1705,7 +1662,7 @@ class AgentComposerService: 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, session=session), + revision=cls._next_revision(session=session, tenant_id=tenant_id, agent_id=agent_id), operation=operation, version_note=version_note, created_by=account_id, @@ -1718,14 +1675,15 @@ class AgentComposerService: def _update_current_version( cls, *, + session: Session, current_snapshot: AgentConfigSnapshot, account_id: str, agent_soul: AgentSoulConfig, operation: AgentConfigRevisionOperation, version_note: str | None, - session: Session, ) -> AgentConfigSnapshot: return cls._create_config_version( + session=session, tenant_id=current_snapshot.tenant_id, agent_id=current_snapshot.agent_id, account_id=account_id, @@ -1733,11 +1691,10 @@ 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, session: Session) -> int: + def _next_revision(cls, *, session: Session, tenant_id: str, agent_id: str) -> int: return ( session.scalar( select(func.max(AgentConfigRevision.revision)).where( @@ -1749,7 +1706,7 @@ class AgentComposerService: ) + 1 @classmethod - def _get_agent_app_agent(cls, *, tenant_id: str, app_id: str, session: Session) -> Agent | None: + def _get_agent_app_agent(cls, *, session: Session, tenant_id: str, app_id: str) -> Agent | None: return session.scalar( select(Agent) .where( @@ -1764,8 +1721,8 @@ class AgentComposerService: ) @classmethod - 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) + def _require_agent_app_agent(cls, *, session: Session, tenant_id: str, app_id: str) -> Agent: + agent = cls._get_agent_app_agent(session=session, tenant_id=tenant_id, app_id=app_id) if agent is None: raise AgentNotFoundError() return agent @@ -1774,11 +1731,11 @@ class AgentComposerService: def _get_agent_draft( cls, *, + session: Session, tenant_id: str, agent_id: str, draft_type: AgentConfigDraftType, account_id: str | None, - session: Session, ) -> AgentConfigDraft | None: stmt = select(AgentConfigDraft).where( AgentConfigDraft.tenant_id == tenant_id, @@ -1795,27 +1752,27 @@ class AgentComposerService: def _get_or_create_agent_draft( cls, *, + session: Session, tenant_id: str, agent: Agent, draft_type: AgentConfigDraftType, account_id: str | None, created_by: str | None, - session: Session, ) -> AgentConfigDraft: draft = cls._get_agent_draft( + session=session, 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 base_snapshot = cls._get_version_if_present( + session=session, 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) @@ -1841,6 +1798,7 @@ class AgentComposerService: def _save_agent_draft( cls, *, + session: Session, tenant_id: str, agent: Agent, draft_type: AgentConfigDraftType, @@ -1848,15 +1806,14 @@ 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( + session=session, tenant_id=tenant_id, agent=agent, 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: @@ -1894,7 +1851,7 @@ class AgentComposerService: } @classmethod - def _get_draft_workflow(cls, *, tenant_id: str, app_id: str, session: Session) -> Workflow: + def _get_draft_workflow(cls, *, session: Session, tenant_id: str, app_id: str) -> Workflow: workflow = session.scalar( select(Workflow) .where( @@ -1910,7 +1867,7 @@ class AgentComposerService: @classmethod def _get_workflow_binding( - cls, *, tenant_id: str, workflow_id: str, node_id: str, session: Session + cls, *, session: Session, tenant_id: str, workflow_id: str, node_id: str ) -> WorkflowAgentNodeBinding | None: # Composer always operates against the draft workflow row, so this lookup # is scoped to ``workflow_version="draft"``. Published bindings are @@ -1934,7 +1891,7 @@ class AgentComposerService: return binding @classmethod - def _require_agent(cls, *, tenant_id: str, agent_id: str | None, session: Session) -> Agent: + def _require_agent(cls, *, session: Session, tenant_id: str, agent_id: str | None) -> Agent: if not agent_id: raise AgentNotFoundError() agent = session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.id == agent_id).limit(1)) @@ -1943,21 +1900,21 @@ class AgentComposerService: return agent @classmethod - def _require_account(cls, *, account_id: str, session: Session) -> Account: + def _require_account(cls, *, session: Session, account_id: str) -> 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, session: Session) -> Agent | None: + def _get_agent_if_present(cls, *, session: Session, tenant_id: str, agent_id: str | None) -> Agent | None: if not agent_id: return None 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, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str | None, version_id: str | None ) -> AgentConfigSnapshot: if not agent_id or not version_id: raise AgentVersionNotFoundError() @@ -1976,7 +1933,7 @@ class AgentComposerService: @classmethod def _get_version_if_present( - cls, *, tenant_id: str, agent_id: str | None, version_id: str | None, session: Session + cls, *, session: Session, tenant_id: str, agent_id: str | None, version_id: str | None ) -> AgentConfigSnapshot | None: if not agent_id or not version_id: return None @@ -2035,11 +1992,11 @@ class AgentComposerService: def _serialize_workflow_state( cls, *, + session: Session, binding: WorkflowAgentNodeBinding, 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] @@ -2054,11 +2011,11 @@ class AgentComposerService: else: save_options.append(ComposerSaveStrategy.SAVE_TO_ROSTER.value) debug_conversation_id = cls._workflow_inline_debug_conversation_id( + session=session, tenant_id=binding.tenant_id, binding=binding, agent=agent, account_id=account_id, - session=session, ) debug_conversation_message_count = ( AgentRosterService(session).count_agent_app_debug_conversation_messages( @@ -2095,7 +2052,7 @@ class AgentComposerService: "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, session=session + session=session, tenant_id=binding.tenant_id, current_snapshot_id=version.id ) if version else None, @@ -2113,12 +2070,13 @@ class AgentComposerService: @staticmethod def _workflow_inline_debug_conversation_id( *, + session: Session, tenant_id: str, binding: WorkflowAgentNodeBinding, agent: Agent | None, account_id: str | None, - session: Session, ) -> str | None: + """Return the editor's inline debug conversation within the caller-owned transaction.""" if ( not account_id or not agent @@ -2133,6 +2091,7 @@ class AgentComposerService: tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, + commit=False, ) @classmethod diff --git a/api/services/agent/dsl_service.py b/api/services/agent/dsl_service.py index 62192d9f6c9..70e7d474f98 100644 --- a/api/services/agent/dsl_service.py +++ b/api/services/agent/dsl_service.py @@ -532,7 +532,11 @@ class AgentDslService: for dataset in knowledge_set.get("datasets", []) if dataset.get("id") ] - existing = get_tenant_knowledge_dataset_rows(tenant_id=tenant_id, dataset_ids=dataset_ids) + existing = get_tenant_knowledge_dataset_rows( + session=self.session, + tenant_id=tenant_id, + dataset_ids=dataset_ids, + ) warnings = [ DslImportWarning( code=f"agent_{asset.kind}_omitted", diff --git a/api/services/agent/knowledge_datasets.py b/api/services/agent/knowledge_datasets.py index 962c562ce15..9db9bbc08c1 100644 --- a/api/services/agent/knowledge_datasets.py +++ b/api/services/agent/knowledge_datasets.py @@ -3,6 +3,8 @@ from __future__ import annotations from typing import Any from uuid import UUID +from sqlalchemy.orm import Session + from models.agent_config_entities import AgentSoulConfig @@ -26,7 +28,7 @@ def list_agent_soul_knowledge_dataset_ids(agent_soul: AgentSoulConfig) -> list[s return dataset_ids -def get_tenant_knowledge_dataset_rows(*, tenant_id: str, dataset_ids: list[str]) -> dict[str, Any]: +def get_tenant_knowledge_dataset_rows(*, session: Session, tenant_id: str, dataset_ids: list[str]) -> dict[str, Any]: """Return tenant-scoped dataset rows for normalized knowledge dataset ids. Knowledge ids come from user-editable config. Malformed ids can never match @@ -46,11 +48,13 @@ def get_tenant_knowledge_dataset_rows(*, tenant_id: str, dataset_ids: list[str]) if not valid_ids: return {} - rows, _ = DatasetService.get_datasets_by_ids(valid_ids, tenant_id) + rows, _ = DatasetService.get_datasets_by_ids(valid_ids, tenant_id, session=session) return {str(row.id): row for row in rows} -def list_missing_tenant_knowledge_dataset_ids(*, tenant_id: str, agent_soul: AgentSoulConfig | None) -> list[str]: +def list_missing_tenant_knowledge_dataset_ids( + *, session: Session, tenant_id: str, agent_soul: AgentSoulConfig | None +) -> list[str]: """Return normalized knowledge dataset ids missing from the tenant scope.""" if agent_soul is None: return [] @@ -59,5 +63,5 @@ def list_missing_tenant_knowledge_dataset_ids(*, tenant_id: str, agent_soul: Age if not dataset_ids: return [] - rows = get_tenant_knowledge_dataset_rows(tenant_id=tenant_id, dataset_ids=dataset_ids) + rows = get_tenant_knowledge_dataset_rows(session=session, tenant_id=tenant_id, dataset_ids=dataset_ids) return [dataset_id for dataset_id in dataset_ids if dataset_id not in rows] diff --git a/api/services/agent/roster_service.py b/api/services/agent/roster_service.py index 2d15fd1b9a2..d1b093a5a65 100644 --- a/api/services/agent/roster_service.py +++ b/api/services/agent/roster_service.py @@ -983,8 +983,8 @@ class AgentRosterService: return icon_type def _copy_app_model_config(self, *, source_app: App, target_app: App, account_id: str) -> None: - source_config = source_app.app_model_config - target_config = target_app.app_model_config + source_config = source_app.app_model_config_with_session(session=self._session) + target_config = target_app.app_model_config_with_session(session=self._session) if source_config is None or target_config is None: return diff --git a/api/services/agent_service.py b/api/services/agent_service.py index a201eeb0485..74c2a9ae56a 100644 --- a/api/services/agent_service.py +++ b/api/services/agent_service.py @@ -12,7 +12,7 @@ from core.plugin.impl.exc import PluginDaemonClientSideError from core.tools.tool_manager import ToolManager from libs.login import current_user from models import Account -from models.model import App, Conversation, EndUser, Message +from models.model import App, Conversation, EndUser, Message, load_annotation_reply_config class AgentService: @@ -48,7 +48,7 @@ class AgentService: if not message: raise ValueError(f"Message not found: {message_id}") - agent_thoughts = message.agent_thoughts + agent_thoughts = message.agent_thoughts_with_session(session=session) if conversation.from_end_user_id: # only select name field @@ -61,7 +61,7 @@ class AgentService: assert current_user.timezone is not None timezone = pytz.timezone(current_user.timezone) - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if not app_model_config: raise ValueError("App model config not found") @@ -76,10 +76,11 @@ class AgentService: "iterations": len(agent_thoughts), }, "iterations": [], - "files": message.message_files, + "files": message.message_files_with_session(session=session), } - agent_config = AgentConfigManager.convert(app_model_config.to_dict()) + annotation_reply = load_annotation_reply_config(session, app_model.id) + agent_config = AgentConfigManager.convert(app_model_config.to_dict(annotation_reply=annotation_reply)) if not agent_config: raise ValueError("Agent config not found") diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index ccca621aab5..c36db8d4b51 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -9,11 +9,11 @@ from werkzeug.datastructures import FileStorage from werkzeug.exceptions import NotFound from core.helper.csv_sanitizer import CSVSanitizer -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 from libs.pagination import paginate_query +from models.dataset import DatasetCollectionBinding from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation from services.app_ref_service import AnnotationRef, AppRef from services.feature_service import FeatureService @@ -103,7 +103,7 @@ class AppAnnotationService: @classmethod def up_insert_app_annotation_from_message( - cls, args: UpsertAnnotationArgs, app_id: str, *, session: Session + cls, args: UpsertAnnotationArgs, app_id: str, session: Session ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() @@ -128,7 +128,9 @@ class AppAnnotationService: question = args.get("question") or message.query or "" - annotation: MessageAnnotation | None = message.annotation + annotation = session.scalar( + select(MessageAnnotation).where(MessageAnnotation.message_id == message.id).limit(1) + ) if annotation: annotation.content = answer annotation.question = question @@ -213,7 +215,7 @@ 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, *, session: Session): + 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 = session.scalar( @@ -243,11 +245,11 @@ class AppAnnotationService: .where(MessageAnnotation.app_id == app_id) .order_by(MessageAnnotation.created_at.desc(), MessageAnnotation.id.desc()) ) - annotations = paginate_query(stmt, page=page, per_page=limit, max_per_page=100) + annotations = paginate_query(stmt, session=session, page=page, per_page=limit, max_per_page=100) return annotations.items, annotations.total or 0 @classmethod - def export_annotation_list_by_app_id(cls, app_id: str, *, session: Session): + def export_annotation_list_by_app_id(cls, app_id: str, session: Session): """ Export all annotations for an app with CSV injection protection. @@ -281,7 +283,7 @@ class AppAnnotationService: @classmethod def insert_app_annotation_directly( - cls, args: InsertAnnotationArgs, app_id: str, *, session: Session + cls, args: InsertAnnotationArgs, app_id: str, session: Session ) -> MessageAnnotation: # get app info current_user, current_tenant_id = current_account_with_tenant() @@ -317,7 +319,10 @@ class AppAnnotationService: @classmethod def update_app_annotation_directly( - cls, args: UpdateAnnotationArgs, annotation_ref: AnnotationRef, session: Session + cls, + args: UpdateAnnotationArgs, + annotation_ref: AnnotationRef, + session: Session, ): annotation = cls._get_annotation_by_ref(annotation_ref, session) @@ -386,7 +391,7 @@ class AppAnnotationService: ) @classmethod - def delete_app_annotations_in_batch(cls, app_ref: AppRef, annotation_ids: list[str], *, session: Session): + 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 = session.execute( select(MessageAnnotation, AppAnnotationSetting) @@ -428,7 +433,7 @@ class AppAnnotationService: return {"deleted_count": deleted_count} @classmethod - def batch_import_app_annotations(cls, app_id: str, file: FileStorage, *, session: Session): + def batch_import_app_annotations(cls, app_id: str, file: FileStorage, session: Session): """ Batch import annotations from CSV file with enhanced security checks. @@ -562,7 +567,7 @@ 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, *, session: 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: @@ -576,11 +581,11 @@ class AppAnnotationService: ) .order_by(AppAnnotationHitHistory.created_at.desc()) ) - annotation_hit_histories = paginate_query(stmt, page=page, per_page=limit, max_per_page=100) + annotation_hit_histories = paginate_query(stmt, session=session, page=page, per_page=limit, max_per_page=100) return annotation_hit_histories.items, annotation_hit_histories.total or 0 @classmethod - def get_annotation_by_id(cls, annotation_id: str, *, session: Session) -> MessageAnnotation | None: + def get_annotation_by_id(cls, annotation_id: str, session: Session) -> MessageAnnotation | None: annotation = session.get(MessageAnnotation, annotation_id) if not annotation: @@ -599,7 +604,6 @@ class AppAnnotationService: message_id: str, from_source: str, score: float, - *, session: Session, ) -> None: # add hit count to annotation @@ -621,11 +625,11 @@ class AppAnnotationService: annotation_content=annotation_content, ) session.add(annotation_hit_history) - session.commit() + session.flush() @classmethod def get_app_annotation_setting_by_app_id( - cls, app_id: str, *, session: Session + cls, app_id: str, session: Session ) -> AnnotationSettingDict | AnnotationSettingDisabledDict: _, current_tenant_id = current_account_with_tenant() # get app info @@ -640,7 +644,7 @@ class AppAnnotationService: select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_id).limit(1) ) if annotation_setting: - collection_binding_detail = annotation_setting.collection_binding_detail + collection_binding_detail = session.get(DatasetCollectionBinding, annotation_setting.collection_binding_id) if collection_binding_detail: return { "id": annotation_setting.id, @@ -662,7 +666,7 @@ class AppAnnotationService: @classmethod def update_app_annotation_setting( - cls, app_id: str, annotation_setting_id: str, args: UpdateAnnotationSettingArgs, *, session: Session + 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 @@ -687,9 +691,9 @@ class AppAnnotationService: annotation_setting.updated_user_id = current_user.id annotation_setting.updated_at = naive_utc_now() session.add(annotation_setting) - session.commit() + session.flush() - collection_binding_detail = annotation_setting.collection_binding_detail + collection_binding_detail = session.get(DatasetCollectionBinding, annotation_setting.collection_binding_id) if collection_binding_detail: return { @@ -710,7 +714,7 @@ class AppAnnotationService: } @classmethod - def clear_all_annotations(cls, app_id: str, *, session: Session): + def clear_all_annotations(cls, app_id: str, session: Session): _, current_tenant_id = current_account_with_tenant() app = session.scalar( select(App).where(App.id == app_id, App.tenant_id == current_tenant_id, App.status == "normal").limit(1) diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index 189d34b3f41..413cc53cc50 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -37,7 +37,7 @@ from graphon.nodes.question_classifier.entities import QuestionClassifierNodeDat from graphon.nodes.tool.entities import ToolNodeData from libs.datetime_utils import naive_utc_now from models import Account, App, AppMode -from models.model import AppModelConfig, AppModelConfigDict, IconType +from models.model import AppModelConfig, AppModelConfigDict, IconType, load_annotation_reply_config from models.workflow import Workflow from services.agent.dsl_service import AgentDslService, AgentPackage from services.agent.workflow_publish_service import WorkflowAgentPublishService @@ -457,7 +457,7 @@ class AppDslService: self._session.add(app) self._session.flush() - app_was_created.send(app, account=account) + app_was_created.send(app, account=account, session=self._session) # save dependencies if dependencies: @@ -537,7 +537,10 @@ class AppDslService: if not model_config or not isinstance(model_config, dict): raise ValueError("Missing model_config for chat/agent-chat/completion app") # Initialize or update model config - if not app.app_model_config: + app_model_config = ( + self._session.get(AppModelConfig, app.app_model_config_id) if app.app_model_config_id else None + ) + if not app_model_config: app_model_config = AppModelConfig( app_id=app.id, created_by=account.id, updated_by=account.id ).from_model_config_dict(cast(AppModelConfigDict, model_config)) @@ -545,9 +548,15 @@ class AppDslService: app.app_model_config_id = app_model_config.id self._session.add(app_model_config) - app_model_config_was_updated.send(app, app_model_config=app_model_config) + # Persist the config and app FK before receivers query them in this transaction. + self._session.flush() + app_model_config_was_updated.send( + app, + app_model_config=app_model_config, + session=self._session, + ) case AppMode.AGENT: - if app.app_model_config is not None: + if app.app_model_config_with_session(session=self._session) is not None: raise ValueError("Agent DSL import only supports creating a new Agent App") agent_data = data.get("agent") raw_agent_packages = data.get("agent_packages") @@ -579,6 +588,7 @@ class AppDslService: """ Export app :param app_model: App instance + :param session: Database session used to load export data :param include_secret: Whether include secret variable :return: """ @@ -621,7 +631,7 @@ class AppDslService: ) ] else: - cls._append_model_config_export_data(export_data, app_model) + cls._append_model_config_export_data(export_data, app_model, session=session) return yaml.dump(export_data, allow_unicode=True) @@ -701,17 +711,21 @@ class AppDslService: return status @classmethod - def _append_model_config_export_data(cls, export_data: dict[str, Any], app_model: App): + def _append_model_config_export_data(cls, export_data: dict[str, Any], app_model: App, *, session: Session) -> None: """ Append model config export data :param export_data: export data :param app_model: App instance + :param session: Database session used to load the model config and annotation reply """ - app_model_config = app_model.app_model_config + app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) if not app_model_config: raise ValueError("Missing app configuration, please check.") - model_config = app_model_config.to_dict() + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + model_config = app_model_config.to_dict(annotation_reply=annotation_reply) # TODO: refactor: we need a better way to filter workspace related data from model config # filter credential id from model config @@ -720,7 +734,7 @@ class AppDslService: export_data["model_config"] = model_config - dependencies = cls._extract_dependencies_from_model_config(app_model_config.to_dict()) + dependencies = cls._extract_dependencies_from_model_config(model_config) export_data["dependencies"] = [ jsonable_encoder(d.model_dump()) for d in DependenciesAnalysisService.generate_dependencies( diff --git a/api/services/app_generate_service.py b/api/services/app_generate_service.py index 7724555b615..80e6c574fd1 100644 --- a/api/services/app_generate_service.py +++ b/api/services/app_generate_service.py @@ -172,7 +172,9 @@ class AppGenerateService: request_id: str, ): effective_mode = ( - AppMode.AGENT_CHAT if app_model.is_agent and app_model.mode != AppMode.AGENT_CHAT else app_model.mode + AppMode.AGENT_CHAT + if app_model.is_agent_with_session(session=session) and app_model.mode != AppMode.AGENT_CHAT + else app_model.mode ) match effective_mode: case AppMode.COMPLETION: @@ -193,7 +195,12 @@ class AppGenerateService: return rate_limit.generate( AgentChatAppGenerator.convert_to_event_stream( AgentChatAppGenerator().generate( - app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=streaming + session=session, + app_model=app_model, + user=user, + args=args, + invoke_from=invoke_from, + streaming=streaming, ), ), request_id, @@ -202,7 +209,12 @@ class AppGenerateService: return rate_limit.generate( AgentAppGenerator.convert_to_event_stream( AgentAppGenerator().generate( - app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=streaming + app_model=app_model, + user=user, + args=args, + invoke_from=invoke_from, + session=session, + streaming=streaming, ), ), request_id, @@ -274,6 +286,7 @@ class AppGenerateService: workflow_run_id=str(uuid.uuid4()), streaming=False, pause_state_config=pause_config, + session=session, ) ), request_id=request_id, @@ -379,6 +392,7 @@ class AppGenerateService: user=user, args=args, streaming=streaming, + session=session, ) ) case AppMode.WORKFLOW: @@ -391,6 +405,7 @@ class AppGenerateService: user=user, args=args, streaming=streaming, + session=session, ) ) case AppMode.CHANNEL | AppMode.RAG_PIPELINE: @@ -422,6 +437,7 @@ class AppGenerateService: user=user, args=args, streaming=streaming, + session=session, ) ) case AppMode.WORKFLOW: @@ -434,6 +450,7 @@ class AppGenerateService: user=user, args=args, streaming=streaming, + session=session, ) ) case AppMode.CHANNEL | AppMode.RAG_PIPELINE: diff --git a/api/services/app_model_config_service.py b/api/services/app_model_config_service.py index ca42e76b00c..726083e1c52 100644 --- a/api/services/app_model_config_service.py +++ b/api/services/app_model_config_service.py @@ -1,5 +1,7 @@ from typing import Any +from sqlalchemy.orm import Session + from core.app.apps.agent_chat.app_config_manager import AgentChatAppConfigManager from core.app.apps.chat.app_config_manager import ChatAppConfigManager from core.app.apps.completion.app_config_manager import CompletionAppConfigManager @@ -8,14 +10,16 @@ from models.model import AppMode, AppModelConfigDict class AppModelConfigService: @classmethod - def validate_configuration(cls, tenant_id: str, config: dict[str, Any], app_mode: AppMode) -> AppModelConfigDict: + def validate_configuration( + cls, tenant_id: str, config: dict[str, Any], app_mode: AppMode, session: Session + ) -> AppModelConfigDict: match app_mode: case AppMode.CHAT: - return ChatAppConfigManager.config_validate(tenant_id, config) + return ChatAppConfigManager.config_validate(tenant_id, config, session) case AppMode.AGENT_CHAT: - return AgentChatAppConfigManager.config_validate(tenant_id, config) + return AgentChatAppConfigManager.config_validate(tenant_id, config, session) case AppMode.COMPLETION: - return CompletionAppConfigManager.config_validate(tenant_id, config) + return CompletionAppConfigManager.config_validate(tenant_id, config, session) case AppMode.WORKFLOW | AppMode.ADVANCED_CHAT | AppMode.CHANNEL | AppMode.RAG_PIPELINE | AppMode.AGENT: # Agent App presentation features go through AgentAppFeatureConfigService, # not this legacy EasyUI model-config validator. diff --git a/api/services/app_service.py b/api/services/app_service.py index dcca5393f6f..a6080879883 100644 --- a/api/services/app_service.py +++ b/api/services/app_service.py @@ -26,8 +26,9 @@ from libs.login import current_user from libs.pagination import PaginatedResult, paginate_query from models import Account, AppStar from models.agent import APP_BACKED_AGENT_SOURCES, Agent, AgentIconType, AgentScope, AgentStatus -from models.model import App, AppMode, AppModelConfig, IconType, Site +from models.model import App, AppMode, AppModelConfig, IconType, Site, load_annotation_reply_config from models.tools import ApiToolProvider +from models.workflow import Workflow from services.agent.errors import AgentNameConflictError from services.billing_service import BillingService from services.enterprise import rbac_service as enterprise_rbac_service @@ -77,6 +78,71 @@ class CreateAppParams(BaseModel): max_active_requests: int | None = None +class AppModelConfigResponseView: + """Expose AppModelConfig response properties through the request session.""" + + def __init__(self, app_model_config: AppModelConfig, *, session: Session) -> None: + self._app_model_config = app_model_config + self._session = session + + def __getattr__(self, name: str) -> Any: + return getattr(self._app_model_config, name) # noqa: no-new-getattr response adapter delegates model fields + + @property + def annotation_reply_dict(self) -> Any: + return load_annotation_reply_config(self._session, self._app_model_config.app_id) + + +class AppResponseView: + """Expose App response properties through one caller-owned database session.""" + + def __init__(self, app: App, *, session: Session) -> None: + self._app = app + self._session = session + + def __getattr__(self, name: str) -> Any: + return getattr(self._app, name) # noqa: no-new-getattr response adapter delegates model fields + + @property + def desc_or_prompt(self) -> str: + return self._app.desc_or_prompt_with_session(session=self._session) + + @property + def site(self) -> Site | None: + return self._app.site_with_session(session=self._session) + + @property + def app_model_config(self) -> AppModelConfigResponseView | None: + app_model_config = self._app.app_model_config_with_session(session=self._session) + if app_model_config is None: + return None + return AppModelConfigResponseView(app_model_config, session=self._session) + + @property + def workflow(self) -> Workflow | None: + return self._app.workflow_with_session(session=self._session) + + @property + def bound_agent_id(self) -> str | None: + return self._app.bound_agent_id_with_session(session=self._session) + + @property + def mode_compatible_with_agent(self) -> str: + return self._app.mode_compatible_with_agent_with_session(session=self._session) + + @property + def deleted_tools(self) -> list[Any]: + return self._app.deleted_tools_with_session(session=self._session) + + @property + def tags(self) -> Sequence[Any]: + return self._app.tags_with_session(session=self._session) + + @property + def author_name(self) -> str | None: + return self._app.author_name_with_session(session=self._session) + + class AppService: @staticmethod def _build_app_list_filters( @@ -153,7 +219,13 @@ class AppService: }[sort_by] @staticmethod - def get_starred_app_ids(*, tenant_id: str, account_id: str, app_ids: Sequence[str], session: Session) -> set[str]: + def get_starred_app_ids( + session: Session, + *, + tenant_id: str, + account_id: str, + app_ids: Sequence[str], + ) -> set[str]: """Return app IDs starred by this account within the tenant.""" if not app_ids: return set() @@ -168,24 +240,38 @@ class AppService: return set(starred_app_ids) @staticmethod - def get_app_by_id(app_id: str, *, session: Session) -> 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(app_id: str, *, session: Session) -> 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(app_ids: Sequence[str], *, session: Session) -> 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(*, name: str, tenant_id: str, session: Session) -> list[App]: + def find_visible_apps_by_name( + session: Session, + *, + name: str, + tenant_id: str, + ) -> list[App]: return list( session.execute( apply_openapi_gate( @@ -199,7 +285,11 @@ class AppService: ) def get_paginate_apps( - self, user_id: str, tenant_id: str, params: AppListParams, session: Session + self, + user_id: str, + tenant_id: str, + params: AppListParams, + session: Session, ) -> PaginatedResult | None: """ Get app list with pagination, filters, and explicit sort order. @@ -223,7 +313,10 @@ class AppService: app_ids = [str(app.id) for app in app_models.items] starred_app_ids = self.get_starred_app_ids( - tenant_id=tenant_id, account_id=user_id, app_ids=app_ids, session=session + session=session, + tenant_id=tenant_id, + account_id=user_id, + app_ids=app_ids, ) for app in app_models.items: app.is_starred = str(app.id) in starred_app_ids @@ -231,7 +324,11 @@ class AppService: return app_models def get_paginate_starred_apps( - self, user_id: str, tenant_id: str, params: StarredAppListParams, session: 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. @@ -422,9 +519,12 @@ class AppService: icon_background=params.icon_background, ) - session.commit() + session.flush() - app_was_created.send(app, account=account) + # Preserve the original commit-before-signal ordering for telemetry. + session.commit() + app_was_created.send(app, account=account, session=session) + session.commit() enterprise_rbac_service.try_sync_creator_access_policy_member_bindings( tenant_id, account.id, @@ -441,15 +541,15 @@ class AppService: return app - def get_app(self, app: App) -> App: + def get_app(self, app: App, *, session: Session) -> App: """ Get App """ assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None # get original app model config - if app.mode == AppMode.AGENT_CHAT or app.is_agent: - model_config = app.app_model_config + if app.mode == AppMode.AGENT_CHAT or app.is_agent_with_session(session=session): + model_config = app.app_model_config_with_session(session=session) if not model_config: return app agent_mode = model_config.agent_mode_dict @@ -773,6 +873,7 @@ class AppService: """ Get app meta info :param app_model: app model + :param session: database session :return: """ app_mode = AppMode.value_of(app_model.mode) @@ -780,7 +881,7 @@ class AppService: meta: dict[str, Any] = {"tool_icons": {}} if app_mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = session.get(Workflow, app_model.workflow_id) if app_model.workflow_id else None if workflow is None: return meta @@ -799,7 +900,9 @@ class AppService: } ) else: - app_model_config: AppModelConfig | None = app_model.app_model_config + app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) if not app_model_config: return meta diff --git a/api/services/audio_service.py b/api/services/audio_service.py index 597b23d412a..0d267e2277c 100644 --- a/api/services/audio_service.py +++ b/api/services/audio_service.py @@ -12,11 +12,10 @@ from werkzeug.datastructures import FileStorage from constants import AUDIO_EXTENSIONS from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features from core.model_manager import ModelManager -from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import ModelType from models.agent_config_entities import AgentSoulConfig from models.enums import MessageStatus -from models.model import App, AppMode, Message +from models.model import App, AppMode, Message, load_annotation_reply_config from services.agent.roster_service import AgentRosterService from services.app_ref_service import MessageRef from services.errors.audio import ( @@ -46,7 +45,14 @@ class AudioService: return session.scalar(stmt.limit(1)) @classmethod - def transcript_asr(cls, app_model: App, file: FileStorage | None, end_user: str | None = None) -> dict[str, str]: + def transcript_asr( + cls, + app_model: App, + file: FileStorage | None, + *, + session: Session, + end_user: str | None = None, + ) -> dict[str, str]: """Transcribe audio after enforcing the effective feature configuration. Published Agent Apps use their active Agent Soul. Historical Agent Apps @@ -56,7 +62,7 @@ class AudioService: SpeechToTextDisabledServiceError: If the effective feature configuration disables STT. """ if app_model.mode == AppMode.AGENT: - agent_soul = AgentRosterService(db.session).get_published_agent_soul_for_app( + agent_soul = AgentRosterService(session).get_published_agent_soul_for_app( tenant_id=app_model.tenant_id, app_id=app_model.id, ) @@ -65,11 +71,12 @@ class AudioService: app_model=app_model, agent_soul=agent_soul, file=file, + session=session, end_user=end_user, ) if app_model.mode in {AppMode.ADVANCED_CHAT, AppMode.WORKFLOW}: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if workflow is None: raise SpeechToTextDisabledServiceError() @@ -77,7 +84,7 @@ class AudioService: if "speech_to_text" not in features_dict or not features_dict["speech_to_text"].get("enabled"): raise SpeechToTextDisabledServiceError() else: - app_model_config = app_model.app_model_config + app_model_config = app_model.app_model_config_with_session(session=session) if not app_model_config: raise SpeechToTextDisabledServiceError() @@ -92,6 +99,8 @@ class AudioService: app_model: App, agent_soul: AgentSoulConfig, file: FileStorage | None, + *, + session: Session, end_user: str | None = None, ) -> dict[str, str]: """Transcribe Agent audio after applying Soul-first runtime feature projection. @@ -99,9 +108,12 @@ class AudioService: Raises: SpeechToTextDisabledServiceError: If the merged Agent feature configuration disables STT. """ + app_model_config = app_model.app_model_config_with_session(session=session) + annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None features = merge_agent_app_features( agent_soul=agent_soul, - app_model_config=app_model.app_model_config, + app_model_config=app_model_config, + annotation_reply=annotation_reply, ) if not features.get("speech_to_text", {}).get("enabled"): raise SpeechToTextDisabledServiceError() @@ -156,7 +168,7 @@ class AudioService: if is_draft: workflow = WorkflowService().get_draft_workflow(app_model=app_model, session=session) else: - workflow = app_model.workflow + workflow = app_model.workflow_with_session(session=session) if ( workflow is None or "text_to_speech" not in workflow.features_dict @@ -167,9 +179,10 @@ class AudioService: voice = workflow.features_dict["text_to_speech"].get("voice") else: if not is_draft: - if app_model.app_model_config is None: + app_model_config = app_model.app_model_config_with_session(session=session) + if app_model_config is None: raise ValueError("AppModelConfig not found") - text_to_speech_dict = app_model.app_model_config.text_to_speech_dict + text_to_speech_dict = app_model_config.text_to_speech_dict if not text_to_speech_dict.get("enabled"): raise ValueError("TTS is not enabled") diff --git a/api/services/data_migration/import_service.py b/api/services/data_migration/import_service.py index b3354413ba1..6b999119b2e 100644 --- a/api/services/data_migration/import_service.py +++ b/api/services/data_migration/import_service.py @@ -238,7 +238,7 @@ class MigrationImportService: raise MigrationDataError(f"Operator account not found: {target.operator_id}") if tenant is None: raise MigrationDataError(f"Target tenant not found: {target.tenant_id}") - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) for workflow_data in package.workflows: app_id = self._optional_string(workflow_data.get("id")) @@ -429,7 +429,7 @@ class MigrationImportService: api_token = ApiToken() api_token.app_id = app_id api_token.tenant_id = tenant_id - api_token.token = ApiToken.generate_api_key("app", 24) + api_token.token = ApiToken.generate_api_key("app", 24, session=session) api_token.type = ApiTokenType.APP session.add(api_token) session.commit() diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index dda5440f772..f595888ff9e 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -26,7 +26,6 @@ from core.rag.retrieval.retrieval_methods import RetrievalMethod from enums.cloud_plan import CloudPlan from events.dataset_event import dataset_was_deleted from events.document_event import document_was_deleted -from extensions.ext_database import db from extensions.ext_redis import redis_client from graphon.file import helpers as file_helpers from graphon.model_runtime.entities.model_entities import ModelFeature, ModelType @@ -247,8 +246,8 @@ class DatasetService: @staticmethod def get_datasets( - page, - per_page, + page: int, + per_page: int, session: Session, tenant_id=None, user=None, @@ -351,7 +350,7 @@ class DatasetService: else: return [], 0 - datasets = paginate_query(query, page=page, per_page=per_page, max_per_page=100) + datasets = paginate_query(query, session=session, page=page, per_page=per_page, max_per_page=100) return datasets.items, datasets.total @@ -374,14 +373,16 @@ class DatasetService: @staticmethod def get_datasets_by_ids( - ids, - tenant_id, + ids: list[str] | None, + tenant_id: str, user=None, accessible_dataset_ids: list[str] | None = None, include_own_datasets: bool = False, + *, + session: Session, ): # Check if ids is not empty to avoid WHERE false condition - if not ids or len(ids) == 0: + if not ids: return [], 0 stmt = select(Dataset).where(Dataset.id.in_(ids), Dataset.tenant_id == tenant_id) @@ -395,7 +396,7 @@ class DatasetService: accessible_filter = sa.or_(Dataset.maintainer == user.id, accessible_filter) stmt = stmt.where(accessible_filter) - datasets = paginate_query(stmt, page=1, per_page=len(ids), max_per_page=len(ids)) + datasets = paginate_query(stmt, session=session, page=1, per_page=len(ids), max_per_page=len(ids)) return datasets.items, datasets.total @@ -550,8 +551,14 @@ class DatasetService: return dataset @staticmethod - def check_doc_form(dataset: Dataset, doc_form: str): - if dataset.doc_form and doc_form != dataset.doc_form: + def get_dataset_for_tenant(dataset_id: str, tenant_id: str, *, session: Session) -> Dataset | None: + """Fetch a dataset only when it belongs to the provided tenant.""" + return session.scalar(select(Dataset).where(Dataset.id == dataset_id, Dataset.tenant_id == tenant_id).limit(1)) + + @staticmethod + def check_doc_form(dataset: Dataset, doc_form: str, *, session: Session): + dataset_doc_form = dataset.get_doc_form(session=session) + if dataset_doc_form and doc_form != dataset_doc_form: raise ValueError("doc_form is different from the dataset doc_form.") @staticmethod @@ -1397,10 +1404,10 @@ class DatasetService: raise NoPermissionError("You do not have permission to access this dataset.") @staticmethod - def get_dataset_queries(dataset_id: str, page: int, per_page: int): - stmt = select(DatasetQuery).filter_by(dataset_id=dataset_id).order_by(db.desc(DatasetQuery.created_at)) + def get_dataset_queries(dataset_id: str, page: int, per_page: int, session: Session): + stmt = select(DatasetQuery).filter_by(dataset_id=dataset_id).order_by(DatasetQuery.created_at.desc()) - dataset_queries = paginate_query(stmt, page=page, per_page=per_page, max_per_page=100) + dataset_queries = paginate_query(stmt, page=page, per_page=per_page, max_per_page=100, session=session) return dataset_queries.items, dataset_queries.total @@ -1422,7 +1429,7 @@ class DatasetService: raise ValueError("Current user or current user id not found") dataset.updated_by = current_user.id dataset.updated_at = naive_utc_now() - session.commit() + session.flush() @staticmethod def get_dataset_auto_disable_logs(dataset_id: str, session: Session) -> AutoDisableLogsDict: @@ -1870,6 +1877,7 @@ class DocumentService: @staticmethod def get_document_by_id(document_id: str, session: Session) -> Document | None: + """Fetch a document by primary key; callers must authorize its dataset before exposing it.""" document = session.get(Document, document_id) return document @@ -2027,7 +2035,7 @@ class DocumentService: .values(name=name) ) - session.commit() + session.flush() return document @@ -2131,7 +2139,7 @@ class DocumentService: session: Session, ) -> tuple[list[Document], str]: # check doc_form - DatasetService.check_doc_form(dataset, knowledge_config.doc_form) + DatasetService.check_doc_form(dataset, knowledge_config.doc_form, session=session) # check document limit assert isinstance(current_user, Account) assert current_user.current_tenant_id is not None @@ -2224,7 +2232,7 @@ class DocumentService: created_by=account.id, ) else: - dataset_process_rule = dataset.latest_process_rule + dataset_process_rule = dataset.get_latest_process_rule(session=session) if not dataset_process_rule: raise ValueError("No process rule found.") elif process_rule.mode == ProcessRuleMode.AUTOMATIC: @@ -2244,9 +2252,9 @@ class DocumentService: session.flush() else: # Fallback when no process_rule provided in knowledge_config: - # 1) reuse dataset.latest_process_rule if present + # 1) reuse the dataset's latest process rule if present # 2) otherwise create an automatic rule - dataset_process_rule = getattr(dataset, "latest_process_rule", None) + dataset_process_rule = dataset.get_latest_process_rule(session=session) if not dataset_process_rule: dataset_process_rule = DatasetProcessRule( dataset_id=dataset.id, @@ -2784,7 +2792,7 @@ class DocumentService: return document @staticmethod - def get_tenant_documents_count(*, session: Session): + def get_tenant_documents_count(session: Session): assert isinstance(current_user, Account) documents_count = ( @@ -3011,7 +3019,7 @@ class DocumentService: cut_name = documents[0].name[:cut_length] dataset.name = cut_name + "..." dataset.description = "useful for when you want to answer queries about the " + documents[0].name - session.commit() + session.flush() return dataset, documents, batch @@ -3398,11 +3406,7 @@ class SegmentService: 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, - session, + keywords_list, [segment_document], dataset, document.doc_form, session=session ) except Exception as e: logger.exception("create segment index failed") @@ -3491,11 +3495,7 @@ class SegmentService: try: # save vector index VectorService.create_segments_vector( - keywords_list, - pre_segment_data_list, - dataset, - document.doc_form, - session, + keywords_list, pre_segment_data_list, dataset, document.doc_form, session=session ) except Exception as e: logger.exception("create segment index failed") @@ -3594,18 +3594,12 @@ 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, - session, - True, + segment, document, dataset, embedding_model_instance, processing_rule, True, session=session ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): if args.enabled or keyword_changed: # update segment vector index - VectorService.update_segment_vector(args.keywords, segment, dataset) + VectorService.update_segment_vector(args.keywords, segment, dataset, session=session) # update summary index if summary is provided and has changed if args.summary is not None: # When user manually provides summary, allow saving even if summary_index_setting doesn't exist @@ -3705,17 +3699,11 @@ 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, - session, - True, + segment, document, dataset, embedding_model_instance, processing_rule, True, session=session ) elif document.doc_form in (IndexStructureType.PARAGRAPH_INDEX, IndexStructureType.QA_INDEX): # update segment vector index - VectorService.update_segment_vector(args.keywords, segment, dataset) + VectorService.update_segment_vector(args.keywords, segment, dataset, session=session) # Handle summary index when content changed if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: from models.dataset import DocumentSegmentSummary @@ -3795,7 +3783,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, session) + VectorService.update_multimodel_vector(segment, args.attachment_ids or [], dataset, session=session) except Exception as e: logger.exception("update segment index failed") segment.enabled = False @@ -4002,7 +3990,7 @@ class SegmentService: session.add(child_chunk) # save vector index try: - VectorService.create_child_chunk_vector(child_chunk, dataset) + VectorService.create_child_chunk_vector(child_chunk, dataset, session=session) except Exception as e: logger.exception("create child chunk index failed") session.rollback() @@ -4077,7 +4065,13 @@ class SegmentService: session.add(child_chunk) session.flush() new_child_chunks.append(child_chunk) - VectorService.update_child_chunk_vector(new_child_chunks, update_child_chunks, delete_child_chunks, dataset) + VectorService.update_child_chunk_vector( + new_child_chunks, + update_child_chunks, + delete_child_chunks, + dataset, + session=session, + ) session.commit() except Exception as e: logger.exception("update child chunk index failed") @@ -4104,7 +4098,7 @@ class SegmentService: child_chunk.updated_at = naive_utc_now() child_chunk.type = SegmentType.CUSTOMIZED session.add(child_chunk) - VectorService.update_child_chunk_vector([], [child_chunk], [], dataset) + VectorService.update_child_chunk_vector([], [child_chunk], [], dataset, session=session) session.commit() except Exception as e: logger.exception("update child chunk index failed") @@ -4116,7 +4110,7 @@ class SegmentService: 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) + VectorService.delete_child_chunk_vector(child_chunk, dataset, session=session) except Exception as e: logger.exception("delete child chunk index failed") session.rollback() @@ -4125,7 +4119,15 @@ class SegmentService: @classmethod def get_child_chunks( - cls, segment_id: str, document_id: str, dataset_id: str, page: int, limit: int, keyword: str | None = None + cls, + segment_id: str, + document_id: str, + dataset_id: str, + page: int, + limit: int, + keyword: str | None = None, + *, + session: Session, ): assert isinstance(current_user, Account) @@ -4142,7 +4144,7 @@ class SegmentService: if keyword: escaped_keyword = helper.escape_like_pattern(keyword) query = query.where(ChildChunk.content.ilike(f"%{escaped_keyword}%", escape="\\")) - return paginate_query(query, page=page, per_page=limit, max_per_page=100) + return paginate_query(query, session=session, 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: Session) -> ChildChunk | None: @@ -4179,6 +4181,8 @@ class SegmentService: keyword: str | None = None, page: int = 1, limit: int = 20, + *, + session: Session, ): """Get segments for a document with optional filtering.""" query = select(DocumentSegment).where( @@ -4194,7 +4198,7 @@ class SegmentService: query = query.where(DocumentSegment.content.ilike(f"%{escaped_keyword}%", escape="\\")) query = query.order_by(DocumentSegment.position.asc(), DocumentSegment.id.asc()) - paginated_segments = paginate_query(query, page=page, per_page=limit, max_per_page=100) + paginated_segments = paginate_query(query, session=session, page=page, per_page=limit, max_per_page=100) return paginated_segments.items, paginated_segments.total @@ -4282,7 +4286,7 @@ class DatasetCollectionBindingService: type=collection_type, ) session.add(dataset_collection_binding) - session.commit() + session.flush() return dataset_collection_binding @classmethod @@ -4316,22 +4320,18 @@ class DatasetPermissionService: @classmethod 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 = [] - for user in user_list: - permission = DatasetPermission( - tenant_id=tenant_id, - dataset_id=dataset_id, - account_id=user["user_id"], - ) - permissions.append(permission) + session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) + permissions = [] + for user in user_list: + permission = DatasetPermission( + tenant_id=tenant_id, + dataset_id=dataset_id, + account_id=user["user_id"], + ) + permissions.append(permission) - session.add_all(permissions) - session.commit() - except Exception as e: - session.rollback() - raise e + session.add_all(permissions) + session.flush() @classmethod def check_permission(cls, user, dataset, requested_permission, requested_partial_member_list, *, session: Session): @@ -4352,9 +4352,5 @@ class DatasetPermissionService: @classmethod def clear_partial_member_list(cls, dataset_id, session: Session): - try: - session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) - session.commit() - except Exception as e: - session.rollback() - raise e + session.execute(delete(DatasetPermission).where(DatasetPermission.dataset_id == dataset_id)) + session.flush() diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index cdd6c48342e..d069c259165 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -30,7 +30,7 @@ from services.errors.knowledge_retrieval import ExternalKnowledgeRetrievalError class ExternalDatasetService: @staticmethod def get_external_knowledge_apis( - page, per_page, tenant_id, search=None + page, per_page, tenant_id, search=None, *, session: Session ) -> tuple[list[ExternalKnowledgeApis], int | None]: query = ( select(ExternalKnowledgeApis) @@ -43,7 +43,7 @@ class ExternalDatasetService: escaped_search = escape_like_pattern(search) query = query.where(ExternalKnowledgeApis.name.ilike(f"%{escaped_search}%", escape="\\")) - external_knowledge_apis = paginate_query(query, page=page, per_page=per_page, max_per_page=100) + external_knowledge_apis = paginate_query(query, session=session, page=page, per_page=per_page, max_per_page=100) return external_knowledge_apis.items, external_knowledge_apis.total @@ -74,7 +74,7 @@ class ExternalDatasetService: ) session.add(external_knowledge_api) - session.commit() + session.flush() return external_knowledge_api @staticmethod @@ -142,7 +142,7 @@ class ExternalDatasetService: 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() + session.flush() return external_knowledge_api @@ -157,7 +157,7 @@ class ExternalDatasetService: raise ValueError("api template not found") session.delete(external_knowledge_api) - session.commit() + session.flush() @staticmethod def external_knowledge_api_use_check( diff --git a/api/services/feedback_service.py b/api/services/feedback_service.py index 62885c901b7..6e60a026ce2 100644 --- a/api/services/feedback_service.py +++ b/api/services/feedback_service.py @@ -90,7 +90,7 @@ class FeedbackService: export_data = [] for feedback, message, conversation, app, account in results: # Get the user query from the message - user_query = message.query or (message.inputs.get("query", "") if message.inputs else "") + user_query = message.query or message.inputs_with_session(session=session).get("query", "") # Format the feedback data feedback_record = { diff --git a/api/services/hit_testing_service.py b/api/services/hit_testing_service.py index 1bfa4025fa0..2f9d5c3677b 100644 --- a/api/services/hit_testing_service.py +++ b/api/services/hit_testing_service.py @@ -235,7 +235,8 @@ class HitTestingService: def compact_retrieve_response( cls, query: str, documents: list[Document], *, session: Session ) -> RetrieveResponseDict: - records = RetrievalService.format_retrieval_documents(documents) + with Session(bind=session.get_bind()) as format_session: + records = RetrievalService.format_retrieval_documents(format_session, documents) return { "query": { diff --git a/api/services/human_input_file_upload_service.py b/api/services/human_input_file_upload_service.py index 3f502c2a9e7..00b4230870c 100644 --- a/api/services/human_input_file_upload_service.py +++ b/api/services/human_input_file_upload_service.py @@ -192,7 +192,7 @@ class HumanInputFileUploadService: # HITL upload runs outside the normal account auth flow, so hydrate the # account tenant context explicitly before delegating to FileService. - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) return account def _resolve_delivery_test_upload_owner( @@ -220,7 +220,7 @@ class HumanInputFileUploadService: if tenant is None: raise InvalidUploadTokenError() - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) if account.current_tenant_id != form_model.tenant_id: raise InvalidUploadTokenError() return account diff --git a/api/services/message_service.py b/api/services/message_service.py index 4fbeb61e1f7..a9658244be2 100644 --- a/api/services/message_service.py +++ b/api/services/message_service.py @@ -192,7 +192,11 @@ class MessageService: 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 + feedback = ( + message.user_feedback_with_session(session=session) + if isinstance(user, EndUser) + else message.admin_feedback_with_session(session=session) + ) if not rating and feedback: session.delete(feedback) @@ -310,7 +314,9 @@ class MessageService: ) # Reuse Conversation.model_config so suggested-questions reads the same # compatibility-normalized config as the rest of the message flow. - app_model_config = app_model_config.from_model_config_dict(conversation.model_config) + app_model_config = app_model_config.from_model_config_dict( + conversation.model_config_with_session(session=session) + ) if not app_model_config: raise ValueError("did not find app model config") diff --git a/api/services/metadata_service.py b/api/services/metadata_service.py index 481eb3b2e29..004a059acc6 100644 --- a/api/services/metadata_service.py +++ b/api/services/metadata_service.py @@ -56,7 +56,7 @@ class MetadataService: created_by=current_user.id, ) session.add(metadata) - session.commit() + session.flush() return metadata @staticmethod @@ -128,7 +128,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def delete_metadata(dataset_id: str, metadata_id: str, *, session: Session): + 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) @@ -174,7 +174,7 @@ class MetadataService: ] @staticmethod - def enable_built_in_field(dataset: Dataset, *, session: Session): + def enable_built_in_field(dataset: Dataset, session: Session): if dataset.built_in_field_enabled: return lock_key = f"dataset_metadata_lock_{dataset.id}" @@ -189,7 +189,7 @@ class MetadataService: else: doc_metadata = copy.deepcopy(document.doc_metadata) doc_metadata[BuiltInField.document_name] = document.name - doc_metadata[BuiltInField.uploader] = document.uploader + doc_metadata[BuiltInField.uploader] = document.get_uploader(session=session) doc_metadata[BuiltInField.upload_date] = document.upload_date.timestamp() doc_metadata[BuiltInField.last_update_date] = document.last_update_date.timestamp() doc_metadata[BuiltInField.source] = MetadataDataSource[document.data_source_type] @@ -203,7 +203,7 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def disable_built_in_field(dataset: Dataset, *, session: Session): + 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}" @@ -260,7 +260,7 @@ class MetadataService: doc_metadata[metadata_value.name] = metadata_value.value if dataset.built_in_field_enabled: doc_metadata[BuiltInField.document_name] = document.name - doc_metadata[BuiltInField.uploader] = document.uploader + doc_metadata[BuiltInField.uploader] = document.get_uploader(session=session) doc_metadata[BuiltInField.upload_date] = document.upload_date.timestamp() doc_metadata[BuiltInField.last_update_date] = document.last_update_date.timestamp() doc_metadata[BuiltInField.source] = MetadataDataSource[document.data_source_type] @@ -319,7 +319,7 @@ class MetadataService: redis_client.set(lock_key, 1, ex=3600) @staticmethod - def get_dataset_metadatas(dataset: Dataset, *, session: Session): + def get_dataset_metadatas(dataset: Dataset, session: Session): return { "doc_metadata": [ { @@ -334,7 +334,7 @@ class MetadataService: ) or 0, } - for item in dataset.doc_metadata or [] + for item in dataset.get_doc_metadata(session=session) if item.get("id") != "built-in" ], "built_in_field_enabled": dataset.built_in_field_enabled, diff --git a/api/services/plugin/plugin_migration.py b/api/services/plugin/plugin_migration.py index 5c2ddb77e0f..d0fc853456f 100644 --- a/api/services/plugin/plugin_migration.py +++ b/api/services/plugin/plugin_migration.py @@ -282,7 +282,9 @@ class PluginMigration: return [] agent_app_model_config_ids = [ - app.app_model_config_id for app in apps if app.is_agent or app.mode == AppMode.AGENT_CHAT + app.app_model_config_id + for app in apps + if app.is_agent_with_session(session=session) or app.mode == AppMode.AGENT_CHAT ] rs = session.scalars(select(AppModelConfig).where(AppModelConfig.id.in_(agent_app_model_config_ids))).all() @@ -455,7 +457,9 @@ class PluginMigration: ) @classmethod - def install_rag_pipeline_plugins(cls, extracted_plugins: str, output_file: str, workers: int = 100) -> None: + def install_rag_pipeline_plugins( + cls, extracted_plugins: str, output_file: str, workers: int = 100, *, session: Session + ) -> None: """ Install rag pipeline plugins. """ @@ -510,7 +514,9 @@ class PluginMigration: total_failed_tenant = 0 while True: # paginate - tenants = paginate_query(sa.select(Tenant).order_by(Tenant.created_at.desc()), page=page, per_page=100) + tenants = paginate_query( + sa.select(Tenant).order_by(Tenant.created_at.desc()), session=session, page=page, per_page=100 + ) if tenants.items is None or len(tenants.items) == 0: break diff --git a/api/services/rag_pipeline/pipeline_generate_service.py b/api/services/rag_pipeline/pipeline_generate_service.py index 276bfaea158..d7ef2205380 100644 --- a/api/services/rag_pipeline/pipeline_generate_service.py +++ b/api/services/rag_pipeline/pipeline_generate_service.py @@ -41,6 +41,7 @@ class PipelineGenerateService: cls.update_document_status(original_document_id, session=session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().generate( + session=session, pipeline=pipeline, workflow=workflow, user=user, @@ -70,7 +71,13 @@ class PipelineGenerateService: 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 + pipeline=pipeline, + workflow=workflow, + node_id=node_id, + user=user, + args=args, + streaming=streaming, + session=session, ) ) @@ -81,7 +88,13 @@ class PipelineGenerateService: 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 + pipeline=pipeline, + workflow=workflow, + node_id=node_id, + user=user, + args=args, + streaming=streaming, + session=session, ) ) diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index d59ac19b593..3e8340c8a11 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -111,6 +111,12 @@ class RagPipelineService: ) self._workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker) + @staticmethod + def get_pipeline_by_id(pipeline_id: str, tenant_id: str, *, session: Session) -> Pipeline | None: + return session.scalar( + select(Pipeline).where(Pipeline.id == pipeline_id, Pipeline.tenant_id == tenant_id).limit(1) + ) + @classmethod def get_pipeline_templates( cls, @@ -1466,6 +1472,7 @@ class RagPipelineService: if not workflow: raise ValueError("Workflow not found") PipelineGenerator().generate( + session=self._session, pipeline=pipeline, workflow=workflow, user=user, diff --git a/api/services/snippet_generate_service.py b/api/services/snippet_generate_service.py index f0cd4d901c3..6fd70de233d 100644 --- a/api/services/snippet_generate_service.py +++ b/api/services/snippet_generate_service.py @@ -376,7 +376,8 @@ class SnippetGenerateService: node_id: str, args: Mapping[str, Any], streaming: bool = True, - session_maker: sessionmaker[Session] | None = None, + *, + session_maker: sessionmaker[Session], ) -> Mapping[str, Any] | Generator[str, None, None]: """ Run a single iteration node in a snippet's draft workflow. @@ -390,6 +391,7 @@ class SnippetGenerateService: :param node_id: ID of the iteration node to run :param args: Dict containing 'inputs' key with iteration input data :param streaming: Whether to stream the response (should be True) + :param session_maker: Factory for the synchronous database work before worker startup :return: SSE streaming generator :raises ValueError: If the snippet has no draft workflow """ @@ -400,16 +402,18 @@ class SnippetGenerateService: app_proxy = cast(App, _SnippetAsApp(snippet)) - return WorkflowAppGenerator.convert_to_event_stream( - WorkflowAppGenerator().single_iteration_generate( - app_model=app_proxy, - workflow=workflow, - node_id=node_id, - user=user, - args=args, - streaming=streaming, + with session_maker() as session: + return WorkflowAppGenerator.convert_to_event_stream( + WorkflowAppGenerator().single_iteration_generate( + app_model=app_proxy, + workflow=workflow, + node_id=node_id, + user=user, + args=args, + streaming=streaming, + session=session, + ) ) - ) @classmethod def generate_single_loop( @@ -419,7 +423,8 @@ class SnippetGenerateService: node_id: str, args: Any, streaming: bool = True, - session_maker: sessionmaker[Session] | None = None, + *, + session_maker: sessionmaker[Session], ) -> Mapping[str, Any] | Generator[str, None, None]: """ Run a single loop node in a snippet's draft workflow. @@ -433,6 +438,7 @@ class SnippetGenerateService: :param node_id: ID of the loop node to run :param args: Pydantic model with 'inputs' attribute containing loop input data :param streaming: Whether to stream the response (should be True) + :param session_maker: Factory for the synchronous database work before worker startup :return: SSE streaming generator :raises ValueError: If the snippet has no draft workflow """ @@ -443,16 +449,18 @@ class SnippetGenerateService: app_proxy = cast(App, _SnippetAsApp(snippet)) - return WorkflowAppGenerator.convert_to_event_stream( - WorkflowAppGenerator().single_loop_generate( - app_model=app_proxy, - workflow=workflow, - node_id=node_id, - user=user, - args=args, # type: ignore[arg-type] - streaming=streaming, + with session_maker() as session: + return WorkflowAppGenerator.convert_to_event_stream( + WorkflowAppGenerator().single_loop_generate( + app_model=app_proxy, + workflow=workflow, + node_id=node_id, + user=user, + args=args, # type: ignore[arg-type] + streaming=streaming, + session=session, + ) ) - ) @staticmethod def parse_files(workflow: Workflow, files: list[dict] | None = None) -> Sequence[File]: diff --git a/api/services/summary_index_service.py b/api/services/summary_index_service.py index 3e8c0ae8340..d82fca196a1 100644 --- a/api/services/summary_index_service.py +++ b/api/services/summary_index_service.py @@ -50,6 +50,8 @@ class SummaryIndexService: segment: DocumentSegment, dataset: Dataset, summary_index_setting: SummaryIndexSettingDict, + *, + session: Session, ) -> tuple[str, LLMUsage]: """ Generate summary for a single segment. @@ -59,6 +61,9 @@ class SummaryIndexService: dataset: Dataset containing the segment summary_index_setting: Summary index configuration + Keyword Args: + session: SQLAlchemy session used to load the segment's document. + Returns: Tuple of (summary_content, llm_usage) where llm_usage is LLMUsage object @@ -72,8 +77,9 @@ class SummaryIndexService: # Get document language to ensure summary is generated in the correct language # This is especially important for image-only chunks where text is empty or minimal document_language = None - if segment.document and segment.document.doc_language: - document_language = segment.document.doc_language + document = segment.get_document(session=session) + if document and document.doc_language: + document_language = document.doc_language summary_content, usage = ParagraphIndexProcessor.generate_summary( tenant_id=dataset.tenant_id, @@ -81,6 +87,7 @@ class SummaryIndexService: summary_index_setting=summary_index_setting, segment_id=segment.id, document_language=document_language, + session=session, ) if not summary_content: @@ -178,6 +185,13 @@ class SummaryIndexService: summary_record_id = summary_record.id # Save the original session parameter for use in error handling original_session = session + + def create_vector() -> Vector: + if original_session is not None: + return Vector(dataset, session=original_session) + with session_factory.create_session() as vector_session: + return Vector(dataset, session=vector_session) + logger.debug( "Starting vectorization for segment %s, summary_record_id=%s, using_provided_session=%s", segment.id, @@ -206,7 +220,7 @@ class SummaryIndexService: # If index_node_id changed, the old vector should have been deleted elsewhere if old_summary_node_id and old_summary_node_id == summary_index_node_id: try: - vector = Vector(dataset) + vector = create_vector() vector.delete_by_ids([old_summary_node_id]) except Exception as e: logger.warning( @@ -258,7 +272,7 @@ class SummaryIndexService: attempt + 1, max_retries, ) - vector = Vector(dataset) + vector = create_vector() # Use duplicate_check=False to ensure re-vectorization even if old vector still exists # The old vector should have been deleted above, but if deletion failed, # we still want to re-vectorize (upsert will overwrite) @@ -695,11 +709,12 @@ class SummaryIndexService: 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 + # Persist GENERATING and release the write transaction before LLM I/O. + session.commit() # Generate summary (returns summary_content and llm_usage) summary_content, llm_usage = SummaryIndexService.generate_summary_for_segment( - segment, dataset, summary_index_setting + segment, dataset, summary_index_setting, session=session ) # Update summary content @@ -718,24 +733,16 @@ class SummaryIndexService: llm_usage.completion_tokens, ) - 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 + 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 e: logger.exception("Failed to generate summary for segment %s", segment.id) + session.rollback() # Update summary record with error status summary_record_in_session = session.scalar( select(DocumentSegmentSummary) @@ -905,7 +912,7 @@ class SummaryIndexService: summary_node_ids = [s.summary_index_node_id for s in summaries if s.summary_index_node_id] if summary_node_ids: try: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.delete_by_ids(summary_node_ids) except Exception as e: logger.warning("Failed to remove summary vectors: %s", str(e)) @@ -1009,6 +1016,8 @@ class SummaryIndexService: def delete_summaries_for_segments( dataset: Dataset, segment_ids: list[str] | None = None, + *, + session: Session, ) -> None: """ Delete summary records and vectors for segments (used only for actual deletion scenarios). @@ -1018,30 +1027,29 @@ class SummaryIndexService: dataset: Dataset containing the segments segment_ids: List of segment IDs to delete summaries for. If None, delete all. """ - with session_factory.create_session() as session: - stmt = select(DocumentSegmentSummary).where(DocumentSegmentSummary.dataset_id == dataset.id) + stmt = select(DocumentSegmentSummary).where(DocumentSegmentSummary.dataset_id == dataset.id) - if segment_ids: - stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) + if segment_ids: + stmt = stmt.where(DocumentSegmentSummary.chunk_id.in_(segment_ids)) - summaries = session.scalars(stmt).all() + summaries = session.scalars(stmt).all() - if not summaries: - return + if not summaries: + return # Delete from vector database - if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: - summary_node_ids = [s.summary_index_node_id for s in summaries if s.summary_index_node_id] - if summary_node_ids: - vector = Vector(dataset) - vector.delete_by_ids(summary_node_ids) + if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: + summary_node_ids = [s.summary_index_node_id for s in summaries if s.summary_index_node_id] + if summary_node_ids: + vector = Vector(dataset, session=session) + vector.delete_by_ids(summary_node_ids) # Delete summary records - for summary in summaries: - session.delete(summary) + for summary in summaries: + session.delete(summary) - session.commit() - logger.info("Deleted %s summary records for dataset %s", len(summaries), dataset.id) + session.flush() + logger.info("Deleted %s summary records for dataset %s", len(summaries), dataset.id) @staticmethod def update_summary_for_segment( @@ -1074,7 +1082,8 @@ class SummaryIndexService: # Vectorization uses dataset.embedding_model, which doesn't require summary_index_setting # Skip qa_model documents - if segment.document and segment.document.doc_form == "qa_model": + document = segment.get_document(session=session) + if document and document.doc_form == "qa_model": return None try: @@ -1095,7 +1104,7 @@ class SummaryIndexService: old_summary_node_id = summary_record.summary_index_node_id if old_summary_node_id: try: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.delete_by_ids([old_summary_node_id]) except Exception as e: logger.warning( @@ -1129,7 +1138,7 @@ class SummaryIndexService: # Delete old vector if exists (before vectorization) if old_summary_node_id: try: - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.delete_by_ids([old_summary_node_id]) except Exception as e: logger.warning( @@ -1157,14 +1166,19 @@ class SummaryIndexService: 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. + if not session.is_active: + session.rollback() + summary_record.summary_content = summary_content summary_record.status = SummaryStatus.ERROR summary_record.error = f"Vectorization failed: {str(e)}" + session.add(summary_record) 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) + session.rollback() # Update summary record with error status if it exists summary_record = session.scalar( select(DocumentSegmentSummary) diff --git a/api/services/vector_service.py b/api/services/vector_service.py index faf4fb085d6..2ecd0d4f2e8 100644 --- a/api/services/vector_service.py +++ b/api/services/vector_service.py @@ -74,8 +74,8 @@ class VectorService: dataset, embedding_model_instance, processing_rule, - session, False, + session=session, ) else: rag_document = Document( @@ -90,7 +90,7 @@ class VectorService: ) documents.append(rag_document) if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodal_document: AttachmentDocument = AttachmentDocument( page_content=attachment["name"], metadata={ @@ -105,12 +105,16 @@ class VectorService: index_processor: BaseIndexProcessor = IndexProcessorFactory(doc_form).init_index_processor() if len(documents) > 0: - index_processor.load(dataset, documents, None, with_keywords=True, keywords_list=keywords_list) + index_processor.load( + dataset, documents, None, with_keywords=True, keywords_list=keywords_list, session=session + ) if len(multimodal_documents) > 0: - index_processor.load(dataset, [], multimodal_documents, with_keywords=False) + index_processor.load(dataset, [], multimodal_documents, with_keywords=False, session=session) @classmethod - def update_segment_vector(cls, keywords: list[str] | None, segment: DocumentSegment, dataset: Dataset): + def update_segment_vector( + cls, keywords: list[str] | None, segment: DocumentSegment, dataset: Dataset, session: Session + ): # update segment index task # format new index @@ -126,19 +130,19 @@ class VectorService: assert segment.index_node_id if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: # update vector index - vector = Vector(dataset=dataset) + vector = Vector(dataset=dataset, session=session) vector.delete_by_ids([segment.index_node_id]) vector.add_texts([document], duplicate_check=True) else: # update keyword index keyword = Keyword(dataset) - keyword.delete_by_ids([segment.index_node_id]) + keyword.delete_by_ids([segment.index_node_id], session) # save keyword index if keywords and len(keywords) > 0: - keyword.add_texts([document], keywords_list=[keywords]) + keyword.add_texts([document], session, keywords_list=[keywords]) else: - keyword.add_texts([document]) + keyword.add_texts([document], session) @classmethod def generate_child_chunks( @@ -148,15 +152,22 @@ class VectorService: dataset: Dataset, embedding_model_instance: ModelInstance, processing_rule: DatasetProcessRule, - session: Session, regenerate: bool = False, + *, + session: Session, ): """Generate child chunks and persist them with the caller's active DB session.""" - index_processor = IndexProcessorFactory(dataset.doc_form).init_index_processor() + index_processor = IndexProcessorFactory(dataset.get_doc_form(session=session)).init_index_processor() assert segment.index_node_id if regenerate: # delete child chunks - index_processor.clean(dataset, [segment.index_node_id], with_keywords=True, delete_child_chunks=True) + index_processor.clean( + dataset, + [segment.index_node_id], + with_keywords=True, + delete_child_chunks=True, + session=session, + ) # generate child chunks document = Document( @@ -179,10 +190,11 @@ class VectorService: process_rule=processing_rule_dict, tenant_id=dataset.tenant_id, doc_language=dataset_document.doc_language, + session=session, ) # save child chunks if documents and documents[0].children: - index_processor.load(dataset, documents) + index_processor.load(dataset, documents, session=session) for position, child_chunk in enumerate(documents[0].children, start=1): child_segment = ChildChunk( @@ -199,10 +211,10 @@ class VectorService: created_by=dataset_document.created_by, ) session.add(child_segment) - session.commit() + session.flush() @classmethod - def create_child_chunk_vector(cls, child_segment: ChildChunk, dataset: Dataset): + def create_child_chunk_vector(cls, child_segment: ChildChunk, dataset: Dataset, *, session: Session): child_document = Document( page_content=child_segment.content, metadata={ @@ -214,7 +226,7 @@ class VectorService: ) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: # save vector index - vector = Vector(dataset=dataset) + vector = Vector(dataset=dataset, session=session) vector.add_texts([child_document], duplicate_check=True) @classmethod @@ -224,6 +236,8 @@ class VectorService: update_child_chunks: list[ChildChunk], delete_child_chunks: list[ChildChunk], dataset: Dataset, + *, + session: Session, ): documents = [] delete_node_ids = [] @@ -256,15 +270,15 @@ class VectorService: delete_node_ids.append(delete_child_chunk.index_node_id) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: # update vector index - vector = Vector(dataset=dataset) + vector = Vector(dataset=dataset, session=session) if delete_node_ids: vector.delete_by_ids(delete_node_ids) if documents: vector.add_texts(documents, duplicate_check=True) @classmethod - def delete_child_chunk_vector(cls, child_chunk: ChildChunk, dataset: Dataset): - vector = Vector(dataset=dataset) + def delete_child_chunk_vector(cls, child_chunk: ChildChunk, dataset: Dataset, *, session: Session): + vector = Vector(dataset=dataset, session=session) assert child_chunk.index_node_id vector.delete_by_ids([child_chunk.index_node_id]) @@ -272,11 +286,10 @@ class VectorService: 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 - attachments = segment.attachments + attachments = segment.get_attachments(session=session) old_attachment_ids = [attachment["id"] for attachment in attachments] if attachments else [] # Check if there's any actual change needed @@ -284,7 +297,7 @@ class VectorService: return try: - vector = Vector(dataset=dataset) + vector = Vector(dataset=dataset, session=session) if dataset.is_multimodal: # Delete old vectors if they exist if old_attachment_ids: @@ -294,14 +307,14 @@ class VectorService: session.execute(delete(SegmentAttachmentBinding).where(SegmentAttachmentBinding.segment_id == segment.id)) if not attachment_ids: - session.commit() + session.flush() return # Bulk fetch upload files - only fetch needed fields upload_file_list = session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all() if not upload_file_list: - session.commit() + session.flush() return # Create a mapping for quick lookup @@ -350,8 +363,7 @@ class VectorService: if documents and dataset.is_multimodal: vector.create_multimodal(documents) - # Single commit for all operations - session.commit() + session.flush() except Exception: logger.exception("Failed to update multimodal vector for segment %s", segment.id) diff --git a/api/services/workflow/workflow_converter.py b/api/services/workflow/workflow_converter.py index 5f787bb51cd..c9d1d0fa46c 100644 --- a/api/services/workflow/workflow_converter.py +++ b/api/services/workflow/workflow_converter.py @@ -26,7 +26,7 @@ from graphon.nodes import BuiltinNodeTypes from graphon.variables.input_entities import VariableEntity from models import Account from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint -from models.model import App, AppMode, AppModelConfig, IconType +from models.model import App, AppMode, AppModelConfig, IconType, load_annotation_reply_config from models.workflow import Workflow, WorkflowType @@ -80,11 +80,14 @@ class WorkflowConverter: :return: new App instance """ # convert app model config - if not app_model.app_model_config: + app_model_config = ( + session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None + ) + if not app_model_config: 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, session=session + app_model=app_model, app_model_config=app_model_config, account_id=account.id, session=session ) # create new app @@ -110,7 +113,8 @@ class WorkflowConverter: workflow.app_id = new_app.id session.commit() - app_was_created.send(new_app, account=account) + app_was_created.send(new_app, account=account, session=session) + session.commit() return new_app @@ -127,7 +131,9 @@ class WorkflowConverter: new_app_mode = self._get_new_app_mode(app_model) # convert app model config - app_config = self._convert_to_app_config(app_model=app_model, app_model_config=app_model_config) + app_config = self._convert_to_app_config( + app_model=app_model, app_model_config=app_model_config, session=session + ) # init workflow graph graph: WorkflowGraph = {"nodes": [], "edges": []} @@ -232,23 +238,38 @@ class WorkflowConverter: return workflow - def _convert_to_app_config(self, app_model: App, app_model_config: AppModelConfig) -> EasyUIBasedAppConfig: + def _convert_to_app_config( + self, app_model: App, app_model_config: AppModelConfig, *, session: Session + ) -> EasyUIBasedAppConfig: app_mode_enum = AppMode.value_of(app_model.mode) app_config: EasyUIBasedAppConfig effective_mode = ( - AppMode.AGENT_CHAT if app_model.is_agent and app_mode_enum != AppMode.AGENT_CHAT else app_mode_enum + AppMode.AGENT_CHAT + if app_model.is_agent_with_session(session=session) and app_mode_enum != AppMode.AGENT_CHAT + else app_mode_enum ) match effective_mode: case AppMode.AGENT_CHAT: app_model.mode = AppMode.AGENT_CHAT + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) app_config = AgentChatAppConfigManager.get_app_config( - app_model=app_model, app_model_config=app_model_config + app_model=app_model, + app_model_config=app_model_config, + annotation_reply=annotation_reply, ) case AppMode.CHAT: - app_config = ChatAppConfigManager.get_app_config(app_model=app_model, app_model_config=app_model_config) + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) + app_config = ChatAppConfigManager.get_app_config( + app_model=app_model, + app_model_config=app_model_config, + annotation_reply=annotation_reply, + ) case AppMode.COMPLETION: + annotation_reply = load_annotation_reply_config(session, app_model_config.app_id) app_config = CompletionAppConfigManager.get_app_config( - app_model=app_model, app_model_config=app_model_config + app_model=app_model, + app_model_config=app_model_config, + annotation_reply=annotation_reply, ) case _: raise ValueError("Invalid app mode") diff --git a/api/services/workflow_run_service.py b/api/services/workflow_run_service.py index 0c5c0876b19..8618b343ae6 100644 --- a/api/services/workflow_run_service.py +++ b/api/services/workflow_run_service.py @@ -82,12 +82,13 @@ class WorkflowRunService: run_ids = [workflow_run.id for workflow_run in workflow_runs] messages_by_run_id: dict[str, Message] = {} if run_ids: - messages = db.session.scalars( - select(Message).where( - Message.app_id == app_model.id, - Message.workflow_run_id.in_(run_ids), - ) - ).all() + with self._session_factory() as session: + messages = session.scalars( + select(Message).where( + Message.app_id == app_model.id, + Message.workflow_run_id.in_(run_ids), + ) + ).all() for loaded_message in messages: run_id = loaded_message.workflow_run_id if run_id is None: diff --git a/api/tasks/add_document_to_index_task.py b/api/tasks/add_document_to_index_task.py index c9d4673c0ad..336fc1000f4 100644 --- a/api/tasks/add_document_to_index_task.py +++ b/api/tasks/add_document_to_index_task.py @@ -44,7 +44,7 @@ def add_document_to_index_task(dataset_document_id: str): indexing_cache_key = f"document_{dataset_document.id}_indexing" try: - dataset = dataset_document.dataset + dataset = dataset_document.get_dataset(session=session) if not dataset: raise Exception(f"Document {dataset_document.id} dataset {dataset_document.dataset_id} doesn't exist.") @@ -70,7 +70,7 @@ def add_document_to_index_task(dataset_document_id: str): }, ) if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -86,7 +86,7 @@ def add_document_to_index_task(dataset_document_id: str): child_documents.append(child_document) document.children = child_documents if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodal_documents.append( AttachmentDocument( page_content=attachment["name"], @@ -101,9 +101,9 @@ def add_document_to_index_task(dataset_document_id: str): ) documents.append(document) - index_type = dataset.doc_form + index_type = dataset.get_doc_form(session=session) index_processor = IndexProcessorFactory(index_type).init_index_processor() - index_processor.load(dataset, documents, multimodal_documents=multimodal_documents) + index_processor.load(dataset, documents, multimodal_documents=multimodal_documents, session=session) # delete auto disable log session.execute( diff --git a/api/tasks/annotation/add_annotation_to_index_task.py b/api/tasks/annotation/add_annotation_to_index_task.py index 81f57d7d6f3..9e0a7d83cc1 100644 --- a/api/tasks/annotation/add_annotation_to_index_task.py +++ b/api/tasks/annotation/add_annotation_to_index_task.py @@ -48,7 +48,8 @@ def add_annotation_to_index_task( document = Document( page_content=question, metadata={"annotation_id": annotation_id, "app_id": app_id, "doc_id": annotation_id} ) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + with session_factory.create_session() as session: + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) vector.create([document], duplicate_check=True) end_at = time.perf_counter() diff --git a/api/tasks/annotation/batch_import_annotations_task.py b/api/tasks/annotation/batch_import_annotations_task.py index 343b8ee2fa3..68f086df8a2 100644 --- a/api/tasks/annotation/batch_import_annotations_task.py +++ b/api/tasks/annotation/batch_import_annotations_task.py @@ -77,7 +77,7 @@ def batch_import_annotations_task(job_id: str, content_list: list[dict], app_id: collection_binding_id=dataset_collection_binding.id, ) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) vector.create(documents, duplicate_check=True) session.commit() diff --git a/api/tasks/annotation/delete_annotation_index_task.py b/api/tasks/annotation/delete_annotation_index_task.py index 79a8db2548b..4e467e2e0ba 100644 --- a/api/tasks/annotation/delete_annotation_index_task.py +++ b/api/tasks/annotation/delete_annotation_index_task.py @@ -34,7 +34,8 @@ def delete_annotation_index_task(annotation_id: str, app_id: str, tenant_id: str ) try: - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + with session_factory.create_session() as session: + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) vector.delete_by_metadata_field("annotation_id", annotation_id) except Exception: logger.exception("Delete annotation index failed when annotation deleted.") diff --git a/api/tasks/annotation/disable_annotation_reply_task.py b/api/tasks/annotation/disable_annotation_reply_task.py index 6a9b52e7e5b..31ffd72a575 100644 --- a/api/tasks/annotation/disable_annotation_reply_task.py +++ b/api/tasks/annotation/disable_annotation_reply_task.py @@ -53,7 +53,7 @@ def disable_annotation_reply_task(job_id: str, app_id: str, tenant_id: str): try: if annotations_exists: - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) vector.delete() except Exception: logger.exception("Delete annotation index failed when annotation deleted.") diff --git a/api/tasks/annotation/enable_annotation_reply_task.py b/api/tasks/annotation/enable_annotation_reply_task.py index 32c010eaef1..237322f6b56 100644 --- a/api/tasks/annotation/enable_annotation_reply_task.py +++ b/api/tasks/annotation/enable_annotation_reply_task.py @@ -73,7 +73,11 @@ def enable_annotation_reply_task( collection_binding_id=old_dataset_collection_binding.id, ) - old_vector = Vector(old_dataset, attributes=["doc_id", "annotation_id", "app_id"]) + old_vector = Vector( + old_dataset, + attributes=["doc_id", "annotation_id", "app_id"], + session=session, + ) try: old_vector.delete() except Exception as e: @@ -109,7 +113,7 @@ def enable_annotation_reply_task( ) documents.append(document) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) try: vector.delete_by_metadata_field("app_id", app_id) except Exception as e: diff --git a/api/tasks/annotation/update_annotation_to_index_task.py b/api/tasks/annotation/update_annotation_to_index_task.py index eecc1f6fc7b..1d4c90155e2 100644 --- a/api/tasks/annotation/update_annotation_to_index_task.py +++ b/api/tasks/annotation/update_annotation_to_index_task.py @@ -49,7 +49,8 @@ def update_annotation_to_index_task( document = Document( page_content=question, metadata={"annotation_id": annotation_id, "app_id": app_id, "doc_id": annotation_id} ) - vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"]) + with session_factory.create_session() as session: + vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session) vector.delete_by_metadata_field("annotation_id", annotation_id) vector.add_texts([document]) end_at = time.perf_counter() diff --git a/api/tasks/app_generate/resume_agent_app_task.py b/api/tasks/app_generate/resume_agent_app_task.py index 88dfbaa04b3..0b884b014bc 100644 --- a/api/tasks/app_generate/resume_agent_app_task.py +++ b/api/tasks/app_generate/resume_agent_app_task.py @@ -52,6 +52,7 @@ def resume_agent_app_execution(*, conversation_id: str, form_id: str) -> None: user=user, conversation_id=conversation_id, invoke_from=_resolve_invoke_from(conversation), + session=db.session(), ) except Exception: logger.exception("Agent App resume failed for conversation %s form %s", conversation_id, form_id) @@ -63,7 +64,7 @@ def _resolve_conversation_user(*, app_model: App, conversation: Conversation) -> if conversation.from_account_id: account = db.session.get(Account, conversation.from_account_id) if account is not None: - account.set_tenant_id(app_model.tenant_id) + account.set_tenant_id_with_session(app_model.tenant_id, session=db.session()) return account if conversation.from_end_user_id: return db.session.get(EndUser, conversation.from_end_user_id) diff --git a/api/tasks/app_generate/workflow_execute_task.py b/api/tasks/app_generate/workflow_execute_task.py index 36bd21e16c1..213f37a8850 100644 --- a/api/tasks/app_generate/workflow_execute_task.py +++ b/api/tasks/app_generate/workflow_execute_task.py @@ -169,13 +169,14 @@ class _AppRunner: user = self._resolve_user() - with self._setup_flask_context(user): + with self._setup_flask_context(user), self._session_factory(expire_on_commit=False) as session: try: response = self._run_app( app=app, workflow=workflow, user=user, pause_state_config=pause_config, + session=session, ) except Exception as exc: if exec_params.streaming: @@ -205,6 +206,7 @@ class _AppRunner: workflow: Workflow, user: Account | EndUser, pause_state_config: PauseStateLayerConfig, + session: Session, ): exec_params = self._exec_params if exec_params.app_mode == AppMode.ADVANCED_CHAT: @@ -217,6 +219,7 @@ class _AppRunner: streaming=exec_params.streaming, workflow_run_id=exec_params.workflow_run_id, pause_state_config=pause_state_config, + session=session, ) if exec_params.app_mode == AppMode.WORKFLOW: return WorkflowAppGenerator().generate( @@ -245,7 +248,7 @@ class _AppRunner: case _Account(): with self._session() as session: user: Account = session.get(Account, user_params.user_id) - user.set_tenant_id(self._exec_params.tenant_id) + user.set_tenant_id_with_session(self._exec_params.tenant_id, session=session) return user case _: raise AssertionError(f"user should only be _Account or _EndUser, got {type(user_params)}") @@ -256,7 +259,7 @@ def _resolve_user_for_run(session: Session, workflow_run: WorkflowRun) -> Accoun if role == CreatorUserRole.ACCOUNT: user = session.get(Account, workflow_run.created_by) if user: - user.set_tenant_id(workflow_run.tenant_id) + user.set_tenant_id_with_session(workflow_run.tenant_id, session=session) return user return session.get(EndUser, workflow_run.created_by) @@ -556,20 +559,22 @@ def _resume_app_execution(payload: dict[str, Any]) -> None: case AdvancedChatAppGenerateEntity(): assert conversation is not None assert message is not None - _resume_advanced_chat( - app_model=app_model, - workflow=workflow, - user=user, - conversation=conversation, - message=message, - generate_entity=generate_entity, - graph_runtime_state=graph_runtime_state, - response_stream_filter=response_stream_filter, - session_factory=session_factory, - pause_state_config=pause_config, - workflow_run_id=workflow_run_id, - workflow_run=workflow_run, - ) + with session_factory() as session: + _resume_advanced_chat( + app_model=app_model, + workflow=workflow, + user=user, + conversation=conversation, + message=message, + generate_entity=generate_entity, + graph_runtime_state=graph_runtime_state, + response_stream_filter=response_stream_filter, + session_factory=session_factory, + pause_state_config=pause_config, + workflow_run_id=workflow_run_id, + workflow_run=workflow_run, + session=session, + ) case WorkflowAppGenerateEntity(): _resume_workflow( app_model=app_model, @@ -601,6 +606,7 @@ def _resume_advanced_chat( pause_state_config: PauseStateLayerConfig, workflow_run_id: str, workflow_run: WorkflowRun, + session: Session, ) -> None: resumed_generate_entity = generate_entity.model_copy(update={"stream": True}) @@ -637,6 +643,7 @@ def _resume_advanced_chat( graph_runtime_state=graph_runtime_state, pause_state_config=pause_state_config, response_stream_filter=response_stream_filter, + session=session, ) except Exception: logger.exception("Failed to resume chatflow execution for workflow run %s", workflow_run_id) diff --git a/api/tasks/async_workflow_tasks.py b/api/tasks/async_workflow_tasks.py index a6cdb0a1f96..319c5379f57 100644 --- a/api/tasks/async_workflow_tasks.py +++ b/api/tasks/async_workflow_tasks.py @@ -314,7 +314,7 @@ def _get_user(session: Session, workflow_run: WorkflowRun | WorkflowTriggerLog) if workflow_run.created_by_role == CreatorUserRole.ACCOUNT: user = session.scalar(select(Account).where(Account.id == workflow_run.created_by)) if user: - user.current_tenant = tenant + user.set_current_tenant_with_session(tenant, session=session) else: # CreatorUserRole.END_USER user = session.scalar(select(EndUser).where(EndUser.id == workflow_run.created_by)) diff --git a/api/tasks/batch_clean_document_task.py b/api/tasks/batch_clean_document_task.py index 57947267168..d243663a428 100644 --- a/api/tasks/batch_clean_document_task.py +++ b/api/tasks/batch_clean_document_task.py @@ -74,14 +74,19 @@ def batch_clean_document_task(document_ids: list[str], dataset_id: str, doc_form if index_node_ids: try: # Fetch dataset in a fresh session to avoid DetachedInstanceError - with session_factory.create_session() as session: + with session_factory.create_session() as session, session.begin(): dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if not dataset: logger.warning("Dataset not found for vector index cleanup, dataset_id: %s", dataset_id) else: index_processor = IndexProcessorFactory(doc_form).init_index_processor() index_processor.clean( - dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, delete_summaries=True + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + delete_summaries=True, + session=session, ) except Exception: logger.exception( diff --git a/api/tasks/batch_create_segment_to_index_task.py b/api/tasks/batch_create_segment_to_index_task.py index 0f92af21dc8..1fbc4add503 100644 --- a/api/tasks/batch_create_segment_to_index_task.py +++ b/api/tasks/batch_create_segment_to_index_task.py @@ -174,10 +174,12 @@ def batch_create_segment_to_index_task( dataset_document.word_count += word_count_change session.add(dataset_document) - with session_factory.create_session() as session: + with session_factory.create_session() as session, session.begin(): dataset = session.get(Dataset, dataset_id) if dataset: - VectorService.create_segments_vector(None, document_segments, dataset, document_config["doc_form"], session) + VectorService.create_segments_vector( + None, document_segments, dataset, document_config["doc_form"], session=session + ) redis_client.setex(indexing_cache_key, 600, "completed") end_at = time.perf_counter() diff --git a/api/tasks/clean_dataset_task.py b/api/tasks/clean_dataset_task.py index 377d0e5cc70..195114499a0 100644 --- a/api/tasks/clean_dataset_task.py +++ b/api/tasks/clean_dataset_task.py @@ -92,7 +92,7 @@ def clean_dataset_task( # This ensures Document/Segment deletion can continue even if vector database cleanup fails try: index_processor = IndexProcessorFactory(doc_form).init_index_processor() - index_processor.clean(dataset, None, with_keywords=True, delete_child_chunks=True) + index_processor.clean(dataset, None, with_keywords=True, delete_child_chunks=True, session=session) logger.info(click.style(f"Successfully cleaned vector database for dataset: {dataset_id}", fg="green")) except Exception: logger.exception(click.style(f"Failed to clean vector database for dataset {dataset_id}", fg="red")) diff --git a/api/tasks/clean_document_task.py b/api/tasks/clean_document_task.py index 869e2b30287..25887c9b704 100644 --- a/api/tasks/clean_document_task.py +++ b/api/tasks/clean_document_task.py @@ -71,11 +71,16 @@ def clean_document_task(document_id: str, dataset_id: str, doc_form: str, file_i # the vector backend or one of its transitive dependencies was unhappy. try: index_processor = IndexProcessorFactory(doc_form).init_index_processor() - with session_factory.create_session() as session: + with session_factory.create_session() as session, session.begin(): dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if dataset: index_processor.clean( - dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, delete_summaries=True + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + delete_summaries=True, + session=session, ) except Exception: logger.exception( diff --git a/api/tasks/clean_notion_document_task.py b/api/tasks/clean_notion_document_task.py index 782d7d02268..e7500b9a8d7 100644 --- a/api/tasks/clean_notion_document_task.py +++ b/api/tasks/clean_notion_document_task.py @@ -25,12 +25,12 @@ def clean_notion_document_task(document_ids: list[str], dataset_id: str): start_at = time.perf_counter() total_index_node_ids = [] - with session_factory.create_session() as session: + with session_factory.create_session() as session, session.begin(): dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if not dataset: raise Exception("Document has no dataset") - index_type = dataset.doc_form + index_type = dataset.get_doc_form(session=session) index_processor = IndexProcessorFactory(index_type).init_index_processor() document_delete_stmt = delete(Document).where(Document.id.in_(document_ids)) @@ -48,11 +48,16 @@ def clean_notion_document_task(document_ids: list[str], dataset_id: str): # exception escaping this task would produce orphans that no later request # can reference back. Mirrors the pattern in ``clean_dataset_task``. try: - with session_factory.create_session() as session: + with session_factory.create_session() as session, session.begin(): dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if dataset: index_processor.clean( - dataset, total_index_node_ids, with_keywords=True, delete_child_chunks=True, delete_summaries=True + dataset, + total_index_node_ids, + with_keywords=True, + delete_child_chunks=True, + delete_summaries=True, + session=session, ) except Exception: logger.exception( diff --git a/api/tasks/create_segment_to_index_task.py b/api/tasks/create_segment_to_index_task.py index 3448325104f..dfa1fa7fd2b 100644 --- a/api/tasks/create_segment_to_index_task.py +++ b/api/tasks/create_segment_to_index_task.py @@ -56,13 +56,13 @@ def create_segment_to_index_task(segment_id: str, keywords: list[str] | None = N }, ) - dataset = segment.dataset + dataset = segment.get_dataset(session=session) if not dataset: logger.info(click.style(f"Segment {segment.id} has no dataset, pass.", fg="cyan")) return - dataset_document = segment.document + dataset_document = segment.get_document(session=session) if not dataset_document: logger.info(click.style(f"Segment {segment.id} has no document, pass.", fg="cyan")) @@ -76,9 +76,9 @@ def create_segment_to_index_task(segment_id: str, keywords: list[str] | None = N logger.info(click.style(f"Segment {segment.id} document status is invalid, pass.", fg="cyan")) return - index_type = dataset.doc_form + index_type = dataset.get_doc_form(session=session) index_processor = IndexProcessorFactory(index_type).init_index_processor() - index_processor.load(dataset, [document]) + index_processor.load(dataset, [document], session=session) # update segment to completed session.execute( diff --git a/api/tasks/deal_dataset_index_update_task.py b/api/tasks/deal_dataset_index_update_task.py index c9b5121a087..63279348455 100644 --- a/api/tasks/deal_dataset_index_update_task.py +++ b/api/tasks/deal_dataset_index_update_task.py @@ -31,7 +31,7 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): if not dataset: raise Exception("Dataset not found") - index_type = dataset.doc_form or IndexStructureType.PARAGRAPH_INDEX + index_type = dataset.get_doc_form(session=session) or IndexStructureType.PARAGRAPH_INDEX index_processor = IndexProcessorFactory(index_type).init_index_processor() if action == "upgrade": dataset_documents = session.scalars( @@ -79,8 +79,14 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): documents.append(document) # save vector index # clean keywords - index_processor.clean(dataset, None, with_keywords=True, delete_child_chunks=False) - index_processor.load(dataset, documents, with_keywords=False) + index_processor.clean( + dataset, + None, + with_keywords=True, + delete_child_chunks=False, + session=session, + ) + index_processor.load(dataset, documents, with_keywords=False, session=session) session.execute( update(DatasetDocument) .where(DatasetDocument.id == dataset_document.id) @@ -115,7 +121,9 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): session.commit() # clean index - index_processor.clean(dataset, None, with_keywords=False, delete_child_chunks=False) + index_processor.clean( + dataset, None, with_keywords=False, delete_child_chunks=False, session=session + ) for dataset_document in dataset_documents: # update from vector index @@ -142,7 +150,7 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): }, ) if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -158,7 +166,7 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): child_documents.append(child_document) document.children = child_documents if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodal_documents.append( AttachmentDocument( page_content=attachment["name"], @@ -174,7 +182,11 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): documents.append(document) # save vector index index_processor.load( - dataset, documents, multimodal_documents=multimodal_documents, with_keywords=False + dataset, + documents, + multimodal_documents=multimodal_documents, + with_keywords=False, + session=session, ) session.execute( update(DatasetDocument) @@ -191,7 +203,9 @@ def deal_dataset_index_update_task(dataset_id: str, action: str): session.commit() else: # clean collection - index_processor.clean(dataset, None, with_keywords=False, delete_child_chunks=False) + index_processor.clean( + dataset, None, with_keywords=False, delete_child_chunks=False, session=session + ) end_at = time.perf_counter() logging.info( diff --git a/api/tasks/deal_dataset_vector_index_task.py b/api/tasks/deal_dataset_vector_index_task.py index 36605359dc5..001436a9af6 100644 --- a/api/tasks/deal_dataset_vector_index_task.py +++ b/api/tasks/deal_dataset_vector_index_task.py @@ -33,10 +33,10 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): if not dataset: raise Exception("Dataset not found") - index_type = dataset.doc_form or IndexStructureType.PARAGRAPH_INDEX + index_type = dataset.get_doc_form(session=session) or IndexStructureType.PARAGRAPH_INDEX index_processor = IndexProcessorFactory(index_type).init_index_processor() if action == "remove": - index_processor.clean(dataset, None, with_keywords=False) + index_processor.clean(dataset, None, with_keywords=False, session=session) elif action == "add": dataset_documents = session.scalars( select(DatasetDocument).where( @@ -82,7 +82,7 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): documents.append(document) # save vector index - index_processor.load(dataset, documents, with_keywords=False) + index_processor.load(dataset, documents, with_keywords=False, session=session) session.execute( update(DatasetDocument) .where(DatasetDocument.id == dataset_document.id) @@ -117,7 +117,9 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): session.commit() # clean index - index_processor.clean(dataset, None, with_keywords=False, delete_child_chunks=False) + index_processor.clean( + dataset, None, with_keywords=False, delete_child_chunks=False, session=session + ) for dataset_document in dataset_documents: # update from vector index @@ -144,7 +146,7 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): }, ) if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -160,7 +162,7 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): child_documents.append(child_document) document.children = child_documents if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodal_documents.append( AttachmentDocument( page_content=attachment["name"], @@ -176,7 +178,11 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): documents.append(document) # save vector index index_processor.load( - dataset, documents, multimodal_documents=multimodal_documents, with_keywords=False + dataset, + documents, + multimodal_documents=multimodal_documents, + with_keywords=False, + session=session, ) session.execute( update(DatasetDocument) @@ -193,7 +199,9 @@ def deal_dataset_vector_index_task(dataset_id: str, action: str): session.commit() else: # clean collection - index_processor.clean(dataset, None, with_keywords=False, delete_child_chunks=False) + index_processor.clean( + dataset, None, with_keywords=False, delete_child_chunks=False, session=session + ) end_at = time.perf_counter() logger.info( diff --git a/api/tasks/delete_segment_from_index_task.py b/api/tasks/delete_segment_from_index_task.py index 306a23aedae..6fc310138e8 100644 --- a/api/tasks/delete_segment_from_index_task.py +++ b/api/tasks/delete_segment_from_index_task.py @@ -56,8 +56,10 @@ def delete_segment_from_index_task( with_keywords=True, delete_child_chunks=True, precomputed_child_node_ids=child_node_ids, - delete_summaries=True, # Actually delete summaries when segment is deleted + delete_summaries=True, # Actually delete summaries when segment is deleted, + session=session, ) + session.commit() if dataset.is_multimodal: # delete segment attachment binding segment_attachment_bindings = session.scalars( @@ -65,7 +67,9 @@ def delete_segment_from_index_task( ).all() if segment_attachment_bindings: attachment_ids = [binding.attachment_id for binding in segment_attachment_bindings] - index_processor.clean(dataset=dataset, node_ids=attachment_ids, with_keywords=False) + index_processor.clean( + session=session, dataset=dataset, node_ids=attachment_ids, with_keywords=False + ) segment_attachment_bind_ids = [i.id for i in segment_attachment_bindings] for i in range(0, len(segment_attachment_bind_ids), 1000): @@ -81,4 +85,5 @@ def delete_segment_from_index_task( end_at = time.perf_counter() logger.info(click.style(f"Segment deleted from index latency: {end_at - start_at}", fg="green")) except Exception: + session.rollback() logger.exception("delete segment from index failed") diff --git a/api/tasks/disable_segment_from_index_task.py b/api/tasks/disable_segment_from_index_task.py index d00e143093b..02141bf7680 100644 --- a/api/tasks/disable_segment_from_index_task.py +++ b/api/tasks/disable_segment_from_index_task.py @@ -38,13 +38,13 @@ def disable_segment_from_index_task(segment_id: str): indexing_cache_key = f"segment_{segment.id}_indexing" try: - dataset = segment.dataset + dataset = segment.get_dataset(session=session) if not dataset: logger.info(click.style(f"Segment {segment.id} has no dataset, pass.", fg="cyan")) return - dataset_document = segment.document + dataset_document = segment.get_document(session=session) if not dataset_document: logger.info(click.style(f"Segment {segment.id} has no document, pass.", fg="cyan")) @@ -61,7 +61,8 @@ def disable_segment_from_index_task(segment_id: str): index_type = dataset_document.doc_form index_processor = IndexProcessorFactory(index_type).init_index_processor() assert segment.index_node_id - index_processor.clean(dataset, [segment.index_node_id]) + index_processor.clean(dataset, [segment.index_node_id], session=session) + session.commit() # Disable summary index for this segment from services.summary_index_service import SummaryIndexService @@ -84,6 +85,7 @@ def disable_segment_from_index_task(segment_id: str): ) except Exception: logger.exception("remove segment from index failed") + session.rollback() segment.enabled = True session.commit() finally: diff --git a/api/tasks/disable_segments_from_index_task.py b/api/tasks/disable_segments_from_index_task.py index cd91ddd0747..aeefa78dfe9 100644 --- a/api/tasks/disable_segments_from_index_task.py +++ b/api/tasks/disable_segments_from_index_task.py @@ -64,7 +64,10 @@ def disable_segments_from_index_task(segment_ids: list, dataset_id: str, documen if segment_attachment_bindings: attachment_ids = [binding.attachment_id for binding in segment_attachment_bindings] index_node_ids.extend(attachment_ids) - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=False) + index_processor.clean( + dataset, index_node_ids, with_keywords=True, delete_child_chunks=False, session=session + ) + session.commit() # Disable summary indexes for these segments from services.summary_index_service import SummaryIndexService @@ -85,6 +88,7 @@ def disable_segments_from_index_task(segment_ids: list, dataset_id: str, documen logger.info(click.style(f"Segments removed from index latency: {end_at - start_at}", fg="green")) except Exception: # update segment error msg + session.rollback() session.execute( update(DocumentSegment) .where( diff --git a/api/tasks/document_indexing_sync_task.py b/api/tasks/document_indexing_sync_task.py index 842e7dcdb2d..7b12c045c50 100644 --- a/api/tasks/document_indexing_sync_task.py +++ b/api/tasks/document_indexing_sync_task.py @@ -111,49 +111,63 @@ def document_indexing_sync_task(dataset_id: str, document_id: str): logger.info(click.style(f"Document {document_id} content changed, starting sync", fg="green")) - try: - index_processor = IndexProcessorFactory(index_type).init_index_processor() - with session_factory.create_session() as session: - dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) - if dataset: - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) - logger.info(click.style(f"Cleaned vector index for document {document_id}", fg="green")) - except Exception: - logger.exception("Failed to clean vector index for document %s", document_id) - - with session_factory.create_session() as session, session.begin(): - document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) - if not document: - logger.warning(click.style(f"Document {document_id} not found during sync", fg="yellow")) - return - - data_source_info = document.data_source_info_dict - data_source_info["last_edited_time"] = last_edited_time - document.data_source_info = json.dumps(data_source_info) - - document.indexing_status = IndexingStatus.PARSING - document.processing_started_at = naive_utc_now() - - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == document_id) - session.execute(segment_delete_stmt) - - logger.info(click.style(f"Deleted segments for document {document_id}", fg="green")) - try: indexing_runner = IndexingRunner() with session_factory.create_session() as session: document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) - if document: - indexing_runner.run([document]) + if not document: + logger.warning(click.style(f"Document {document_id} not found during sync", fg="yellow")) + return + + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + # End the read transaction before the external vector cleanup. + session.commit() + if dataset: + try: + index_processor = IndexProcessorFactory(index_type).init_index_processor() + index_processor.clean( + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + session=session, + ) + session.commit() + logger.info(click.style(f"Cleaned vector index for document {document_id}", fg="green")) + except Exception: + logger.exception("Failed to clean vector index for document %s", document_id) + session.rollback() + document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) + if not document: + logger.warning(click.style(f"Document {document_id} not found during sync", fg="yellow")) + return + + data_source_info = document.data_source_info_dict + data_source_info["last_edited_time"] = last_edited_time + document.data_source_info = json.dumps(data_source_info) + + document.indexing_status = IndexingStatus.PARSING + document.processing_started_at = naive_utc_now() + + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == document_id) + session.execute(segment_delete_stmt) + # Make the source update and segment deletion visible before extraction. + session.commit() + + logger.info(click.style(f"Deleted segments for document {document_id}", fg="green")) + + indexing_runner.run([document], session) + session.commit() end_at = time.perf_counter() logger.info(click.style(f"Sync completed for document {document_id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: logger.info(click.style(str(ex), fg="yellow")) except Exception as e: logger.exception("document_indexing_sync_task failed for document_id: %s", document_id) - with session_factory.create_session() as session, session.begin(): + with session_factory.create_session() as session: document = session.scalar(select(Document).where(Document.id == document_id).limit(1)) if document: document.indexing_status = IndexingStatus.ERROR document.error = str(e) document.stopped_at = naive_utc_now() + session.commit() diff --git a/api/tasks/document_indexing_task.py b/api/tasks/document_indexing_task.py index c173606d74d..5d8e6dd701c 100644 --- a/api/tasks/document_indexing_task.py +++ b/api/tasks/document_indexing_task.py @@ -91,7 +91,7 @@ def _document_indexing(dataset_id: str, document_ids: Sequence[str]): session.commit() return - # Phase 1: Update status to parsing (short transaction) + # Phase 1: Persist parsing status before slow extraction and vector operations. with session_factory.create_session() as session, session.begin(): documents: list[Document] = list( session.scalars( @@ -100,17 +100,29 @@ def _document_indexing(dataset_id: str, document_ids: Sequence[str]): ) for document in documents: - if document: - document.indexing_status = IndexingStatus.PARSING - document.processing_started_at = naive_utc_now() - session.add(document) - # Transaction committed and closed + document.indexing_status = IndexingStatus.PARSING + document.processing_started_at = naive_utc_now() + session.add(document) - # Phase 2: Execute indexing (no transaction - IndexingRunner creates its own sessions) + # Phase 2: Execute indexing without holding locks from the parsing-status update. has_error = False try: indexing_runner = IndexingRunner() - indexing_runner.run(documents) + with session_factory.create_session() as session: + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if not dataset: + logger.info(click.style(f"Dataset is not found: {dataset_id}", fg="yellow")) + return + + documents = list( + session.scalars( + select(Document).where(Document.id.in_(document_ids), Document.dataset_id == dataset_id) + ).all() + ) + + indexing_runner.run(documents, session) + session.commit() + end_at = time.perf_counter() logger.info(click.style(f"Processed dataset: {dataset_id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: @@ -122,9 +134,6 @@ def _document_indexing(dataset_id: str, document_ids: Sequence[str]): if not has_error: with session_factory.create_session() as session: - # Trigger summary index generation for completed documents if enabled - # Only generate for high_quality indexing technique and when summary_index_setting is enabled - # Re-query dataset to get latest summary_index_setting (in case it was updated) dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if not dataset: logger.warning("Dataset %s not found after indexing", dataset_id) @@ -133,8 +142,6 @@ def _document_indexing(dataset_id: str, document_ids: Sequence[str]): if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: summary_index_setting = dataset.summary_index_setting if summary_index_setting and summary_index_setting.get("enable"): - # expire all session to get latest document's indexing status - session.expire_all() # Check each document's indexing status and trigger summary generation if completed documents = list( diff --git a/api/tasks/document_indexing_update_task.py b/api/tasks/document_indexing_update_task.py index ba26c203317..59ffae8d5b8 100644 --- a/api/tasks/document_indexing_update_task.py +++ b/api/tasks/document_indexing_update_task.py @@ -29,55 +29,75 @@ def document_indexing_update_task(dataset_id: str, document_id: str): logger.info(click.style(f"Start update document: {document_id}", fg="green")) start_at = time.perf_counter() - with session_factory.create_session() as session, session.begin(): - document = session.scalar( - select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) - ) - - if not document: - logger.info(click.style(f"Document not found: {document_id}", fg="red")) - return - - document.indexing_status = IndexingStatus.PARSING - document.processing_started_at = naive_utc_now() - - dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) - if not dataset: - return - - index_type = document.doc_form - segments = session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document_id)).all() - index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] - - clean_success = False - try: - index_processor = IndexProcessorFactory(index_type).init_index_processor() - if index_node_ids: - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) - end_at = time.perf_counter() - logger.info( - click.style( - "Cleaned document when document update data source or process rule: {} latency: {}".format( - document_id, end_at - start_at - ), - fg="green", - ) - ) - clean_success = True - except Exception: - logger.exception("Failed to clean document index during update, document_id: %s", document_id) - - if clean_success: - with session_factory.create_session() as session, session.begin(): - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == document_id) - session.execute(segment_delete_stmt) - has_error = False try: - indexing_runner = IndexingRunner() - indexing_runner.run([document]) - end_at = time.perf_counter() - logger.info(click.style(f"update document: {document.id} latency: {end_at - start_at}", fg="green")) + with session_factory.create_session() as session: + document = session.scalar( + select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) + ) + + if not document: + logger.info(click.style(f"Document not found: {document_id}", fg="red")) + return + + document.indexing_status = IndexingStatus.PARSING + document.processing_started_at = naive_utc_now() + + dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) + if not dataset: + return + + index_type = document.doc_form + segments = session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document_id)).all() + index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] + # Persist the parsing status before vector cleanup and extraction. + session.commit() + + clean_success = False + try: + index_processor = IndexProcessorFactory(index_type).init_index_processor() + if index_node_ids: + index_processor.clean( + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + session=session, + ) + end_at = time.perf_counter() + logger.info( + click.style( + "Cleaned document when document update data source or process rule: {} latency: {}".format( + document_id, end_at - start_at + ), + fg="green", + ) + ) + clean_success = True + except Exception: + logger.exception("Failed to clean document index during update, document_id: %s", document_id) + session.rollback() + document = session.scalar( + select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) + ) + if not document: + logger.info(click.style(f"Document not found: {document_id}", fg="red")) + return + document.indexing_status = IndexingStatus.PARSING + document.processing_started_at = naive_utc_now() + session.commit() + + if clean_success: + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == document_id) + session.execute(segment_delete_stmt) + session.commit() + + indexing_runner = IndexingRunner() + indexing_runner.run([document], session) + session.commit() + + end_at = time.perf_counter() + logger.info(click.style(f"update document: {document.id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: logger.info(click.style(str(ex), fg="yellow")) has_error = True @@ -91,6 +111,9 @@ def document_indexing_update_task(dataset_id: str, document_id: str): # Trigger summary index generation for the updated document if enabled. # Only generate for high_quality indexing technique and when summary_index_setting is enabled. with session_factory.create_session() as session: + document = session.scalar( + select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) + ) dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if not dataset: logger.warning("Dataset %s not found after update indexing", dataset_id) @@ -99,10 +122,6 @@ def document_indexing_update_task(dataset_id: str, document_id: str): if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: summary_index_setting = dataset.summary_index_setting if summary_index_setting and summary_index_setting.get("enable"): - session.expire_all() - document = session.scalar( - select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) - ) if ( document and document.indexing_status == IndexingStatus.COMPLETED diff --git a/api/tasks/duplicate_document_indexing_task.py b/api/tasks/duplicate_document_indexing_task.py index 69747ba497d..10f37db7494 100644 --- a/api/tasks/duplicate_document_indexing_task.py +++ b/api/tasks/duplicate_document_indexing_task.py @@ -113,7 +113,7 @@ def _duplicate_document_indexing_task(dataset_id: str, document_ids: Sequence[st ).all() ) for document in documents: - if document: + if document is not None: document.indexing_status = IndexingStatus.ERROR document.error = str(e) document.stopped_at = naive_utc_now() @@ -141,25 +141,36 @@ def _duplicate_document_indexing_task(dataset_id: str, document_ids: Sequence[st index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] # delete from vector index - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) + index_processor.clean( + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + session=session, + ) segment_ids = [segment.id for segment in segments] - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) - session.execute(segment_delete_stmt) + if segment_ids: + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) + session.execute(segment_delete_stmt) session.commit() document.indexing_status = IndexingStatus.PARSING document.processing_started_at = naive_utc_now() session.add(document) + # Do not keep segment deletions or parsing status changes open during extraction. session.commit() indexing_runner = IndexingRunner() - indexing_runner.run(list(documents)) + indexing_runner.run(list(documents), session) + session.commit() end_at = time.perf_counter() logger.info(click.style(f"Processed dataset: {dataset_id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: + session.rollback() logger.info(click.style(str(ex), fg="yellow")) except Exception: + session.rollback() logger.exception("duplicate_document_indexing_task failed, dataset_id: %s", dataset_id) diff --git a/api/tasks/enable_segment_to_index_task.py b/api/tasks/enable_segment_to_index_task.py index 8334ca25881..c9653eb285f 100644 --- a/api/tasks/enable_segment_to_index_task.py +++ b/api/tasks/enable_segment_to_index_task.py @@ -52,13 +52,13 @@ def enable_segment_to_index_task(segment_id: str): }, ) - dataset = segment.dataset + dataset = segment.get_dataset(session=session) if not dataset: logger.info(click.style(f"Segment {segment.id} has no dataset, pass.", fg="cyan")) return - dataset_document = segment.document + dataset_document = segment.get_document(session=session) if not dataset_document: logger.info(click.style(f"Segment {segment.id} has no document, pass.", fg="cyan")) @@ -74,7 +74,7 @@ def enable_segment_to_index_task(segment_id: str): index_processor = IndexProcessorFactory(dataset_document.doc_form).init_index_processor() if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -91,7 +91,7 @@ def enable_segment_to_index_task(segment_id: str): document.children = child_documents multimodel_documents = [] if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodel_documents.append( AttachmentDocument( page_content=attachment["name"], @@ -106,7 +106,8 @@ def enable_segment_to_index_task(segment_id: str): ) # save vector index - index_processor.load(dataset, [document], multimodal_documents=multimodel_documents) + index_processor.load(dataset, [document], multimodal_documents=multimodel_documents, session=session) + session.commit() # Enable summary index for this segment from services.summary_index_service import SummaryIndexService @@ -123,6 +124,7 @@ def enable_segment_to_index_task(segment_id: str): logger.info(click.style(f"Segment enabled to index: {segment.id} latency: {end_at - start_at}", fg="green")) except Exception as e: logger.exception("enable segment to index failed") + session.rollback() segment.enabled = False segment.disabled_at = naive_utc_now() segment.status = SegmentStatus.ERROR diff --git a/api/tasks/enable_segments_to_index_task.py b/api/tasks/enable_segments_to_index_task.py index 603abf62fe3..b33de4ebc1e 100644 --- a/api/tasks/enable_segments_to_index_task.py +++ b/api/tasks/enable_segments_to_index_task.py @@ -72,7 +72,7 @@ def enable_segments_to_index_task(segment_ids: list, dataset_id: str, document_i ) if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: - child_chunks = segment.get_child_chunks() + child_chunks = segment.get_child_chunks(session=session) if child_chunks: child_documents = [] for child_chunk in child_chunks: @@ -89,7 +89,7 @@ def enable_segments_to_index_task(segment_ids: list, dataset_id: str, document_i document.children = child_documents if dataset.is_multimodal: - for attachment in segment.attachments: + for attachment in segment.get_attachments(session=session): multimodal_documents.append( AttachmentDocument( page_content=attachment["name"], @@ -104,7 +104,9 @@ def enable_segments_to_index_task(segment_ids: list, dataset_id: str, document_i ) documents.append(document) # save vector index - index_processor.load(dataset, documents, multimodal_documents=multimodal_documents) + index_processor.load(dataset, documents, multimodal_documents=multimodal_documents, session=session) + + session.commit() # Enable summary indexes for these segments from services.summary_index_service import SummaryIndexService @@ -123,6 +125,7 @@ def enable_segments_to_index_task(segment_ids: list, dataset_id: str, document_i except Exception as e: logger.exception("enable segments to index failed") # update segment error msg + session.rollback() session.execute( update(DocumentSegment) .where( diff --git a/api/tasks/rag_pipeline/priority_rag_pipeline_run_task.py b/api/tasks/rag_pipeline/priority_rag_pipeline_run_task.py index d8fa73b42df..d78c30ebfee 100644 --- a/api/tasks/rag_pipeline/priority_rag_pipeline_run_task.py +++ b/api/tasks/rag_pipeline/priority_rag_pipeline_run_task.py @@ -126,7 +126,7 @@ def run_single_rag_pipeline_task(rag_pipeline_invoke_entity: Mapping[str, Any], tenant = session.scalar(select(Tenant).where(Tenant.id == tenant_id).limit(1)) if not tenant: raise ValueError(f"Tenant {tenant_id} not found") - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) pipeline = session.scalar(select(Pipeline).where(Pipeline.id == pipeline_id).limit(1)) if not pipeline: @@ -172,19 +172,21 @@ def run_single_rag_pipeline_task(rag_pipeline_invoke_entity: Mapping[str, Any], pipeline_generator = PipelineGenerator() # Using protected method intentionally for async execution - pipeline_generator._generate( # type: ignore[attr-defined] - flask_app=flask_app, - context=context, - pipeline=pipeline, - workflow_id=workflow_id, - user=account, - application_generate_entity=entity, - invoke_from=InvokeFrom.PUBLISHED_PIPELINE, - workflow_execution_repository=workflow_execution_repository, - workflow_node_execution_repository=workflow_node_execution_repository, - streaming=streaming, - workflow_thread_pool_id=workflow_thread_pool_id, - ) + with Session(db.engine, expire_on_commit=False) as session: + pipeline_generator._generate( # type: ignore[attr-defined] + session=session, + flask_app=flask_app, + context=context, + pipeline=pipeline, + workflow_id=workflow_id, + user=account, + application_generate_entity=entity, + invoke_from=InvokeFrom.PUBLISHED_PIPELINE, + workflow_execution_repository=workflow_execution_repository, + workflow_node_execution_repository=workflow_node_execution_repository, + streaming=streaming, + workflow_thread_pool_id=workflow_thread_pool_id, + ) except Exception: logging.exception("Error in priority pipeline task") raise diff --git a/api/tasks/rag_pipeline/rag_pipeline_run_task.py b/api/tasks/rag_pipeline/rag_pipeline_run_task.py index 8e1e096ed03..6b51a03c1ad 100644 --- a/api/tasks/rag_pipeline/rag_pipeline_run_task.py +++ b/api/tasks/rag_pipeline/rag_pipeline_run_task.py @@ -140,7 +140,7 @@ def run_single_rag_pipeline_task(rag_pipeline_invoke_entity: Mapping[str, Any], tenant = session.scalar(select(Tenant).where(Tenant.id == tenant_id).limit(1)) if not tenant: raise ValueError(f"Tenant {tenant_id} not found") - account.current_tenant = tenant + account.set_current_tenant_with_session(tenant, session=session) pipeline = session.scalar(select(Pipeline).where(Pipeline.id == pipeline_id).limit(1)) if not pipeline: @@ -187,6 +187,7 @@ def run_single_rag_pipeline_task(rag_pipeline_invoke_entity: Mapping[str, Any], pipeline_generator = PipelineGenerator() # Using protected method intentionally for async execution pipeline_generator._generate( # type: ignore[attr-defined] + session=session, flask_app=flask_app, context=context, pipeline=pipeline, diff --git a/api/tasks/recover_document_indexing_task.py b/api/tasks/recover_document_indexing_task.py index 73b121961c5..4606981c44b 100644 --- a/api/tasks/recover_document_indexing_task.py +++ b/api/tasks/recover_document_indexing_task.py @@ -36,11 +36,12 @@ def recover_document_indexing_task(dataset_id: str, document_id: str): try: indexing_runner = IndexingRunner() if document.indexing_status in {"waiting", "parsing", "cleaning"}: - indexing_runner.run([document]) + indexing_runner.run([document], session) elif document.indexing_status == "splitting": - indexing_runner.run_in_splitting_status(document) + indexing_runner.run_in_splitting_status(document, session) elif document.indexing_status == "indexing": - indexing_runner.run_in_indexing_status(document) + indexing_runner.run_in_indexing_status(document, session) + session.commit() end_at = time.perf_counter() logger.info(click.style(f"Processed document: {document.id} latency: {end_at - start_at}", fg="green")) except DocumentIsPausedError as ex: diff --git a/api/tasks/regenerate_summary_index_task.py b/api/tasks/regenerate_summary_index_task.py index 5cb8d4281f0..7a4a6331350 100644 --- a/api/tasks/regenerate_summary_index_task.py +++ b/api/tasks/regenerate_summary_index_task.py @@ -152,7 +152,7 @@ def regenerate_summary_index_task( try: from core.rag.datasource.vdb.vector_factory import Vector - vector = Vector(dataset) + vector = Vector(dataset, session=session) vector.delete_by_ids([summary_record.summary_index_node_id]) except Exception as e: logger.warning( @@ -162,7 +162,7 @@ def regenerate_summary_index_task( ) # Re-vectorize with new embedding model - SummaryIndexService.vectorize_summary(summary_record, segment, dataset) + SummaryIndexService.vectorize_summary(summary_record, segment, dataset, session=session) session.commit() total_segments_processed += 1 diff --git a/api/tasks/remove_document_from_index_task.py b/api/tasks/remove_document_from_index_task.py index 2314d32232c..2d3f8ea2a7a 100644 --- a/api/tasks/remove_document_from_index_task.py +++ b/api/tasks/remove_document_from_index_task.py @@ -38,7 +38,7 @@ def remove_document_from_index_task(document_id: str): indexing_cache_key = f"document_{document.id}_indexing" try: - dataset = document.dataset + dataset = document.get_dataset(session=session) if not dataset: raise Exception("Document has no dataset") @@ -64,7 +64,13 @@ def remove_document_from_index_task(document_id: str): index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] if index_node_ids: try: - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=False) + index_processor.clean( + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=False, + session=session, + ) except Exception: logger.exception("clean dataset %s from index failed", dataset.id) # update segment to disable diff --git a/api/tasks/retry_document_indexing_task.py b/api/tasks/retry_document_indexing_task.py index dddb7715d22..f8430cc206a 100644 --- a/api/tasks/retry_document_indexing_task.py +++ b/api/tasks/retry_document_indexing_task.py @@ -43,7 +43,7 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ tenant = session.scalar(select(Tenant).where(Tenant.id == dataset.tenant_id).limit(1)) if not tenant: raise ValueError("Tenant not found") - user.current_tenant = tenant + user.set_current_tenant_with_session(tenant, session=session) for document_id in document_ids: retry_indexing_cache_key = f"document_{document_id}_is_retried" @@ -88,16 +88,24 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ if segments: index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] # delete from vector index - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) + index_processor.clean( + dataset, + index_node_ids, + with_keywords=True, + delete_child_chunks=True, + session=session, + ) segment_ids = [segment.id for segment in segments] - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) - session.execute(segment_delete_stmt) + if segment_ids: + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) + session.execute(segment_delete_stmt) session.commit() document.indexing_status = IndexingStatus.PARSING document.processing_started_at = naive_utc_now() session.add(document) + # The runner performs slow extraction/indexing in a separate transaction phase. session.commit() if dataset.runtime_mode == "rag_pipeline": @@ -106,14 +114,20 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ rag_pipeline_service.retry_error_document(dataset, document, user) else: indexing_runner = IndexingRunner() - indexing_runner.run([document]) + indexing_runner.run([document], session) + session.commit() redis_client.delete(retry_indexing_cache_key) except Exception as ex: - document.indexing_status = IndexingStatus.ERROR - document.error = str(ex) - document.stopped_at = naive_utc_now() - session.add(document) - session.commit() + session.rollback() + document = session.scalar( + select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) + ) + if document: + document.indexing_status = IndexingStatus.ERROR + document.error = str(ex) + document.stopped_at = naive_utc_now() + session.add(document) + session.commit() logger.info(click.style(str(ex), fg="yellow")) redis_client.delete(retry_indexing_cache_key) logger.exception("retry_document_indexing_task failed, document_id: %s", document_id) diff --git a/api/tasks/sync_website_document_indexing_task.py b/api/tasks/sync_website_document_indexing_task.py index 5bdeac7e9df..dbe029b6016 100644 --- a/api/tasks/sync_website_document_indexing_task.py +++ b/api/tasks/sync_website_document_indexing_task.py @@ -73,27 +73,37 @@ def sync_website_document_indexing_task(dataset_id: str, document_id: str): if segments: index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] # delete from vector index - index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True) + index_processor.clean( + dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session + ) segment_ids = [segment.id for segment in segments] - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) - session.execute(segment_delete_stmt) + if segment_ids: + segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) + session.execute(segment_delete_stmt) session.commit() document.indexing_status = IndexingStatus.PARSING document.processing_started_at = naive_utc_now() session.add(document) + # Release document/segment locks before extraction starts. session.commit() indexing_runner = IndexingRunner() - indexing_runner.run([document]) + indexing_runner.run([document], session) + session.commit() redis_client.delete(sync_indexing_cache_key) except Exception as ex: - document.indexing_status = IndexingStatus.ERROR - document.error = str(ex) - document.stopped_at = naive_utc_now() - session.add(document) - session.commit() + session.rollback() + document = session.scalar( + select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) + ) + if document: + document.indexing_status = IndexingStatus.ERROR + document.error = str(ex) + document.stopped_at = naive_utc_now() + session.add(document) + session.commit() logger.info(click.style(str(ex), fg="yellow")) redis_client.delete(sync_indexing_cache_key) logger.exception("sync_website_document_indexing_task failed, document_id: %s", document_id) diff --git a/api/tests/integration_tests/controllers/console/app/test_model_config_permissions.py b/api/tests/integration_tests/controllers/console/app/test_model_config_permissions.py index 3634034c815..1237abc387b 100644 --- a/api/tests/integration_tests/controllers/console/app/test_model_config_permissions.py +++ b/api/tests/integration_tests/controllers/console/app/test_model_config_permissions.py @@ -108,13 +108,6 @@ class TestModelConfigResourcePermissions: ) monkeypatch.setattr(AppModelConfigService, "validate_configuration", mock_validate_config) - # Mock database operations - mock_db_session = mock.Mock() - mock_db_session.add = mock.Mock() - mock_db_session.flush = mock.Mock() - mock_db_session.commit = mock.Mock() - monkeypatch.setattr(model_config_api.db, "session", mock_db_session) - # Mock app_model_config_was_updated event mock_event = mock.Mock() mock_event.send = mock.Mock() diff --git a/api/tests/test_containers_integration_tests/controllers/console/app/test_conversation_read_timestamp.py b/api/tests/test_containers_integration_tests/controllers/console/app/test_conversation_read_timestamp.py index 05c8452da80..c7625580d30 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/app/test_conversation_read_timestamp.py +++ b/api/tests/test_containers_integration_tests/controllers/console/app/test_conversation_read_timestamp.py @@ -41,7 +41,7 @@ def test_get_conversation_mark_read_keeps_updated_at_unchanged( return_value=read_at, autospec=True, ): - loaded = _get_conversation(account, app, conversation.id) + loaded = _get_conversation(db_session_with_containers, account, app, conversation.id) db_session_with_containers.refresh(conversation) @@ -58,4 +58,4 @@ def test_get_conversation_raises_not_found_for_missing_conversation( app = create_console_app(db_session_with_containers, tenant.id, account.id, AppMode.CHAT) with pytest.raises(NotFound): - _get_conversation(account, app, "00000000-0000-0000-0000-000000000000") + _get_conversation(db_session_with_containers, account, app, "00000000-0000-0000-0000-000000000000") diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py index 4a41aa352d2..f8abd102143 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py @@ -450,6 +450,7 @@ class TestPipelineRunApis: pipeline = make_pipeline() user = make_account() + session = MagicMock(spec=Session) payload = { "inputs": empty_mapping(), @@ -461,16 +462,23 @@ class TestPipelineRunApis: with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", + return_value=pipeline, + ) as load_pipeline, patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", return_value=MagicMock(), - ), + ) as generate, patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response", return_value={"ok": True}, ), ): - assert method(api, MagicMock(), user, pipeline) == {"ok": True} + assert method(api, session, user, pipeline.id) == {"ok": True} + + load_pipeline.assert_called_once_with(session, pipeline.id) + assert generate.call_args.kwargs["session"] is session def test_draft_run_rate_limit(self, app: Flask) -> None: api = DraftRagPipelineRunApi() @@ -478,6 +486,7 @@ class TestPipelineRunApis: pipeline = make_pipeline() user = make_account() + session = MagicMock(spec=Session) payload: dict[str, object] = { "inputs": empty_mapping(), "datasource_type": "x", @@ -492,13 +501,17 @@ class TestPipelineRunApis: "payload", payload, ), + patch( + "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", + return_value=pipeline, + ), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", side_effect=InvokeRateLimitError("limit"), ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), user, pipeline) + method(api, session, user, pipeline.id) class TestDraftNodeRun: @@ -638,6 +651,7 @@ class TestPublishedRagPipelineRunApi: pipeline = make_pipeline() user = make_account() + session = MagicMock(spec=Session) payload = { "inputs": empty_mapping(), @@ -650,24 +664,32 @@ class TestPublishedRagPipelineRunApi: with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", + return_value=pipeline, + ) as load_pipeline, patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", return_value=MagicMock(), - ), + ) as generate, patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response", return_value={"ok": True}, ), ): - result = method(api, MagicMock(), user, pipeline) + result = method(api, session, user, pipeline.id) assert result == {"ok": True} + load_pipeline.assert_called_once_with(session, pipeline.id) + assert generate.call_args.kwargs["session"] is session + def test_published_run_rate_limit(self, app: Flask) -> None: api = PublishedRagPipelineRunApi() method = unwrap(api.post) pipeline = make_pipeline() user = make_account() + session = MagicMock(spec=Session) payload = { "inputs": empty_mapping(), @@ -679,13 +701,17 @@ class TestPublishedRagPipelineRunApi: with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", + return_value=pipeline, + ), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", side_effect=InvokeRateLimitError("limit"), ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), user, pipeline) + method(api, session, user, pipeline.id) class TestDefaultBlockConfigApi: diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py index aab79538f3c..7c63cee4107 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py @@ -129,96 +129,67 @@ class TestDataSourceApi: assert status == 200 assert response["data"] == [] - def test_patch_enable_binding(self, app: Flask, mock_engine: None) -> None: + def test_patch_enable_binding(self, app: Flask) -> None: api = DataSourceApi() method = inspect.unwrap(api.patch) binding = MagicMock(id="b1", disabled=True) + session = MagicMock() + session.scalar.return_value = binding - with ( - app.test_request_context("/"), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, - patch("controllers.console.datasets.data_source.db.session.add"), - patch("controllers.console.datasets.data_source.db.session.commit"), - ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.execute.return_value.scalar_one_or_none.return_value = binding - - response, status = method(api, "tenant-1", "b1", "enable") + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "b1", "enable") assert status == 200 assert binding.disabled is False - def test_patch_disable_binding(self, app: Flask, mock_engine: None) -> None: + def test_patch_disable_binding(self, app: Flask) -> None: api = DataSourceApi() method = inspect.unwrap(api.patch) binding = MagicMock(id="b1", disabled=False) + session = MagicMock() + session.scalar.return_value = binding - with ( - app.test_request_context("/"), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, - patch("controllers.console.datasets.data_source.db.session.add"), - patch("controllers.console.datasets.data_source.db.session.commit"), - ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.execute.return_value.scalar_one_or_none.return_value = binding - - response, status = method(api, "tenant-1", "b1", "disable") + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "b1", "disable") assert status == 200 assert binding.disabled is True - def test_patch_binding_not_found(self, app: Flask, mock_engine: None) -> None: + def test_patch_binding_not_found(self, app: Flask) -> None: api = DataSourceApi() method = inspect.unwrap(api.patch) + session = MagicMock() + session.scalar.return_value = None - with ( - app.test_request_context("/"), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, - ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.execute.return_value.scalar_one_or_none.return_value = None - + with app.test_request_context("/"): with pytest.raises(NotFound): - method(api, "tenant-1", "b1", "enable") + method(api, session, "tenant-1", "b1", "enable") - def test_patch_enable_already_enabled(self, app: Flask, mock_engine: None) -> None: + def test_patch_enable_already_enabled(self, app: Flask) -> None: api = DataSourceApi() method = inspect.unwrap(api.patch) binding = MagicMock(id="b1", disabled=False) + session = MagicMock() + session.scalar.return_value = binding - with ( - app.test_request_context("/"), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, - ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.execute.return_value.scalar_one_or_none.return_value = binding - + with app.test_request_context("/"): with pytest.raises(ValueError): - method(api, "tenant-1", "b1", "enable") + method(api, session, "tenant-1", "b1", "enable") - def test_patch_disable_already_disabled(self, app: Flask, mock_engine: None) -> None: + def test_patch_disable_already_disabled(self, app: Flask) -> None: api = DataSourceApi() method = inspect.unwrap(api.patch) binding = MagicMock(id="b1", disabled=True) + session = MagicMock() + session.scalar.return_value = binding - with ( - app.test_request_context("/"), - patch("controllers.console.datasets.data_source.sessionmaker") as mock_session_class, - ): - mock_session = MagicMock() - mock_session_class.return_value.begin.return_value.__enter__.return_value = mock_session - mock_session.execute.return_value.scalar_one_or_none.return_value = binding - + with app.test_request_context("/"): with pytest.raises(ValueError): - method(api, "tenant-1", "b1", "disable") + method(api, session, "tenant-1", "b1", "disable") class TestDataSourceNotionListApi: @@ -238,7 +209,7 @@ class TestDataSourceNotionListApi: ), ): with pytest.raises(NotFound): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) def test_get_success_no_dataset_id(self, app: Flask, current_user: Account, mock_engine: None) -> None: api = DataSourceNotionListApi() @@ -277,7 +248,7 @@ class TestDataSourceNotionListApi: ), ), ): - response, status = method(api, "tenant-1", current_user) + response, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 200 @@ -343,7 +314,7 @@ class TestDataSourceNotionListApi: ), ), ): - response, status = method(api, tenant_id, current_user) + response, status = method(api, db_session_with_containers, tenant_id, current_user) assert status == 200 @@ -363,10 +334,9 @@ class TestDataSourceNotionListApi: "controllers.console.datasets.data_source.DatasetService.get_dataset", return_value=dataset, ), - patch("controllers.console.datasets.data_source.sessionmaker"), ): with pytest.raises(ValueError): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) class TestDataSourceNotionPreviewApi: @@ -429,7 +399,7 @@ class TestDataSourceNotionIndexingEstimateApi: return_value=MagicMock(model_dump=lambda: {"total_pages": 1}), ), ): - response, status = method(api, "tenant-1") + response, status = method(api, MagicMock(), "tenant-1") assert status == 200 @@ -458,7 +428,7 @@ class TestDataSourceNotionDatasetSyncApi: return_value=None, ), ): - response, status = method(api, "ds-1") + response, status = method(api, MagicMock(), "ds-1") assert status == 200 @@ -474,7 +444,7 @@ class TestDataSourceNotionDatasetSyncApi: ), ): with pytest.raises(NotFound): - method(api, "ds-1") + method(api, MagicMock(), "ds-1") class TestDataSourceNotionDocumentSyncApi: @@ -501,7 +471,7 @@ class TestDataSourceNotionDocumentSyncApi: return_value=None, ), ): - response, status = method(api, "ds-1", "doc-1") + response, status = method(api, MagicMock(), "ds-1", "doc-1") assert status == 200 @@ -521,4 +491,4 @@ class TestDataSourceNotionDocumentSyncApi: ), ): with pytest.raises(NotFound): - method(api, "ds-1", "doc-1") + method(api, MagicMock(), "ds-1", "doc-1") diff --git a/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py b/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py index 3d5fce4b6ca..14e34689ebb 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py +++ b/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py @@ -9,6 +9,7 @@ from unittest.mock import patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound import controllers.console.explore.conversation as conversation_module @@ -27,6 +28,10 @@ from services.errors.conversation import ( class InstalledAppCarrier: app: App | None + def app_with_session(self, *, session: Session) -> App | None: + del session + return self.app + @pytest.fixture def chat_app() -> InstalledApp: diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_account.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_account.py index 7b5bef7b613..0a04c4fdad0 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_account.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_account.py @@ -13,14 +13,16 @@ from tests.test_containers_integration_tests.controllers.openapi.conftest import class TestAccountInfo: - def test_returns_account_and_owner_workspace(self, app: Flask, make_account: Callable[..., Account]) -> None: + def test_returns_account_and_owner_workspace( + self, app: Flask, db_session_with_containers: Session, make_account: Callable[..., Account] + ) -> None: account = make_account() owner_tenant = account.current_tenant assert owner_tenant is not None api = AccountApi() with app.test_request_context("/openapi/v1/account"): - result = unwrap(api.get)(api, auth_data=auth_for(account)) + result = unwrap(api.get)(api, db_session_with_containers, auth_data=auth_for(account)) assert result.subject_type == "account" assert result.subject_email == account.email @@ -45,7 +47,7 @@ class TestAccountInfo: api = AccountApi() with app.test_request_context("/openapi/v1/account"): - result = unwrap(api.get)(api, auth_data=auth_for(account)) + result = unwrap(api.get)(api, db_session_with_containers, auth_data=auth_for(account)) assert {w.id for w in result.workspaces} == {owner_tenant.id, second.id} roles = {w.id: w.role for w in result.workspaces} 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 4222f49a28a..9e6acab5543 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 @@ -53,7 +53,10 @@ class TestSessionList: with app.test_request_context("/openapi/v1/account/sessions"): with account_auth_context(account, token_id=mint.token_id): result = unwrap(api.get)( - api, auth_data=auth_for(account, token_id=mint.token_id), query=SessionListQuery() + api, + db_session_with_containers, + auth_data=auth_for(account, token_id=mint.token_id), + query=SessionListQuery(), ) assert result.total == 1 @@ -75,7 +78,10 @@ class TestSessionList: with app.test_request_context("/openapi/v1/account/sessions"): with account_auth_context(account, token_id=mine.token_id): result = unwrap(api.get)( - api, auth_data=auth_for(account, token_id=mine.token_id), query=SessionListQuery() + api, + db_session_with_containers, + auth_data=auth_for(account, token_id=mine.token_id), + query=SessionListQuery(), ) assert {row.id for row in result.data} == {str(mine.token_id)} @@ -91,7 +97,9 @@ class TestSessionRevoke: revoke_api = AccountSessionsSelfApi() with app.test_request_context("/openapi/v1/account/sessions/self", method="DELETE"): with account_auth_context(account, token_id=mint.token_id): - result = unwrap(revoke_api.delete)(revoke_api, auth_data=auth_for(account, token_id=mint.token_id)) + result = unwrap(revoke_api.delete)( + revoke_api, db_session_with_containers, auth_data=auth_for(account, token_id=mint.token_id) + ) assert result.status == "revoked" @@ -100,7 +108,10 @@ class TestSessionRevoke: with app.test_request_context("/openapi/v1/account/sessions"): with account_auth_context(account, token_id=mint.token_id): listing = unwrap(list_api.get)( - list_api, auth_data=auth_for(account, token_id=mint.token_id), query=SessionListQuery() + list_api, + db_session_with_containers, + auth_data=auth_for(account, token_id=mint.token_id), + query=SessionListQuery(), ) assert listing.total == 0 @@ -115,7 +126,10 @@ class TestSessionRevoke: with app.test_request_context(f"/openapi/v1/account/sessions/{session_id}", method="DELETE"): with account_auth_context(account, token_id=mint.token_id): result = unwrap(api.delete)( - api, session_id=session_id, auth_data=auth_for(account, token_id=mint.token_id) + api, + db_session_with_containers, + session_id=session_id, + auth_data=auth_for(account, token_id=mint.token_id), ) assert result.status == "revoked" @@ -134,4 +148,9 @@ class TestSessionRevoke: with app.test_request_context(f"/openapi/v1/account/sessions/{session_id}", method="DELETE"): with account_auth_context(outsider, token_id=uuid4()): with pytest.raises(NotFound): - unwrap(api.delete)(api, session_id=session_id, auth_data=auth_for(outsider, token_id=uuid4())) + unwrap(api.delete)( + api, + db_session_with_containers, + session_id=session_id, + auth_data=auth_for(outsider, token_id=uuid4()), + ) 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 ce1425d9e61..52f70b9d590 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 @@ -61,7 +61,12 @@ class TestAppList: api = AppListApi() with app.test_request_context(f"/openapi/v1/apps?workspace_id={tenant.id}"): - result = unwrap(api.get)(api, auth_data=auth_for(account), query=AppListQuery(workspace_id=str(tenant.id))) + result = unwrap(api.get)( + api, + db_session_with_containers, + auth_data=auth_for(account), + query=AppListQuery(workspace_id=str(tenant.id)), + ) # The api-disabled app is gated out, so it counts neither in `data` # nor in `total` (the gate is pushed into the query for stable paging). @@ -81,6 +86,7 @@ class TestAppList: with app.test_request_context(f"/openapi/v1/apps?workspace_id={tenant.id}&name={target.id}"): result = unwrap(api.get)( api, + db_session_with_containers, auth_data=auth_for(account), query=AppListQuery(workspace_id=str(tenant.id), name=str(target.id)), ) @@ -104,6 +110,7 @@ class TestAppList: with app.test_request_context(f"/openapi/v1/apps?workspace_id={outsider_tenant.id}&name={foreign_app.id}"): result = unwrap(api.get)( api, + db_session_with_containers, auth_data=auth_for(outsider), query=AppListQuery(workspace_id=str(outsider_tenant.id), name=str(foreign_app.id)), ) @@ -122,7 +129,11 @@ class TestAppDescribe: api = AppDescribeApi() with app.test_request_context(f"/openapi/v1/apps/{app_model.id}?fields=info"): result = unwrap(api.get)( - api, app_id=app_model.id, auth_data=auth_for(account), query=AppDescribeQuery(fields="info") + api, + db_session_with_containers, + app_id=app_model.id, + auth_data=auth_for(account), + query=AppDescribeQuery(fields="info"), ) assert result.info is not None @@ -133,14 +144,22 @@ class TestAppDescribe: assert result.parameters is None assert result.input_schema is None - def test_describe_unknown_app_is_404(self, app: Flask, make_account: Callable[..., Account]) -> None: + def test_describe_unknown_app_is_404( + self, app: Flask, db_session_with_containers: Session, make_account: Callable[..., Account] + ) -> None: account = make_account() missing_id = str(uuid4()) api = AppDescribeApi() with app.test_request_context(f"/openapi/v1/apps/{missing_id}"): with pytest.raises(NotFound): - unwrap(api.get)(api, app_id=missing_id, auth_data=auth_for(account), query=AppDescribeQuery()) + unwrap(api.get)( + api, + db_session_with_containers, + app_id=missing_id, + auth_data=auth_for(account), + query=AppDescribeQuery(), + ) def test_describe_api_disabled_app_is_404( self, app: Flask, db_session_with_containers: Session, make_account: Callable[..., Account] @@ -153,4 +172,10 @@ class TestAppDescribe: api = AppDescribeApi() with app.test_request_context(f"/openapi/v1/apps/{hidden.id}"): with pytest.raises(NotFound): - unwrap(api.get)(api, app_id=hidden.id, auth_data=auth_for(account), query=AppDescribeQuery()) + unwrap(api.get)( + api, + db_session_with_containers, + app_id=hidden.id, + auth_data=auth_for(account), + query=AppDescribeQuery(), + ) diff --git a/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py b/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py index 5e794ae1982..1346e013e66 100644 --- a/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py +++ b/api/tests/test_containers_integration_tests/controllers/openapi/test_workspaces.py @@ -24,7 +24,7 @@ class TestWorkspacesList: api = WorkspacesApi() with app.test_request_context("/openapi/v1/workspaces"): - result = unwrap(api.get)(api, auth_data=auth_for(account)) + result = unwrap(api.get)(api, db_session_with_containers, auth_data=auth_for(account)) ids = {w.id for w in result.workspaces} assert ids == {owner_tenant.id} @@ -45,7 +45,7 @@ class TestWorkspacesList: api = WorkspacesApi() with app.test_request_context("/openapi/v1/workspaces"): - result = unwrap(api.get)(api, auth_data=auth_for(account)) + result = unwrap(api.get)(api, db_session_with_containers, auth_data=auth_for(account)) assert {w.id for w in result.workspaces} == {owner_tenant.id, second.id} @@ -60,7 +60,9 @@ class TestWorkspaceDetail: api = WorkspaceByIdApi() with app.test_request_context(f"/openapi/v1/workspaces/{tenant.id}"): - detail = unwrap(api.get)(api, workspace_id=tenant.id, auth_data=auth_for(account)) + detail = unwrap(api.get)( + api, db_session_with_containers, workspace_id=tenant.id, auth_data=auth_for(account) + ) assert detail.id == tenant.id assert detail.role == TenantAccountRole.OWNER.value @@ -80,7 +82,9 @@ class TestWorkspaceDetail: api = WorkspaceByIdApi() with app.test_request_context(f"/openapi/v1/workspaces/{someone_elses_ws.id}"): with pytest.raises(NotFound): - unwrap(api.get)(api, workspace_id=someone_elses_ws.id, auth_data=auth_for(outsider)) + unwrap(api.get)( + api, db_session_with_containers, workspace_id=someone_elses_ws.id, auth_data=auth_for(outsider) + ) class TestWorkspaceSwitch: @@ -95,8 +99,10 @@ class TestWorkspaceSwitch: ) api = WorkspaceSwitchApi() - with app.test_request_context(f"/openapi/v1/workspaces/{target.id}:switch", method="POST"): - detail = unwrap(api.post)(api, workspace_id=target.id, auth_data=auth_for(account)) + with app.test_request_context(f"/openapi/v1/workspaces/{target.id}/switch", method="POST"): + detail = unwrap(api.post)( + api, db_session_with_containers, workspace_id=target.id, auth_data=auth_for(account) + ) # Response reflects the post-switch state. assert detail.id == target.id @@ -105,7 +111,9 @@ class TestWorkspaceSwitch: # And the switch persisted: the previously-current owner workspace is no # longer current (verified through the real read path). with app.test_request_context("/openapi/v1/workspaces"): - listing = unwrap(WorkspacesApi().get)(WorkspacesApi(), auth_data=auth_for(account)) + listing = unwrap(WorkspacesApi().get)( + WorkspacesApi(), db_session_with_containers, auth_data=auth_for(account) + ) by_id = {w.id: w for w in listing.workspaces} assert by_id[target.id].current is True assert by_id[owner_tenant.id].current is False @@ -120,4 +128,6 @@ class TestWorkspaceSwitch: api = WorkspaceSwitchApi() with app.test_request_context(f"/openapi/v1/workspaces/{outsider_ws.id}:switch", method="POST"): with pytest.raises(NotFound): - unwrap(api.post)(api, workspace_id=outsider_ws.id, auth_data=auth_for(account)) + unwrap(api.post)( + api, db_session_with_containers, workspace_id=outsider_ws.id, auth_data=auth_for(account) + ) 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 d670425be0c..e58bed04270 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 @@ -16,7 +16,7 @@ since these test controller-level behavior. import uuid from contextlib import ExitStack from datetime import UTC, datetime -from unittest.mock import ANY, Mock, PropertyMock, patch +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask @@ -270,25 +270,25 @@ def mock_dataset(): @pytest.fixture(autouse=True) -def dataset_model_property_defaults(): - properties: dict[str, object] = { - "app_count": 0, - "document_count": 0, - "word_count": 0, - "author_name": None, - "tags": [], - "doc_form": None, - "external_knowledge_info": None, - "doc_metadata": [], - "is_published": False, - "total_documents": 0, - "total_available_documents": 0, +def dataset_model_getter_defaults(): + getters: dict[str, object] = { + "get_app_count": 0, + "get_document_count": 0, + "get_word_count": 0, + "get_author_name": None, + "get_tags": [], + "get_doc_form": None, + "get_external_knowledge_info": None, + "get_doc_metadata": [], + "get_is_published": False, + "get_total_documents": 0, + "get_total_available_documents": 0, } with ExitStack() as stack: - for name, value in properties.items(): - property_mock = stack.enter_context(patch.object(Dataset, name, new_callable=PropertyMock)) - property_mock.return_value = value + for name, value in getters.items(): + getter_mock = stack.enter_context(patch.object(Dataset, name, autospec=True)) + getter_mock.return_value = value yield @@ -770,7 +770,7 @@ class TestDatasetApiDelete: method="DELETE", ): api = DatasetApi() - result = unwrap(api.delete)(api, _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) + result = unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) assert result == ("", 204) @@ -793,7 +793,7 @@ class TestDatasetApiDelete: ): api = DatasetApi() with pytest.raises(NotFound): - unwrap(api.delete)(api, _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) + unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) @patch("controllers.service_api.dataset.dataset.current_user") @patch("controllers.service_api.dataset.dataset.DatasetService") @@ -814,7 +814,7 @@ class TestDatasetApiDelete: ): api = DatasetApi() with pytest.raises(DatasetInUseError): - unwrap(api.delete)(api, _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) + unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) # --------------------------------------------------------------------------- diff --git a/api/tests/test_containers_integration_tests/models/test_account.py b/api/tests/test_containers_integration_tests/models/test_account.py index 22388694bc9..1f1c4a4ede1 100644 --- a/api/tests/test_containers_integration_tests/models/test_account.py +++ b/api/tests/test_containers_integration_tests/models/test_account.py @@ -214,7 +214,7 @@ class TestTenantGetAccounts(_DBTrackingTestBase): self._create_join(db_session_with_containers, tenant.id, account1.id, TenantAccountRole.OWNER, current=False) self._create_join(db_session_with_containers, tenant.id, account2.id, TenantAccountRole.NORMAL, current=False) - accounts = tenant.get_accounts() + accounts = tenant.get_accounts(session=db_session_with_containers) assert len(accounts) == 2 account_ids = {a.id for a in accounts} diff --git a/api/tests/test_containers_integration_tests/models/test_conversation_status_count.py b/api/tests/test_containers_integration_tests/models/test_conversation_status_count.py index 6352f815df7..e706c07c820 100644 --- a/api/tests/test_containers_integration_tests/models/test_conversation_status_count.py +++ b/api/tests/test_containers_integration_tests/models/test_conversation_status_count.py @@ -273,7 +273,7 @@ class TestSiteGenerateCode: def test_generate_code_returns_string_of_correct_length(self, db_session_with_containers: Session) -> None: """Site.generate_code returns a code string of the requested length.""" - code = Site.generate_code(8) + code = Site.generate_code(8, session=db_session_with_containers) assert isinstance(code, str) assert len(code) == 8 @@ -307,7 +307,7 @@ class TestSiteGenerateCode: db_session_with_containers.add(site) db_session_with_containers.flush() - code = Site.generate_code(8) + code = Site.generate_code(8, session=db_session_with_containers) assert isinstance(code, str) assert len(code) == 8 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 2710df5e56c..0aa6cc820a5 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 @@ -243,9 +243,7 @@ class TestAnnotationService: } with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly( - annotation_args, app.id, session=db_session_with_containers - ) + AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id, db_session_with_containers) def test_insert_app_annotation_directly_app_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -443,11 +441,7 @@ class TestAnnotationService: # Get annotation list annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, - page=1, - limit=10, - keyword="", - session=db_session_with_containers, + app.id, page=1, limit=10, keyword="", session=db_session_with_containers ) # Verify results @@ -474,22 +468,18 @@ 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, session=db_session_with_containers) + AppAnnotationService.insert_app_annotation_directly(annotation_args, app.id, 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, session=db_session_with_containers) + AppAnnotationService.insert_app_annotation_directly(other_args, app.id, 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, - session=db_session_with_containers, + app.id, page=1, limit=10, keyword=unique_keyword, session=db_session_with_containers ) # Verify only matching annotations are returned @@ -516,9 +506,7 @@ class TestAnnotationService: "question": "Question with 50% discount", "answer": "Answer about 50% discount offer", } - AppAnnotationService.insert_app_annotation_directly( - annotation_with_percent, app.id, session=db_session_with_containers - ) + AppAnnotationService.insert_app_annotation_directly(annotation_with_percent, app.id, db_session_with_containers) annotation_with_underscore = { "question": "Question with test_data", @@ -541,17 +529,11 @@ class TestAnnotationService: "question": "Question with 100% different", "answer": "Answer about 100% different content", } - AppAnnotationService.insert_app_annotation_directly( - annotation_no_match, app.id, session=db_session_with_containers - ) + AppAnnotationService.insert_app_annotation_directly(annotation_no_match, app.id, 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%", - session=db_session_with_containers, + app.id, page=1, limit=10, keyword="50%", session=db_session_with_containers ) assert total == 1 assert len(annotation_list) == 1 @@ -559,11 +541,7 @@ 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", - session=db_session_with_containers, + app.id, page=1, limit=10, keyword="test_data", session=db_session_with_containers ) assert total == 1 assert len(annotation_list) == 1 @@ -571,11 +549,7 @@ 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", - session=db_session_with_containers, + app.id, page=1, limit=10, keyword="path\\to\\file", session=db_session_with_containers ) assert total == 1 assert len(annotation_list) == 1 @@ -583,11 +557,7 @@ 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%", - session=db_session_with_containers, + 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 @@ -945,9 +915,7 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations( - app.id, file_storage, session=db_session_with_containers - ) + result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage, db_session_with_containers) # Verify result structure assert "job_id" in result @@ -987,9 +955,7 @@ class TestAnnotationService: mock_pd.read_csv.return_value = mock_df # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations( - app.id, file_storage, session=db_session_with_containers - ) + result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage, db_session_with_containers) # Verify error result assert "error_msg" in result @@ -1035,9 +1001,7 @@ class TestAnnotationService: ].get_features.return_value.annotation_quota_limit.size = 0 # Batch import annotations - result = AppAnnotationService.batch_import_app_annotations( - app.id, file_storage, session=db_session_with_containers - ) + result = AppAnnotationService.batch_import_app_annotations(app.id, file_storage, db_session_with_containers) # Verify error result assert "error_msg" in result @@ -1079,7 +1043,7 @@ class TestAnnotationService: db_session_with_containers.commit() # Get annotation setting - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=db_session_with_containers) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, db_session_with_containers) # Verify result structure assert result["enabled"] is True @@ -1098,7 +1062,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, session=db_session_with_containers) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, db_session_with_containers) # Verify result structure assert result["enabled"] is False @@ -1180,9 +1144,7 @@ class TestAnnotationService: annotations.append(annotation) # Export annotation list - exported_annotations = AppAnnotationService.export_annotation_list_by_app_id( - app.id, session=db_session_with_containers - ) + exported_annotations = AppAnnotationService.export_annotation_list_by_app_id(app.id, db_session_with_containers) # Verify results assert len(exported_annotations) == 3 @@ -1209,9 +1171,7 @@ 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, session=db_session_with_containers - ) + AppAnnotationService.export_annotation_list_by_app_id(non_existent_app_id, db_session_with_containers) def test_insert_app_annotation_directly_with_setting_success( self, db_session_with_containers: Session, mock_external_service_dependencies 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 1f87acfc2cb..6a9d7979931 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 @@ -820,17 +820,19 @@ class TestAppDslService: db_session_with_containers.commit() service = AppDslService(db_session_with_containers) - service._create_or_update_app( - app=app, - data={ - "app": {"mode": AppMode.CHAT}, - "model_config": {"model": {"provider": "openai"}}, - }, - account=account, - ) + with patch("services.app_dsl_service.app_model_config_was_updated") as signal: + service._create_or_update_app( + app=app, + data={ + "app": {"mode": AppMode.CHAT}, + "model_config": {"model": {"provider": "openai"}}, + }, + account=account, + ) db_session_with_containers.expire_all() assert app.app_model_config_id is not None + assert signal.send.call_args.kwargs["session"] is db_session_with_containers def test_create_or_update_app_invalid_mode_raises(self, db_session_with_containers: Session): service = AppDslService(db_session_with_containers) @@ -1245,17 +1247,29 @@ class TestAppDslService: ) monkeypatch.setattr(app_dsl_service, "jsonable_encoder", lambda x: x) - app_model_config = SimpleNamespace(to_dict=lambda: {"agent_mode": {"tools": [{"credential_id": "secret"}]}}) - app_model = _app_stub(app_model_config=app_model_config) + app_model_config = MagicMock(app_id="app-1") + app_model_config.to_dict.return_value = {"agent_mode": {"tools": [{"credential_id": "secret"}]}} + app_model = _app_stub(id="app-1", app_model_config_id="config-1") + session = MagicMock(spec=Session) + session.get.return_value = app_model_config + annotation_reply = {"enabled": False} + monkeypatch.setattr(app_dsl_service, "load_annotation_reply_config", lambda *_args: annotation_reply) export_data: dict = {} - AppDslService._append_model_config_export_data(export_data, app_model) + AppDslService._append_model_config_export_data(export_data, app_model, session=session) assert export_data["model_config"]["agent_mode"]["tools"] == [{}] assert export_data["dependencies"] == [{"tenant": _DEFAULT_TENANT_ID, "dep": "dep-1"}] + session.get.assert_called_once_with(AppModelConfig, "config-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) def test_append_model_config_export_data_requires_app_config(self): + session = MagicMock(spec=Session) + session.get.return_value = None with pytest.raises(ValueError, match="Missing app configuration"): - AppDslService._append_model_config_export_data({}, _app_stub(app_model_config=None)) + AppDslService._append_model_config_export_data( + {}, _app_stub(app_model_config_id="config-1"), session=session + ) + session.get.assert_called_once_with(AppModelConfig, "config-1") # ── Dependency Extraction ───────────────────────────────────────── 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 8deaf6d462d..2875f014380 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 @@ -191,7 +191,7 @@ class TestAppService: mock_current_user.current_tenant_id = account.current_tenant_id with patch("services.app_service.current_user", mock_current_user): - retrieved_app = app_service.get_app(created_app) + retrieved_app = app_service.get_app(created_app, session=db_session_with_containers) # Verify retrieved app matches created app assert retrieved_app.id == created_app.id @@ -1595,13 +1595,13 @@ class TestAppService: def test_get_app_meta_returns_empty_when_workflow_missing( self, db_session_with_containers: Session, mock_external_service_dependencies ): - """Test get_app_meta returns empty tool_icons when workflow is None.""" + """Test get_app_meta returns empty tool_icons when the workflow ID is absent.""" from types import SimpleNamespace from services.app_service import AppService app_service = AppService() - workflow_app = SimpleNamespace(mode="workflow", workflow=None) + workflow_app = SimpleNamespace(mode="workflow", workflow_id=None) meta = app_service.get_app_meta(workflow_app, session=db_session_with_containers) assert meta == {"tool_icons": {}} @@ -1609,13 +1609,13 @@ class TestAppService: def test_get_app_meta_returns_empty_when_model_config_missing( self, db_session_with_containers: Session, mock_external_service_dependencies ): - """Test get_app_meta returns empty tool_icons when app_model_config is None.""" + """Test get_app_meta returns empty tool_icons when the model config ID is absent.""" from types import SimpleNamespace from services.app_service import AppService app_service = AppService() - chat_app = SimpleNamespace(mode="chat", app_model_config=None) + chat_app = SimpleNamespace(mode="chat", app_model_config_id=None) 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_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_conversation_service.py index 60858df08e8..09411f0b0e1 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 @@ -853,11 +853,7 @@ class TestConversationServiceMessageAnnotation: # Act result_items, result_total = AppAnnotationService.get_annotation_list_by_app_id( - app_id=app_model.id, - page=1, - limit=10, - keyword="", - session=db_session_with_containers, + app_id=app_model.id, page=1, limit=10, keyword="", session=db_session_with_containers ) # Assert @@ -934,9 +930,7 @@ class TestConversationServiceMessageAnnotation: } # Act - result = AppAnnotationService.insert_app_annotation_directly( - args, app_model.id, session=db_session_with_containers - ) + result = AppAnnotationService.insert_app_annotation_directly(args, app_model.id, db_session_with_containers) # Assert assert result.question == args["question"] @@ -1011,7 +1005,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, session=db_session_with_containers) + result = AppAnnotationService.export_annotation_list_by_app_id(app_model.id, db_session_with_containers) # Assert assert len(result) == 10 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py index 0df4ece5d2d..a20367a9b4f 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py @@ -288,10 +288,10 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: ) assert result == [] - def test_update_partial_member_list_database_error_rollback(self, db_session_with_containers: Session): - """ - Test error handling and rollback on database error. - """ + def test_update_partial_member_list_database_error_requires_caller_rollback( + self, db_session_with_containers: Session + ): + """The transaction owner rolls back a failed partial-member update.""" # Arrange owner, tenant = DatasetPermissionTestDataFactory.create_account_with_tenant(role=TenantAccountRole.OWNER) existing_member, _ = DatasetPermissionTestDataFactory.create_account_with_tenant( @@ -309,32 +309,28 @@ class TestDatasetPermissionServiceUpdatePartialMemberList: DatasetPermissionTestDataFactory.build_user_list_payload([existing_member.id]), session=db_session_with_containers, ) + db_session_with_containers.commit() user_list = DatasetPermissionTestDataFactory.build_user_list_payload([replacement_member.id]) - rollback_called = {"count": 0} - original_rollback = db_session_with_containers.rollback + original_flush = db_session_with_containers.flush # Act / Assert with pytest.MonkeyPatch.context() as mp: - def _raise_commit(): + def _raise_flush(): raise Exception("Database connection error") - def _rollback_and_mark(): - rollback_called["count"] += 1 - original_rollback() - - mp.setattr(db_session_with_containers, "commit", _raise_commit) - mp.setattr(db_session_with_containers, "rollback", _rollback_and_mark) + mp.setattr(db_session_with_containers, "flush", _raise_flush) with pytest.raises(Exception, match="Database connection error"): DatasetPermissionService.update_partial_member_list( tenant.id, dataset.id, user_list, session=db_session_with_containers ) + mp.setattr(db_session_with_containers, "flush", original_flush) # Assert + db_session_with_containers.rollback() result = DatasetPermissionService.get_dataset_partial_member_list( dataset.id, session=db_session_with_containers ) - assert rollback_called["count"] == 1 assert result == [existing_member.id] assert db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).count() == 1 @@ -388,10 +384,10 @@ class TestDatasetPermissionServiceClearPartialMemberList: ) assert result == [] - def test_clear_partial_member_list_database_error_rollback(self, db_session_with_containers: Session): - """ - Test error handling and rollback on database error. - """ + def test_clear_partial_member_list_database_error_requires_caller_rollback( + self, db_session_with_containers: Session + ): + """The transaction owner rolls back a failed partial-member clear.""" # Arrange owner, tenant = DatasetPermissionTestDataFactory.create_account_with_tenant(role=TenantAccountRole.OWNER) member_1, _ = DatasetPermissionTestDataFactory.create_account_with_tenant( @@ -407,29 +403,25 @@ class TestDatasetPermissionServiceClearPartialMemberList: DatasetPermissionService.update_partial_member_list( tenant.id, dataset.id, users, session=db_session_with_containers ) - rollback_called = {"count": 0} - original_rollback = db_session_with_containers.rollback + db_session_with_containers.commit() + original_flush = db_session_with_containers.flush # Act / Assert with pytest.MonkeyPatch.context() as mp: - def _raise_commit(): + def _raise_flush(): raise Exception("Database connection error") - def _rollback_and_mark(): - rollback_called["count"] += 1 - original_rollback() - - mp.setattr(db_session_with_containers, "commit", _raise_commit) - mp.setattr(db_session_with_containers, "rollback", _rollback_and_mark) + mp.setattr(db_session_with_containers, "flush", _raise_flush) with pytest.raises(Exception, match="Database connection error"): DatasetPermissionService.clear_partial_member_list(dataset.id, session=db_session_with_containers) + mp.setattr(db_session_with_containers, "flush", original_flush) # Assert + db_session_with_containers.rollback() result = DatasetPermissionService.get_dataset_partial_member_list( dataset.id, session=db_session_with_containers ) - assert rollback_called["count"] == 1 assert set(result) == {member_1.id, member_2.id} assert db_session_with_containers.query(DatasetPermission).filter_by(dataset_id=dataset.id).count() == 2 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_get_segments.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_get_segments.py index bd8f5371b8e..b4fe36d07c4 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_get_segments.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_get_segments.py @@ -173,7 +173,9 @@ class TestSegmentServiceGetSegments: ) # Act - items, total = SegmentService.get_segments(document_id=document.id, tenant_id=tenant.id, page=1, limit=20) + items, total = SegmentService.get_segments( + document_id=document.id, tenant_id=tenant.id, page=1, limit=20, session=db_session_with_containers + ) # Assert assert len(items) == 2 @@ -226,7 +228,10 @@ class TestSegmentServiceGetSegments: # Act items, total = SegmentService.get_segments( - document_id=document.id, tenant_id=tenant.id, status_list=["completed", "indexing"] + document_id=document.id, + tenant_id=tenant.id, + status_list=["completed", "indexing"], + session=db_session_with_containers, ) # Assert @@ -270,7 +275,9 @@ class TestSegmentServiceGetSegments: ) # Act - items, total = SegmentService.get_segments(document_id=document.id, tenant_id=tenant.id, status_list=[]) + items, total = SegmentService.get_segments( + document_id=document.id, tenant_id=tenant.id, status_list=[], session=db_session_with_containers + ) # Assert — empty status_list should return all segments (no status filter applied) assert len(items) == 2 @@ -311,7 +318,9 @@ class TestSegmentServiceGetSegments: ) # Act - items, total = SegmentService.get_segments(document_id=document.id, tenant_id=tenant.id, keyword="search term") + items, total = SegmentService.get_segments( + document_id=document.id, tenant_id=tenant.id, keyword="search term", session=db_session_with_containers + ) # Assert assert len(items) == 1 @@ -364,7 +373,9 @@ class TestSegmentServiceGetSegments: ) # Act - items, total = SegmentService.get_segments(document_id=document.id, tenant_id=tenant.id) + items, total = SegmentService.get_segments( + document_id=document.id, tenant_id=tenant.id, session=db_session_with_containers + ) # Assert — segments should be ordered by position ASC assert len(items) == 3 @@ -386,7 +397,9 @@ class TestSegmentServiceGetSegments: non_existent_doc_id = str(uuid4()) # Act - items, total = SegmentService.get_segments(document_id=non_existent_doc_id, tenant_id=tenant.id) + items, total = SegmentService.get_segments( + document_id=non_existent_doc_id, tenant_id=tenant.id, session=db_session_with_containers + ) # Assert assert items == [] @@ -447,6 +460,7 @@ class TestSegmentServiceGetSegments: keyword="important", page=1, limit=10, + session=db_session_with_containers, ) # Assert — only the first segment matches both filters @@ -494,6 +508,7 @@ class TestSegmentServiceGetSegments: document_id=document.id, tenant_id=tenant.id, status_list=None, + session=db_session_with_containers, ) # Assert — None status_list should return all segments @@ -532,6 +547,7 @@ class TestSegmentServiceGetSegments: document_id=document.id, tenant_id=tenant.id, limit=200, + session=db_session_with_containers, ) # Assert — total is 105, but items per page capped at 100 diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py index b6768d2ca26..2ee72415137 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_retrieval.py @@ -582,7 +582,9 @@ class TestDatasetServiceGetDatasetsByIds: dataset_ids = [dataset.id for dataset in datasets] # Act - result_datasets, total = DatasetService.get_datasets_by_ids(dataset_ids, tenant.id) + result_datasets, total = DatasetService.get_datasets_by_ids( + dataset_ids, tenant.id, session=db_session_with_containers + ) # Assert assert len(result_datasets) == 3 @@ -596,7 +598,7 @@ class TestDatasetServiceGetDatasetsByIds: dataset_ids = [] # Act - datasets, total = DatasetService.get_datasets_by_ids(dataset_ids, tenant_id) + datasets, total = DatasetService.get_datasets_by_ids(dataset_ids, tenant_id, session=db_session_with_containers) # Assert assert datasets == [] @@ -608,7 +610,7 @@ class TestDatasetServiceGetDatasetsByIds: tenant_id = str(uuid4()) # Act - datasets, total = DatasetService.get_datasets_by_ids(None, tenant_id) + datasets, total = DatasetService.get_datasets_by_ids(None, tenant_id, session=db_session_with_containers) # Assert assert datasets == [] @@ -684,7 +686,7 @@ class TestDatasetServiceGetDatasetQueries: ) # Act - queries, total = DatasetService.get_dataset_queries(dataset.id, page, per_page) + queries, total = DatasetService.get_dataset_queries(dataset.id, page, per_page, db_session_with_containers) # Assert assert len(queries) == 3 @@ -702,7 +704,7 @@ class TestDatasetServiceGetDatasetQueries: per_page = 20 # Act - queries, total = DatasetService.get_dataset_queries(dataset.id, page, per_page) + queries, total = DatasetService.get_dataset_queries(dataset.id, page, per_page, db_session_with_containers) # Assert assert queries == [] 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 4106b68545c..73a79f1ae50 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 @@ -173,7 +173,9 @@ class TestHitTestingService: assert response.query.content == query assert len(response.records) == 1 assert response.records[0].content == "formatted content" - mock_format.assert_called_once_with([mock_doc]) + mock_format.assert_called_once() + assert mock_format.call_args.args[0] is not db_session_with_containers + assert mock_format.call_args.args[1] == [mock_doc] def test_compact_external_retrieve_response_should_return_records_for_external_provider( self, db_session_with_containers: Session 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 d397c62b6a8..7ea52ff6fef 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 @@ -157,21 +157,22 @@ class TestOAuthServerServiceTokenOperations: ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN, ) - def test_validate_access_token_returns_none_when_not_found(self, mock_redis): + def test_validate_access_token_returns_none_when_not_found(self, mock_redis, db_session_with_containers: Session): mock_redis.get.return_value = None session = MagicMock() - result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session) + result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", db_session_with_containers) assert result is None - def test_validate_access_token_loads_user_when_exists(self, mock_redis): + def test_validate_access_token_loads_user_when_exists(self, mock_redis, db_session_with_containers: Session): 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", session) + result = OAuthServerService.validate_oauth_access_token( + "client-1", "access-token", db_session_with_containers + ) assert result is expected_user - mock_load.assert_called_once_with("user-88", session) + mock_load.assert_called_once_with("user-88", db_session_with_containers) diff --git a/api/tests/test_containers_integration_tests/tasks/test_clean_notion_document_task.py b/api/tests/test_containers_integration_tests/tasks/test_clean_notion_document_task.py index f6e03b84c90..b8327db67f4 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_clean_notion_document_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_clean_notion_document_task.py @@ -239,6 +239,7 @@ class TestCleanNotionDocumentTask: # args: (dataset, total_index_node_ids) assert isinstance(args[0], Dataset) assert args[1] == [] + assert kwargs["session"] is not None def test_clean_notion_document_task_with_different_index_types( self, db_session_with_containers: Session, mock_index_processor_factory, mock_external_service_dependencies @@ -898,8 +899,10 @@ class TestCleanNotionDocumentTask: # Verify all data exists before cleanup # Note: There may be documents from previous tests, so we check for at least 3 - assert db_session_with_containers.scalar(select(func.count()).select_from(Document)) >= 3 - assert db_session_with_containers.scalar(select(func.count()).select_from(DocumentSegment)) >= 9 + document_count = db_session_with_containers.scalar(select(func.count()).select_from(Document)) or 0 + segment_count = db_session_with_containers.scalar(select(func.count()).select_from(DocumentSegment)) or 0 + assert document_count >= 3 + assert segment_count >= 9 # Clean up documents from only the first dataset target_dataset = datasets[0] diff --git a/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py index 5287cd06dbc..7f102f7375f 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py @@ -213,6 +213,10 @@ class TestDatasetIndexingTaskIntegration: assert len(opened) >= 2 assert opened_ids <= closed_ids + def _runner_documents_arg(self, patched_external_dependencies) -> Sequence[Document]: + """Return the document batch passed to the runner.""" + return patched_external_dependencies["indexing_runner_instance"].run.call_args.args[0] + def test_legacy_document_indexing_task_still_works( self, db_session_with_containers: Session, patched_external_dependencies ): @@ -241,7 +245,7 @@ class TestDatasetIndexingTaskIntegration: # Assert patched_external_dependencies["indexing_runner_instance"].run.assert_called_once() - run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0] + run_args = self._runner_documents_arg(patched_external_dependencies) assert len(run_args) == len(document_ids) self._assert_documents_parsing(db_session_with_containers, document_ids) @@ -298,7 +302,8 @@ class TestDatasetIndexingTaskIntegration: _document_indexing(dataset.id, []) # Assert - patched_external_dependencies["indexing_runner_instance"].run.assert_called_once_with([]) + patched_external_dependencies["indexing_runner_instance"].run.assert_called_once() + assert self._runner_documents_arg(patched_external_dependencies) == [] def test_tenant_queue_dispatches_next_task_after_completion( self, db_session_with_containers: Session, patched_external_dependencies @@ -512,7 +517,7 @@ class TestDatasetIndexingTaskIntegration: _document_indexing(dataset.id, mixed_ids) # Assert - run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0] + run_args = self._runner_documents_arg(patched_external_dependencies) assert len(run_args) == 2 self._assert_documents_parsing(db_session_with_containers, existing_ids) @@ -602,7 +607,7 @@ class TestDatasetIndexingTaskIntegration: _document_indexing(dataset.id, large_document_ids) # Assert - run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0] + run_args = self._runner_documents_arg(patched_external_dependencies) assert len(run_args) == 100 self._assert_documents_parsing(db_session_with_containers, large_document_ids) @@ -662,7 +667,7 @@ class TestDatasetIndexingTaskIntegration: _document_indexing(dataset.id, [document_id]) # Assert - run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0] + run_args = self._runner_documents_arg(patched_external_dependencies) assert len(run_args) == 1 self._assert_documents_parsing(db_session_with_containers, [document_id]) @@ -746,6 +751,6 @@ class TestDatasetIndexingTaskIntegration: _document_indexing(dataset.id, document_ids) # Assert - run_args = patched_external_dependencies["indexing_runner_instance"].run.call_args[0][0] + run_args = self._runner_documents_arg(patched_external_dependencies) assert len(run_args) == batch_limit self._assert_documents_parsing(db_session_with_containers, document_ids) diff --git a/api/tests/test_containers_integration_tests/tasks/test_deal_dataset_vector_index_task.py b/api/tests/test_containers_integration_tests/tasks/test_deal_dataset_vector_index_task.py index fbdd81aed5d..b0704ae9303 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_deal_dataset_vector_index_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_deal_dataset_vector_index_task.py @@ -7,7 +7,7 @@ add, update, and remove actions. """ import uuid -from unittest.mock import ANY, Mock, patch +from unittest.mock import Mock, patch import pytest from faker import Faker @@ -330,7 +330,13 @@ class TestDealDatasetVectorIndexTask: # Verify index processor clean and load methods were called mock_factory = mock_index_processor_factory.return_value mock_processor = mock_factory.init_index_processor.return_value - mock_processor.clean.assert_called_once_with(ANY, None, with_keywords=False, delete_child_chunks=False) + mock_processor.clean.assert_called_once() + clean_args, clean_kwargs = mock_processor.clean.call_args + assert clean_args[0].id == dataset.id + assert clean_args[1] is None + assert clean_kwargs["with_keywords"] is False + assert clean_kwargs["delete_child_chunks"] is False + assert clean_kwargs["session"] is not None mock_processor.load.assert_called_once() def test_deal_dataset_vector_index_task_dataset_not_found_error( @@ -470,7 +476,13 @@ class TestDealDatasetVectorIndexTask: # Verify that index processor clean was called but no load mock_factory = mock_index_processor_factory.return_value mock_processor = mock_factory.init_index_processor.return_value - mock_processor.clean.assert_called_once_with(ANY, None, with_keywords=False, delete_child_chunks=False) + mock_processor.clean.assert_called_once() + clean_args, clean_kwargs = mock_processor.clean.call_args + assert clean_args[0].id == dataset.id + assert clean_args[1] is None + assert clean_kwargs["with_keywords"] is False + assert clean_kwargs["delete_child_chunks"] is False + assert clean_kwargs["session"] is not None mock_processor.load.assert_not_called() def test_deal_dataset_vector_index_task_add_action_with_exception_handling( diff --git a/api/tests/test_containers_integration_tests/tasks/test_delete_segment_from_index_task.py b/api/tests/test_containers_integration_tests/tasks/test_delete_segment_from_index_task.py index a7edf4f77a2..edd23ca2b68 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_delete_segment_from_index_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_delete_segment_from_index_task.py @@ -11,6 +11,7 @@ import logging from unittest.mock import MagicMock, patch from faker import Faker +from sqlalchemy import inspect from sqlalchemy.orm import Session from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType @@ -502,6 +503,7 @@ class TestDeleteSegmentFromIndexTask: index_node_ids = [segment.index_node_id for segment in segments] segment_ids = [segment.id for segment in segments] + expected_dataset_id = dataset.id # Mock the index processor to raise an exception mock_processor = MagicMock() @@ -517,7 +519,7 @@ class TestDeleteSegmentFromIndexTask: # Verify index processor clean method was called assert mock_processor.clean.call_count == 1 call_args = mock_processor.clean.call_args - assert call_args[0][0].id == dataset.id # Verify dataset ID matches + assert inspect(call_args[0][0]).identity == (expected_dataset_id,) assert call_args[0][1] == index_node_ids # Verify index node IDs match assert call_args[1]["with_keywords"] is True assert call_args[1]["delete_child_chunks"] is True diff --git a/api/tests/test_containers_integration_tests/tasks/test_disable_segment_from_index_task.py b/api/tests/test_containers_integration_tests/tasks/test_disable_segment_from_index_task.py index 34e2ce4e80d..a5577217e81 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_disable_segment_from_index_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_disable_segment_from_index_task.py @@ -456,8 +456,8 @@ class TestDisableSegmentFromIndexTask: # Verify index processor was called mock_index_processor.clean.assert_called_once() call_args = mock_index_processor.clean.call_args - # Check that the call was made with the correct parameters - assert len(call_args[0]) == 2 # Check two arguments were passed + # Check that the call was made with the correct parameters. + assert len(call_args[0]) == 2 assert call_args[0][1] == [segment.index_node_id] # Check index node IDs # Verify segment was re-enabled diff --git a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py index 6c1454b6d87..b58b2f01da4 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py @@ -51,6 +51,19 @@ class TestDocumentIndexingTasks: "features": mock_features, } + def _runner_documents_arg(self, mock_external_service_dependencies) -> list[Document]: + """Return the document batch passed to the runner.""" + return mock_external_service_dependencies["indexing_runner_instance"].run.call_args.args[0] + + def _assert_documents_parsing(self, db_session_with_containers: Session, document_ids: list[str]) -> None: + """Assert the short status transaction remains committed when the runner exits early.""" + db_session_with_containers.expire_all() + for doc_id in document_ids: + updated_document = db_session_with_containers.query(Document).where(Document.id == doc_id).first() + assert updated_document is not None + assert updated_document.indexing_status == IndexingStatus.PARSING + assert updated_document.processing_started_at is not None + def _create_test_dataset_and_documents( self, db_session_with_containers: Session, mock_external_service_dependencies, document_count=3 ): @@ -261,7 +274,7 @@ class TestDocumentIndexingTasks: # Verify the run method was called with correct documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] # First argument should be documents list + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 3 def test_document_indexing_task_dataset_not_found( @@ -331,7 +344,7 @@ class TestDocumentIndexingTasks: # Verify the run method was called with only existing documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] # First argument should be documents list + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 2 # Only existing documents def test_document_indexing_task_indexing_runner_exception( @@ -368,12 +381,8 @@ class TestDocumentIndexingTasks: mock_external_service_dependencies["indexing_runner"].assert_called_once() mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once() - # Verify documents were still updated to parsing status before the exception - # Re-query documents from database since _document_indexing close the session - for doc_id in document_ids: - updated_document = db_session_with_containers.query(Document).where(Document.id == doc_id).first() - assert updated_document.indexing_status == IndexingStatus.PARSING - assert updated_document.processing_started_at is not None + # Parsing status is committed before the runner starts, so runner failure cannot roll it back. + self._assert_documents_parsing(db_session_with_containers, document_ids) def test_document_indexing_task_mixed_document_states( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -455,7 +464,7 @@ class TestDocumentIndexingTasks: # Verify the run method was called with all documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] # First argument should be documents list + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 4 def test_document_indexing_task_billing_sandbox_plan_batch_limit( @@ -592,12 +601,8 @@ class TestDocumentIndexingTasks: mock_external_service_dependencies["indexing_runner"].assert_called_once() mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once() - # Verify documents were still updated to parsing status before the exception - # Re-query documents from database since _document_indexing uses a different session - for doc_id in document_ids: - updated_document = db_session_with_containers.query(Document).where(Document.id == doc_id).first() - assert updated_document.indexing_status == IndexingStatus.PARSING - assert updated_document.processing_started_at is not None + # The pause stops the runner but does not undo the earlier parsing-status transaction. + self._assert_documents_parsing(db_session_with_containers, document_ids) # ==================== NEW TESTS FOR REFACTORED FUNCTIONS ==================== def test_old_document_indexing_task_success( @@ -715,7 +720,7 @@ class TestDocumentIndexingTasks: # Verify the run method was called with correct documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 2 # Verify task function was not called (no waiting tasks) @@ -830,12 +835,8 @@ class TestDocumentIndexingTasks: mock_external_service_dependencies["indexing_runner"].assert_called_once() mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once() - # Verify documents were still updated to parsing status before the exception - # Re-query documents from database since _document_indexing uses a different session - for doc_id in document_ids: - updated_document = db_session_with_containers.query(Document).where(Document.id == doc_id).first() - assert updated_document.indexing_status == IndexingStatus.PARSING - assert updated_document.processing_started_at is not None + # The core indexing error does not undo the earlier parsing status, and the tenant queue still advances. + self._assert_documents_parsing(db_session_with_containers, document_ids) # Verify waiting task was still processed despite core processing error mock_task_func.apply_async.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_update_task.py b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_update_task.py index 208fc1aa1d7..bda54389419 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_update_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_update_task.py @@ -139,9 +139,9 @@ class TestDocumentIndexingUpdateTask: clean_call = mock_external_dependencies["processor"].clean.call_args assert clean_call is not None args, kwargs = clean_call - # args[0] is a Dataset instance (from another session) — validate by id + # args[0] is a Dataset instance (from another session), so validate by id. assert getattr(args[0], "id", None) == dataset.id - # args[1] should contain our node_ids + # args[1] should contain our node_ids. assert set(args[1]) == set(node_ids) assert kwargs.get("with_keywords") is True assert kwargs.get("delete_child_chunks") is True diff --git a/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py index e1c7e3e09a6..74199255e96 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py @@ -62,6 +62,19 @@ class TestDuplicateDocumentIndexingTasks: "index_processor": mock_processor, } + def _runner_documents_arg(self, mock_external_service_dependencies) -> list[Document]: + """Return the document batch passed to the runner.""" + return mock_external_service_dependencies["indexing_runner_instance"].run.call_args.args[0] + + def _assert_documents_parsing(self, db_session_with_containers: Session, document_ids: list[str]) -> None: + """Assert the short status transaction remains committed when the runner exits early.""" + db_session_with_containers.expire_all() + for doc_id in document_ids: + updated_document = db_session_with_containers.scalar(select(Document).where(Document.id == doc_id).limit(1)) + assert updated_document is not None + assert updated_document.indexing_status == IndexingStatus.PARSING + assert updated_document.processing_started_at is not None + def _create_test_dataset_and_documents( self, db_session_with_containers: Session, mock_external_service_dependencies, document_count=3 ): @@ -329,7 +342,7 @@ class TestDuplicateDocumentIndexingTasks: # Verify the run method was called with correct documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] # First argument should be documents list + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 3 def _test_duplicate_document_indexing_task_with_segment_cleanup( @@ -450,7 +463,7 @@ class TestDuplicateDocumentIndexingTasks: # Verify the run method was called with only existing documents call_args = mock_external_service_dependencies["indexing_runner_instance"].run.call_args assert call_args is not None - processed_documents = call_args[0][0] # First argument should be documents list + processed_documents = self._runner_documents_arg(mock_external_service_dependencies) assert len(processed_documents) == 2 # Only existing documents def _test_duplicate_document_indexing_task_indexing_runner_exception( @@ -487,12 +500,8 @@ class TestDuplicateDocumentIndexingTasks: mock_external_service_dependencies["indexing_runner"].assert_called_once() mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once() - # Verify documents were still updated to parsing status before the exception - # Re-query documents from database since _duplicate_document_indexing_task close the session - for doc_id in document_ids: - updated_document = db_session_with_containers.scalar(select(Document).where(Document.id == doc_id).limit(1)) - assert updated_document.indexing_status == IndexingStatus.PARSING - assert updated_document.processing_started_at is not None + # Parsing status is committed before the runner starts, so runner failure cannot roll it back. + self._assert_documents_parsing(db_session_with_containers, document_ids) def _test_duplicate_document_indexing_task_billing_sandbox_plan_batch_limit( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -621,9 +630,10 @@ class TestDuplicateDocumentIndexingTasks: _duplicate_document_indexing_task(dataset.id, document_ids) # Assert: Verify IndexingRunner was called with empty list - # Note: The actual implementation does call run([]) with empty list + # Note: The actual implementation does call run([]) with an empty list. mock_external_service_dependencies["indexing_runner"].assert_called_once() - mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once_with([]) + mock_external_service_dependencies["indexing_runner_instance"].run.assert_called_once() + assert mock_external_service_dependencies["indexing_runner_instance"].run.call_args.args[0] == [] def test_deprecated_duplicate_document_indexing_task_delegates_to_core( self, db_session_with_containers: Session, mock_external_service_dependencies diff --git a/api/tests/unit_tests/commands/test_fix_app_site_missing.py b/api/tests/unit_tests/commands/test_fix_app_site_missing.py new file mode 100644 index 00000000000..a7b05e3bbb1 --- /dev/null +++ b/api/tests/unit_tests/commands/test_fix_app_site_missing.py @@ -0,0 +1,77 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from sqlalchemy.orm import Session + +from commands import system as system_commands + + +def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytest.MonkeyPatch) -> None: + account = object() + tenant = MagicMock() + tenant.get_accounts.return_value = [account] + app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + + session = Session() + phase_events: list[str] = [] + scalar = MagicMock(return_value=app) + get = MagicMock(return_value=tenant) + commit = MagicMock(side_effect=lambda: phase_events.append("commit")) + monkeypatch.setattr(session, "scalar", scalar) + monkeypatch.setattr(session, "get", get) + monkeypatch.setattr(session, "commit", commit) + + scoped_session = MagicMock(return_value=session) + scoped_session.scalar.return_value = app + + connection = MagicMock() + connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] + engine = MagicMock() + engine.begin.return_value.__enter__.return_value = connection + + monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=scoped_session)) + send = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal")) + monkeypatch.setattr(system_commands.app_was_created, "send", send) + + system_commands.fix_app_site_missing.callback() + + scoped_session.assert_called_once_with() + scalar.assert_called_once() + get.assert_called_once_with(system_commands.Tenant, app.tenant_id) + tenant.get_accounts.assert_called_once_with(session=session) + send.assert_called_once_with(app, account=account, session=session) + commit.assert_called_once_with() + assert phase_events == ["signal", "commit"] + assert isinstance(send.call_args.kwargs["session"], Session) + + +def test_fix_app_site_missing_rolls_back_when_signal_fails(monkeypatch: pytest.MonkeyPatch) -> None: + account = object() + tenant = MagicMock() + tenant.get_accounts.return_value = [account] + app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + session = MagicMock() + phase_events: list[str] = [] + session.scalar.return_value = app + session.get.return_value = tenant + session.rollback.side_effect = lambda: phase_events.append("rollback") + + connection = MagicMock() + connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] + engine = MagicMock() + engine.begin.return_value.__enter__.return_value = connection + + monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=MagicMock(return_value=session))) + + def fail_signal(*_args, **_kwargs) -> None: + phase_events.append("signal") + raise RuntimeError("failed") + + monkeypatch.setattr(system_commands.app_was_created, "send", MagicMock(side_effect=fail_signal)) + + system_commands.fix_app_site_missing.callback() + + session.rollback.assert_called_once_with() + session.commit.assert_not_called() + assert phase_events == ["signal", "rollback"] diff --git a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py index b4639d54d2f..a57ce6655b9 100644 --- a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py +++ b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py @@ -1,16 +1,25 @@ from types import SimpleNamespace +from unittest.mock import MagicMock import pytest -from controllers.common import agent_app_parameters from controllers.common.agent_app_parameters import get_published_agent_app_feature_dict_and_user_input_form from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError -def test_published_agent_app_parameters_use_soul_file_upload(monkeypatch): +def _app_model(*, bound_agent_id: str | None, app_model_config=None): + return SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + bound_agent_id=bound_agent_id, + app_model_config_with_session=lambda *, session: app_model_config, + ) + + +def test_published_agent_app_parameters_use_soul_file_upload(): app_model_config = SimpleNamespace( - to_dict=lambda: { + to_dict=lambda **_kwargs: { "opening_statement": "Hi from legacy presentation config", "file_upload": { "enabled": False, @@ -18,11 +27,7 @@ def test_published_agent_app_parameters_use_soul_file_upload(monkeypatch): }, } ) - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=app_model_config, - ) + app_model = _app_model(bound_agent_id="agent-1", app_model_config=app_model_config) agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="snapshot-1", @@ -43,10 +48,13 @@ def test_published_agent_app_parameters_use_soul_file_upload(monkeypatch): "app_variables": [{"name": "topic", "type": "string", "required": True}], } ) - query_results = iter([agent, snapshot]) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: next(query_results)) + session = MagicMock() + session.scalar.side_effect = [agent, snapshot, None] - features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(app_model) + features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form( + app_model, + session=session, + ) parameters = get_parameters_from_feature_dict(features_dict=features_dict, user_input_form=user_input_form) assert parameters["opening_statement"] == "Hi from legacy presentation config" @@ -62,26 +70,19 @@ def test_published_agent_app_parameters_use_soul_file_upload(monkeypatch): def test_published_agent_app_parameters_requires_bound_agent(): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id=None, - app_model_config=None, - ) + app_model = _app_model(bound_agent_id=None) with pytest.raises(AgentAppGeneratorError, match="no bound Agent"): - get_published_agent_app_feature_dict_and_user_input_form(app_model) + get_published_agent_app_feature_dict_and_user_input_form(app_model, session=MagicMock()) -def test_published_agent_app_parameters_requires_existing_active_agent(monkeypatch): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=None, - ) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: None) +def test_published_agent_app_parameters_requires_existing_active_agent(): + app_model = _app_model(bound_agent_id="agent-1") + session = MagicMock() + session.scalar.return_value = None with pytest.raises(AgentAppGeneratorError, match="no bound Agent"): - get_published_agent_app_feature_dict_and_user_input_form(app_model) + get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session) @pytest.mark.parametrize( @@ -91,78 +92,69 @@ def test_published_agent_app_parameters_requires_existing_active_agent(monkeypat False, ], ) -def test_published_agent_app_parameters_requires_published_agent(monkeypatch, active_config_is_published): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=None, - ) +def test_published_agent_app_parameters_requires_published_agent(active_config_is_published): + app_model = _app_model(bound_agent_id="agent-1") agent = SimpleNamespace( id="agent-1", active_config_snapshot_id=None, active_config_is_published=active_config_is_published, ) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: agent) + session = MagicMock() + session.scalar.return_value = agent with pytest.raises(AgentAppNotPublishedError, match="not been published"): - get_published_agent_app_feature_dict_and_user_input_form(app_model) + get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session) -def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot(monkeypatch): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=None, - ) +def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot(): + app_model = _app_model(bound_agent_id="agent-1") agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="snapshot-1", active_config_is_published=False, ) snapshot = SimpleNamespace(config_snapshot_dict={}) - query_results = iter([agent, snapshot]) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: next(query_results)) + session = MagicMock() + session.scalar.side_effect = [agent, snapshot] - features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(app_model) + features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form( + app_model, + session=session, + ) assert features_dict["file_upload"]["enabled"] is True assert user_input_form == [] -def test_published_agent_app_parameters_requires_published_snapshot(monkeypatch): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=None, - ) +def test_published_agent_app_parameters_requires_published_snapshot(): + app_model = _app_model(bound_agent_id="agent-1") agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="snapshot-1", active_config_is_published=True, ) - query_results = iter([agent, None]) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: next(query_results)) + session = MagicMock() + session.scalar.side_effect = [agent, None] with pytest.raises(AgentAppGeneratorError, match="published version not found"): - get_published_agent_app_feature_dict_and_user_input_form(app_model) + get_published_agent_app_feature_dict_and_user_input_form(app_model, session=session) -def test_published_agent_app_parameters_allows_missing_legacy_app_model_config(monkeypatch): - app_model = SimpleNamespace( - tenant_id="tenant-1", - bound_agent_id="agent-1", - app_model_config=None, - ) +def test_published_agent_app_parameters_allows_missing_legacy_app_model_config(): + app_model = _app_model(bound_agent_id="agent-1") agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="snapshot-1", active_config_is_published=True, ) snapshot = SimpleNamespace(config_snapshot_dict={}) - query_results = iter([agent, snapshot]) - monkeypatch.setattr(agent_app_parameters.db.session, "scalar", lambda _: next(query_results)) + session = MagicMock() + session.scalar.side_effect = [agent, snapshot] - features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form(app_model) + features_dict, user_input_form = get_published_agent_app_feature_dict_and_user_input_form( + app_model, + session=session, + ) assert features_dict["file_upload"] == { "allowed_file_extensions": ["JPG", "JPEG", "PNG", "GIF", "WEBP", "SVG"], 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 60a576346a0..df5debb24d6 100644 --- a/api/tests/unit_tests/controllers/common/test_app_access.py +++ b/api/tests/unit_tests/controllers/common/test_app_access.py @@ -2,6 +2,8 @@ from __future__ import annotations +from unittest.mock import MagicMock + import pytest from controllers.common.app_access import ( @@ -108,7 +110,7 @@ class TestResolveAppAccessFilter: self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True)) permissions = _permissions(app_default_keys=["app.preview"]) - flt = resolve_app_access_filter("tenant-1", "acc-1", permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) assert flt.accessible_app_ids is None assert flt.can_manage_own_apps is False @@ -119,7 +121,7 @@ class TestResolveAppAccessFilter: workspace_keys=["app.full_access", "app.create_and_management"], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) # Workspace-level preview grant defeats the whitelist restriction. assert flt.accessible_app_ids is None @@ -134,7 +136,7 @@ class TestResolveAppAccessFilter: ], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) assert flt.accessible_app_ids == {"app-1"} @@ -144,18 +146,26 @@ class TestResolveAppAccessFilter: app_overrides=[ResourcePermissionKeys(resource_id="app-1", permission_keys=["app.acl.preview"])], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) assert flt.accessible_app_ids == {"app-1", "app-5"} def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[])) + session = MagicMock() + captured: dict[str, object] = {} + + def get_permissions(tenant_id: str, account_id: str, *, session: object): + captured.update(tenant_id=tenant_id, account_id=account_id, session=session) + return _permissions(workspace_keys=["app.create_and_management"]) + monkeypatch.setattr( f"{_RBAC_MODULE}.RBACService.MyPermissions.get", - lambda tenant_id, account_id, session: _permissions(workspace_keys=["app.create_and_management"]), + get_permissions, ) - flt = resolve_app_access_filter("tenant-1", "acc-1") + flt = resolve_app_access_filter("tenant-1", "acc-1", session=session) assert flt.accessible_app_ids == set() assert flt.can_manage_own_apps is True + assert captured == {"tenant_id": "tenant-1", "account_id": "acc-1", "session": session} diff --git a/api/tests/unit_tests/controllers/common/test_session.py b/api/tests/unit_tests/controllers/common/test_session.py index 05b96059572..9da06133a6c 100644 --- a/api/tests/unit_tests/controllers/common/test_session.py +++ b/api/tests/unit_tests/controllers/common/test_session.py @@ -1,6 +1,8 @@ from __future__ import annotations import pytest +from sqlalchemy import Engine, literal, select +from sqlalchemy.orm import Session from controllers.common import session as session_module @@ -22,32 +24,6 @@ class FakeSession: self.rolled_back = True -class FakeSessionBegin: - session: FakeSession - entered: bool - exited: bool - exc_type: object | None - - def __init__(self, session: FakeSession) -> None: - self.session = session - self.entered = False - self.exited = False - self.exc_type = None - - def __enter__(self) -> FakeSession: - self.entered = True - return self.session - - def __exit__(self, exc_type: object | None, *_args: object) -> None: - self.exited = True - self.exc_type = exc_type - if exc_type is None: - self.session.commit() - else: - self.session.rollback() - self.session.closed = True - - class FakeSessionContext: session: FakeSession entered: bool @@ -70,20 +46,10 @@ class FakeSessionContext: self.session.closed = True -class FakeSessionMaker: - begin_context: FakeSessionBegin - - def __init__(self, session: FakeSession) -> None: - self.begin_context = FakeSessionBegin(session) - - def begin(self) -> FakeSessionBegin: - return self.begin_context - - def test_with_session_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + session_context = FakeSessionContext(session) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) class Handler: @session_module.with_session(write=True) @@ -96,15 +62,15 @@ def test_with_session_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) assert session.closed assert session.committed assert not session.rolled_back - assert session_maker.begin_context.entered - assert session_maker.begin_context.exited - assert session_maker.begin_context.exc_type is None + assert session_context.entered + assert session_context.exited + assert session_context.exc_type is None def test_with_session_default_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + session_context = FakeSessionContext(session) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) class Handler: @session_module.with_session @@ -119,8 +85,8 @@ def test_with_session_default_write_commits_on_success(monkeypatch: pytest.Monke def test_with_session_write_rolls_back_on_error(monkeypatch: pytest.MonkeyPatch) -> None: session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + session_context = FakeSessionContext(session) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) class Handler: @session_module.with_session(write=True) @@ -133,9 +99,23 @@ def test_with_session_write_rolls_back_on_error(monkeypatch: pytest.MonkeyPatch) assert session.closed assert not session.committed assert session.rolled_back - assert session_maker.begin_context.entered - assert session_maker.begin_context.exited - assert session_maker.begin_context.exc_type is RuntimeError + assert session_context.entered + assert session_context.exited + assert session_context.exc_type is RuntimeError + + +def test_with_session_write_allows_commit_then_more_database_work( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine +) -> None: + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: Session(sqlite_engine)) + + class Handler: + @session_module.with_session + def post(self, session: Session): + session.commit() + return session.scalar(select(literal(1))) + + assert Handler().post() == 1 def test_with_session_read_mode_does_not_commit(monkeypatch: pytest.MonkeyPatch) -> None: @@ -161,8 +141,8 @@ def test_with_session_read_mode_does_not_commit(monkeypatch: pytest.MonkeyPatch) def test_with_session_preserves_wrapped_metadata(monkeypatch: pytest.MonkeyPatch) -> None: session = FakeSession() - session_maker = FakeSessionMaker(session) - monkeypatch.setattr(session_module.session_factory, "get_session_maker", lambda: session_maker) + session_context = FakeSessionContext(session) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) class Handler: @session_module.with_session 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 51737a859ba..e25f5bea5bf 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 @@ -1,7 +1,7 @@ from inspect import getsource, unwrap from types import SimpleNamespace from typing import Any, cast -from unittest.mock import Mock +from unittest.mock import MagicMock, Mock, call import pytest from flask import Flask @@ -45,11 +45,7 @@ from controllers.console.agent.roster import ( ) from controllers.console.app import completion as completion_controller from controllers.console.app import message as message_controller -from controllers.console.app.completion import ( - AgentBuildChatFinalizeApi, - AgentChatMessageApi, - AgentChatMessageStopApi, -) +from controllers.console.app.completion import AgentBuildChatFinalizeApi, AgentChatMessageApi, AgentChatMessageStopApi from controllers.console.app.error import CompletionRequestError from controllers.console.app.message import ( AgentChatMessageListApi, @@ -152,7 +148,6 @@ def _candidates_response(variant: str) -> dict: def test_agent_v2_console_routes_are_agent_id_first() -> None: paths = {route for item in console_ns.resources for route in item.urls} - for route in ( "/agent", "/agent/", @@ -190,7 +185,6 @@ def test_agent_v2_console_routes_are_agent_id_first() -> None: "/agent/invite-options", ): assert route in paths - for route in ( "/agents", "/agents/invite-options", @@ -223,7 +217,7 @@ def test_agent_app_list_and_create_use_agent_route( captured: dict[str, object] = {} class FakeAppService: - def get_app(self, app_obj: object) -> object: + def get_app(self, app_obj: object, *, session: object) -> object: return app_obj def get_paginate_apps(self, user_id: str, tenant_id: str, params, session) -> object: @@ -303,26 +297,27 @@ def test_agent_app_list_and_create_use_agent_route( lambda _self, **kwargs: {"agent-list": "debug-conversation-list"}, ) monkeypatch.setattr( - roster_controller.AgentRosterService, - "count_agent_app_debug_conversation_messages", - lambda _self, **kwargs: 0, + roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0 ) + + def get_or_create_debug_conversation(_self: object, **kwargs: object) -> str: + captured["get_or_create_debug_conversation"] = kwargs + return "debug-conversation-created" + monkeypatch.setattr( roster_controller.AgentRosterService, "get_or_create_agent_app_debug_conversation_id", - lambda _self, **kwargs: "debug-conversation-created", + get_or_create_debug_conversation, ) monkeypatch.setattr( roster_controller.FeatureService, "get_system_features", lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - with app.test_request_context( "/console/api/agent?page=1&limit=10&mode=workflow&sort_by=recently_created&is_created_by_me=true" ): - listed = unwrap(AgentAppListApi.get)(AgentAppListApi(), "tenant-1", SimpleNamespace(id=account_id)) - + listed = unwrap(AgentAppListApi.get)(AgentAppListApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id)) assert listed["page"] == 1 assert listed["limit"] == 10 assert listed["total"] == 1 @@ -348,19 +343,13 @@ def test_agent_app_list_and_create_use_agent_route( assert list_params.sort_by == "recently_created" assert list_params.is_created_by_me is True assert list_params.status == "normal" - with app.test_request_context( "/console/api/agent", - json={ - "name": "Iris", - "description": "Agent app", - "role": "Coordinator", - "icon_type": "emoji", - "icon": "robot", - }, + json={"name": "Iris", "description": "Agent app", "role": "Coordinator", "icon_type": "emoji", "icon": "robot"}, ): - created, status = unwrap(AgentAppListApi.post)(AgentAppListApi(), "tenant-1", SimpleNamespace(id=account_id)) - + created, status = unwrap(AgentAppListApi.post)( + AgentAppListApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id) + ) assert status == 201 assert created["id"] == "agent-created" assert created["app_id"] == "app-created" @@ -372,6 +361,12 @@ def test_agent_app_list_and_create_use_agent_route( create_params = cast(Any, create_call["params"]) assert create_params.mode == "agent" assert create_params.agent_role == "Coordinator" + assert captured["get_or_create_debug_conversation"] == { + "tenant_id": "tenant-1", + "agent_id": "agent-created", + "account_id": account_id, + "commit": False, + } def test_agent_app_create_payload_allows_optional_role() -> None: @@ -381,7 +376,6 @@ def test_agent_app_create_payload_allows_optional_role() -> None: blank = roster_controller.AgentAppCreatePayload.model_validate( {"name": "Iris", "description": "Agent app", "role": " ", "icon_type": "emoji", "icon": "robot"} ) - assert omitted.role is None assert blank.role == "" @@ -401,21 +395,14 @@ def test_agent_app_create_omits_optional_role_as_empty_string( monkeypatch.setattr( roster_controller, "_serialize_agent_app_detail", - lambda app_model, **_kwargs: {"id": "agent-created", "app_id": app_model.id}, + lambda _session, app_model, **_kwargs: {"id": "agent-created", "app_id": app_model.id}, ) - current_user = SimpleNamespace(id=account_id) with app.test_request_context( "/console/api/agent", - json={ - "name": "No-role Iris", - "description": "Agent app", - "icon_type": "emoji", - "icon": "robot", - }, + json={"name": "No-role Iris", "description": "Agent app", "icon_type": "emoji", "icon": "robot"}, ): - created, status = unwrap(AgentAppListApi.post)(AgentAppListApi(), "tenant-1", current_user) - + created, status = unwrap(AgentAppListApi.post)(AgentAppListApi(), MagicMock(), "tenant-1", current_user) assert status == 201 assert created == {"id": "agent-created", "app_id": "app-created"} create_call = cast(dict[str, object], captured["create"]) @@ -439,28 +426,16 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( active_config_snapshot_id=None, ) captured: dict[str, object] = {} - - monkeypatch.setattr( - roster_controller.AgentRosterService, - "get_agent_app_model", - lambda _self, **kwargs: app_model, - ) - monkeypatch.setattr(roster_controller, "resolve_agent_runtime_app_model", lambda **kwargs: app_model) - monkeypatch.setattr(roster_controller.db.session, "scalar", lambda _stmt: agent) - monkeypatch.setattr( - roster_controller.AgentRosterService, - "get_app_backing_agent", - lambda _self, **kwargs: agent, - ) + monkeypatch.setattr(roster_controller.AgentRosterService, "get_agent_app_model", lambda _self, **kwargs: app_model) + monkeypatch.setattr(roster_controller, "_resolve_agent_runtime_app_model", lambda _session, **kwargs: app_model) + monkeypatch.setattr(roster_controller.AgentRosterService, "get_app_backing_agent", lambda _self, **kwargs: agent) monkeypatch.setattr( roster_controller.AgentRosterService, "get_or_create_agent_app_debug_conversation_id", lambda _self, **kwargs: "debug-conversation-detail", ) monkeypatch.setattr( - roster_controller.AgentRosterService, - "count_agent_app_debug_conversation_messages", - lambda _self, **kwargs: 2, + roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 2 ) monkeypatch.setattr( roster_controller.FeatureService, @@ -469,8 +444,8 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( ) class FakeAppService: - def get_app(self, app_obj: object) -> object: - captured["get_app"] = app_obj + def get_app(self, app_obj: object, *, session: object) -> object: + captured["get_app"] = {"app": app_obj, "session": session} return app_obj def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: @@ -481,8 +456,9 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( captured["delete"] = app_obj monkeypatch.setattr(roster_controller, "AppService", FakeAppService) - - detail = unwrap(AgentAppApi.get)(AgentAppApi(), "tenant-1", SimpleNamespace(id=account_id), agent_id) + session = Mock() + session.scalar.return_value = agent + detail = unwrap(AgentAppApi.get)(AgentAppApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id) assert detail["id"] == agent_id assert detail["app_id"] == "app-1" assert detail["debug_conversation_id"] == "debug-conversation-detail" @@ -491,13 +467,12 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( assert detail["role"] == "Resolved role" assert detail["active_config_is_published"] is False assert "bound_agent_id" not in detail - + assert captured["get_app"] == {"app": app_model, "session": session} with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001", json={"name": "Renamed", "description": "", "role": "Reviewer", "icon_type": "emoji", "icon": "R"}, ): - updated = unwrap(AgentAppApi.put)(AgentAppApi(), "tenant-1", SimpleNamespace(id=account_id), agent_id) - + updated = unwrap(AgentAppApi.put)(AgentAppApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id) assert updated["name"] == "Renamed" assert updated["id"] == agent_id assert updated["app_id"] == "app-1" @@ -510,8 +485,7 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( update_call = cast(dict[str, object], captured["update"]) assert update_call["app"] is app_model assert cast(dict[str, object], update_call["args"])["role"] == "Reviewer" - - deleted, status = unwrap(AgentAppApi.delete)(AgentAppApi(), "tenant-1", agent_id) + deleted, status = unwrap(AgentAppApi.delete)(AgentAppApi(), session, "tenant-1", agent_id) assert (deleted, status) == ("", 204) assert captured["delete"] is app_model @@ -529,13 +503,12 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( captured.update(kwargs) return copied_app - monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda: FakeRosterService()) + monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda *_args: FakeRosterService()) monkeypatch.setattr( roster_controller, "_serialize_agent_app_detail", - lambda app_model, **_kwargs: {"id": "copied-agent", "app_id": app_model.id, "name": app_model.name}, + lambda _session, app_model, **_kwargs: {"id": "copied-agent", "app_id": app_model.id, "name": app_model.name}, ) - with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/copy", json={ @@ -547,8 +520,9 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( "icon_background": "#fff", }, ): - copied, status = unwrap(AgentAppCopyApi.post)(AgentAppCopyApi(), "tenant-1", current_user, agent_id) - + copied, status = unwrap(AgentAppCopyApi.post)( + AgentAppCopyApi(), MagicMock(), "tenant-1", current_user, agent_id + ) assert status == 201 assert copied == {"id": "copied-agent", "app_id": "copied-app", "name": "Iris"} assert captured == { @@ -575,29 +549,19 @@ def test_agent_debug_conversation_refresh_uses_current_user( captured.update(kwargs) return "new-debug-conversation-id" - monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda: FakeRosterService()) - + monkeypatch.setattr(roster_controller, "_agent_roster_service", lambda *_args: FakeRosterService()) with app.test_request_context( - "/console/api/agent/00000000-0000-0000-0000-000000000001/debug-conversation/refresh", - method="POST", + "/console/api/agent/00000000-0000-0000-0000-000000000001/debug-conversation/refresh", method="POST" ): response = unwrap(AgentDebugConversationRefreshApi.post)( - AgentDebugConversationRefreshApi(), - "tenant-1", - SimpleNamespace(id=account_id), - agent_id, + AgentDebugConversationRefreshApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id), agent_id ) - assert response == { "debug_conversation_id": "new-debug-conversation-id", "debug_conversation_has_messages": False, "debug_conversation_message_count": 0, } - assert captured == { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - } + assert captured == {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id} def test_agent_publish_and_build_draft_routes_call_composer_service( @@ -605,7 +569,7 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( ) -> None: agent_id = "00000000-0000-0000-0000-000000000001" current_user = SimpleNamespace(id=account_id) - captured: dict[str, object] = {} + captured: dict[str, dict[str, object]] = {} def publish_agent_app_draft(**kwargs: object) -> dict[str, object]: captured["publish"] = kwargs @@ -631,125 +595,91 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( captured["discard"] = kwargs return {"result": "success"} + monkeypatch.setattr(roster_controller.AgentComposerService, "publish_agent_app_draft", publish_agent_app_draft) monkeypatch.setattr( - roster_controller.AgentComposerService, - "publish_agent_app_draft", - publish_agent_app_draft, + roster_controller.AgentComposerService, "checkout_agent_app_build_draft", checkout_agent_app_build_draft ) monkeypatch.setattr( - roster_controller.AgentComposerService, - "checkout_agent_app_build_draft", - checkout_agent_app_build_draft, + roster_controller.AgentComposerService, "load_agent_app_build_draft", load_agent_app_build_draft ) monkeypatch.setattr( - roster_controller.AgentComposerService, - "load_agent_app_build_draft", - load_agent_app_build_draft, + roster_controller.AgentComposerService, "save_agent_app_build_draft", save_agent_app_build_draft ) monkeypatch.setattr( - roster_controller.AgentComposerService, - "save_agent_app_build_draft", - save_agent_app_build_draft, + roster_controller.AgentComposerService, "apply_agent_app_build_draft", apply_agent_app_build_draft ) monkeypatch.setattr( - roster_controller.AgentComposerService, - "apply_agent_app_build_draft", - apply_agent_app_build_draft, + roster_controller.AgentComposerService, "discard_agent_app_build_draft", discard_agent_app_build_draft ) - monkeypatch.setattr( - roster_controller.AgentComposerService, - "discard_agent_app_build_draft", - 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"}, + "/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) + published = unwrap(AgentPublishApi.post)(AgentPublishApi(), MagicMock(), "tenant-1", current_user, agent_id) assert published["active_config_snapshot_id"] == "version-1" - assert_call_without_session( - "publish", - { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "version_note": "publish v1", - }, - ) - + captured["publish"].pop("session", None) + assert captured["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", - json={"force": True}, + "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/checkout", json={"force": True} ): checked_out = unwrap(AgentBuildDraftCheckoutApi.post)( - AgentBuildDraftCheckoutApi(), "tenant-1", current_user, agent_id + AgentBuildDraftCheckoutApi(), MagicMock(), "tenant-1", current_user, agent_id ) assert checked_out["draft"]["id"] == "build-draft-1" - assert_call_without_session( - "checkout", - { - "tenant_id": "tenant-1", - "agent_id": agent_id, - "account_id": account_id, - "force": True, - }, - ) - + captured["checkout"].pop("session", None) + assert captured["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) + loaded = unwrap(AgentBuildDraftApi.get)(AgentBuildDraftApi(), MagicMock(), "tenant-1", current_user, agent_id) assert loaded["draft"]["id"] == "build-draft-1" - assert_call_without_session("load", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) - + captured["load"].pop("session", None) + assert captured["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", json={"variant": "agent_app", "save_strategy": "save_to_current_version", "agent_soul": {}}, ): - saved = unwrap(AgentBuildDraftApi.put)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) + saved = unwrap(AgentBuildDraftApi.put)(AgentBuildDraftApi(), MagicMock(), "tenant-1", current_user, agent_id) assert saved["draft"]["id"] == "build-draft-1" assert captured["save"]["tenant_id"] == "tenant-1" assert captured["save"]["agent_id"] == agent_id assert captured["save"]["account_id"] == account_id - assert captured["save"]["payload"].variant == ComposerVariant.AGENT_APP - + assert cast(Any, captured["save"]["payload"]).variant == ComposerVariant.AGENT_APP with app.test_request_context( - "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/apply", - method="POST", + "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/apply", method="POST" ): - applied = unwrap(AgentBuildDraftApplyApi.post)(AgentBuildDraftApplyApi(), "tenant-1", current_user, agent_id) + applied = unwrap(AgentBuildDraftApplyApi.post)( + AgentBuildDraftApplyApi(), MagicMock(), "tenant-1", current_user, agent_id + ) assert applied == {"result": "success", "draft": {"id": "draft-1"}} - assert_call_without_session("apply", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) - + captured["apply"].pop("session", None) + assert captured["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", - method="DELETE", + "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", method="DELETE" ): - discarded = unwrap(AgentBuildDraftApi.delete)(AgentBuildDraftApi(), "tenant-1", current_user, agent_id) + discarded = unwrap(AgentBuildDraftApi.delete)( + AgentBuildDraftApi(), MagicMock(), "tenant-1", current_user, agent_id + ) assert discarded == {"result": "success"} - assert_call_without_session("discard", {"tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id}) + captured["discard"].pop("session", None) + assert captured["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( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata(monkeypatch: pytest.MonkeyPatch) -> None: agent_id = "00000000-0000-0000-0000-000000000001" app_model = SimpleNamespace( - id="app-1", - enable_api=True, - api_base_url="https://api.example.test/v1", - api_rpm=60, - api_rph=600, + id="app-1", enable_api=True, api_base_url="https://api.example.test/v1", api_rpm=60, api_rph=600 ) - monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", lambda **kwargs: app_model) - monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda app_id: 2) - - response = unwrap(AgentApiAccessApi.get)(AgentApiAccessApi(), "tenant-1", agent_id) - + monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", lambda _session, **kwargs: app_model) + monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, app_id: 2) + response = unwrap(AgentApiAccessApi.get)(AgentApiAccessApi(), MagicMock(), "tenant-1", agent_id) assert response == { "enabled": True, "service_api_base_url": "https://api.example.test/v1", @@ -768,23 +698,17 @@ def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata( } -def test_agent_api_status_and_key_routes_resolve_backing_app( - app: Flask, - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_agent_api_status_and_key_routes_resolve_backing_app(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: agent_id = "00000000-0000-0000-0000-000000000001" api_key_id = "00000000-0000-0000-0000-000000000002" app_model = SimpleNamespace( - id="app-1", - enable_api=False, - api_base_url="https://api.example.test/v1", - api_rpm=0, - api_rph=0, + id="app-1", enable_api=False, api_base_url="https://api.example.test/v1", api_rpm=0, api_rph=0 ) captured: dict[str, object] = {} - - monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", lambda **kwargs: app_model) - monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda app_id: 1) + session = MagicMock() + resolve_app = Mock(return_value=app_model) + monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", resolve_app) + monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, app_id: 1) class FakeAppService: def update_app_api_status(self, app_obj: object, enable_api: bool, *, session: object) -> object: @@ -794,22 +718,25 @@ def test_agent_api_status_and_key_routes_resolve_backing_app( monkeypatch.setattr(roster_controller, "AppService", FakeAppService) - def fake_get_api_key_list(self, resource_id: str, tenant_id: str): - captured["list_keys"] = {"resource_id": resource_id, "tenant_id": tenant_id} + def fake_get_api_key_list(self, resource_id: str, tenant_id: str, *, session: object): + captured["list_keys"] = {"session": session, "resource_id": resource_id, "tenant_id": tenant_id} return roster_controller.ApiKeyList(data=[]) - def fake_create_api_key(self, resource_id: str, tenant_id: str): - captured["create_key"] = {"resource_id": resource_id, "tenant_id": tenant_id} - return SimpleNamespace( - id=api_key_id, - type="app", - token="app-test-token", - last_used_at=None, - created_at=None, - ) + def fake_create_api_key(self, resource_id: str, tenant_id: str, *, session: object): + captured["create_key"] = {"session": session, "resource_id": resource_id, "tenant_id": tenant_id} + return SimpleNamespace(id=api_key_id, type="app", token="app-test-token", last_used_at=None, created_at=None) - def fake_delete_api_key(self, resource_id: str, key_id: str, tenant_id: str, current_user: object) -> None: + def fake_delete_api_key( + self, + resource_id: str, + key_id: str, + tenant_id: str, + current_user: object, + *, + session: object, + ) -> None: captured["delete_key"] = { + "session": session, "resource_id": resource_id, "api_key_id": key_id, "tenant_id": tenant_id, @@ -819,52 +746,45 @@ def test_agent_api_status_and_key_routes_resolve_backing_app( monkeypatch.setattr(AgentApiKeyListApi, "_get_api_key_list", fake_get_api_key_list) monkeypatch.setattr(AgentApiKeyListApi, "_create_api_key", fake_create_api_key) monkeypatch.setattr(AgentApiKeyApi, "_delete_api_key", fake_delete_api_key) - with app.test_request_context( - "/console/api/agent/00000000-0000-0000-0000-000000000001/api-enable", - json={"enable_api": True}, + "/console/api/agent/00000000-0000-0000-0000-000000000001/api-enable", json={"enable_api": True} ): - enabled = unwrap(AgentApiStatusApi.post)(AgentApiStatusApi(), "tenant-1", agent_id) + enabled = unwrap(AgentApiStatusApi.post)(AgentApiStatusApi(), session, "tenant-1", agent_id) assert enabled["enabled"] is True assert captured["enable"] == {"app": app_model, "enable_api": True} - - keys = unwrap(AgentApiKeyListApi.get)(AgentApiKeyListApi(), "tenant-1", agent_id) + keys = unwrap(AgentApiKeyListApi.get)(AgentApiKeyListApi(), session, "tenant-1", agent_id) assert keys == {"data": []} - assert captured["list_keys"] == {"resource_id": "app-1", "tenant_id": "tenant-1"} - - created, status = unwrap(AgentApiKeyListApi.post)(AgentApiKeyListApi(), "tenant-1", agent_id) + assert captured["list_keys"] == {"session": session, "resource_id": "app-1", "tenant_id": "tenant-1"} + created, status = unwrap(AgentApiKeyListApi.post)(AgentApiKeyListApi(), session, "tenant-1", agent_id) assert status == 201 assert created["id"] == api_key_id assert created["token"] == "app-test-token" - assert captured["create_key"] == {"resource_id": "app-1", "tenant_id": "tenant-1"} - + assert captured["create_key"] == {"session": session, "resource_id": "app-1", "tenant_id": "tenant-1"} current_user = SimpleNamespace(id="account-1", is_admin_or_owner=True) deleted, delete_status = unwrap(AgentApiKeyApi.delete)( - AgentApiKeyApi(), - "tenant-1", - current_user, - agent_id, - api_key_id, + AgentApiKeyApi(), session, "tenant-1", current_user, agent_id, api_key_id ) assert (deleted, delete_status) == ("", 204) assert captured["delete_key"] == { + "session": session, "resource_id": "app-1", "api_key_id": api_key_id, "tenant_id": "tenant-1", "current_user": current_user, } + assert resolve_app.call_args_list == [ + call(session, tenant_id="tenant-1", agent_id=agent_id), + call(session, tenant_id="tenant-1", agent_id=agent_id), + call(session, tenant_id="tenant-1", agent_id=agent_id), + call(session, tenant_id="tenant-1", agent_id=agent_id), + ] def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: agent_id = "00000000-0000-0000-0000-000000000001" app_model = _app_detail_obj(id="app-1", bound_agent_id=agent_id) captured: dict[str, object] = {} - - monkeypatch.setattr( - roster_controller.AgentRosterService, - "get_agent_app_model", - lambda _self, **kwargs: app_model, - ) + monkeypatch.setattr(roster_controller.AgentRosterService, "get_agent_app_model", lambda _self, **kwargs: app_model) monkeypatch.setattr( roster_controller.AgentRosterService, "get_app_backing_agent", @@ -883,14 +803,10 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk lambda _self, **kwargs: "debug-conversation-detail", ) monkeypatch.setattr( - roster_controller.AgentRosterService, - "count_agent_app_debug_conversation_messages", - lambda _self, **kwargs: 0, + roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0 ) monkeypatch.setattr( - roster_controller.AgentRosterService, - "active_config_is_published", - lambda _self, **kwargs: False, + roster_controller.AgentRosterService, "active_config_is_published", lambda _self, **kwargs: False ) monkeypatch.setattr( roster_controller.FeatureService, @@ -899,7 +815,7 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk ) class FakeAppService: - def get_app(self, app_obj: object) -> object: + def get_app(self, app_obj: object, *, session: object) -> object: return app_obj def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: @@ -907,13 +823,13 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) monkeypatch.setattr(roster_controller, "AppService", FakeAppService) - with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001", json={"name": "Renamed", "description": "", "role": "", "icon_type": "emoji", "icon": "R"}, ): - updated = unwrap(AgentAppApi.put)(AgentAppApi(), "tenant-1", SimpleNamespace(id="account-1"), agent_id) - + updated = unwrap(AgentAppApi.put)( + AgentAppApi(), MagicMock(), "tenant-1", SimpleNamespace(id="account-1"), agent_id + ) assert updated["role"] == "" update_call = cast(dict[str, object], captured["update"]) assert cast(dict[str, object], update_call["args"])["role"] == "" @@ -927,10 +843,8 @@ def test_invite_options_get_parses_app_id(app: Flask, monkeypatch: pytest.Monkey return {"data": [], "page": kwargs["page"], "limit": kwargs["limit"], "total": 0, "has_more": False} monkeypatch.setattr(roster_controller.AgentRosterService, "list_invite_options", list_invite_options) - with app.test_request_context("/console/api/agent/invite-options?page=1&limit=10&app_id=app-1"): - result = unwrap(AgentInviteOptionsApi.get)(AgentInviteOptionsApi(), "tenant-1") - + result = unwrap(AgentInviteOptionsApi.get)(AgentInviteOptionsApi(), MagicMock(), "tenant-1") assert result == {"data": [], "page": 1, "limit": 10, "total": 0, "has_more": False} assert captured == {"tenant_id": "tenant-1", "page": 1, "limit": 10, "keyword": None, "app_id": "app-1"} @@ -939,9 +853,7 @@ def test_agent_versions_call_services(app: Flask, monkeypatch: pytest.MonkeyPatc agent_id = "00000000-0000-0000-0000-000000000001" version_id = "00000000-0000-0000-0000-000000000002" monkeypatch.setattr( - roster_controller.AgentRosterService, - "list_agent_versions", - lambda _self, **kwargs: [_version_response()], + roster_controller.AgentRosterService, "list_agent_versions", lambda _self, **kwargs: [_version_response()] ) monkeypatch.setattr( roster_controller.AgentRosterService, @@ -972,18 +884,17 @@ def test_agent_versions_call_services(app: Flask, monkeypatch: pytest.MonkeyPatc return {"result": "success", "active_config_snapshot_id": kwargs["version_id"]} monkeypatch.setattr(roster_controller.AgentRosterService, "restore_agent_version", restore_agent_version) - assert ( - unwrap(AgentRosterVersionsApi.get)(AgentRosterVersionsApi(), "tenant-1", agent_id)["data"][0]["id"] + unwrap(AgentRosterVersionsApi.get)(AgentRosterVersionsApi(), MagicMock(), "tenant-1", agent_id)["data"][0]["id"] == "version-1" ) version_detail = unwrap(AgentRosterVersionDetailApi.get)( - AgentRosterVersionDetailApi(), "tenant-1", agent_id, version_id + AgentRosterVersionDetailApi(), MagicMock(), "tenant-1", agent_id, version_id ) assert version_detail["id"] == version_id assert version_detail["agent_id"] == agent_id restored = unwrap(AgentRosterVersionRestoreApi.post)( - AgentRosterVersionRestoreApi(), "tenant-1", SimpleNamespace(id="account-1"), agent_id, version_id + AgentRosterVersionRestoreApi(), MagicMock(), "tenant-1", SimpleNamespace(id="account-1"), agent_id, version_id ) assert restored == { "result": "success", @@ -1126,17 +1037,13 @@ def test_agent_observability_routes_resolve_app_from_agent_id( }, } - monkeypatch.setattr(roster_controller, "resolve_agent_runtime_app_model", lambda **kwargs: app_model) - monkeypatch.setattr(roster_controller, "_agent_observability_service", lambda: FakeObservabilityService()) - + monkeypatch.setattr(roster_controller, "_resolve_agent_runtime_app_model", lambda _session, **kwargs: app_model) + monkeypatch.setattr(roster_controller, "_agent_observability_service", lambda *_args: FakeObservabilityService()) account = SimpleNamespace(id=account_id, timezone="UTC") with app.test_request_context( - "/console/api/agent/00000000-0000-0000-0000-000000000001/logs" - "?page=2&limit=5&keyword=hello&statuses=success&statuses=failed&sources=webapp:app-1" - "&sources=workflow:app-2:workflow-1:v1:node-1&sort_by=created_at&sort_order=asc" + "/console/api/agent/00000000-0000-0000-0000-000000000001/logs?page=2&limit=5&keyword=hello&statuses=success&statuses=failed&sources=webapp:app-1&sources=workflow:app-2:workflow-1:v1:node-1&sort_by=created_at&sort_order=asc" ): - logs = unwrap(AgentLogsApi.get)(AgentLogsApi(), "tenant-1", account, agent_id) - + logs = unwrap(AgentLogsApi.get)(AgentLogsApi(), MagicMock(), "tenant-1", account, agent_id) assert logs["data"][0]["id"] == "conversation-1" assert logs["data"][0]["source"]["id"] == "webapp:app-1" logs_call = cast(dict[str, object], captured["logs"]) @@ -1150,18 +1057,12 @@ def test_agent_observability_routes_resolve_app_from_agent_id( assert logs_params.sources == ("webapp:app-1", "workflow:app-2:workflow-1:v1:node-1") assert logs_params.sort_by == "created_at" assert logs_params.sort_order == "asc" - with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/logs/00000000-0000-0000-0000-000000000002/messages" ): messages = unwrap(AgentLogMessagesApi.get)( - AgentLogMessagesApi(), - "tenant-1", - account, - agent_id, - "00000000-0000-0000-0000-000000000002", + AgentLogMessagesApi(), MagicMock(), "tenant-1", account, agent_id, "00000000-0000-0000-0000-000000000002" ) - assert messages["data"][0]["id"] == "message-1" messages_call = cast(dict[str, object], captured["messages"]) assert messages_call["app"] is app_model @@ -1170,20 +1071,18 @@ def test_agent_observability_routes_resolve_app_from_agent_id( messages_params = cast(Any, messages_call["params"]) assert messages_params.sources == () assert messages_params.statuses == () - with app.test_request_context("/console/api/agent/00000000-0000-0000-0000-000000000001/log-sources"): - sources = unwrap(AgentLogSourcesApi.get)(AgentLogSourcesApi(), "tenant-1", account, agent_id) - + sources = unwrap(AgentLogSourcesApi.get)(AgentLogSourcesApi(), MagicMock(), "tenant-1", account, agent_id) assert sources["data"][0]["id"] == "webapp:app-1" sources_call = cast(dict[str, object], captured["sources"]) assert sources_call["app"] is app_model assert sources_call["agent_id"] == agent_id - with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/statistics/summary?source=api" ): - statistics = unwrap(AgentStatisticsSummaryApi.get)(AgentStatisticsSummaryApi(), "tenant-1", account, agent_id) - + statistics = unwrap(AgentStatisticsSummaryApi.get)( + AgentStatisticsSummaryApi(), MagicMock(), "tenant-1", account, agent_id + ) assert statistics["summary"]["total_messages"] == 1 stats_call = cast(dict[str, object], captured["statistics"]) assert stats_call["app"] is app_model @@ -1232,37 +1131,40 @@ def test_workflow_composer_get_put_validate_candidates_impact_and_save( "bindings": [], }, ) - with app.test_request_context("?snapshot_id=preview-version"): workflow_state = unwrap(WorkflowAgentComposerApi.get)( - WorkflowAgentComposerApi(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" ) assert workflow_state["node_id"] == "node-1" assert captured_load["account_id"] == account_id assert captured_load["snapshot_id"] == "preview-version" with app.test_request_context(json=payload): saved_state = unwrap(WorkflowAgentComposerApi.put)( - WorkflowAgentComposerApi(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" ) assert saved_state["save_options"] == ["node_job_only"] assert unwrap(WorkflowAgentComposerValidateApi.post)( - WorkflowAgentComposerValidateApi(), "tenant-1", app_model, "node-1" + WorkflowAgentComposerValidateApi(), MagicMock(), "tenant-1", app_model, "node-1" ) == {"result": "success", "errors": [], "warnings": [], "knowledge_retrieval_placeholder": []} assert ( unwrap(WorkflowAgentComposerCandidatesApi.get)( - WorkflowAgentComposerCandidatesApi(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerCandidatesApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" )["variant"] == "workflow" ) with app.test_request_context(json=payload): assert unwrap(WorkflowAgentComposerImpactApi.post)( - WorkflowAgentComposerImpactApi(), "tenant-1", app_model, "node-1" + WorkflowAgentComposerImpactApi(), MagicMock(), "tenant-1", app_model, "node-1" ) == {"current_snapshot_id": "version-1", "workflow_node_count": 1, "bindings": []} assert unwrap(WorkflowAgentComposerSaveToRosterApi.post)( - WorkflowAgentComposerSaveToRosterApi(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerSaveToRosterApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" )["save_options"] == ["node_job_only"] +def test_workflow_composer_get_uses_write_transaction() -> None: + assert "@with_session\n @get_app_model" in getsource(WorkflowAgentComposerApi) + + def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str) -> None: app_model = SimpleNamespace(id="app-1") captured: dict[str, object] = {} @@ -1291,7 +1193,6 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk monkeypatch.setattr( composer_controller.AgentComposerService, "copy_workflow_composer_from_roster", fake_copy_from_roster ) - with app.test_request_context( json={ "source_agent_id": "roster-agent-1", @@ -1300,11 +1201,10 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk } ): result = unwrap(WorkflowAgentComposerCopyFromRosterApi.post)( - WorkflowAgentComposerCopyFromRosterApi(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerCopyFromRosterApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" ) - assert result["binding"]["binding_type"] == "inline_agent" - assert captured.pop("session") is not None + captured.pop("session", None) assert captured == { "tenant_id": "tenant-1", "app_id": "app-1", @@ -1318,12 +1218,10 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk def test_workflow_impact_returns_empty_without_version(app: Flask) -> None: payload = {"variant": ComposerVariant.WORKFLOW.value, "save_strategy": ComposerSaveStrategy.NODE_JOB_ONLY.value} - with app.test_request_context(json=payload): result = unwrap(WorkflowAgentComposerImpactApi.post)( - WorkflowAgentComposerImpactApi(), "tenant-1", SimpleNamespace(id="app-1"), "node-1" + WorkflowAgentComposerImpactApi(), MagicMock(), "tenant-1", SimpleNamespace(id="app-1"), "node-1" ) - assert result == {"current_snapshot_id": None, "workflow_node_count": 0, "bindings": []} @@ -1354,45 +1252,31 @@ def test_agent_composer_routes_resolve_app_from_agent_id( captured["candidates"] = kwargs return _candidates_response("agent_app") - monkeypatch.setattr( - composer_controller.AgentComposerService, - "load_agent_composer", - load_agent_composer, - ) - monkeypatch.setattr( - composer_controller.AgentComposerService, - "save_agent_composer", - save_agent_composer, - ) + monkeypatch.setattr(composer_controller.AgentComposerService, "load_agent_composer", load_agent_composer) + monkeypatch.setattr(composer_controller.AgentComposerService, "save_agent_composer", save_agent_composer) monkeypatch.setattr(composer_controller.ComposerConfigValidator, "validate_publish_payload", lambda payload: None) monkeypatch.setattr( - composer_controller.AgentComposerService, - "collect_validation_findings", - collect_validation_findings, + composer_controller.AgentComposerService, "collect_validation_findings", collect_validation_findings ) - monkeypatch.setattr( - composer_controller.AgentComposerService, - "get_agent_app_candidates", - get_agent_app_candidates, - ) - - assert unwrap(AgentComposerApi.get)(AgentComposerApi(), "tenant-1", agent_id)["variant"] == "agent_app" + monkeypatch.setattr(composer_controller.AgentComposerService, "get_agent_app_candidates", get_agent_app_candidates) + assert unwrap(AgentComposerApi.get)(AgentComposerApi(), MagicMock(), "tenant-1", agent_id)["variant"] == "agent_app" assert cast(dict[str, object], captured["load"])["agent_id"] == agent_id - with app.test_request_context(json=payload): assert ( - unwrap(AgentComposerApi.put)(AgentComposerApi(), "tenant-1", account_id, agent_id)["variant"] == "agent_app" + unwrap(AgentComposerApi.put)(AgentComposerApi(), MagicMock(), "tenant-1", account_id, agent_id)["variant"] + == "agent_app" ) assert cast(dict[str, object], captured["save"])["agent_id"] == agent_id - assert unwrap(AgentComposerValidateApi.post)(AgentComposerValidateApi(), "tenant-1", agent_id) == { + assert unwrap(AgentComposerValidateApi.post)(AgentComposerValidateApi(), MagicMock(), "tenant-1", agent_id) == { "result": "success", "errors": [], "warnings": [], "knowledge_retrieval_placeholder": [], } assert cast(dict[str, object], captured["validate"])["agent_id"] == agent_id - - candidates = unwrap(AgentComposerCandidatesApi.get)(AgentComposerCandidatesApi(), "tenant-1", account_id, agent_id) + candidates = unwrap(AgentComposerCandidatesApi.get)( + AgentComposerCandidatesApi(), MagicMock(), "tenant-1", account_id, agent_id + ) assert candidates["variant"] == "agent_app" assert cast(dict[str, object], captured["candidates"])["agent_id"] == agent_id @@ -1405,6 +1289,11 @@ def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( captured: dict[str, object] = {} def resolve_agent_app_model(**kwargs: object) -> object: + captured["stop_resolve"] = kwargs + return app_model + + def resolve_agent_app_model_with_session(_self, **kwargs: object) -> object: + captured["resolve_session"] = _self._session captured["resolve"] = kwargs return app_model @@ -1414,33 +1303,41 @@ def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( def stop_chat_message(**kwargs: object) -> tuple[dict[str, object], int]: captured["stop"] = kwargs - return {"result": "success"}, 200 + return ({"result": "success"}, 200) monkeypatch.setattr(completion_controller, "resolve_agent_runtime_app_model", resolve_agent_app_model) + monkeypatch.setattr( + completion_controller.AgentRosterService, + "get_agent_runtime_app_model", + resolve_agent_app_model_with_session, + ) monkeypatch.setattr(completion_controller, "_create_chat_message", create_chat_message) monkeypatch.setattr(completion_controller, "_stop_chat_message", stop_chat_message) - session = Mock() - with app.test_request_context(json={"inputs": {}, "query": "hello"}): assert unwrap(AgentChatMessageApi.post)( AgentChatMessageApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id ) == {"result": "generated"} - assert cast(dict[str, object], captured["resolve"]) == {"tenant_id": "tenant-1", "agent_id": agent_id} + assert captured["resolve_session"] is session create_call = cast(dict[str, object], captured["create"]) assert create_call["session"] is session assert create_call["app_model"] is app_model assert cast(SimpleNamespace, create_call["current_user"]).id == account_id - assert unwrap(AgentChatMessageStopApi.post)( - AgentChatMessageStopApi(), "tenant-1", account_id, agent_id, "task-1" + AgentChatMessageStopApi(), session, "tenant-1", account_id, agent_id, "task-1" ) == ({"result": "success"}, 200) + assert captured["stop_resolve"] == { + "session": session, + "tenant_id": "tenant-1", + "agent_id": agent_id, + } stop_call = cast(dict[str, object], captured["stop"]) assert stop_call == {"current_user_id": account_id, "app_model": app_model, "task_id": "task-1"} def test_agent_chat_stream_preflight_raises_first_error_event() -> None: + class ClosableStream: def __init__(self) -> None: self.closed = False @@ -1464,25 +1361,17 @@ def test_agent_chat_stream_preflight_raises_first_error_event() -> None: self.closed = True stream = ClosableStream() - with pytest.raises(CompletionRequestError) as exc_info: completion_controller._raise_agent_stream_error_before_response(stream) - assert "Incorrect API key provided" in exc_info.value.description assert stream.closed is True def test_agent_chat_stream_preflight_preserves_first_normal_event() -> None: stream = iter( - [ - "event: ping\n\n", - 'data: {"event":"message","answer":"hello"}\n\n', - 'data: {"event":"message_end"}\n\n', - ] + ["event: ping\n\n", 'data: {"event":"message","answer":"hello"}\n\n', 'data: {"event":"message_end"}\n\n'] ) - wrapped = completion_controller._raise_agent_stream_error_before_response(stream) - assert list(wrapped) == [ "event: ping\n\n", 'data: {"event":"message","answer":"hello"}\n\n', @@ -1505,16 +1394,17 @@ def test_agent_build_chat_finalize_route_resolves_app_from_agent_id( captured["finalize"] = kwargs return {"result": "generated"} - monkeypatch.setattr(completion_controller, "resolve_agent_runtime_app_model", resolve_agent_app_model) + monkeypatch.setattr( + completion_controller.AgentRosterService, + "get_agent_runtime_app_model", + lambda _self, **kwargs: resolve_agent_app_model(**kwargs), + ) monkeypatch.setattr(completion_controller, "_create_build_chat_finalization_message", create_finalization_message) - session = Mock() - with app.test_request_context(): assert unwrap(AgentBuildChatFinalizeApi.post)( AgentBuildChatFinalizeApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id ) == {"result": "generated"} - assert cast(dict[str, object], captured["resolve"]) == {"tenant_id": "tenant-1", "agent_id": agent_id} finalize_call = cast(dict[str, object], captured["finalize"]) assert finalize_call["session"] is session @@ -1537,31 +1427,25 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( def generate(**kwargs: object) -> object: captured["generate"] = kwargs return iter( - [ - "event: ping\n\n", - 'data: {"event":"message","answer":"working"}\n\n', - 'data: {"event":"message_end"}\n\n', - ] + ["event: ping\n\n", 'data: {"event":"message","answer":"working"}\n\n', 'data: {"event":"message_end"}\n\n'] ) monkeypatch.setattr( - completion_controller, - "_resolve_current_user_agent_debug_conversation_id", - resolve_debug_conversation, + completion_controller, "_resolve_current_user_agent_debug_conversation_id", resolve_debug_conversation ) monkeypatch.setattr(completion_controller.AppGenerateService, "generate", generate) - + session = Mock() with app.test_request_context(headers={"X-Trace-Id": "trace-1"}): result = completion_controller._create_build_chat_finalization_message( current_tenant_id="tenant-1", current_user=SimpleNamespace(id=account_id), app_model=app_model, agent_id="agent-1", - session=Mock(), + session=session, ) - assert result == ({"result": "success"}, 200) assert captured["resolve_debug_conversation"] == { + "session": session, "current_tenant_id": "tenant-1", "current_user": SimpleNamespace(id=account_id), "app_model": app_model, @@ -1581,6 +1465,7 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( def test_drain_streaming_generate_response_returns_on_message_end() -> None: + class ClosableResponse: def __init__(self) -> None: self._chunks = iter( @@ -1602,21 +1487,18 @@ def test_drain_streaming_generate_response_returns_on_message_end() -> None: self.closed = True response = ClosableResponse() - assert completion_controller._drain_streaming_generate_response(response) is None assert response.closed is True def test_drain_streaming_generate_response_maps_error_event() -> None: response = iter(['data: {"event":"error","message":"backend failed"}\n\n']) - with pytest.raises(CompletionRequestError, match="backend failed"): completion_controller._drain_streaming_generate_response(response) def test_drain_streaming_generate_response_raises_when_stream_ends_early() -> None: response = iter(['data: {"event":"message","answer":"working"}\n\n']) - with pytest.raises(CompletionRequestError, match="did not complete"): completion_controller._drain_streaming_generate_response(response) @@ -1639,21 +1521,14 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace( lambda **kwargs: "debug-conversation-1", ) monkeypatch.setattr( - completion_controller.helper, - "compact_generate_response", - lambda response: {"response": response}, + completion_controller.helper, "compact_generate_response", lambda response: {"response": response} ) - with app.test_request_context( - json={"inputs": {}, "query": "hello", "response_mode": "streaming"}, - headers={"X-Trace-Id": "trace-1"}, + json={"inputs": {}, "query": "hello", "response_mode": "streaming"}, headers={"X-Trace-Id": "trace-1"} ): result = completion_controller._create_chat_message( - current_user=current_user, - app_model=app_model, - session=Mock(), + current_user=current_user, app_model=app_model, session=Mock() ) - assert result == {"response": {"answer": "ok"}} assert captured["app_model"] is app_model assert captured["user"] is current_user @@ -1711,18 +1586,14 @@ def test_agent_chat_helper_ignores_private_exit_intent_payload_key( def test_agent_chat_helper_rejects_foreign_debug_conversation( - app: Flask, - monkeypatch: pytest.MonkeyPatch, - account_id: str, + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str ) -> None: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") - monkeypatch.setattr( completion_controller, "_resolve_current_user_agent_debug_conversation_id", lambda **kwargs: "owned-conversation", ) - with app.test_request_context( json={ "inputs": {}, @@ -1759,21 +1630,20 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app return SimpleNamespace(id="backing-agent") monkeypatch.setattr(completion_controller, "AgentRosterService", FakeRosterService) - monkeypatch.setattr(completion_controller, "db", SimpleNamespace(session="session-1")) - explicit_id = completion_controller._resolve_current_user_agent_debug_conversation_id( + session="session-1", # type: ignore[arg-type] current_tenant_id="tenant-1", current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"), agent_id="agent-1", ) fallback_id = completion_controller._resolve_current_user_agent_debug_conversation_id( + session="session-1", # type: ignore[arg-type] current_tenant_id="tenant-1", current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"), agent_id=None, ) - assert explicit_id == "debug-agent-1" assert fallback_id == "debug-backing-agent" assert calls[1] == {"get_or_create": {"tenant_id": "tenant-1", "agent_id": "agent-1", "account_id": "account-1"}} @@ -1810,20 +1680,14 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app ], ) def test_agent_chat_helper_maps_generation_errors( - app: Flask, - monkeypatch: pytest.MonkeyPatch, - error: Exception, - expected: type[Exception], + app: Flask, monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception] ) -> None: app_model = SimpleNamespace(id="app-1", mode="chat") monkeypatch.setattr(completion_controller.AppGenerateService, "generate", lambda **_: (_ for _ in ()).throw(error)) - with app.test_request_context(json={"inputs": {}, "query": "hello"}): with pytest.raises(expected): completion_controller._create_chat_message( - current_user=SimpleNamespace(id="account-1"), - app_model=app_model, - session=Mock(), + current_user=SimpleNamespace(id="account-1"), app_model=app_model, session=Mock() ) @@ -1833,9 +1697,11 @@ def test_agent_chat_message_routes_resolve_app_from_agent_id(app: Flask, monkeyp app_model = SimpleNamespace(id="app-1", mode="agent") current_user = SimpleNamespace(id="account-1") captured: dict[str, object] = {} + resolver_calls: list[dict[str, object]] = [] + session = Mock() def resolve_agent_app_model(**kwargs: object) -> object: - captured["resolve"] = kwargs + resolver_calls.append(kwargs) return app_model def list_chat_messages(**kwargs: object) -> dict[str, object]: @@ -1859,31 +1725,39 @@ def test_agent_chat_message_routes_resolve_app_from_agent_id(app: Flask, monkeyp monkeypatch.setattr(message_controller, "_update_message_feedback", update_message_feedback) monkeypatch.setattr(message_controller, "_get_message_suggested_questions", get_message_suggested_questions) monkeypatch.setattr(message_controller, "_get_message_detail", get_message_detail) - - assert unwrap(AgentChatMessageListApi.get)(AgentChatMessageListApi(), "tenant-1", current_user, agent_id) == { - "data": [] - } - assert cast(dict[str, object], captured["list"])["app_model"] is app_model - + assert unwrap(AgentChatMessageListApi.get)( + AgentChatMessageListApi(), session, "tenant-1", current_user, agent_id + ) == {"data": []} + list_call = cast(dict[str, object], captured["list"]) + assert list_call["session"] is session + assert list_call["app_model"] is app_model with app.test_request_context(json={"message_id": message_id, "rating": "like"}): - assert unwrap(AgentMessageFeedbackApi.post)(AgentMessageFeedbackApi(), "tenant-1", current_user, agent_id) == { - "result": "success" - } + assert unwrap(AgentMessageFeedbackApi.post)( + AgentMessageFeedbackApi(), session, "tenant-1", current_user, agent_id + ) == {"result": "success"} feedback_call = cast(dict[str, object], captured["feedback"]) + assert feedback_call["session"] is session assert feedback_call["app_model"] is app_model assert feedback_call["current_user"] is current_user - assert unwrap(AgentMessageSuggestedQuestionApi.get)( - AgentMessageSuggestedQuestionApi(), "tenant-1", current_user, agent_id, message_id + AgentMessageSuggestedQuestionApi(), session, "tenant-1", current_user, agent_id, message_id ) == {"data": ["next"]} suggested_call = cast(dict[str, object], captured["suggested"]) + assert suggested_call["session"] is session assert suggested_call["app_model"] is app_model assert suggested_call["current_user"] is current_user assert suggested_call["message_id"] == message_id - - assert unwrap(AgentMessageApi.get)(AgentMessageApi(), "tenant-1", agent_id, message_id) == {"id": message_id} + assert unwrap(AgentMessageApi.get)(AgentMessageApi(), session, "tenant-1", agent_id, message_id) == { + "id": message_id + } detail_call = cast(dict[str, object], captured["detail"]) - assert detail_call == {"app_model": app_model, "message_id": message_id} + assert detail_call == {"session": session, "app_model": app_model, "message_id": message_id} + assert resolver_calls == [ + {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, + ] def test_list_chat_messages_supports_first_id_pagination(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -1895,10 +1769,7 @@ def test_list_chat_messages_supports_first_id_pagination(app: Flask, monkeypatch older_message = SimpleNamespace(id=older_message_id, created_at=1) scalar_values = iter([conversation, first_message, True]) scalars_result = SimpleNamespace(all=lambda: [older_message]) - session = SimpleNamespace( - scalar=lambda _stmt: next(scalar_values), - scalars=lambda _stmt: scalars_result, - ) + session = SimpleNamespace(scalar=lambda _stmt: next(scalar_values), scalars=lambda _stmt: scalars_result) class FakeMessagePaginationResponse: @classmethod @@ -1911,23 +1782,18 @@ def test_list_chat_messages_supports_first_id_pagination(app: Flask, monkeypatch } ) - monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=session)) monkeypatch.setattr(message_controller, "attach_message_extra_contents", lambda messages: None) monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) - with app.test_request_context( - "/console/api/agent/agent-1/chat-messages" - f"?conversation_id={conversation_id}&first_id={first_message_id}&limit=1" + f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}&first_id={first_message_id}&limit=1" ): - result = message_controller._list_chat_messages(app_model=SimpleNamespace(id="app-1", mode="chat")) - + result = message_controller._list_chat_messages( + session=session, app_model=SimpleNamespace(id="app-1", mode="chat") + ) assert result == {"data": [older_message_id], "limit": 1, "has_more": True} -def test_list_agent_chat_messages_uses_current_user_conversation( - app: Flask, - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_list_agent_chat_messages_uses_current_user_conversation(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: conversation_id = "00000000-0000-0000-0000-000000000010" message_id = "00000000-0000-0000-0000-000000000011" conversation = SimpleNamespace(id=conversation_id) @@ -1935,10 +1801,7 @@ def test_list_agent_chat_messages_uses_current_user_conversation( current_user = SimpleNamespace(id="account-1") app_model = SimpleNamespace(id="app-1", mode="agent") captured: dict[str, object] = {} - session = SimpleNamespace( - scalar=lambda _stmt: False, - scalars=lambda _stmt: SimpleNamespace(all=lambda: [message]), - ) + session = SimpleNamespace(scalar=lambda _stmt: False, scalars=lambda _stmt: SimpleNamespace(all=lambda: [message])) class FakeMessagePaginationResponse: @classmethod @@ -1955,63 +1818,53 @@ 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=SessionProxy())) monkeypatch.setattr(message_controller, "attach_message_extra_contents", lambda messages: None) monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) - with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): - result = message_controller._list_chat_messages(app_model=app_model, current_user=current_user) - + result = message_controller._list_chat_messages(session=session, 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} -def test_list_agent_chat_messages_rejects_foreign_conversation( - app: Flask, - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_list_agent_chat_messages_rejects_foreign_conversation(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: conversation_id = "00000000-0000-0000-0000-000000000010" monkeypatch.setattr( message_controller.ConversationService, "get_conversation", lambda **kwargs: (_ for _ in ()).throw(message_controller.ConversationNotExistsError()), ) - with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): with pytest.raises(NotFound): message_controller._list_chat_messages( + session=Mock(), app_model=SimpleNamespace(id="app-1", mode="agent"), current_user=SimpleNamespace(id="account-1"), ) def test_update_message_feedback_rejects_empty_rating_without_existing_feedback( - app: Flask, monkeypatch: pytest.MonkeyPatch + app: Flask, ) -> None: message_id = "00000000-0000-0000-0000-000000000002" - message = SimpleNamespace(id=message_id, app_id="app-1", admin_feedback=None) - session = SimpleNamespace(scalar=lambda _stmt: message) - monkeypatch.setattr(message_controller, "db", SimpleNamespace(session=session)) - + message = SimpleNamespace( + id=message_id, + app_id="app-1", + admin_feedback_with_session=MagicMock(return_value=None), + ) + session = MagicMock() + session.scalar.return_value = message with app.test_request_context(json={"message_id": message_id, "rating": None}): with pytest.raises(ValueError, match="rating cannot be None"): message_controller._update_message_feedback( + session=session, current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"), ) + message.admin_feedback_with_session.assert_called_once_with(session=session) + @pytest.mark.parametrize( ("error", "expected"), @@ -2033,18 +1886,22 @@ def test_update_message_feedback_rejects_empty_rating_without_existing_feedback( ], ) def test_get_message_suggested_questions_maps_service_errors( - monkeypatch: pytest.MonkeyPatch, - error: Exception, - expected: type[Exception], + monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception] ) -> None: + session = Mock() + + def raise_error(**kwargs: object) -> None: + assert kwargs["session"] is session + raise error + monkeypatch.setattr( message_controller.MessageService, "get_suggested_questions_after_answer", - lambda **_: (_ for _ in ()).throw(error), + raise_error, ) - with pytest.raises(expected): message_controller._get_message_suggested_questions( + session=session, current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"), message_id="00000000-0000-0000-0000-000000000002", diff --git a/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py b/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py new file mode 100644 index 00000000000..e05d57bd73a --- /dev/null +++ b/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py @@ -0,0 +1,50 @@ +from unittest.mock import MagicMock +from uuid import UUID + +import pytest + +from controllers.console.agent import app_helpers + + +def test_resolve_agent_app_model_reuses_caller_session(monkeypatch: pytest.MonkeyPatch) -> None: + session = MagicMock() + app = MagicMock() + service = MagicMock() + service.get_agent_app_model.return_value = app + service_factory = MagicMock(return_value=service) + monkeypatch.setattr(app_helpers, "AgentRosterService", service_factory) + + result = app_helpers.resolve_agent_app_model( + session=session, + tenant_id="tenant-1", + agent_id=UUID("00000000-0000-0000-0000-000000000001"), + ) + + assert result is app + service_factory.assert_called_once_with(session) + service.get_agent_app_model.assert_called_once_with( + tenant_id="tenant-1", + agent_id="00000000-0000-0000-0000-000000000001", + ) + + +def test_resolve_agent_runtime_app_model_reuses_caller_session(monkeypatch: pytest.MonkeyPatch) -> None: + session = MagicMock() + app = MagicMock() + service = MagicMock() + service.get_agent_runtime_app_model.return_value = app + service_factory = MagicMock(return_value=service) + monkeypatch.setattr(app_helpers, "AgentRosterService", service_factory) + + result = app_helpers.resolve_agent_runtime_app_model( + session=session, + tenant_id="tenant-1", + agent_id=UUID("00000000-0000-0000-0000-000000000001"), + ) + + assert result is app + service_factory.assert_called_once_with(session) + service.get_agent_runtime_app_model.assert_called_once_with( + tenant_id="tenant-1", + agent_id="00000000-0000-0000-0000-000000000001", + ) 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 0ab8814f368..265c542e635 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 @@ -2,6 +2,7 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from dify_agent.client import DifyAgentClientError, DifyAgentHTTPError, DifyAgentTimeoutError @@ -123,8 +124,10 @@ def test_handle_maps_sandbox_and_agent_backend_errors() -> None: def test_agent_app_sandbox_resources_proxy_service(monkeypatch: pytest.MonkeyPatch) -> None: service = _AgentAppService() + session = MagicMock() + resolver = MagicMock(return_value=_app_model()) monkeypatch.setattr(module, "AgentAppSandboxService", lambda: service) - monkeypatch.setattr(module, "resolve_agent_runtime_app_model", lambda *, tenant_id, agent_id: _app_model()) + monkeypatch.setattr(module, "resolve_agent_runtime_app_model", resolver) monkeypatch.setattr( module, "query_params_from_request", @@ -136,10 +139,10 @@ def test_agent_app_sandbox_resources_proxy_service(monkeypatch: pytest.MonkeyPat SimpleNamespace(get_json=lambda silent=True: {"conversation_id": "conv-1", "path": "report.txt"}), ) - info = unwrap(module.AgentAppSandboxInfoResource.get)(object(), "tenant-1", "agent-1") - listing = unwrap(module.AgentAppSandboxListResource.get)(object(), "tenant-1", "agent-1") - preview = unwrap(module.AgentAppSandboxReadResource.get)(object(), "tenant-1", "agent-1") - upload = unwrap(module.AgentAppSandboxUploadResource.post)(object(), "tenant-1", "agent-1") + info = unwrap(module.AgentAppSandboxInfoResource.get)(object(), session, "tenant-1", "agent-1") + listing = unwrap(module.AgentAppSandboxListResource.get)(object(), session, "tenant-1", "agent-1") + preview = unwrap(module.AgentAppSandboxReadResource.get)(object(), session, "tenant-1", "agent-1") + upload = unwrap(module.AgentAppSandboxUploadResource.post)(object(), session, "tenant-1", "agent-1") assert info == {"session_id": "abc1234", "workspace_cwd": "~/workspace/abc1234"} assert listing["path"] == "sub/report.txt" @@ -151,6 +154,7 @@ def test_agent_app_sandbox_resources_proxy_service(monkeypatch: pytest.MonkeyPat ("read", "tenant-1", "app-1", "conv-1", "sub/report.txt"), ("upload", "tenant-1", "app-1", "conv-1", "report.txt"), ] + assert all(call.kwargs["session"] is session for call in resolver.call_args_list) def test_agent_app_sandbox_resource_returns_normalized_errors(monkeypatch: pytest.MonkeyPatch) -> None: @@ -162,16 +166,17 @@ def test_agent_app_sandbox_resource_returns_normalized_errors(monkeypatch: pytes raise AgentSandboxInspectorError("no_active_session", "no active session", status_code=404) monkeypatch.setattr(module, "AgentAppSandboxService", FailingService) - monkeypatch.setattr(module, "resolve_agent_runtime_app_model", lambda *, tenant_id, agent_id: _app_model()) + session = MagicMock() + monkeypatch.setattr(module, "resolve_agent_runtime_app_model", MagicMock(return_value=_app_model())) monkeypatch.setattr( module, "query_params_from_request", lambda model: SimpleNamespace(conversation_id="conv-1", path=".") ) - assert unwrap(module.AgentAppSandboxInfoResource.get)(object(), "tenant-1", "agent-1") == ( + assert unwrap(module.AgentAppSandboxInfoResource.get)(object(), session, "tenant-1", "agent-1") == ( {"code": "no_active_session", "message": "no active session"}, 404, ) - assert unwrap(module.AgentAppSandboxListResource.get)(object(), "tenant-1", "agent-1") == ( + assert unwrap(module.AgentAppSandboxListResource.get)(object(), session, "tenant-1", "agent-1") == ( {"code": "no_active_session", "message": "no active session"}, 404, ) diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py b/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py index 0b9ca517e13..1d857e7284c 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py @@ -6,9 +6,9 @@ workflow-node agent binding, and service delegation for the new config surface. from __future__ import annotations -import inspect +from inspect import unwrap from types import SimpleNamespace -from unittest.mock import PropertyMock, patch +from unittest.mock import MagicMock, PropertyMock, patch from flask import Flask @@ -35,13 +35,28 @@ app = Flask(__name__) def _raw(method): - return inspect.unwrap(method) + return unwrap(method) -_APP = SimpleNamespace(id="app-1", tenant_id="tenant-1", bound_agent_id="agent-1") +_APP = SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + bound_agent_id_with_session=lambda *, session: "agent-1", +) _USER = SimpleNamespace(id="acct-1") +def test_resolve_bound_agent_uses_injected_session(): + session = MagicMock() + resolver = MagicMock(return_value="agent-1") + app_model = SimpleNamespace(bound_agent_id_with_session=resolver) + result = inspector._resolve_agent_id(session, app_model, None) + + assert result == "agent-1" + resolver.assert_called_once_with(session=session) + assert resolver.call_args.kwargs["session"] is session + + def test_manifest_by_agent_resolves_build_draft_version(): raw = _raw(AgentConfigManifestByAgentApi.get) with app.test_request_context("/?draft_type=debug_build"): @@ -59,8 +74,7 @@ def test_manifest_by_agent_resolves_build_draft_version(): "env_keys": [], "note": "", } - body = raw(AgentConfigManifestByAgentApi(), "tenant-1", _USER, "agent-1") - + body = raw(AgentConfigManifestByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1") assert body["config_version"]["kind"] == "build_draft" assert config_service.return_value.manifest.call_args.kwargs["config_version_id"] == "build-draft-1" assert config_service.return_value.manifest.call_args.kwargs["config_version_kind"].value == "build_draft" @@ -69,10 +83,7 @@ def test_manifest_by_agent_resolves_build_draft_version(): def test_manifest_resolves_workflow_node_agent_and_normal_draft(): raw = _raw(AgentConfigManifestApi.get) with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.manifest.return_value = { @@ -83,31 +94,27 @@ def test_manifest_resolves_workflow_node_agent_and_normal_draft(): "env_keys": [], "note": "", } - body = raw(AgentConfigManifestApi(), _USER, _APP) - + body = raw(AgentConfigManifestApi(), MagicMock(), _USER, _APP) assert body["agent_id"] == "wf-agent-9" assert composer.resolve_workflow_node_agent_id.call_args.kwargs["node_id"] == "node-1" assert config_service.return_value.manifest.call_args.kwargs["config_version_kind"].value == "draft" def test_normal_draft_resolution_commits_created_draft_before_service_session() -> None: - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.db") as db, - ): + session = MagicMock() + with patch(f"{_MOD}.AgentComposerService") as composer: composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} - version_id, version_kind = inspector._resolve_console_version( + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="acct-1", version_id=None, draft_type="draft", ) - assert version_id == "draft-1" assert version_kind.value == "draft" - db.session.commit.assert_called_once() + session.commit.assert_called_once() def test_skill_inspect_by_agent_returns_strict_json_response(): @@ -128,13 +135,7 @@ def test_skill_inspect_by_agent_returns_strict_json_response(): "hash": "sha256:abc", "source": "config_skill_zip", "files": [ - { - "path": "SKILL.md", - "name": "SKILL.md", - "type": "file", - "previewable": True, - "downloadable": True, - } + {"path": "SKILL.md", "name": "SKILL.md", "type": "file", "previewable": True, "downloadable": True} ], "skill_md": { "path": "SKILL.md", @@ -145,8 +146,9 @@ def test_skill_inspect_by_agent_returns_strict_json_response(): }, "warnings": [], } - response = raw(AgentConfigSkillInspectByAgentApi(), "tenant-1", _USER, "agent-1", "pdf-toolkit") - + response = raw( + AgentConfigSkillInspectByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1", "pdf-toolkit" + ) assert response.status_code == 200 assert response.get_json()["name"] == "pdf-toolkit" assert b"PDF Toolkit" in response.get_data() @@ -155,10 +157,7 @@ def test_skill_inspect_by_agent_returns_strict_json_response(): def test_file_preview_api_passes_through_and_maps_errors(): raw = _raw(AgentConfigFilePreviewApi.get) with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.preview_file.return_value = { @@ -168,20 +167,16 @@ def test_file_preview_api_passes_through_and_maps_errors(): "binary": False, "text": "hello", } - body = raw(AgentConfigFilePreviewApi(), _USER, _APP, "sample.txt") + body = raw(AgentConfigFilePreviewApi(), MagicMock(), _USER, _APP, "sample.txt") assert body["text"] == "hello" - with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.preview_file.side_effect = AgentConfigServiceError( "config_file_not_found", "missing", status_code=404 ) - body, status = raw(AgentConfigFilePreviewApi(), _USER, _APP, "missing.txt") + body, status = raw(AgentConfigFilePreviewApi(), MagicMock(), _USER, _APP, "missing.txt") assert status == 404 assert body["code"] == "config_file_not_found" @@ -204,8 +199,7 @@ def test_skill_upload_by_agent_delegates_after_version_resolution(): ) as upload_skill, ): composer.load_agent_app_build_draft.return_value = {"draft": {"id": "build-draft-1"}} - body, status = raw(AgentConfigSkillUploadByAgentApi(), "tenant-1", _USER, "agent-1") - + body, status = raw(AgentConfigSkillUploadByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1") assert status == 201 assert body["skill"]["name"] == "alpha" assert upload_skill.call_args.kwargs["version_id"] == "build-draft-1" @@ -219,10 +213,7 @@ def test_file_upload_by_agent_delegates_to_service_owned_upload_lookup(): patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP), patch(f"{_MOD}.AgentComposerService") as composer, patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value={"upload_file_id": "upload-1"}, + type(console_ns), "payload", new_callable=PropertyMock, return_value={"upload_file_id": "upload-1"} ), patch(f"{_MOD}.AgentConfigService") as config_service, ): @@ -231,8 +222,7 @@ def test_file_upload_by_agent_delegates_to_service_owned_upload_lookup(): "file": {"id": "guide.txt", "name": "guide.txt", "file_id": "upload-1"}, "config_version": {"id": "build-draft-1", "kind": "build_draft", "writable": True}, } - body, status = raw(AgentConfigFilesByAgentApi(), "tenant-1", _USER, "agent-1") - + body, status = raw(AgentConfigFilesByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1") assert status == 201 assert body["file"]["name"] == "guide.txt" assert config_service.return_value.push_file_for_console.call_args.kwargs["upload_file_id"] == "upload-1" @@ -245,10 +235,7 @@ def test_file_upload_by_agent_delegates_to_service_owned_upload_lookup(): def test_skill_list_api_uses_config_list_shape() -> None: raw = _raw(AgentConfigSkillsApi.get) with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.list_skills.return_value = { @@ -256,8 +243,7 @@ def test_skill_list_api_uses_config_list_shape() -> None: "config_version": {"id": "draft-1", "kind": "draft", "writable": True}, "items": [{"id": "alpha", "name": "alpha", "file_id": "tool-file-1", "description": "Alpha"}], } - body = raw(AgentConfigSkillsApi(), _USER, _APP) - + body = raw(AgentConfigSkillsApi(), MagicMock(), _USER, _APP) assert body["items"][0]["name"] == "alpha" assert body["items"][0]["file_id"] == "tool-file-1" assert config_service.return_value.list_skills.call_args.kwargs["agent_id"] == "wf-agent-9" @@ -266,10 +252,7 @@ def test_skill_list_api_uses_config_list_shape() -> None: def test_file_list_api_uses_config_list_shape() -> None: raw = _raw(AgentConfigFilesApi.get) with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.list_files.return_value = { @@ -286,8 +269,7 @@ def test_file_list_api_uses_config_list_shape() -> None: } ], } - body = raw(AgentConfigFilesApi(), _USER, _APP) - + body = raw(AgentConfigFilesApi(), MagicMock(), _USER, _APP) assert body == { "agent_id": "wf-agent-9", "config_version": {"id": "draft-1", "kind": "draft", "writable": True}, @@ -321,8 +303,7 @@ def test_skill_file_preview_by_agent_reads_path_query() -> None: "binary": False, "text": "hello world", } - body = raw(AgentConfigSkillFilePreviewByAgentApi(), "tenant-1", _USER, "agent-1", "alpha") - + body = raw(AgentConfigSkillFilePreviewByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1", "alpha") assert body["path"] == "references/guide.md" assert config_service.return_value.preview_skill_file.call_args.kwargs["path"] == "references/guide.md" @@ -341,8 +322,7 @@ def test_skill_file_download_by_agent_returns_proxy_url() -> None: ): composer.load_agent_app_build_draft.return_value = {"draft": {"id": "build-draft-1"}} config_service.return_value.resolve_skill_file_member_path.return_value = "references/guide.md" - response = raw(AgentConfigSkillFileDownloadByAgentApi(), "tenant-1", _USER, "agent-1", "alpha") - + response = raw(AgentConfigSkillFileDownloadByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1", "alpha") assert response.status_code == 200 assert response.get_json()["url"].endswith( "/agent/agent-1/config/skills/alpha/files/content?path=references%2Fguide.md&draft_type=debug_build" @@ -361,8 +341,9 @@ def test_skill_file_download_by_agent_validates_member_path() -> None: config_service.return_value.resolve_skill_file_member_path.side_effect = AgentConfigServiceError( "config_skill_file_not_found", "missing", status_code=404 ) - body, status = raw(AgentConfigSkillFileDownloadByAgentApi(), "tenant-1", _USER, "agent-1", "alpha") - + body, status = raw( + AgentConfigSkillFileDownloadByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1", "alpha" + ) assert status == 404 assert body["code"] == "config_skill_file_not_found" @@ -375,21 +356,16 @@ def test_skill_file_download_api_forwards_workflow_node_and_draft_type() -> None patch(f"{_MOD}.AgentConfigService") as config_service, patch( f"{_MOD}.url_for", - return_value=( - "/console/api/apps/app-1/agent/config/skills/alpha/files/content" - "?node_id=node-1&draft_type=debug_build&path=references%2Fguide.md" - ), + return_value="/console/api/apps/app-1/agent/config/skills/alpha/files/content?node_id=node-1&draft_type=debug_build&path=references%2Fguide.md", ) as url_for_mock, ): composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_app_build_draft.return_value = {"draft": {"id": "build-draft-1"}} config_service.return_value.resolve_skill_file_member_path.return_value = "references/guide.md" - response = raw(AgentConfigSkillFileDownloadApi(), _USER, _APP, "alpha") - + response = raw(AgentConfigSkillFileDownloadApi(), MagicMock(), _USER, _APP, "alpha") assert response.status_code == 200 assert response.get_json()["url"].endswith( - "/apps/app-1/agent/config/skills/alpha/files/content" - "?node_id=node-1&draft_type=debug_build&path=references%2Fguide.md" + "/apps/app-1/agent/config/skills/alpha/files/content?node_id=node-1&draft_type=debug_build&path=references%2Fguide.md" ) assert url_for_mock.call_args.kwargs == { "_external": False, @@ -404,17 +380,13 @@ def test_skill_file_download_api_forwards_workflow_node_and_draft_type() -> None def test_skill_file_download_api_propagates_member_lookup_404s() -> None: raw = _raw(AgentConfigSkillFileDownloadApi.get) with app.test_request_context("/?node_id=node-1&path=references/missing.md"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.resolve_skill_file_member_path.side_effect = AgentConfigServiceError( "config_skill_file_not_found", "missing", status_code=404 ) - body, status = raw(AgentConfigSkillFileDownloadApi(), _USER, _APP, "alpha") - + body, status = raw(AgentConfigSkillFileDownloadApi(), MagicMock(), _USER, _APP, "alpha") assert status == 404 assert body["code"] == "config_skill_file_not_found" @@ -422,14 +394,10 @@ def test_skill_file_download_api_propagates_member_lookup_404s() -> None: def test_file_download_api_returns_signed_url_json() -> None: raw = _raw(AgentConfigFileDownloadApi.get) with app.test_request_context("/?node_id=node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentConfigService") as config_service, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentConfigService") as config_service: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} config_service.return_value.download_file_url.return_value = "https://example.com/guide.txt" - response = raw(AgentConfigFileDownloadApi(), _USER, _APP, "guide.txt") - + response = raw(AgentConfigFileDownloadApi(), MagicMock(), _USER, _APP, "guide.txt") assert response.status_code == 200 assert response.get_json() == {"url": "https://example.com/guide.txt"} diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py b/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py index c8b28ef29f0..453e2f29ff0 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py @@ -7,12 +7,13 @@ resolution, query handling, and error mapping, not auth. from __future__ import annotations -import inspect +from inspect import unwrap from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch from flask import Flask +from controllers.console.app import agent_drive_inspector as inspector from controllers.console.app.agent_drive_inspector import ( AgentDriveDownloadApi, AgentDriveDownloadByAgentApi, @@ -32,10 +33,25 @@ app = Flask(__name__) def _raw(method): - return inspect.unwrap(method) + return unwrap(method) -_APP = SimpleNamespace(id="app-1", tenant_id="tenant-1", bound_agent_id="agent-1") +_APP = SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + bound_agent_id_with_session=lambda *, session: "agent-1", +) + + +def test_resolve_bound_agent_uses_injected_session(): + session = MagicMock() + resolver = MagicMock(return_value="agent-1") + app_model = SimpleNamespace(bound_agent_id_with_session=resolver) + result = inspector._resolve_agent_id(session, app_model, None) + + assert result == "agent-1" + resolver.assert_called_once_with(session=session) + assert resolver.call_args.kwargs["session"] is session def test_list_filters_value_pointers_out_of_console_payload(): @@ -53,8 +69,7 @@ def test_list_filters_value_pointers_out_of_console_payload(): "created_at": 1718000000, } ] - body = raw(AgentDriveListApi(), _APP) - + body = raw(AgentDriveListApi(), MagicMock(), _APP) assert body["items"][0]["key"] == "pdf-toolkit/SKILL.md" assert "file_id" not in body["items"][0] assert drive.return_value.manifest.call_args.kwargs["prefix"] == "pdf-toolkit/" @@ -62,6 +77,7 @@ def test_list_filters_value_pointers_out_of_console_payload(): def test_list_by_agent_filters_value_pointers_out_of_console_payload(): raw = _raw(AgentDriveListByAgentApi.get) + session = MagicMock() with app.test_request_context("/?prefix=pdf-toolkit/"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -78,31 +94,28 @@ def test_list_by_agent_filters_value_pointers_out_of_console_payload(): "created_at": 1718000000, } ] - body = raw(AgentDriveListByAgentApi(), "tenant-1", "agent-1") - + body = raw(AgentDriveListByAgentApi(), session, "tenant-1", "agent-1") assert body["items"][0]["key"] == "pdf-toolkit/SKILL.md" assert "file_id" not in body["items"][0] - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") assert drive.return_value.manifest.call_args.kwargs["agent_id"] == "agent-1" + assert drive.return_value.manifest.call_args.kwargs["session"] is session def test_list_resolves_workflow_node_binding_agent(): raw = _raw(AgentDriveListApi.get) with app.test_request_context("/?node_id=agent-node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.manifest.return_value = [] - raw(AgentDriveListApi(), _APP) - + raw(AgentDriveListApi(), MagicMock(), _APP) assert drive.return_value.manifest.call_args.kwargs["agent_id"] == "wf-agent-9" assert composer.resolve_workflow_node_agent_id.call_args.kwargs["node_id"] == "agent-node-1" def test_skill_list_by_agent_calls_service(): raw = _raw(AgentDriveSkillListByAgentApi.get) + session = MagicMock() with app.test_request_context("/"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -121,30 +134,27 @@ def test_skill_list_by_agent_calls_service(): "created_at": 1718000000, } ] - body = raw(AgentDriveSkillListByAgentApi(), "tenant-1", "agent-1") - + body = raw(AgentDriveSkillListByAgentApi(), session, "tenant-1", "agent-1") assert body["items"][0]["path"] == "pdf-toolkit" - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") assert drive.return_value.list_skills.call_args.kwargs["agent_id"] == "agent-1" + assert drive.return_value.list_skills.call_args.kwargs["session"] is session def test_skill_list_resolves_workflow_node_binding_agent(): raw = _raw(AgentDriveSkillListApi.get) with app.test_request_context("/?node_id=agent-node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.list_skills.return_value = [] - body = raw(AgentDriveSkillListApi(), _APP) - + body = raw(AgentDriveSkillListApi(), MagicMock(), _APP) assert body == {"items": []} assert drive.return_value.list_skills.call_args.kwargs["agent_id"] == "wf-agent-9" def test_skill_inspect_by_agent_returns_strict_json_response(): raw = _raw(AgentDriveSkillInspectByAgentApi.get) + session = MagicMock() payload = { "path": "pdf-toolkit", "skill_md_key": "pdf-toolkit/SKILL.md", @@ -181,11 +191,11 @@ def test_skill_inspect_by_agent_returns_strict_json_response(): patch(f"{_MOD}.AgentDriveService") as drive, ): drive.return_value.inspect_skill.return_value = payload - response = raw(AgentDriveSkillInspectByAgentApi(), "tenant-1", "agent-1", "pdf-toolkit") - + response = raw(AgentDriveSkillInspectByAgentApi(), session, "tenant-1", "agent-1", "pdf-toolkit") assert response.status_code == 200 assert response.get_json()["skill_md"]["text"] == "# PDF Toolkit\nUse it.\n" assert b"# PDF Toolkit\\nUse it.\\n" in response.get_data() + assert drive.return_value.inspect_skill.call_args.kwargs["session"] is session def test_skill_inspect_resolves_workflow_node_binding_agent(): @@ -207,25 +217,24 @@ def test_skill_inspect_resolves_workflow_node_binding_agent(): "warnings": [], } with app.test_request_context("/?node_id=agent-node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.inspect_skill.return_value = payload - response = raw(AgentDriveSkillInspectApi(), _APP, "pdf-toolkit") - + response = raw(AgentDriveSkillInspectApi(), MagicMock(), _APP, "pdf-toolkit") assert response.get_json()["path"] == "pdf-toolkit" assert drive.return_value.inspect_skill.call_args.kwargs["agent_id"] == "wf-agent-9" def test_list_400_when_no_agent_bound(): raw = _raw(AgentDriveListApi.get) - app_without_agent = SimpleNamespace(id="app-1", tenant_id="tenant-1", bound_agent_id=None) + resolver = MagicMock(return_value=None) + app_without_agent = SimpleNamespace(bound_agent_id_with_session=resolver) + session = MagicMock() with app.test_request_context("/"): - body, status = raw(AgentDriveListApi(), app_without_agent) + body, status = raw(AgentDriveListApi(), session, app_without_agent) assert status == 400 assert body["code"] == "agent_not_bound" + resolver.assert_called_once_with(session=session) def test_preview_passes_through_and_maps_errors(): @@ -239,21 +248,21 @@ def test_preview_passes_through_and_maps_errors(): "binary": False, "text": "# hi", } - body = raw(AgentDrivePreviewApi(), _APP) + body = raw(AgentDrivePreviewApi(), MagicMock(), _APP) assert body["text"] == "# hi" - with app.test_request_context("/?key=ghost/SKILL.md"): with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.preview.side_effect = AgentDriveError( "drive_key_not_found", "no drive entry", status_code=404 ) - body, status = raw(AgentDrivePreviewApi(), _APP) + body, status = raw(AgentDrivePreviewApi(), MagicMock(), _APP) assert status == 404 assert body["code"] == "drive_key_not_found" def test_preview_by_agent_passes_through_and_maps_errors(): raw = _raw(AgentDrivePreviewByAgentApi.get) + session = MagicMock() with app.test_request_context("/?key=pdf-toolkit/SKILL.md"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -266,10 +275,10 @@ def test_preview_by_agent_passes_through_and_maps_errors(): "binary": False, "text": "# hi", } - body = raw(AgentDrivePreviewByAgentApi(), "tenant-1", "agent-1") + body = raw(AgentDrivePreviewByAgentApi(), session, "tenant-1", "agent-1") assert body["text"] == "# hi" - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") - + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + assert drive.return_value.preview.call_args.kwargs["session"] is session with app.test_request_context("/?key=ghost/SKILL.md"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP), @@ -278,7 +287,7 @@ def test_preview_by_agent_passes_through_and_maps_errors(): drive.return_value.preview.side_effect = AgentDriveError( "drive_key_not_found", "no drive entry", status_code=404 ) - body, status = raw(AgentDrivePreviewByAgentApi(), "tenant-1", "agent-1") + body, status = raw(AgentDrivePreviewByAgentApi(), session, "tenant-1", "agent-1") assert status == 404 assert body["code"] == "drive_key_not_found" @@ -288,18 +297,20 @@ def test_download_returns_signed_url_json(): with app.test_request_context("/?key=pdf-toolkit/.DIFY-SKILL-FULL.zip"): with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.download_url.return_value = "https://signed.example/zip" - body = raw(AgentDriveDownloadApi(), _APP) + body = raw(AgentDriveDownloadApi(), MagicMock(), _APP) assert body == {"url": "https://signed.example/zip"} def test_download_by_agent_returns_signed_url_json(): raw = _raw(AgentDriveDownloadByAgentApi.get) + session = MagicMock() with app.test_request_context("/?key=pdf-toolkit/.DIFY-SKILL-FULL.zip"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): drive.return_value.download_url.return_value = "https://signed.example/zip" - body = raw(AgentDriveDownloadByAgentApi(), "tenant-1", "agent-1") + body = raw(AgentDriveDownloadByAgentApi(), session, "tenant-1", "agent-1") assert body == {"url": "https://signed.example/zip"} - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + assert drive.return_value.download_url.call_args.kwargs["session"] is session diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_skills.py b/api/tests/unit_tests/controllers/console/app/test_agent_skills.py index 79a363e74ca..e555e9c5094 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_skills.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_skills.py @@ -7,14 +7,15 @@ bare Flask request context with the services mocked — covering request handlin from __future__ import annotations -import inspect import io +from inspect import unwrap from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest from flask import Flask +from controllers.console.app import agent as agent_controller from controllers.console.app.agent import ( AgentDriveFilesByAgentApi, AgentSkillByAgentApi, @@ -31,7 +32,7 @@ app = Flask(__name__) def _raw(method): - return inspect.unwrap(method) + return unwrap(method) def _file_ctx(*, files: dict[str, bytes] | None = None): @@ -40,23 +41,40 @@ def _file_ctx(*, files: dict[str, bytes] | None = None): _USER = SimpleNamespace(id="user-1") -_APP = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT, bound_agent_id="agent-1") -_WORKFLOW_APP = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.WORKFLOW, bound_agent_id=None) +_APP = SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + mode=AppMode.AGENT, + bound_agent_id_with_session=lambda *, session: "agent-1", +) +_WORKFLOW_APP = SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + mode=AppMode.WORKFLOW, + bound_agent_id_with_session=lambda *, session: None, +) + + +def test_resolve_bound_agent_uses_injected_session(): + session = MagicMock() + resolver = MagicMock(return_value="agent-1") + app_model = SimpleNamespace(bound_agent_id_with_session=resolver) + result = agent_controller._resolve_agent_id(session, app_model, None) + + assert result == "agent-1" + resolver.assert_called_once_with(session=session) + assert resolver.call_args.kwargs["session"] is session def test_upload_standardizes_into_drive_and_returns_skill_ref(): raw = _raw(AgentSkillUploadApi.post) - with _file_ctx(files={"file": b"zip-bytes"}): - with ( - patch(f"{_MOD}.SkillStandardizeService") as svc, - ): + with patch(f"{_MOD}.SkillStandardizeService") as svc: svc.return_value.standardize.return_value = { "skill": {"path": "skill-a", "skill_md_key": "skill-a/SKILL.md"}, "manifest": {"name": "Skill A"}, } - body, status = raw(AgentSkillUploadApi(), _USER, _APP) - + body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) assert status == 201 assert body["skill"] == {"path": "skill-a", "skill_md_key": "skill-a/SKILL.md"} assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "agent-1" @@ -64,25 +82,24 @@ def test_upload_standardizes_into_drive_and_returns_skill_ref(): def test_upload_by_agent_resolves_app_and_standardizes_into_drive(): raw = _raw(AgentSkillUploadByAgentApi.post) - with _file_ctx(files={"file": b"zip-bytes"}): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.SkillStandardizeService") as svc, ): + session = MagicMock() svc.return_value.standardize.return_value = {"skill": {"path": "skill-a"}, "manifest": {}} - body, status = raw(AgentSkillUploadByAgentApi(), "tenant-1", _USER, "agent-1") - + body, status = raw(AgentSkillUploadByAgentApi(), session, "tenant-1", _USER, "agent-1") assert status == 201 assert body["skill"] == {"path": "skill-a"} - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "agent-1" def test_upload_no_file_is_400(): raw = _raw(AgentSkillUploadApi.post) with _file_ctx(files={}): - body, status = raw(AgentSkillUploadApi(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) assert status == 400 assert body["code"] == "no_file" @@ -94,18 +111,21 @@ def test_upload_maps_package_error(): svc.return_value.standardize.side_effect = SkillPackageError( "missing_skill_md", "no SKILL.md", status_code=400 ) - body, status = raw(AgentSkillUploadApi(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) assert status == 400 assert body["code"] == "missing_skill_md" def test_upload_no_bound_agent_is_400(): raw = _raw(AgentSkillUploadApi.post) - app_without_agent = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT, bound_agent_id=None) + resolver = MagicMock(return_value=None) + app_without_agent = SimpleNamespace(bound_agent_id_with_session=resolver) + session = MagicMock() with _file_ctx(files={"file": b"zip"}): - body, status = raw(AgentSkillUploadApi(), _USER, app_without_agent) + body, status = raw(AgentSkillUploadApi(), session, _USER, app_without_agent) assert status == 400 assert body["code"] == "agent_not_bound" + resolver.assert_called_once_with(session=session) def test_upload_resolves_workflow_node_agent(): @@ -113,14 +133,10 @@ def test_upload_resolves_workflow_node_agent(): with app.test_request_context( "/?node_id=agent-node-1", method="POST", data={"file": (io.BytesIO(b"zip"), "skill.zip")} ): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.SkillStandardizeService") as svc, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.SkillStandardizeService") as svc: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" svc.return_value.standardize.return_value = {"skill": {"path": "s"}, "manifest": {}} - body, status = raw(AgentSkillUploadApi(), _USER, _WORKFLOW_APP) - + body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _WORKFLOW_APP) assert status == 201 assert body["skill"] == {"path": "s"} assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "wf-agent-1" @@ -131,14 +147,11 @@ def test_upload_maps_drive_error(): with _file_ctx(files={"file": b"zip"}): with patch(f"{_MOD}.SkillStandardizeService") as svc: svc.return_value.standardize.side_effect = AgentDriveError("source_not_found", "nope", status_code=404) - body, status = raw(AgentSkillUploadApi(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) assert status == 404 assert body["code"] == "source_not_found" -# ── ENG-625: drive files commit + delete endpoints ──────────────────────────── - - def _json_ctx(payload: dict | None = None, *, method: str = "POST", query_string: str = ""): return app.test_request_context(f"/?{query_string}", method=method, json=payload or {}) @@ -149,18 +162,14 @@ def test_files_commit_validates_upload_and_returns_drive_ref(): raw = _raw(AgentDriveFilesApi.post) upload = SimpleNamespace(id="uf-1", name="sample qna.pdf") with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}): - with ( - patch(f"{_MOD}.console_ns") as ns, - patch(f"{_MOD}.db") as db_mock, - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.console_ns") as ns, patch(f"{_MOD}.AgentDriveService") as drive: + session = MagicMock() ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - db_mock.session.scalar.return_value = upload + session.scalar.return_value = upload drive.return_value.commit.return_value = [ {"key": "files/sample qna.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesApi(), _USER, _APP) - + body, status = raw(AgentDriveFilesApi(), session, _USER, _APP) assert status == 201 assert body["file"]["drive_key"] == "files/sample qna.pdf" assert body["file"]["file_id"] == "uf-1" @@ -176,18 +185,17 @@ def test_files_by_agent_commit_uses_agent_route_and_ignores_node_id(): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.console_ns") as ns, - patch(f"{_MOD}.db") as db_mock, patch(f"{_MOD}.AgentDriveService") as drive, ): + session = MagicMock() ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - db_mock.session.scalar.return_value = upload + session.scalar.return_value = upload drive.return_value.commit.return_value = [ {"key": "files/sample.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesByAgentApi(), "tenant-1", _USER, "agent-1") - + body, status = raw(AgentDriveFilesByAgentApi(), session, "tenant-1", _USER, "agent-1") assert status == 201 - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") def test_files_commit_404_when_upload_not_in_tenant(): @@ -195,13 +203,11 @@ def test_files_commit_404_when_upload_not_in_tenant(): raw = _raw(AgentDriveFilesApi.post) with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}): - with ( - patch(f"{_MOD}.console_ns") as ns, - patch(f"{_MOD}.db") as db_mock, - ): + with patch(f"{_MOD}.console_ns") as ns: + session = MagicMock() ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - db_mock.session.scalar.return_value = None - body, status = raw(AgentDriveFilesApi(), _USER, _APP) + session.scalar.return_value = None + body, status = raw(AgentDriveFilesApi(), session, _USER, _APP) assert status == 404 assert body["code"] == "upload_file_not_found" @@ -214,18 +220,17 @@ def test_files_commit_resolves_workflow_node_agent(): with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}, query_string="node_id=agent-node-1"): with ( patch(f"{_MOD}.console_ns") as ns, - patch(f"{_MOD}.db") as db_mock, patch(f"{_MOD}.AgentDriveService") as drive, patch(f"{_MOD}.AgentComposerService") as composer, ): + session = MagicMock() ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - db_mock.session.scalar.return_value = upload + session.scalar.return_value = upload composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" drive.return_value.commit.return_value = [ {"key": "files/sample.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesApi(), _USER, _WORKFLOW_APP) - + body, status = raw(AgentDriveFilesApi(), session, _USER, _WORKFLOW_APP) assert status == 201 assert drive.return_value.commit.call_args.kwargs["agent_id"] == "wf-agent-1" @@ -236,14 +241,11 @@ def test_files_delete_updates_soul_then_drive(): raw = _raw(AgentDriveFilesApi.delete) calls: list[str] = [] with _json_ctx(method="DELETE", query_string="key=files/sample.pdf"): - with ( - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.commit.side_effect = lambda **kw: ( calls.append("drive") or [{"key": "files/sample.pdf", "removed": True}] ) - body = raw(AgentDriveFilesApi(), _USER, _APP) - + body = raw(AgentDriveFilesApi(), MagicMock(), _USER, _APP) assert calls == ["drive"] assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} @@ -255,11 +257,11 @@ def test_files_by_agent_delete_uses_agent_route_and_ignores_node_id(): patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): + session = MagicMock() drive.return_value.commit.return_value = [{"key": "files/sample.pdf", "removed": True}] - body = raw(AgentDriveFilesByAgentApi(), "tenant-1", _USER, "agent-1") - + body = raw(AgentDriveFilesByAgentApi(), session, "tenant-1", _USER, "agent-1") assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") def test_files_delete_resolves_workflow_node_agent(): @@ -267,14 +269,10 @@ def test_files_delete_resolves_workflow_node_agent(): raw = _raw(AgentDriveFilesApi.delete) with _json_ctx(method="DELETE", query_string="key=files/sample.pdf&node_id=agent-node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" drive.return_value.commit.return_value = [{"key": "files/sample.pdf", "removed": True}] - body = raw(AgentDriveFilesApi(), _USER, _WORKFLOW_APP) - + body = raw(AgentDriveFilesApi(), MagicMock(), _USER, _WORKFLOW_APP) assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} assert drive.return_value.commit.call_args.kwargs["agent_id"] == "wf-agent-1" @@ -284,12 +282,10 @@ def test_files_delete_survives_drive_failure(): raw = _raw(AgentDriveFilesApi.delete) with _json_ctx(method="DELETE", query_string="key=files/sample.pdf"): - with ( - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.commit.side_effect = RuntimeError("storage down") with pytest.raises(RuntimeError, match="storage down"): - raw(AgentDriveFilesApi(), _USER, _APP) + raw(AgentDriveFilesApi(), MagicMock(), _USER, _APP) def test_skill_delete_uses_slug_prefix_and_is_idempotent(): @@ -297,15 +293,12 @@ def test_skill_delete_uses_slug_prefix_and_is_idempotent(): raw = _raw(AgentSkillApi.delete) with _json_ctx(method="DELETE"): - with ( - patch(f"{_MOD}.AgentDriveService") as drive, - ): + with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.commit.return_value = [ {"key": "tender-analyzer/SKILL.md", "removed": True}, {"key": "tender-analyzer/.DIFY-SKILL-FULL.zip", "removed": True}, ] - body = raw(AgentSkillApi(), _USER, _APP, "tender-analyzer") - + body = raw(AgentSkillApi(), MagicMock(), _USER, _APP, "tender-analyzer") assert body == { "result": "success", "removed_keys": ["tender-analyzer/SKILL.md", "tender-analyzer/.DIFY-SKILL-FULL.zip"], @@ -319,11 +312,11 @@ def test_skill_delete_by_agent_uses_agent_route(): patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): + session = MagicMock() drive.return_value.commit.return_value = [{"key": "tender-analyzer/SKILL.md", "removed": True}] - body = raw(AgentSkillByAgentApi(), "tenant-1", _USER, "agent-1", "tender-analyzer") - + body = raw(AgentSkillByAgentApi(), session, "tenant-1", _USER, "agent-1", "tender-analyzer") assert body == {"result": "success", "removed_keys": ["tender-analyzer/SKILL.md"]} - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") def test_skill_delete_rejects_path_like_slug(): @@ -331,14 +324,11 @@ def test_skill_delete_rejects_path_like_slug(): raw = _raw(AgentSkillApi.delete) with _json_ctx(method="DELETE"): - body, status = raw(AgentSkillApi(), _USER, _APP, "a/b") + body, status = raw(AgentSkillApi(), MagicMock(), _USER, _APP, "a/b") assert status == 400 assert body["code"] == "drive_key_invalid" -# ── ENG-371: infer-tools endpoint ───────────────────────────────────────────── - - def test_infer_tools_returns_draft_suggestions(): from controllers.console.app.agent import AgentSkillInferToolsApi @@ -350,8 +340,7 @@ def test_infer_tools_returns_draft_suggestions(): "cli_tools": [{"name": "ffmpeg", "inferred_from": "audio-transcribe"}], "reason": None, } - body = raw(AgentSkillInferToolsApi(), _APP, "audio-transcribe") - + body = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "audio-transcribe") assert body["inferable"] is True assert svc.return_value.infer.call_args.kwargs["slug"] == "audio-transcribe" @@ -363,11 +352,11 @@ def test_infer_tools_by_agent_uses_agent_route(): patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.SkillToolInferenceService") as svc, ): + session = MagicMock() svc.return_value.infer.return_value = {"inferable": True, "cli_tools": [], "reason": None} - body = raw(AgentSkillInferToolsByAgentApi(), "tenant-1", "agent-1", "audio-transcribe") - + body = raw(AgentSkillInferToolsByAgentApi(), session, "tenant-1", "agent-1", "audio-transcribe") assert body["inferable"] is True - resolve_app.assert_called_once_with(tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") assert svc.return_value.infer.call_args.kwargs["agent_id"] == "agent-1" @@ -376,14 +365,10 @@ def test_infer_tools_resolves_workflow_node_agent(): raw = _raw(AgentSkillInferToolsApi.post) with _json_ctx(query_string="node_id=agent-node-1"): - with ( - patch(f"{_MOD}.AgentComposerService") as composer, - patch(f"{_MOD}.SkillToolInferenceService") as svc, - ): + with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.SkillToolInferenceService") as svc: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" svc.return_value.infer.return_value = {"inferable": False, "cli_tools": [], "reason": "none"} - body = raw(AgentSkillInferToolsApi(), _WORKFLOW_APP, "audio-transcribe") - + body = raw(AgentSkillInferToolsApi(), MagicMock(), _WORKFLOW_APP, "audio-transcribe") assert body["inferable"] is False assert svc.return_value.infer.call_args.kwargs["agent_id"] == "wf-agent-1" @@ -398,7 +383,7 @@ def test_infer_tools_maps_inference_errors(): svc.return_value.infer.side_effect = SkillToolInferenceError( "default_model_not_configured", "no model", status_code=400 ) - body, status = raw(AgentSkillInferToolsApi(), _APP, "audio-transcribe") + body, status = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "audio-transcribe") assert status == 400 assert body["code"] == "default_model_not_configured" @@ -408,10 +393,9 @@ def test_infer_tools_rejects_path_like_slug_and_unbound_app(): raw = _raw(AgentSkillInferToolsApi.post) with _json_ctx(): - body, status = raw(AgentSkillInferToolsApi(), _APP, "a/b") + body, status = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "a/b") assert (status, body["code"]) == (400, "drive_key_invalid") - - app_without_agent = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT, bound_agent_id=None) + app_without_agent = SimpleNamespace(bound_agent_id_with_session=MagicMock(return_value=None)) with _json_ctx(): - body, status = raw(AgentSkillInferToolsApi(), app_without_agent, "x") + body, status = raw(AgentSkillInferToolsApi(), MagicMock(), app_without_agent, "x") assert (status, body["code"]) == (400, "agent_not_bound") 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 cc95f7f8a94..b12a015217e 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 ANY, Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -110,16 +110,17 @@ def test_annotation_file_payload_valid(): def test_get_app_ref_raises_not_found_when_app_is_not_in_current_tenant(): + session = MagicMock() + session.scalar.return_value = None with ( patch.object( annotation_module, "current_account_with_tenant", return_value=(SimpleNamespace(id="account-1"), "tenant-1"), ), - patch.object(annotation_module.db.session, "scalar", return_value=None), ): with pytest.raises(NotFound): - annotation_module._get_app_ref("app-1") + annotation_module._get_app_ref(session, "app-1") class TestConsoleAnnotationRefBoundaries: @@ -127,6 +128,8 @@ class TestConsoleAnnotationRefBoundaries: api = annotation_module.AnnotationApi() handler = unwrap(api.delete) delete_mock = Mock() + session = MagicMock() + session.scalar.return_value = _app_model() with ( app.test_request_context("/?annotation_id=ann-1&annotation_id=ann-2", method="DELETE"), @@ -135,20 +138,21 @@ class TestConsoleAnnotationRefBoundaries: "current_account_with_tenant", return_value=(SimpleNamespace(id="account-1"), "tenant-1"), ), - patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), patch.object(annotation_module.AppAnnotationService, "delete_app_annotations_in_batch", delete_mock), ): - response, status = handler(api, "app-1") + response, status = handler(api, session, "app-1") assert response == "" assert status == 204 - delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], session=ANY) + delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], session) def test_update_uses_annotation_ref(self, app: Flask): api = annotation_module.AnnotationUpdateDeleteApi() handler = unwrap(api.post) update_mock = Mock(return_value=_annotation_model()) payload = {"question": "updated"} + session = MagicMock() + session.scalar.return_value = _app_model() with ( app.test_request_context("/annotations/ann-1", method="POST", json=payload), @@ -158,19 +162,21 @@ class TestConsoleAnnotationRefBoundaries: "current_account_with_tenant", return_value=(SimpleNamespace(id="account-1"), "tenant-1"), ), - patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), patch.object(annotation_module.AppAnnotationService, "update_app_annotation_directly", update_mock), ): - response = handler(api, "app-1", "ann-1") + response = handler(api, session, "app-1", "ann-1") assert response["question"] == "q" update_mock.assert_called_once() assert update_mock.call_args.args[1] == AnnotationRef("tenant-1", "app-1", "ann-1") + assert update_mock.call_args.args[2] is session def test_delete_uses_annotation_ref(self, app: Flask): api = annotation_module.AnnotationUpdateDeleteApi() handler = unwrap(api.delete) delete_mock = Mock() + session = MagicMock() + session.scalar.return_value = _app_model() with ( app.test_request_context("/annotations/ann-1", method="DELETE"), @@ -179,15 +185,15 @@ class TestConsoleAnnotationRefBoundaries: "current_account_with_tenant", return_value=(SimpleNamespace(id="account-1"), "tenant-1"), ), - patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), patch.object(annotation_module.AppAnnotationService, "delete_app_annotation", delete_mock), ): - response, status = handler(api, "app-1", "ann-1") + response, status = handler(api, session, "app-1", "ann-1") assert response == "" assert status == 204 delete_mock.assert_called_once() assert delete_mock.call_args.args[0] == AnnotationRef("tenant-1", "app-1", "ann-1") + assert delete_mock.call_args.args[1] is session def test_hit_history_uses_annotation_ref(self, app: Flask): api = annotation_module.AnnotationHitHistoryListApi() @@ -202,6 +208,8 @@ class TestConsoleAnnotationRefBoundaries: created_at=None, ) hit_history_mock = Mock(return_value=([history], 1)) + session = MagicMock() + session.scalar.return_value = _app_model() with ( app.test_request_context("/hit-histories?page=2&limit=5", method="GET"), @@ -210,10 +218,9 @@ class TestConsoleAnnotationRefBoundaries: "current_account_with_tenant", return_value=(SimpleNamespace(id="account-1"), "tenant-1"), ), - patch.object(annotation_module.db.session, "scalar", return_value=_app_model()), patch.object(annotation_module.AppAnnotationService, "get_annotation_hit_histories", hit_history_mock), ): - response = handler(api, "app-1", "ann-1") + response = handler(api, session, "app-1", "ann-1") assert response["total"] == 1 - hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5, session=ANY) + hit_history_mock.assert_called_once_with(AnnotationRef("tenant-1", "app-1", "ann-1"), 2, 5, session) 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 6a22d8769bc..2c8dd997761 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,6 +193,7 @@ class TestAnnotationImportServiceValidation: @pytest.fixture def mock_db_session(self): + """Mock database session.""" return MagicMock() def test_max_records_limit_enforced(self, mock_app, mock_db_session): @@ -212,7 +213,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, session=mock_db_session) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, mock_db_session) # Should return error about too many records assert "error_msg" in result @@ -229,7 +230,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, session=mock_db_session) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, mock_db_session) # Should return error about insufficient records assert "error_msg" in result @@ -248,7 +249,7 @@ class TestAnnotationImportServiceValidation: ): mock_auth.return_value = (MagicMock(id="user_id"), "tenant_id") - result = AppAnnotationService.batch_import_app_annotations("app_id", file, session=mock_db_session) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, mock_db_session) assert "error_msg" in result assert "malformed" in result["error_msg"].lower() @@ -269,9 +270,7 @@ 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, session=mock_db_session - ) + result = AppAnnotationService.batch_import_app_annotations("app_id", file, mock_db_session) # Should return success response assert "job_id" in result diff --git a/api/tests/unit_tests/controllers/console/app/test_app_apis.py b/api/tests/unit_tests/controllers/console/app/test_app_apis.py index 22b11dad6a2..36fa0723cc9 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_apis.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_apis.py @@ -227,7 +227,7 @@ class TestAppEndpoints: app.test_request_context("/console/api/apps/app-1", method="PUT", json=payload), patch.object(type(console_ns), "payload", payload), ): - response = method(api, app_model=_make_app(icon_type=app_module.IconType.EMOJI)) + response = method(api, MagicMock(spec=Session), app_model=_make_app(icon_type=app_module.IconType.EMOJI)) assert response == {"id": "app-1"} assert app_service.update_app.call_args.args[1]["icon_type"] is None @@ -264,7 +264,7 @@ class TestAppEndpoints: app.test_request_context("/console/api/apps/app-1/icon", method="POST", json=payload), patch.object(type(console_ns), "payload", payload), ): - response = method(api, app_model=_make_app()) + response = method(api, MagicMock(spec=Session), app_model=_make_app()) assert response == {"id": "app-1"} assert app_service.update_app_icon.call_args.args[1:] == ( @@ -367,7 +367,7 @@ class TestSiteEndpoints: site = self._add_site(db.session) with database_app.test_request_context("/", json={"title": "My Site", "input_placeholder": "Ask me anything"}): - result = method(api, _make_account(), app_model=_make_app()) + result = method(api, db.session, _make_account(), app_model=_make_app()) db.session.refresh(site) assert isinstance(result, dict) @@ -386,7 +386,7 @@ class TestSiteEndpoints: monkeypatch.setattr(site_module.Site, "generate_code", lambda *_args, **_kwargs: "code") with database_app.test_request_context("/"): - result = method(api, _make_account(), app_model=_make_app()) + result = method(api, db.session, _make_account(), app_model=_make_app()) db.session.refresh(site) assert isinstance(result, dict) 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 d4935ba2efd..6b1dce534ed 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 ANY, MagicMock +from unittest.mock import MagicMock import pytest from flask import Flask @@ -334,7 +334,7 @@ def test_create_app_endpoint_rejects_agent_mode(app_module, monkeypatch: pytest. app_module.console_ns.payload = payload try: with pytest.raises(ValidationError): - _unwrap(app_module.AppListApi().post)("tenant-1", SimpleNamespace(id="account-1")) + _unwrap(app_module.AppListApi().post)(MagicMock(), "tenant-1", SimpleNamespace(id="account-1")) finally: app_module.console_ns.payload = None @@ -439,6 +439,54 @@ def test_app_detail_with_site_includes_nested_serialization(app_models): assert "role" not in serialized +def test_app_response_view_uses_the_caller_session_for_query_backed_fields(app_module, monkeypatch): + session = MagicMock() + app_obj = MagicMock() + app_model_config = SimpleNamespace(app_id="app-1") + app_obj.desc_or_prompt_with_session.return_value = "Description" + app_obj.site_with_session.return_value = SimpleNamespace(id="site-1") + app_obj.app_model_config_with_session.return_value = app_model_config + app_obj.workflow_with_session.return_value = SimpleNamespace(id="workflow-1") + app_obj.bound_agent_id_with_session.return_value = "agent-1" + app_obj.mode_compatible_with_agent_with_session.return_value = "agent" + app_obj.deleted_tools_with_session.return_value = [] + app_obj.tags_with_session.return_value = [] + app_obj.author_name_with_session.return_value = "Author" + load_annotation_reply = MagicMock(return_value={"enabled": False}) + monkeypatch.setattr("services.app_service.load_annotation_reply_config", load_annotation_reply) + + view = app_module.AppResponseView(app_obj, session=session) + site = view.site + workflow = view.workflow + model_config = view.app_model_config + + assert view.desc_or_prompt == "Description" + assert site is not None + assert site.id == "site-1" + assert workflow is not None + assert workflow.id == "workflow-1" + assert view.bound_agent_id == "agent-1" + assert view.mode_compatible_with_agent == "agent" + assert view.deleted_tools == [] + assert view.tags == [] + assert view.author_name == "Author" + assert model_config is not None + assert model_config.annotation_reply_dict == {"enabled": False} + for method in ( + app_obj.desc_or_prompt_with_session, + app_obj.site_with_session, + app_obj.app_model_config_with_session, + app_obj.workflow_with_session, + app_obj.bound_agent_id_with_session, + app_obj.mode_compatible_with_agent_with_session, + app_obj.deleted_tools_with_session, + app_obj.tags_with_session, + app_obj.author_name_with_session, + ): + method.assert_called_once_with(session=session) + load_annotation_reply.assert_called_once_with(session, "app-1") + + def test_app_pagination_aliases_per_page_and_has_next(app_models): AppPagination = app_models.AppPagination item_one = SimpleNamespace( @@ -513,10 +561,8 @@ def test_app_list_uses_injected_session_for_draft_workflows( "FeatureService", SimpleNamespace(get_system_features=lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False))), ) - monkeypatch.setattr( - app_module.enterprise_rbac_service.RBACService.MyPermissions, - "get", - lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( + get_permissions = MagicMock( + return_value=app_module.enterprise_rbac_service.MyPermissionsResponse( app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( @@ -525,7 +571,12 @@ def test_app_list_uses_injected_session_for_draft_workflows( ) ] ) - ), + ) + ) + monkeypatch.setattr( + app_module.enterprise_rbac_service.RBACService.MyPermissions, + "get", + get_permissions, ) monkeypatch.setattr(app_module, "db", SimpleNamespace(session=scoped_session)) @@ -536,6 +587,7 @@ def test_app_list_uses_injected_session_for_draft_workflows( assert response["data"][0]["has_draft_trigger"] is True session.execute.assert_called_once() scoped_session.execute.assert_not_called() + get_permissions.assert_called_once_with("tenant-1", "user-1", session=session) assert response["data"][0]["permission_keys"] == ["app.acl.edit"] @@ -585,7 +637,7 @@ def test_app_create_api_attaches_permission_keys(app, app_module): replace_whitelist, ) - resp, status = method(app_module.AppListApi(), "tenant-1", SimpleNamespace(id="acct-1")) + resp, status = method(app_module.AppListApi(), MagicMock(), "tenant-1", SimpleNamespace(id="acct-1")) assert status == 201 assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"] @@ -810,7 +862,8 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): with app.test_request_context("/apps/app-1"): with pytest.MonkeyPatch.context() as monkeypatch: monkeypatch.setattr(dify_config, "RBAC_ENABLED", True) - monkeypatch.setattr(app_module, "AppService", lambda: SimpleNamespace(get_app=lambda app_model: app_obj)) + get_app = MagicMock(return_value=app_obj) + monkeypatch.setattr(app_module, "AppService", lambda: SimpleNamespace(get_app=get_app)) monkeypatch.setattr( app_module.FeatureService, "get_system_features", @@ -834,9 +887,17 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): get_permissions, ) - resp = method(app_module.AppApi(), "tenant-1", SimpleNamespace(id="acct-1"), app_model=app_obj) + session = MagicMock() + resp = method( + app_module.AppApi(), + session, + "tenant-1", + SimpleNamespace(id="acct-1"), + app_model=app_obj, + ) - get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=ANY) + get_app.assert_called_once_with(app_obj, session=session) + get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=session) assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit", "app.acl.monitor"] diff --git a/api/tests/unit_tests/controllers/console/app/test_audio.py b/api/tests/unit_tests/controllers/console/app/test_audio.py index 8084e7c883e..3a23ccf7f09 100644 --- a/api/tests/unit_tests/controllers/console/app/test_audio.py +++ b/api/tests/unit_tests/controllers/console/app/test_audio.py @@ -114,7 +114,7 @@ def test_agent_console_audio_api_uses_agent_draft(app: Flask, monkeypatch: pytes ) assert response == {"text": "agent transcript"} - assert calls["resolver"] == {"tenant_id": "tenant-1", "agent_id": agent_id} + assert calls["resolver"] == {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id} assert calls["rbac"] == { "tenant_id": "tenant-1", "account_id": "account-1", @@ -133,6 +133,7 @@ def test_agent_console_audio_api_uses_agent_draft(app: Flask, monkeypatch: pytes "app_model": app_model, "agent_soul": agent_soul, "file": calls["asr"]["file"], + "session": session, "end_user": None, } diff --git a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py index 2d2d5b4f361..1f2819e6d39 100644 --- a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py @@ -20,10 +20,8 @@ def _make_account(): def test_completion_conversation_list_returns_paginated_result(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: api = conversation_module.CompletionConversationApi() method = unwrap(api.get) - account = _make_account() monkeypatch.setattr(conversation_module, "parse_time_range", lambda *_args, **_kwargs: (None, None)) - paginate_result = MagicMock() paginate_result.page = 1 paginate_result.per_page = 20 @@ -31,40 +29,32 @@ def test_completion_conversation_list_returns_paginated_result(app: Flask, monke paginate_result.has_next = False paginate_result.items = [] monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) - with app.test_request_context("/console/api/apps/app-1/completion-conversations", method="GET"): - response = method(api, account, app_model=SimpleNamespace(id="app-1")) - + response = method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1")) assert response == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []} def test_completion_conversation_list_invalid_time_range(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: api = conversation_module.CompletionConversationApi() method = unwrap(api.get) - account = _make_account() monkeypatch.setattr( conversation_module, "parse_time_range", lambda *_args, **_kwargs: (_ for _ in ()).throw(ValueError("bad range")), ) - with app.test_request_context( - "/console/api/apps/app-1/completion-conversations", - method="GET", - query_string={"start": "bad"}, + "/console/api/apps/app-1/completion-conversations", method="GET", query_string={"start": "bad"} ): with pytest.raises(BadRequest): - method(api, account, app_model=SimpleNamespace(id="app-1")) + method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1")) def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: api = conversation_module.ChatConversationApi() method = unwrap(api.get) - account = _make_account() monkeypatch.setattr(conversation_module, "parse_time_range", lambda *_args, **_kwargs: (None, None)) - paginate_result = MagicMock() paginate_result.page = 1 paginate_result.per_page = 20 @@ -72,48 +62,88 @@ def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeyp paginate_result.has_next = False paginate_result.items = [] monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) - with app.test_request_context("/console/api/apps/app-1/chat-conversations", method="GET"): - response = method(api, account, app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT)) - + response = method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT)) assert response == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []} def test_get_conversation_updates_read_at(monkeypatch: pytest.MonkeyPatch) -> None: conversation = SimpleNamespace(id="c1", app_id="app-1") - session = MagicMock() session.scalar.return_value = conversation - - monkeypatch.setattr(conversation_module.db, "session", session) - - result = conversation_module._get_conversation(_make_account(), SimpleNamespace(id="app-1"), "c1") - + result = conversation_module._get_conversation(session, _make_account(), SimpleNamespace(id="app-1"), "c1") assert result is conversation session.execute.assert_called_once() - session.commit.assert_called_once() + session.flush.assert_called_once() session.refresh.assert_called_once_with(conversation) def test_get_conversation_missing_raises_not_found(monkeypatch: pytest.MonkeyPatch) -> None: session = MagicMock() session.scalar.return_value = None - - monkeypatch.setattr(conversation_module.db, "session", session) - with pytest.raises(NotFound): - conversation_module._get_conversation(_make_account(), SimpleNamespace(id="app-1"), "missing") + conversation_module._get_conversation(session, _make_account(), SimpleNamespace(id="app-1"), "missing") + + +def test_conversation_response_source_uses_caller_session() -> None: + session = MagicMock() + account = object() + annotation = MagicMock() + annotation.account_with_session.return_value = account + message = MagicMock() + conversation = MagicMock() + conversation.inputs_with_session.return_value = {"topic": "support"} + conversation.model_config_with_session.return_value = {"model_id": "model-1"} + conversation.summary_or_query_with_session.return_value = "summary" + conversation.annotated_with_session.return_value = True + conversation.annotation_with_session.return_value = annotation + conversation.message_count_with_session.return_value = 3 + conversation.user_feedback_stats_with_session.return_value = {"like": 2, "dislike": 1} + conversation.admin_feedback_stats_with_session.return_value = {"like": 1, "dislike": 0} + conversation.status_count_with_session.return_value = {"success": 1, "failed": 0} + conversation.first_message_with_session.return_value = message + conversation.from_end_user_session_id_with_session.return_value = "end-user-session" + conversation.from_account_name_with_session.return_value = "Account" + + source = conversation_module.ConversationResponseSource(conversation, session=session) + + assert source.inputs == {"topic": "support"} + assert source.model_config == {"model_id": "model-1"} + assert source.summary_or_query == "summary" + assert source.annotated is True + annotation_source = source.annotation + assert annotation_source is not None + assert annotation_source.account is account + assert source.message_count == 3 + assert source.user_feedback_stats == {"like": 2, "dislike": 1} + assert source.admin_feedback_stats == {"like": 1, "dislike": 0} + assert source.status_count == {"success": 1, "failed": 0} + assert source.first_message is not None + assert source.from_end_user_session_id == "end-user-session" + assert source.from_account_name == "Account" + conversation.model_config_with_session.assert_called_once_with(session=session) + conversation.inputs_with_session.assert_called_once_with(session=session) + conversation.summary_or_query_with_session.assert_called_once_with(session=session) + conversation.annotated_with_session.assert_called_once_with(session=session) + conversation.annotation_with_session.assert_called_once_with(session=session) + conversation.message_count_with_session.assert_called_once_with(session=session) + conversation.user_feedback_stats_with_session.assert_called_once_with(session=session) + conversation.admin_feedback_stats_with_session.assert_called_once_with(session=session) + conversation.status_count_with_session.assert_called_once_with(session=session) + conversation.first_message_with_session.assert_called_once_with(session=session) + conversation.from_end_user_session_id_with_session.assert_called_once_with(session=session) + conversation.from_account_name_with_session.assert_called_once_with(session=session) + annotation.account_with_session.assert_called_once_with(session=session) def test_completion_conversation_delete_maps_not_found(monkeypatch: pytest.MonkeyPatch) -> None: api = conversation_module.CompletionConversationDetailApi() method = unwrap(api.delete) - monkeypatch.setattr( conversation_module.ConversationService, "delete", lambda *_args, **_kwargs: (_ for _ in ()).throw(ConversationNotExistsError()), ) - + session = MagicMock() with pytest.raises(NotFound): - method(api, _make_account(), app_model=SimpleNamespace(id="app-1"), conversation_id="c1") + method(api, session, _make_account(), app_model=SimpleNamespace(id="app-1"), conversation_id="c1") diff --git a/api/tests/unit_tests/controllers/console/app/test_message_api.py b/api/tests/unit_tests/controllers/console/app/test_message_api.py index 20187da6159..0c2efc12198 100644 --- a/api/tests/unit_tests/controllers/console/app/test_message_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_message_api.py @@ -1,6 +1,9 @@ from __future__ import annotations from datetime import UTC, datetime +from inspect import unwrap +from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from flask import Flask @@ -8,6 +11,131 @@ from flask import Flask from controllers.console.app import message as message_module +def test_app_message_routes_pass_injected_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + session = MagicMock() + current_user = SimpleNamespace(id="account-1") + app_model = SimpleNamespace(id="app-1", mode="chat") + message_id = "550e8400-e29b-41d4-a716-446655440000" + list_messages = MagicMock(return_value={"data": []}) + update_feedback = MagicMock(return_value={"result": "success"}) + get_suggested_questions = MagicMock(return_value={"data": ["next"]}) + get_message_detail = MagicMock(return_value={"id": message_id}) + monkeypatch.setattr(message_module, "_list_chat_messages", list_messages) + monkeypatch.setattr(message_module, "_update_message_feedback", update_feedback) + monkeypatch.setattr(message_module, "_get_message_suggested_questions", get_suggested_questions) + monkeypatch.setattr(message_module, "_get_message_detail", get_message_detail) + + assert unwrap(message_module.ChatMessageListApi.get)( + message_module.ChatMessageListApi(), session, current_user, app_model + ) == {"data": []} + assert unwrap(message_module.MessageFeedbackApi.post)( + message_module.MessageFeedbackApi(), session, current_user, app_model + ) == {"result": "success"} + assert unwrap(message_module.MessageSuggestedQuestionApi.get)( + message_module.MessageSuggestedQuestionApi(), session, current_user, app_model, message_id + ) == {"data": ["next"]} + assert unwrap(message_module.MessageApi.get)(message_module.MessageApi(), session, app_model, message_id) == { + "id": message_id + } + + assert list_messages.call_args.kwargs["session"] is session + assert update_feedback.call_args.kwargs["session"] is session + assert get_suggested_questions.call_args.kwargs["session"] is session + assert get_message_detail.call_args.kwargs["session"] is session + + +def test_update_message_feedback_commits_injected_session(app: Flask) -> None: + message_id = "550e8400-e29b-41d4-a716-446655440000" + feedback = SimpleNamespace(rating="dislike", content=None) + get_admin_feedback = MagicMock(return_value=feedback) + message = SimpleNamespace( + id=message_id, + conversation_id="conversation-1", + admin_feedback_with_session=get_admin_feedback, + ) + session = MagicMock() + session.scalar.return_value = message + + with app.test_request_context(json={"message_id": message_id, "rating": "like", "content": "helpful"}): + result = message_module._update_message_feedback( + session=session, + current_user=SimpleNamespace(id="account-1"), + app_model=SimpleNamespace(id="app-1"), + ) + + assert result == {"result": "success"} + assert feedback.rating == "like" + assert feedback.content == "helpful" + get_admin_feedback.assert_called_once_with(session=session) + session.commit.assert_called_once_with() + + +def test_get_message_detail_uses_injected_session(monkeypatch: pytest.MonkeyPatch) -> None: + message_id = "550e8400-e29b-41d4-a716-446655440000" + message = SimpleNamespace(id=message_id) + response_source = object() + response_source_factory = MagicMock(return_value=response_source) + session = MagicMock() + session.scalar.return_value = message + monkeypatch.setattr(message_module, "attach_message_extra_contents", MagicMock()) + monkeypatch.setattr(message_module, "MessageResponseSource", response_source_factory) + monkeypatch.setattr(message_module, "dump_response", lambda _model, value: value) + + result = message_module._get_message_detail( + session=session, + app_model=SimpleNamespace(id="app-1"), + message_id=message_id, + ) + + assert result is response_source + response_source_factory.assert_called_once_with(message, session=session) + session.scalar.assert_called_once() + + +def test_message_response_source_uses_caller_session_for_nested_fields() -> None: + session = MagicMock() + account = object() + feedback = MagicMock() + feedback.from_account_with_session.return_value = account + annotation = MagicMock() + annotation.account_with_session.return_value = account + annotation.annotation_create_account_with_session.return_value = account + thought = object() + message_file = {"id": "file-1"} + message = MagicMock() + message.inputs_with_session.return_value = {"topic": "support"} + message.user_feedback_with_session.return_value = feedback + message.feedbacks_with_session.return_value = [feedback] + message.annotation_with_session.return_value = annotation + message.annotation_hit_history_with_session.return_value = annotation + message.agent_thoughts_with_session.return_value = [thought] + message.message_files_with_session.return_value = [message_file] + + source = message_module.MessageResponseSource(message, session=session) + + assert source.inputs == {"topic": "support"} + assert source.user_feedback is feedback + assert source.feedbacks[0].from_account is account + annotation_source = source.annotation + assert annotation_source is not None + assert annotation_source.account is account + annotation_hit_history_source = source.annotation_hit_history + assert annotation_hit_history_source is not None + assert annotation_hit_history_source.annotation_create_account is account + assert source.agent_thoughts == [thought] + assert source.message_files == [message_file] + message.inputs_with_session.assert_called_once_with(session=session) + message.user_feedback_with_session.assert_called_once_with(session=session) + message.feedbacks_with_session.assert_called_once_with(session=session) + message.annotation_with_session.assert_called_once_with(session=session) + message.annotation_hit_history_with_session.assert_called_once_with(session=session) + message.agent_thoughts_with_session.assert_called_once_with(session=session) + message.message_files_with_session.assert_called_once_with(session=session) + feedback.from_account_with_session.assert_called_once_with(session=session) + annotation.account_with_session.assert_called_once_with(session=session) + annotation.annotation_create_account_with_session.assert_called_once_with(session=session) + + def test_chat_messages_query_valid(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: """Test valid ChatMessagesQuery with all fields.""" query = message_module.ChatMessagesQuery( diff --git a/api/tests/unit_tests/controllers/console/app/test_model_config_api.py b/api/tests/unit_tests/controllers/console/app/test_model_config_api.py index 714eae618ed..8257605fed4 100644 --- a/api/tests/unit_tests/controllers/console/app/test_model_config_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_model_config_api.py @@ -1,36 +1,58 @@ from __future__ import annotations +import importlib import json from inspect import unwrap -from types import SimpleNamespace +from typing import Never from unittest.mock import MagicMock +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy import func, select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import object_session, sessionmaker +from controllers.common import session as controller_session from controllers.console.app import model_config as model_config_module -from models.model import AppMode, AppModelConfig +from models.model import App, AppMode, AppModelConfig + +app_wraps_module = importlib.import_module("controllers.console.app.wraps") -def test_post_updates_app_model_config_for_chat(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def _poison_implicit_app_config_properties(monkeypatch: pytest.MonkeyPatch) -> None: + def fail(_app: App) -> Never: + raise AssertionError("implicit App model-config property was accessed") + + monkeypatch.setattr(App, "app_model_config", property(fail)) + monkeypatch.setattr(App, "is_agent", property(fail)) + + +@pytest.mark.parametrize("app_mode", [AppMode.CHAT, AppMode.COMPLETION]) +def test_post_updates_non_agent_model_config_without_implicit_properties( + app: Flask, + monkeypatch: pytest.MonkeyPatch, + app_mode: AppMode, +) -> None: api = model_config_module.ModelConfigResource() method = unwrap(api.post) - app_model = SimpleNamespace( + app_model = App( id="app-1", - mode=AppMode.CHAT.value, - is_agent=False, - app_model_config_id=None, + mode=app_mode, + app_model_config_id="config-0", updated_by=None, updated_at=None, ) + original_config = AppModelConfig(app_id="app-1", created_by="u1", updated_by="u1") + original_config.agent_mode = None + _poison_implicit_app_config_properties(monkeypatch) monkeypatch.setattr( model_config_module.AppModelConfigService, "validate_configuration", lambda **_kwargs: {"pre_prompt": "hi"}, ) session = MagicMock() - monkeypatch.setattr(model_config_module.db, "session", session) def _from_model_config_dict(self, model_config): self.pre_prompt = model_config["pre_prompt"] @@ -40,30 +62,116 @@ def test_post_updates_app_model_config_for_chat(app: Flask, monkeypatch: pytest. monkeypatch.setattr(AppModelConfig, "from_model_config_dict", _from_model_config_dict) send_mock = MagicMock() monkeypatch.setattr(model_config_module.app_model_config_was_updated, "send", send_mock) + session.get.return_value = original_config with app.test_request_context("/console/api/apps/app-1/model-config", method="POST", json={"pre_prompt": "hi"}): - response = method(api, "t1", "u1", app_model=app_model) + response = method(api, session, "t1", "u1", app_model=app_model) + session.get.assert_called_once_with(AppModelConfig, "config-0") session.add.assert_called_once() session.flush.assert_called_once() - session.commit.assert_called_once() - send_mock.assert_called_once() + session.commit.assert_not_called() + assert send_mock.call_args.kwargs["session"] is session assert app_model.app_model_config_id == "config-1" + assert app_model.mode == app_mode assert response["result"] == "success" +def test_post_uses_one_session_and_rolls_back_when_signal_fails( + app: Flask, + sqlite_engine: Engine, + monkeypatch: pytest.MonkeyPatch, +) -> None: + App.metadata.create_all(sqlite_engine, tables=[App.__table__, AppModelConfig.__table__]) + make_session = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + app_id = str(uuid4()) + tenant_id = str(uuid4()) + user_id = str(uuid4()) + + with make_session.begin() as setup_session: + original_config = AppModelConfig(app_id=app_id, created_by=user_id, updated_by=user_id) + original_config.agent_mode = json.dumps({"tools": []}) + setup_session.add(original_config) + setup_session.flush() + original_config_id = original_config.id + setup_session.add( + App( + id=app_id, + tenant_id=tenant_id, + name="Atomic app", + description="", + mode=AppMode.AGENT_CHAT, + icon_type=None, + icon=None, + icon_background=None, + app_model_config_id=original_config_id, + enable_site=True, + enable_api=True, + max_active_requests=None, + created_by=user_id, + ) + ) + + monkeypatch.setattr(controller_session.session_factory, "create_session", make_session) + monkeypatch.setattr( + model_config_module.AppModelConfigService, + "validate_configuration", + lambda **_kwargs: {"agent_mode": {"tools": []}}, + ) + + captured: dict[str, object] = {} + + def load_app_model(session, requested_app_id: str): + loaded_app = session.get(App, requested_app_id) + captured["load_session"] = session + return loaded_app + + def fail_signal(sender: App, **kwargs: object) -> None: + signal_session = kwargs["session"] + assert object_session(sender) is signal_session + assert signal_session is captured["load_session"] + raise RuntimeError("signal failed") + + monkeypatch.setattr(app_wraps_module, "_load_app_model", load_app_model) + monkeypatch.setattr(model_config_module.app_model_config_was_updated, "send", fail_signal) + + method = model_config_module.ModelConfigResource.post + while not method.__code__.co_filename.endswith("controllers/common/session.py"): + method = method.__wrapped__ + assert method.__wrapped__.__code__.co_filename.endswith("controllers/console/app/wraps.py") + + api = model_config_module.ModelConfigResource() + with ( + app.test_request_context(f"/console/api/apps/{app_id}/model-config", method="POST", json={}), + pytest.raises(RuntimeError, match="signal failed"), + ): + method( + api, + current_tenant_id=tenant_id, + current_user_id=user_id, + app_id=app_id, + ) + + with make_session() as verification_session: + persisted_app = verification_session.get(App, app_id) + assert persisted_app is not None + assert persisted_app.app_model_config_id == original_config_id + config_count = verification_session.scalar(select(func.count()).select_from(AppModelConfig)) + assert config_count == 1 + + def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: api = model_config_module.ModelConfigResource() method = unwrap(api.post) - app_model = SimpleNamespace( + app_model = App( id="app-1", - mode=AppMode.AGENT_CHAT.value, - is_agent=True, + mode=AppMode.AGENT_CHAT, app_model_config_id="config-0", updated_by=None, updated_at=None, ) + _poison_implicit_app_config_properties(monkeypatch) original_config = AppModelConfig(app_id="app-1", created_by="u1", updated_by="u1") original_config.agent_mode = json.dumps( @@ -83,8 +191,7 @@ def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.Mon ) session = MagicMock() - session.get.return_value = original_config - monkeypatch.setattr(model_config_module.db, "session", session) + session.scalar.return_value = original_config monkeypatch.setattr( model_config_module.AppModelConfigService, @@ -129,9 +236,11 @@ def test_post_encrypts_agent_tool_parameters(app: Flask, monkeypatch: pytest.Mon monkeypatch.setattr(model_config_module.app_model_config_was_updated, "send", send_mock) with app.test_request_context("/console/api/apps/app-1/model-config", method="POST", json={"pre_prompt": "hi"}): - response = method(api, "t1", "u1", app_model=app_model) + response = method(api, session, "t1", "u1", app_model=app_model) stored_config = session.add.call_args[0][0] stored_agent_mode = json.loads(stored_config.agent_mode) + session.scalar.assert_called_once() + assert app_model.mode == AppMode.AGENT_CHAT assert stored_agent_mode["tools"][0]["tool_parameters"]["secret"] == "encrypted" assert response["result"] == "success" 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 2d811deb916..fd07c42f752 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -613,6 +613,36 @@ def test_advanced_chat_run_conversation_not_exists(app: Flask, monkeypatch: pyte handler(api, Mock(), "t1", app_model=SimpleNamespace(id="app")) +@pytest.mark.parametrize( + ("resource", "payload"), + [ + (workflow_module.DraftWorkflowTriggerRunApi, {"node_id": "node-1"}), + (workflow_module.DraftWorkflowTriggerRunAllApi, {"node_ids": ["node-1"]}), + ], +) +def test_trigger_run_loads_draft_with_request_session( + app: Flask, + monkeypatch: pytest.MonkeyPatch, + resource: type, + payload: dict[str, object], +) -> None: + get_draft_workflow = Mock(return_value=None) + monkeypatch.setattr( + workflow_module, + "WorkflowService", + lambda: SimpleNamespace(get_draft_workflow=get_draft_workflow), + ) + session = Mock() + app_model = SimpleNamespace(id="app-1") + handler = inspect.unwrap(resource.post) + + with app.test_request_context("/", method="POST", json=payload): + with pytest.raises(ValueError, match="Workflow not found"): + handler(resource(), session, SimpleNamespace(id="account-1"), app_model) + + get_draft_workflow.assert_called_once_with(app_model, session=session) + + def test_workflow_online_users_filters_inaccessible_workflow(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: app_id_1 = "11111111-1111-1111-1111-111111111111" app_id_2 = "22222222-2222-2222-2222-222222222222" diff --git a/api/tests/unit_tests/controllers/console/app/test_wraps.py b/api/tests/unit_tests/controllers/console/app/test_wraps.py index c7af2f26411..002ce2e3df8 100644 --- a/api/tests/unit_tests/controllers/console/app/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/app/test_wraps.py @@ -1,5 +1,7 @@ from __future__ import annotations +from contextlib import nullcontext +from inspect import getsource from types import SimpleNamespace from unittest.mock import MagicMock @@ -7,7 +9,10 @@ import pytest from sqlalchemy import Select from sqlalchemy.orm import Session +from controllers.common import session as session_module from controllers.common.session import with_session +from controllers.console.app import completion as completion_module +from controllers.console.app import workflow as workflow_module from controllers.console.app import wraps as wraps_module from controllers.console.app.error import AppNotFoundError from models.model import App, AppMode, TrialApp @@ -59,6 +64,7 @@ def test_get_app_model_rejects_wrong_mode(monkeypatch: pytest.MonkeyPatch) -> No def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: pytest.MonkeyPatch) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") + session = FakeSession() def scalar(statement: Select[tuple[App]]) -> object | None: has_trial_app_join = any( @@ -66,54 +72,54 @@ def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: p ) return None if has_trial_app_join else app_model - scoped_session = MagicMock() - scoped_session.scalar.side_effect = scalar + monkeypatch.setattr(session, "scalar", scalar) recommended_get_app = MagicMock(return_value=None) - monkeypatch.setattr(wraps_module.db, "session", scoped_session) monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app) - @wraps_module.get_app_model_with_trial - def handler(app_model): - return app_model.id + class Handler: + @wraps_module.get_app_model_with_trial + def get(self, _injected_session, app_model): + return app_model.id with pytest.raises(AppNotFoundError): - handler(app_id="app-1") + Handler().get(session, app_id="app-1") - recommended_get_app.assert_called_once_with("app-1", session=scoped_session.return_value) + recommended_get_app.assert_called_once_with("app-1", session=session) def test_get_app_model_with_trial_falls_back_to_recommended_app(monkeypatch: pytest.MonkeyPatch) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") session = MagicMock(spec=Session) - scoped_session = MagicMock(return_value=session) trial_app_loader = MagicMock(return_value=None) recommended_get_app = MagicMock(return_value=app_model) - monkeypatch.setattr(wraps_module.db, "session", scoped_session) monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader) monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app) - @wraps_module.get_app_model_with_trial - def handler(app_model): - return app_model.id + class Handler: + @wraps_module.get_app_model_with_trial + def get(self, _injected_session, app_model): + return app_model.id - assert handler(app_id="app-1") == "app-1" - trial_app_loader.assert_called_once_with("app-1") + assert Handler().get(session, app_id="app-1") == "app-1" + trial_app_loader.assert_called_once_with(session, "app-1") recommended_get_app.assert_called_once_with("app-1", session=session) def test_get_app_model_with_trial_prefers_trial_registration(monkeypatch: pytest.MonkeyPatch) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") + session = MagicMock(spec=Session) trial_app_loader = MagicMock(return_value=app_model) recommended_get_app = MagicMock() monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader) monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app) - @wraps_module.get_app_model_with_trial - def handler(app_model): - return app_model.id + class Handler: + @wraps_module.get_app_model_with_trial + def get(self, _injected_session, app_model): + return app_model.id - assert handler(app_id="app-1") == "app-1" - trial_app_loader.assert_called_once_with("app-1") + assert Handler().get(session, app_id="app-1") == "app-1" + trial_app_loader.assert_called_once_with(session, "app-1") recommended_get_app.assert_not_called() @@ -147,3 +153,48 @@ def test_get_app_model_prefers_injected_session(monkeypatch: pytest.MonkeyPatch) assert Handler().get(session, app_id="app-1") == "app-1" assert session.scalar_called + + +def test_get_app_model_with_trial_prefers_injected_session(monkeypatch: pytest.MonkeyPatch) -> None: + app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal") + session = FakeSession(app_model) + monkeypatch.setattr( + wraps_module.db, + "session", + SimpleNamespace(scalar=lambda *_args, **_kwargs: pytest.fail("db.session should not be used")), + ) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: nullcontext(session)) + + class Handler: + @with_session(write=False) + @wraps_module.get_app_model_with_trial(None) + def get(self, injected_session, app_model): + assert injected_session is session + return app_model.id + + assert Handler().get(app_id="app-1") == "app-1" + assert session.scalar_called + + +def test_get_app_model_with_trial_requires_injected_session() -> None: + @wraps_module.get_app_model_with_trial(None) + def handler(app_model): + return app_model.id + + with pytest.raises(RuntimeError, match="requires @with_session"): + handler(app_id="app-1") + + +@pytest.mark.parametrize( + "resource", + [ + completion_module.CompletionMessageApi, + completion_module.ChatMessageApi, + workflow_module.AdvancedChatDraftWorkflowRunApi, + workflow_module.DraftWorkflowRunApi, + workflow_module.DraftWorkflowTriggerRunApi, + workflow_module.DraftWorkflowTriggerRunAllApi, + ], +) +def test_migrated_handlers_open_session_before_app_lookup(resource: type) -> None: + assert "@with_session\n @get_app_model" in getsource(resource) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_data_source.py b/api/tests/unit_tests/controllers/console/datasets/test_data_source.py index b608e8a73fe..a6cb79417a7 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_data_source.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_data_source.py @@ -5,6 +5,7 @@ from collections.abc import Callable from datetime import UTC, datetime from typing import cast from unittest.mock import MagicMock, PropertyMock, patch +from uuid import uuid4 import pytest from flask import Flask @@ -106,6 +107,22 @@ def test_get_data_source_integrates_preserves_empty_list_when_no_binding(flask_a assert response == {"data": []} +def test_patch_data_source_binding_uses_injected_session(flask_app: Flask) -> None: + binding = MagicMock(disabled=True) + session = MagicMock() + session.scalar.return_value = binding + + with flask_app.test_request_context("/"): + response, status = unwrap(DataSourceApi().patch)(DataSourceApi(), session, "tenant-1", uuid4(), "enable") + + assert status == 200 + assert response == {"result": "success"} + assert binding.disabled is False + session.scalar.assert_called_once() + session.add.assert_not_called() + session.commit.assert_not_called() + + def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask, current_user: Account) -> None: page = MagicMock( page_id="page-1", @@ -128,6 +145,7 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask get_online_document_pages=MagicMock(return_value=iter([online_document_message])), datasource_provider_type=MagicMock(return_value="online_document"), ) + session = MagicMock() with ( flask_app.test_request_context("/?credential_id=credential-1"), @@ -137,10 +155,11 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask return_value={"token": "token"}, ), patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()), - patch.object(module, "sessionmaker"), patch("core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", return_value=runtime), ): - response, status = unwrap(DataSourceNotionListApi().get)(DataSourceNotionListApi(), "tenant-1", current_user) + response, status = unwrap(DataSourceNotionListApi().get)( + DataSourceNotionListApi(), session, "tenant-1", current_user + ) assert status == 200 assert response == { 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 dad78d3d4ef..541c4b3f392 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 ANY, MagicMock, PropertyMock, patch +from unittest.mock import ANY, MagicMock, PropertyMock, call, patch import pytest from flask import Flask @@ -46,24 +46,23 @@ from services.enterprise import rbac_service as enterprise_rbac_service @pytest.fixture(autouse=True) def dataset_model_property_defaults(): - properties: dict[str, object] = { - "app_count": 0, - "document_count": 0, - "word_count": 0, - "author_name": None, - "tags": [], - "doc_form": None, - "external_knowledge_info": None, - "doc_metadata": [], - "is_published": False, - "total_documents": 0, - "total_available_documents": 0, + getter_values: dict[str, object] = { + "get_app_count": 0, + "get_document_count": 0, + "get_word_count": 0, + "get_author_name": None, + "get_tags": [], + "get_doc_form": None, + "get_external_knowledge_info": None, + "get_doc_metadata": [], + "get_is_published": False, + "get_total_documents": 0, + "get_total_available_documents": 0, } - + getters = {} with ExitStack() as stack: - for name, value in properties.items(): - property_mock = stack.enter_context(patch.object(Dataset, name, new_callable=PropertyMock)) - property_mock.return_value = value + for name, value in getter_values.items(): + getters[name] = stack.enter_context(patch.object(Dataset, name, autospec=True, return_value=value)) stack.enter_context( patch( "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.MyPermissions.get", @@ -76,7 +75,7 @@ def dataset_model_property_defaults(): return_value={}, ) ) - yield + yield getters def make_dataset(**overrides) -> Dataset: @@ -171,25 +170,14 @@ class TestDatasetList: def test_get_success_basic(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [make_dataset(icon_info={"icon": "📙", "icon_type": "emoji"})] - with app.test_request_context("/datasets"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 200 assert resp["total"] == 1 assert resp["data"][0]["embedding_available"] is True @@ -200,28 +188,33 @@ class TestDatasetList: "icon_url": None, } + def test_get_serializes_database_fields_with_caller_session(self, app: Flask, dataset_model_property_defaults): + api = DatasetListApi() + method = unwrap(api.get) + current_user = self._mock_user() + dataset = make_dataset() + session = MagicMock() + with app.test_request_context("/datasets"): + with ( + patch.object(DatasetService, "get_datasets", return_value=([dataset], 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), + ): + method(api, session, "tenant-1", current_user) + + for getter in dataset_model_property_defaults.values(): + getter.assert_called_once_with(dataset, session=session) + def test_get_with_ids_filter(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [make_dataset()] - with app.test_request_context("/datasets?ids=1&ids=2"): with ( - patch.object( - DatasetService, - "get_datasets_by_ids", - return_value=(datasets, 2), - ) as by_ids_mock, - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets_by_ids", return_value=(datasets, 2)) as by_ids_mock, + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) by_ids_mock.assert_called_once() assert status == 200 assert resp["total"] == 2 @@ -236,28 +229,21 @@ class TestDatasetList: default_permission_keys=["dataset.acl.readonly"], overrides=[ enterprise_rbac_service.ResourcePermissionKeys( - resource_id="dataset-1", - permission_keys=["dataset.acl.readonly", "dataset.acl.edit"], + resource_id="dataset-1", permission_keys=["dataset.acl.readonly", "dataset.acl.edit"] ) ], ) ) - with app.test_request_context("/datasets"): with ( patch.object(DatasetService, "get_datasets", return_value=([dataset], 1)), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), patch( "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.MyPermissions.get", return_value=permissions, ) as get_permissions, ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) 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"] @@ -271,7 +257,6 @@ class TestDatasetList: permission_keys=["dataset.create_and_management"] ) ) - with app.test_request_context("/datasets"): with ( patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), @@ -284,14 +269,9 @@ class TestDatasetList: "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetAccess.whitelist_resources", return_value=SimpleNamespace(resource_ids=[]), ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - method(api, "tenant-1", current_user) - + method(api, MagicMock(), "tenant-1", current_user) assert get_datasets.call_args.kwargs["accessible_dataset_ids"] == [] assert get_datasets.call_args.kwargs["include_own_datasets"] is True @@ -302,7 +282,6 @@ class TestDatasetList: permissions = enterprise_rbac_service.MyPermissionsResponse( dataset=enterprise_rbac_service.ResourcePermissionSnapshot(default_permission_keys=["dataset.preview"]) ) - with app.test_request_context("/datasets"): with ( patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), @@ -315,14 +294,9 @@ class TestDatasetList: "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetAccess.whitelist_resources", return_value=SimpleNamespace(unrestricted=True, resource_ids=[]), ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - method(api, "tenant-1", current_user) - + method(api, MagicMock(), "tenant-1", current_user) assert get_datasets.call_args.kwargs["accessible_dataset_ids"] is None def test_get_limits_to_dataset_read_overrides(self, app: Flask): @@ -333,25 +307,18 @@ class TestDatasetList: dataset=enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ enterprise_rbac_service.ResourcePermissionKeys( - resource_id="dataset-acl-shared", - permission_keys=["dataset.acl.preview"], + resource_id="dataset-acl-shared", permission_keys=["dataset.acl.preview"] ), enterprise_rbac_service.ResourcePermissionKeys( - resource_id="dataset-full", - permission_keys=["dataset.full_access"], + resource_id="dataset-full", permission_keys=["dataset.full_access"] ), enterprise_rbac_service.ResourcePermissionKeys( - resource_id="dataset-shared", - permission_keys=["dataset.preview"], - ), - enterprise_rbac_service.ResourcePermissionKeys( - resource_id="dataset-hidden", - permission_keys=[], + resource_id="dataset-shared", permission_keys=["dataset.preview"] ), + enterprise_rbac_service.ResourcePermissionKeys(resource_id="dataset-hidden", permission_keys=[]), ] ) ) - with app.test_request_context("/datasets"): with ( patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), @@ -363,22 +330,12 @@ class TestDatasetList: patch( "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetAccess.whitelist_resources", return_value=SimpleNamespace( - resource_ids=[ - "dataset-shared", - "dataset-acl-shared", - "dataset-full", - "dataset-whitelist-only", - ] + resource_ids=["dataset-shared", "dataset-acl-shared", "dataset-full", "dataset-whitelist-only"] ), ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - method(api, "tenant-1", current_user) - + method(api, MagicMock(), "tenant-1", current_user) assert get_datasets.call_args.kwargs["accessible_dataset_ids"] == [ "dataset-acl-shared", "dataset-full", @@ -392,7 +349,6 @@ class TestDatasetList: method = unwrap(api.get) current_user = self._mock_user() permissions = enterprise_rbac_service.MyPermissionsResponse() - with app.test_request_context("/datasets?ids=dataset-1"): with ( patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), @@ -405,50 +361,35 @@ class TestDatasetList: "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.DatasetAccess.whitelist_resources", return_value=SimpleNamespace(resource_ids=[]), ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - method(api, "tenant-1", current_user) - - get_datasets_by_ids.assert_called_once_with( - ["dataset-1"], - "tenant-1", - user=current_user, - accessible_dataset_ids=[], - include_own_datasets=False, - ) + method(api, MagicMock(), "tenant-1", current_user) + session = get_datasets_by_ids.call_args.kwargs["session"] + assert isinstance(session, MagicMock) + assert get_datasets_by_ids.call_args.args == (["dataset-1"], "tenant-1") + assert get_datasets_by_ids.call_args.kwargs == { + "user": current_user, + "accessible_dataset_ids": [], + "include_own_datasets": False, + "session": session, + } def test_get_with_tag_ids(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [make_dataset()] - with app.test_request_context("/datasets?tag_ids=tag1"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 200 def test_get_allows_legacy_weighted_score_without_weight_type(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [ make_dataset( @@ -471,47 +412,26 @@ class TestDatasetList: } ) ] - with app.test_request_context("/datasets"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 200 assert resp["data"][0]["retrieval_model_dict"]["weights"]["weight_type"] is None def test_get_merges_partial_retrieval_model_defaults(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [make_dataset(retrieval_model={"top_k": 4, "score_threshold_enabled": False})] - with app.test_request_context("/datasets"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 200 retrieval_model = resp["data"][0]["retrieval_model_dict"] assert retrieval_model["search_method"] == "semantic_search" @@ -522,62 +442,35 @@ class TestDatasetList: def test_embedding_available_false(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [ make_dataset( - indexing_technique="high_quality", - embedding_model="text-embed", - embedding_model_provider="openai", + indexing_technique="high_quality", embedding_model="text-embed", embedding_model_provider="openai" ) ] - config = MagicMock() - config.get_models.return_value = [] # model not available - + config.get_models.return_value = [] with app.test_request_context("/datasets"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=config, - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=config), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, MagicMock(), "tenant-1", current_user) assert resp["data"][0]["embedding_available"] is False def test_partial_members_permission(self, app: Flask): api = DatasetListApi() method = unwrap(api.get) - current_user = self._mock_user() datasets = [make_dataset(permission="partial_members")] - + session = MagicMock() + session.execute.return_value.all.return_value = [("ds-1", "u1")] with app.test_request_context("/datasets"): with ( - patch.object( - DatasetService, - "get_datasets", - return_value=(datasets, 1), - ), - patch( - "controllers.console.datasets.datasets.db.session.execute", - return_value=MagicMock(all=lambda: [("ds-1", "u1")]), - ), - patch.object( - ProviderManager, - "get_configurations", - return_value=MagicMock(get_models=lambda **_: []), - ), + patch.object(DatasetService, "get_datasets", return_value=(datasets, 1)), + patch.object(ProviderManager, "get_configurations", return_value=MagicMock(get_models=lambda **_: [])), ): - resp, status = method(api, "tenant-1", current_user) - + resp, status = method(api, session, "tenant-1", current_user) assert resp["data"][0]["partial_member_list"] == ["u1"] @@ -585,61 +478,36 @@ class TestDatasetListApiPost: def test_post_success(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - - payload = { - "name": "My Dataset", - "description": "desc", - "indexing_technique": "economy", - "provider": "vendor", - } - + payload = {"name": "My Dataset", "description": "desc", "indexing_technique": "economy", "provider": "vendor"} user = make_account() - dataset = make_dataset(name=payload["name"], description=payload["description"]) - with ( app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload), - patch.object( - DatasetService, - "create_empty_dataset", - return_value=dataset, - ), + patch.object(DatasetService, "create_empty_dataset", return_value=dataset), ): _, status = method(api, MagicMock(), "tenant-1", user) - assert status == 201 def test_post_forbidden(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - payload = {"name": "test"} - user = make_account(TenantAccountRole.NORMAL) - - with ( - app.test_request_context("/datasets", json=payload), - patch.object(type(console_ns), "payload", payload), - ): + with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(Forbidden): method(api, MagicMock(), "tenant-1", user) def test_post_duplicate_name(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - payload = {"name": "duplicate"} - user = make_account() - with ( app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload), patch.object( - DatasetService, - "create_empty_dataset", - side_effect=services.errors.dataset.DatasetNameDuplicateError(), + DatasetService, "create_empty_dataset", side_effect=services.errors.dataset.DatasetNameDuplicateError() ), ): with pytest.raises(DatasetNameDuplicateError): @@ -648,7 +516,6 @@ class TestDatasetListApiPost: def test_post_invalid_payload_missing_name(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - with app.test_request_context("/datasets", json={}), patch.object(type(console_ns), "payload", {}): with pytest.raises(ValueError): method(api, MagicMock(), "tenant-1", make_account()) @@ -656,12 +523,7 @@ class TestDatasetListApiPost: def test_post_invalid_indexing_technique(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - - payload = { - "name": "bad", - "indexing_technique": "invalid-tech", - } - + payload = {"name": "bad", "indexing_technique": "invalid-tech"} with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError, match="Invalid indexing technique"): method(api, MagicMock(), "tenant-1", make_account()) @@ -669,12 +531,7 @@ class TestDatasetListApiPost: def test_post_invalid_provider(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) - - payload = { - "name": "bad", - "provider": "unknown", - } - + payload = {"name": "bad", "provider": "unknown"} with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError, match="Invalid provider"): method(api, MagicMock(), "tenant-1", make_account()) @@ -684,66 +541,40 @@ class TestDatasetApiGet: def test_get_success_basic(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "123e4567-e89b-12d3-a456-426614174000" - user = make_account() tenant_id = "tenant-1" - dataset = make_dataset(id=dataset_id) - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock, ): - # embedding models exist → embedding_available stays True provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = [] - - data, status = method(api, tenant_id, user, dataset_id) - + data, status = method(api, MagicMock(), tenant_id, user, dataset_id) assert status == 200 assert data["embedding_available"] is True def test_get_attaches_permission_keys_when_rbac_enabled(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "123e4567-e89b-12d3-a456-426614174000" user = MagicMock(id="account-1") tenant_id = "tenant-1" dataset = make_dataset(id=dataset_id) - with ( app.test_request_context(f"/datasets/{dataset_id}"), patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), patch( "controllers.console.datasets.datasets.enterprise_rbac_service.RBACService.MyPermissions.get", return_value=enterprise_rbac_service.MyPermissionsResponse( dataset=enterprise_rbac_service.ResourcePermissionSnapshot( overrides=[ enterprise_rbac_service.ResourcePermissionKeys( - resource_id=dataset_id, - permission_keys=["dataset.acl.readonly", "dataset.acl.edit"], + resource_id=dataset_id, permission_keys=["dataset.acl.readonly", "dataset.acl.edit"] ) ] ) @@ -752,9 +583,7 @@ class TestDatasetApiGet: patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock, ): provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = [] - - data, status = method(api, tenant_id, user, dataset_id) - + data, status = method(api, MagicMock(), tenant_id, user, 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"] @@ -762,66 +591,38 @@ class TestDatasetApiGet: def test_get_uses_default_external_retrieval_model(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "dataset-id" dataset = make_dataset(id=dataset_id, retrieval_model=None) - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock, ): provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = [] - - data, status = method(api, "tenant", make_account(), dataset_id) - + data, status = method(api, MagicMock(), "tenant", make_account(), dataset_id) assert status == 200 - assert data["external_retrieval_model"] == { - "top_k": 2, - "score_threshold": 0.0, - "score_threshold_enabled": None, - } + assert data["external_retrieval_model"] == {"top_k": 2, "score_threshold": 0.0, "score_threshold_enabled": None} def test_get_dataset_not_found(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "missing-id" - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, "tenant", make_account(), dataset_id) + method(api, MagicMock(), "tenant", make_account(), dataset_id) def test_get_permission_denied(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "dataset-id" dataset = make_dataset(id=dataset_id) - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object( DatasetService, "check_dataset_permission", @@ -829,77 +630,45 @@ class TestDatasetApiGet: ), ): with pytest.raises(Forbidden, match="no access"): - method(api, "tenant", make_account(), dataset_id) + method(api, MagicMock(), "tenant", make_account(), dataset_id) def test_get_high_quality_embedding_unavailable(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "dataset-id" user = make_account() tenant_id = "tenant-1" - dataset = make_dataset( id=dataset_id, indexing_technique="high_quality", embedding_model="text-embedding", embedding_model_provider="openai", ) - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock, ): - # embedding model NOT configured provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = [] - - data, _ = method(api, tenant_id, user, dataset_id) - + data, _ = method(api, MagicMock(), tenant_id, user, dataset_id) assert data["embedding_available"] is False def test_get_partial_members_permission(self, app: Flask): api = DatasetApi() method = unwrap(api.get) - dataset_id = "dataset-id" - dataset = make_dataset(id=dataset_id, permission="partial_members") - partial_members = ["u1", "u2"] - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), - patch.object( - DatasetPermissionService, - "get_dataset_partial_member_list", - return_value=partial_members, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), + patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=partial_members), patch("controllers.console.datasets.datasets.create_plugin_provider_manager") as provider_manager_mock, ): provider_manager_mock.return_value.get_configurations.return_value.get_models.return_value = [] - - data, _ = method(api, "tenant", make_account(), dataset_id) - + data, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) assert data["partial_member_list"] == partial_members @@ -907,59 +676,29 @@ class TestDatasetApiPatch: def test_patch_success_basic(self, app: Flask): api = DatasetApi() method = unwrap(api.patch) - dataset_id = "dataset-id" - - payload = { - "name": "updated-name", - "description": "updated description", - } - + payload = {"name": "updated-name", "description": "updated description"} user = make_account() tenant_id = "tenant-1" - dataset = make_dataset(id=dataset_id, tenant_id=tenant_id) - with ( app.test_request_context(f"/datasets/{dataset_id}"), patch.object(type(console_ns), "payload", payload), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "check_permission", - return_value=None, - ), - patch.object( - DatasetService, - "update_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "get_dataset_partial_member_list", - return_value=[], - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "check_permission", return_value=None), + patch.object(DatasetService, "update_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]), ): result, status = method(api, MagicMock(), tenant_id, user, dataset_id) - assert status == 200 assert result["partial_member_list"] == [] def test_patch_dataset_not_found(self, app: Flask): api = DatasetApi() method = unwrap(api.patch) - with ( app.test_request_context("/datasets/missing"), - patch.object( - DatasetService, - "get_dataset", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): method(api, MagicMock(), "tenant-1", make_account(), "missing") @@ -967,25 +706,14 @@ class TestDatasetApiPatch: def test_patch_permission_denied(self, app: Flask): api = DatasetApi() method = unwrap(api.patch) - dataset_id = "dataset-id" dataset = make_dataset(id=dataset_id) - payload = {"name": "x"} - with ( app.test_request_context(f"/datasets/{dataset_id}"), patch.object(type(console_ns), "payload", payload), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "check_permission", - side_effect=Forbidden("no permission"), - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "check_permission", side_effect=Forbidden("no permission")), ): with pytest.raises(Forbidden): method(api, MagicMock(), "tenant", make_account(), dataset_id) @@ -993,92 +721,37 @@ class TestDatasetApiPatch: def test_patch_partial_members_update(self, app: Flask): api = DatasetApi() method = unwrap(api.patch) - dataset_id = "dataset-id" - - payload = { - "permission": "partial_members", - "partial_member_list": [{"user_id": "u1"}, {"user_id": "u2"}], - } - + payload = {"permission": "partial_members", "partial_member_list": [{"user_id": "u1"}, {"user_id": "u2"}]} dataset = make_dataset(id=dataset_id, permission="partial_members") - with ( app.test_request_context(f"/datasets/{dataset_id}"), patch.object(type(console_ns), "payload", payload), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "check_permission", - return_value=None, - ), - patch.object( - DatasetService, - "update_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "update_partial_member_list", - return_value=None, - ), - patch.object( - DatasetPermissionService, - "get_dataset_partial_member_list", - return_value=["u1", "u2"], - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "check_permission", return_value=None), + patch.object(DatasetService, "update_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "update_partial_member_list", return_value=None), + patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["u1", "u2"]), ): result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) - assert result["partial_member_list"] == ["u1", "u2"] def test_patch_clear_partial_members(self, app: Flask): api = DatasetApi() method = unwrap(api.patch) - dataset_id = "dataset-id" - - payload = { - "permission": "only_me", - } - + payload = {"permission": "only_me"} dataset = make_dataset(id=dataset_id) - with ( app.test_request_context(f"/datasets/{dataset_id}"), patch.object(type(console_ns), "payload", payload), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "check_permission", - return_value=None, - ), - patch.object( - DatasetService, - "update_dataset", - return_value=dataset, - ), - patch.object( - DatasetPermissionService, - "clear_partial_member_list", - return_value=None, - ), - patch.object( - DatasetPermissionService, - "get_dataset_partial_member_list", - return_value=[], - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "check_permission", return_value=None), + patch.object(DatasetService, "update_dataset", return_value=dataset), + patch.object(DatasetPermissionService, "clear_partial_member_list", return_value=None), + patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]), ): result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) - assert result["partial_member_list"] == [] @@ -1086,112 +759,73 @@ class TestDatasetApiDelete: def test_delete_success(self, app: Flask): api = DatasetApi() method = unwrap(api.delete) - dataset_id = "dataset-id" user = make_account() - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "delete_dataset", - return_value=True, - ), - patch.object( - DatasetPermissionService, - "clear_partial_member_list", - return_value=None, - ), + patch.object(DatasetService, "delete_dataset", return_value=True), + patch.object(DatasetPermissionService, "clear_partial_member_list", return_value=None), ): - result, status = method(api, user, dataset_id) - + result, status = method(api, MagicMock(), user, dataset_id) assert status == 204 assert result == "" def test_delete_forbidden_no_permission(self, app: Flask): api = DatasetApi() method = unwrap(api.delete) - dataset_id = "dataset-id" user = make_account(TenantAccountRole.NORMAL) - with app.test_request_context(f"/datasets/{dataset_id}"): with pytest.raises(Forbidden): - method(api, user, dataset_id) + method(api, MagicMock(), user, dataset_id) def test_delete_dataset_not_found(self, app: Flask): api = DatasetApi() method = unwrap(api.delete) - dataset_id = "missing-dataset" user = make_account() - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "delete_dataset", - return_value=False, - ), + patch.object(DatasetService, "delete_dataset", return_value=False), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, user, dataset_id) + method(api, MagicMock(), user, dataset_id) def test_delete_dataset_in_use(self, app: Flask): api = DatasetApi() method = unwrap(api.delete) - dataset_id = "dataset-id" user = make_account() - with ( app.test_request_context(f"/datasets/{dataset_id}"), - patch.object( - DatasetService, - "delete_dataset", - side_effect=services.errors.dataset.DatasetInUseError(), - ), + patch.object(DatasetService, "delete_dataset", side_effect=services.errors.dataset.DatasetInUseError()), ): with pytest.raises(DatasetInUseError): - method(api, user, dataset_id) + method(api, MagicMock(), user, dataset_id) class TestDatasetUseCheckApi: def test_get_use_check_true(self, app: Flask): api = DatasetUseCheckApi() method = unwrap(api.get) - dataset_id = "dataset-id" - with ( app.test_request_context(f"/datasets/{dataset_id}/use-check"), - patch.object( - DatasetService, - "dataset_use_check", - return_value=True, - ), + patch.object(DatasetService, "dataset_use_check", return_value=True), ): - result, status = method(api, dataset_id) - + result, status = method(api, MagicMock(), dataset_id) assert status == 200 assert result == {"is_using": True} def test_get_use_check_false(self, app: Flask): api = DatasetUseCheckApi() method = unwrap(api.get) - dataset_id = "dataset-id" - with ( app.test_request_context(f"/datasets/{dataset_id}/use-check"), - patch.object( - DatasetService, - "dataset_use_check", - return_value=False, - ), + patch.object(DatasetService, "dataset_use_check", return_value=False), ): - result, status = method(api, dataset_id) - + result, status = method(api, MagicMock(), dataset_id) assert status == 200 assert result == {"is_using": False} @@ -1200,15 +834,7 @@ class TestDatasetQueryApi: def _query_record(self, index: int = 1) -> DatasetQuery: query = DatasetQuery( dataset_id="dataset-id", - content=json.dumps( - [ - { - "content_type": "text_query", - "content": f"question {index}", - "file_info": None, - } - ] - ), + content=json.dumps([{"content_type": "text_query", "content": f"question {index}", "file_info": None}]), source="hit_testing", source_app_id=None, created_by_role=CreatorUserRole.ACCOUNT, @@ -1221,35 +847,17 @@ class TestDatasetQueryApi: def test_get_queries_success(self, app: Flask): api = DatasetQueryApi() method = unwrap(api.get) - dataset_id = "dataset-id" - current_user = make_account() - dataset = make_dataset(id=dataset_id) - queries = [self._query_record(1), self._query_record(2)] - with ( app.test_request_context("/datasets/queries?page=1&limit=20"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), - patch.object( - DatasetService, - "get_dataset_queries", - return_value=(queries, 2), - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), + patch.object(DatasetService, "get_dataset_queries", return_value=(queries, 2)), ): - response, status = method(api, current_user, dataset_id) - + response, status = method(api, MagicMock(), current_user, dataset_id) assert status == 200 assert response["total"] == 2 assert response["page"] == 1 @@ -1258,13 +866,7 @@ class TestDatasetQueryApi: assert len(response["data"]) == 2 assert response["data"][0] == { "id": "query-1", - "queries": [ - { - "content_type": "text_query", - "content": "question 1", - "file_info": None, - } - ], + "queries": [{"content_type": "text_query", "content": "question 1", "file_info": None}], "source": "hit_testing", "source_app_id": None, "created_by_role": "account", @@ -1272,40 +874,71 @@ class TestDatasetQueryApi: "created_at": 1704110400, } + def test_get_image_query_uses_caller_session(self, app: Flask): + api = DatasetQueryApi() + method = unwrap(api.get) + dataset = make_dataset(id="dataset-id") + query = self._query_record() + query.content = json.dumps([{"content_type": "image_query", "content": "file-1"}]) + upload_file = SimpleNamespace( + id="file-1", + name="image.png", + size=10, + extension="png", + mime_type="image/png", + ) + session = MagicMock() + session.scalar.return_value = upload_file + with ( + app.test_request_context("/datasets/queries"), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), + patch.object(DatasetService, "get_dataset_queries", return_value=([query], 1)), + patch("models.dataset.db") as db_mock, + patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"), + ): + db_mock.session.scalar.return_value = upload_file + response, status = method(api, session, make_account(), "dataset-id") + + assert status == 200 + assert response["data"][0]["queries"] == [ + { + "content_type": "image_query", + "content": "file-1", + "file_info": { + "id": "file-1", + "name": "image.png", + "size": 10, + "extension": "png", + "mime_type": "image/png", + "source_url": "signed-url", + }, + } + ] + session.scalar.assert_called_once() + db_mock.session.scalar.assert_not_called() + def test_get_queries_dataset_not_found(self, app: Flask): api = DatasetQueryApi() method = unwrap(api.get) - dataset_id = "dataset-id" current_user = make_account() - with ( app.test_request_context("/datasets/queries"), - patch.object( - DatasetService, - "get_dataset", - return_value=None, - ), + patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, current_user, dataset_id) + method(api, MagicMock(), current_user, dataset_id) def test_get_queries_permission_denied(self, app: Flask): api = DatasetQueryApi() method = unwrap(api.get) - dataset_id = "dataset-id" current_user = make_account() - dataset = make_dataset(id=dataset_id) - with ( app.test_request_context("/datasets/queries"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object( DatasetService, "check_dataset_permission", @@ -1313,39 +946,22 @@ class TestDatasetQueryApi: ), ): with pytest.raises(Forbidden): - method(api, current_user, dataset_id) + method(api, MagicMock(), current_user, dataset_id) def test_get_queries_pagination_has_more(self, app: Flask): api = DatasetQueryApi() method = unwrap(api.get) - dataset_id = "dataset-id" current_user = make_account() - dataset = make_dataset(id=dataset_id) - queries = [self._query_record(index) for index in range(1, 21)] - with ( app.test_request_context("/datasets/queries?page=1&limit=20"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - return_value=None, - ), - patch.object( - DatasetService, - "get_dataset_queries", - return_value=(queries, 40), - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", return_value=None), + patch.object(DatasetService, "get_dataset_queries", return_value=(queries, 40)), ): - response, status = method(api, current_user, dataset_id) - + response, status = method(api, MagicMock(), current_user, dataset_id) assert status == 200 assert response["has_more"] is True assert len(response["data"]) == 20 @@ -1371,12 +987,7 @@ class TestDatasetIndexingEstimateApi: def _base_payload(self): return { - "info_list": { - "data_source_type": "upload_file", - "file_info_list": { - "file_ids": ["file-1"], - }, - }, + "info_list": {"data_source_type": "upload_file", "file_info_list": {"file_ids": ["file-1"]}}, "process_rule": {"chunk_size": 100}, "indexing_technique": "high_quality", "doc_form": IndexStructureType.PARAGRAPH_INDEX, @@ -1387,36 +998,20 @@ class TestDatasetIndexingEstimateApi: def test_post_success_upload_file(self, app: Flask): api = DatasetIndexingEstimateApi() method = unwrap(api.post) - payload = self._base_payload() - mock_file = self._upload_file() + session = MagicMock() + session.scalars.return_value.all.return_value = [mock_file] mock_response = IndexingEstimate(total_segments=100, preview=[]) with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch( - "controllers.console.datasets.datasets.DocumentService.estimate_args_validate", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [mock_file]), - ), - patch( - "controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", - return_value=mock_response, - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), + patch("controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", return_value=mock_response), ): - response, status = method(api, "tenant-1") - + response, status = method(api, session, "tenant-1") assert status == 200 assert response == { "tokens": 0, @@ -1429,153 +1024,101 @@ class TestDatasetIndexingEstimateApi: def test_post_file_not_found(self, app: Flask): api = DatasetIndexingEstimateApi() method = unwrap(api.post) - payload = self._base_payload() - + session = MagicMock() + session.scalars.return_value.all.return_value = None with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch( - "controllers.console.datasets.datasets.DocumentService.estimate_args_validate", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: None), - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1") + method(api, session, "tenant-1") def test_post_llm_bad_request_error(self, app: Flask): api = DatasetIndexingEstimateApi() method = unwrap(api.post) mock_file = self._upload_file() - payload = self._base_payload() - + session = MagicMock() + session.scalars.return_value.all.return_value = [mock_file] with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch( - "controllers.console.datasets.datasets.DocumentService.estimate_args_validate", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [mock_file]), - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), patch( "controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", side_effect=LLMBadRequestError(), ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1") + method(api, session, "tenant-1") def test_post_provider_token_not_init(self, app: Flask): api = DatasetIndexingEstimateApi() method = unwrap(api.post) mock_file = self._upload_file() - payload = self._base_payload() - + session = MagicMock() + session.scalars.return_value.all.return_value = [mock_file] with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch( - "controllers.console.datasets.datasets.DocumentService.estimate_args_validate", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [mock_file]), - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), patch( "controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", side_effect=ProviderTokenNotInitError("token missing"), ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1") + method(api, session, "tenant-1") def test_post_generic_exception(self, app: Flask): api = DatasetIndexingEstimateApi() method = unwrap(api.post) mock_file = self._upload_file() - payload = self._base_payload() - + session = MagicMock() + session.scalars.return_value.all.return_value = [mock_file] with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), patch( - "controllers.console.datasets.datasets.DocumentService.estimate_args_validate", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [mock_file]), - ), - patch( - "controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", - side_effect=Exception("boom"), + "controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", side_effect=Exception("boom") ), ): with pytest.raises(IndexingEstimateError): - method(api, "tenant-1") + method(api, session, "tenant-1") class TestDatasetRelatedAppListApi: def test_get_success(self, app: Flask): api = DatasetRelatedAppListApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") - app1 = make_related_app(id="app-1", name="App 1") app2 = make_related_app(id="app-2", name="App 2") - - join1 = MagicMock(app=app1) - join2 = MagicMock(app=app2) - + join1 = MagicMock(app_id="app-1") + join2 = MagicMock(app_id="app-2") + session = MagicMock() with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets.DatasetService.check_dataset_permission", return_value=None), + patch("controllers.console.datasets.datasets.DatasetService.get_related_apps", return_value=[join1, join2]), patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.check_dataset_permission", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.get_related_apps", - return_value=[join1, join2], - ), + "controllers.console.datasets.datasets.AppService.get_app_by_id", + side_effect=[app1, app2], + ) as get_app_by_id, + patch.object( + App, + "mode_compatible_with_agent_with_session", + autospec=True, + side_effect=lambda app_model, *, session: str(app_model.mode), + ) as compatible_mode, ): - response, status = method(api, make_account(), "dataset-1") - + response, status = method(api, session, make_account(), "dataset-1") assert status == 200 assert response["total"] == 2 assert response["data"] == [ @@ -1600,69 +1143,53 @@ class TestDatasetRelatedAppListApi: "icon_url": None, }, ] + assert compatible_mode.call_args_list == [call(app1, session=session), call(app2, session=session)] + assert get_app_by_id.call_args_list == [call("app-1", session), call("app-2", session)] def test_get_dataset_not_found(self, app: Flask): api = DatasetRelatedAppListApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, make_account(), "dataset-1") + method(api, MagicMock(), make_account(), "dataset-1") def test_get_permission_denied(self, app: Flask): api = DatasetRelatedAppListApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets.DatasetService.check_dataset_permission", side_effect=services.errors.account.NoPermissionError("no permission"), ), ): with pytest.raises(Forbidden): - method(api, make_account(), "dataset-1") + method(api, MagicMock(), make_account(), "dataset-1") def test_get_filters_none_apps(self, app: Flask): api = DatasetRelatedAppListApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") - app1 = make_related_app() - - join1 = MagicMock(app=app1) - join2 = MagicMock(app=None) - + join1 = MagicMock(app_id="app-1") + join2 = MagicMock(app_id="app-2") + session = MagicMock() with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets.DatasetService.check_dataset_permission", return_value=None), + patch("controllers.console.datasets.datasets.DatasetService.get_related_apps", return_value=[join1, join2]), patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.check_dataset_permission", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.get_related_apps", - return_value=[join1, join2], + "controllers.console.datasets.datasets.AppService.get_app_by_id", + side_effect=[app1, None], ), ): - response, status = method(api, make_account(), "dataset-1") - + response, status = method(api, session, make_account(), "dataset-1") assert status == 200 assert response["total"] == 1 assert response["data"] == [ @@ -1683,7 +1210,6 @@ class TestDatasetIndexingStatusApi: def test_get_success_with_documents(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) - document = MagicMock() document.id = "doc-1" document.indexing_status = "completed" @@ -1695,24 +1221,14 @@ class TestDatasetIndexingStatusApi: document.paused_at = None document.error = None document.stopped_at = None - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [document]), - ), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - return_value=3, - ), - ): - response, status = method(api, "tenant-1", "dataset-1") - + session = MagicMock() + session.scalars.return_value.all.return_value = [document] + session.scalar.return_value = 3 + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "dataset-1") assert status == 200 assert "data" in response assert len(response["data"]) == 1 - item = response["data"][0] assert item["completed_segments"] == 3 assert item["total_segments"] == 3 @@ -1720,23 +1236,16 @@ class TestDatasetIndexingStatusApi: def test_get_success_no_documents(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: []), - ), - ): - response, status = method(api, "tenant-1", "dataset-1") - + session = MagicMock() + session.scalars.return_value.all.return_value = [] + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "dataset-1") assert status == 200 assert response == {"data": []} def test_segment_counts_different_values(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) - document = MagicMock() document.id = "doc-1" document.indexing_status = "indexing" @@ -1748,20 +1257,11 @@ class TestDatasetIndexingStatusApi: document.paused_at = None document.error = None document.stopped_at = None - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [document]), - ), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - side_effect=[2, 5], - ), - ): - response, status = method(api, "tenant-1", "dataset-1") - + session = MagicMock() + session.scalars.return_value.all.return_value = [document] + session.scalar.side_effect = [2, 5] + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "dataset-1") assert status == 200 item = response["data"][0] assert item["completed_segments"] == 2 @@ -1772,7 +1272,6 @@ class TestDatasetApiKeyApi: def test_get_api_keys_success(self, app: Flask): api = DatasetApiKeyApi() method = unwrap(api.get) - mock_key_1 = MagicMock(spec=ApiToken) mock_key_1.id = "key-1" mock_key_1.type = "dataset" @@ -1785,16 +1284,10 @@ class TestDatasetApiKeyApi: mock_key_2.token = "ds-def" mock_key_2.last_used_at = None mock_key_2.created_at = None - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalars", - return_value=MagicMock(all=lambda: [mock_key_1, mock_key_2]), - ), - ): - response = method(api, "tenant-1") - + session = MagicMock() + session.scalars.return_value.all.return_value = [mock_key_1, mock_key_2] + with app.test_request_context("/"): + response = method(api, session, "tenant-1") assert "data" in response assert len(response["data"]) == 2 assert response["data"][0]["id"] == "key-1" @@ -1805,58 +1298,33 @@ class TestDatasetApiKeyApi: def test_post_create_api_key_success(self, app: Flask): api = DatasetApiKeyApi() method = unwrap(api.post) - mock_token = MagicMock() mock_token.id = "new-key-id" mock_token.last_used_at = None mock_token.created_at = datetime.datetime(2024, 1, 1, 0, 0, 0, tzinfo=datetime.UTC) - mock_api_token_cls = MagicMock() mock_api_token_cls.return_value = mock_token mock_api_token_cls.generate_api_key.return_value = "dataset-abc123" - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - return_value=3, - ), - patch( - "controllers.console.datasets.datasets.ApiToken", - mock_api_token_cls, - ), - patch( - "controllers.console.datasets.datasets.db.session.add", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.commit", - return_value=None, - ), - ): - response, status = method(api, "tenant-1") - + session = MagicMock() + session.scalar.return_value = 3 + with app.test_request_context("/"), patch("controllers.console.datasets.datasets.ApiToken", mock_api_token_cls): + response, status = method(api, session, "tenant-1") assert status == 200 assert isinstance(response, dict) assert response["id"] == "new-key-id" assert response["token"] == "dataset-abc123" assert response["type"] == "dataset" assert response["created_at"] is not None + mock_api_token_cls.generate_api_key.assert_called_once_with("dataset-", 24, session=session) def test_post_exceed_max_keys(self, app: Flask): api = DatasetApiKeyApi() method = unwrap(api.post) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - return_value=10, - ), - ): + session = MagicMock() + session.scalar.return_value = 10 + with app.test_request_context("/"): with pytest.raises(BadRequest) as exc_info: - method(api, "tenant-1") - + method(api, session, "tenant-1") assert exc_info.value.code == 400 assert vars(exc_info.value)["data"] == { "message": "Cannot create more than 10 API keys for this resource type.", @@ -1868,74 +1336,44 @@ class TestDatasetApiDeleteApi: def test_delete_success(self, app: Flask): api = DatasetApiDeleteApi() method = unwrap(api.delete) - mock_key = MagicMock() - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - return_value=mock_key, - ), - patch( - "controllers.console.datasets.datasets.db.session.commit", - return_value=None, - ), - patch( - "controllers.console.datasets.datasets.db.session.delete", - return_value=None, - ), - ): - response, status = method(api, "tenant-1", "api-key-id") - + session = MagicMock() + session.scalar.return_value = mock_key + with app.test_request_context("/"): + response, status = method(api, session, "tenant-1", "api-key-id") assert status == 204 assert response == "" def test_delete_key_not_found(self, app: Flask): api = DatasetApiDeleteApi() method = unwrap(api.delete) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.db.session.scalar", - return_value=None, - ), - ): + session = MagicMock() + session.scalar.return_value = None + with app.test_request_context("/"): with pytest.raises(NotFound): - method(api, "tenant-1", "api-key-id") + method(api, session, "tenant-1", "api-key-id") class TestDatasetEnableApiApi: def test_enable_api(self, app: Flask): api = DatasetEnableApiApi() method = unwrap(api.post) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None), ): - response, status = method(api, "dataset-1", "enable") - + response, status = method(api, MagicMock(), "dataset-1", "enable") assert status == 200 assert response["result"] == "success" def test_disable_api(self, app: Flask): api = DatasetEnableApiApi() method = unwrap(api.post) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None), ): - response, status = method(api, "dataset-1", "disable") - + response, status = method(api, MagicMock(), "dataset-1", "disable") assert status == 200 assert response["result"] == "success" @@ -1944,46 +1382,31 @@ class TestDatasetApiBaseUrlApi: def test_get_api_base_url_from_config(self, app: Flask): api = DatasetApiBaseUrlApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", - "https://example.com", - ), + patch("controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", "https://example.com"), ): response = method(api) - assert response["api_base_url"] == "https://example.com/v1" def test_get_api_base_url_from_request(self, app: Flask): api = DatasetApiBaseUrlApi() method = unwrap(api.get) - with ( app.test_request_context("http://localhost:5000/"), - patch( - "controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", - None, - ), + patch("controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", None), ): response = method(api) - assert response["api_base_url"] == "http://localhost:5000/v1" def test_get_api_base_url_no_double_v1(self, app: Flask): api = DatasetApiBaseUrlApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", - "https://example.com/v1", - ), + patch("controllers.console.datasets.datasets.dify_config.SERVICE_API_URL", "https://example.com/v1"), ): response = method(api) - assert response["api_base_url"] == "https://example.com/v1" @@ -1991,20 +1414,15 @@ class TestDatasetRetrievalSettingApi: def test_get_success(self, app: Flask): api = DatasetRetrievalSettingApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.dify_config.VECTOR_STORE", - "qdrant", - ), + patch("controllers.console.datasets.datasets.dify_config.VECTOR_STORE", "qdrant"), patch( "controllers.console.datasets.datasets._get_retrieval_methods_by_vector_type", return_value={"retrieval_method": ["semantic", "hybrid"]}, ), ): response = method(api) - assert "retrieval_method" in response @@ -2012,7 +1430,6 @@ class TestDatasetRetrievalSettingMockApi: def test_get_success(self, app: Flask): api = DatasetRetrievalSettingMockApi() method = unwrap(api.get) - with ( app.test_request_context("/"), patch( @@ -2021,7 +1438,6 @@ class TestDatasetRetrievalSettingMockApi: ), ): response = method(api, "milvus") - assert response["retrieval_method"] == ["semantic"] @@ -2029,124 +1445,89 @@ class TestDatasetErrorDocs: def test_get_success(self, app: Flask): api = DatasetErrorDocs() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") error_doc = make_document_status(id="error-doc", indexing_status=IndexingStatus.ERROR, error="failed") - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets.DocumentService.get_error_documents_by_dataset_id", return_value=[error_doc], ), ): - response, status = method(api, "dataset-1") - + response, status = method(api, MagicMock(), "dataset-1") assert status == 200 assert response["total"] == 1 def test_get_dataset_not_found(self, app: Flask): api = DatasetErrorDocs() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "dataset-1") + method(api, MagicMock(), "dataset-1") class TestDatasetPermissionUserListApi: def test_get_success(self, app: Flask): api = DatasetPermissionUserListApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") users = ["u1", "u2"] - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.check_dataset_permission", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets.DatasetService.check_dataset_permission", return_value=None), patch( "controllers.console.datasets.datasets.DatasetPermissionService.get_dataset_partial_member_list", return_value=users, ), ): - response, status = method(api, make_account(), "dataset-1") - + response, status = method(api, MagicMock(), make_account(), "dataset-1") assert status == 200 assert response["data"] == users def test_get_permission_denied(self, app: Flask): api = DatasetPermissionUserListApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets.DatasetService.check_dataset_permission", side_effect=services.errors.account.NoPermissionError("no permission"), ), ): with pytest.raises(Forbidden): - method(api, make_account(), "dataset-1") + method(api, MagicMock(), make_account(), "dataset-1") class TestDatasetAutoDisableLogApi: def test_get_success(self, app: Flask): api = DatasetAutoDisableLogApi() method = unwrap(api.get) - dataset = make_dataset(id="dataset-1") logs = {"document_ids": ["doc-1"], "count": 1} - with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset_auto_disable_logs", - return_value=logs, + "controllers.console.datasets.datasets.DatasetService.get_dataset_auto_disable_logs", return_value=logs ), ): - response, status = method(api, "dataset-1") - + response, status = method(api, MagicMock(), "dataset-1") assert status == 200 assert response == logs def test_get_dataset_not_found(self, app: Flask): api = DatasetAutoDisableLogApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "dataset-1") + method(api, MagicMock(), "dataset-1") diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py index 7c415c01cdc..85ffd1b2d20 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py @@ -1,5 +1,6 @@ import datetime -import inspect +from inspect import unwrap +from types import SimpleNamespace from unittest.mock import ANY, MagicMock, patch import pytest @@ -73,8 +74,20 @@ def make_serializable_document(**overrides): "total_segments": None, } attrs.update(overrides) - document = MagicMock(spec_set=list(attrs)) + document = MagicMock( + spec_set=[ + *attrs, + "get_data_source_detail_dict", + "get_dataset_process_rule", + "get_doc_metadata_details", + "get_hit_count", + ] + ) document.configure_mock(**attrs) + document.get_data_source_detail_dict.return_value = attrs["data_source_detail_dict"] + document.get_dataset_process_rule.return_value = None + document.get_doc_metadata_details.return_value = attrs["doc_metadata_details"] + document.get_hit_count.return_value = attrs["hit_count"] return document @@ -112,8 +125,22 @@ def make_document_detail(**overrides): "need_summary": False, } attrs.update(overrides) - document = MagicMock(spec_set=list(attrs)) + document = MagicMock( + spec_set=[ + *attrs, + "get_data_source_detail_dict", + "get_dataset_process_rule", + "get_doc_metadata_details", + "get_hit_count", + "get_segment_count", + ] + ) document.configure_mock(**attrs) + document.get_data_source_detail_dict.return_value = attrs["data_source_detail_dict"] + document.get_dataset_process_rule.return_value = attrs["dataset_process_rule"] + document.get_doc_metadata_details.return_value = attrs["doc_metadata_details"] + document.get_hit_count.return_value = attrs["hit_count"] + document.get_segment_count.return_value = attrs["segment_count"] return document @@ -186,18 +213,14 @@ def document(): @pytest.fixture def patch_dataset(dataset): - with patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=dataset, - ): + with patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset): yield @pytest.fixture def patch_permission(): with patch( - "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", - return_value=None, + "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", return_value=None ): yield @@ -205,21 +228,21 @@ def patch_permission(): class TestGetProcessRuleApi: def test_get_default_success(self, app: Flask, patch_tenant): api = GetProcessRuleApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - with app.test_request_context("/"): - response = method(api, user) - + response = method(api, MagicMock(), user) assert "rules" in response def test_get_with_document_preserves_legacy_segmentation_delimiter(self, app: Flask, patch_tenant): api = GetProcessRuleApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant document = MagicMock(dataset_id="ds-1") - process_rule = MagicMock( + session = MagicMock() + dataset = MagicMock() + dataset.get_latest_process_rule.return_value = SimpleNamespace( mode="custom", rules_dict={"segmentation": {"delimiter": "---", "max_tokens": 123}}, ) @@ -227,76 +250,89 @@ class TestGetProcessRuleApi: with ( app.test_request_context("/?document_id=doc-1"), patch( - "controllers.console.datasets.datasets_document.db.get_or_404", + "controllers.console.datasets.datasets_document.DocumentService.get_document_by_id", return_value=document, - ), + ) as mock_get_document, patch( "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), + return_value=dataset, ), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=process_rule, - ), ): - response = method(api, user) + response = method(api, session, user) + mock_get_document.assert_called_once_with("doc-1", session) + dataset.get_latest_process_rule.assert_called_once_with(session=session) assert response["rules"]["segmentation"]["separator"] == "---" assert response["rules"]["segmentation"]["max_tokens"] == 123 assert "delimiter" not in response["rules"]["segmentation"] - def test_get_with_document_dataset_not_found(self, app: Flask, patch_tenant): + def test_get_with_document_preserves_null_rules(self, app: Flask, patch_tenant): api = GetProcessRuleApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - - document = MagicMock(dataset_id="ds-1") - + session = MagicMock() + dataset = MagicMock() + dataset.get_latest_process_rule.return_value = SimpleNamespace(mode="custom", rules_dict=None) with ( app.test_request_context("/?document_id=doc-1"), patch( - "controllers.console.datasets.datasets_document.db.get_or_404", - return_value=document, + "controllers.console.datasets.datasets_document.DocumentService.get_document_by_id", + return_value=MagicMock(dataset_id="ds-1"), ), patch( "controllers.console.datasets.datasets_document.DatasetService.get_dataset", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", return_value=None, ), + ): + response = method(api, session, user) + + assert response["mode"] == "custom" + assert response["rules"] is None + + def test_get_with_document_dataset_not_found(self, app: Flask, patch_tenant): + api = GetProcessRuleApi() + method = unwrap(api.get) + user, _ = patch_tenant + document = MagicMock(dataset_id="ds-1") + session = MagicMock() + with ( + app.test_request_context("/?document_id=doc-1"), + patch( + "controllers.console.datasets.datasets_document.DocumentService.get_document_by_id", + return_value=document, + ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, user) + method(api, session, user) class TestDatasetDocumentListApi: def test_get_with_fetch_true_counts_segments(self, app: Flask, patch_tenant, patch_dataset, patch_permission): api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - doc = make_serializable_document() pagination = MagicMock(items=[doc], total=1) - + session = MagicMock() + session.scalar.return_value = 2 with ( app.test_request_context("/?fetch=true"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=2, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.enrich_documents_with_summary_index_status", return_value=None, ), ): - resp = method(api, tenant_id, user, "ds-1") - + resp = method(api, session, tenant_id, user, "ds-1") assert resp["data"][0]["id"] == "doc-1" assert resp["data"][0]["completed_segments"] == 2 assert resp["data"][0]["total_segments"] == 2 @@ -305,17 +341,12 @@ class TestDatasetDocumentListApi: self, app: Flask, patch_tenant, patch_dataset, patch_permission ): api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - pagination = MagicMock(items=[make_serializable_document()], total=1) - with ( app.test_request_context("/?keyword=test&status=enabled&sort=created_at"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.apply_display_status_filter", side_effect=lambda q, s: q, @@ -325,44 +356,43 @@ class TestDatasetDocumentListApi: return_value=None, ), ): - resp = method(api, tenant_id, user, "ds-1") - + resp = method(api, MagicMock(), tenant_id, user, "ds-1") assert resp["total"] == 1 def test_get_success(self, app: Flask, patch_tenant, patch_dataset, patch_permission): api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - - pagination = MagicMock(items=[make_serializable_document()], total=1) - + document = make_serializable_document() + pagination = MagicMock(items=[document], total=1) + session = MagicMock() with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.enrich_documents_with_summary_index_status", return_value=None, ), ): - response = method(api, tenant_id, user, "ds-1") - + response = method(api, session, tenant_id, user, "ds-1") assert response["total"] == 1 assert response["data"][0]["id"] == "doc-1" assert "completed_segments" not in response["data"][0] assert "total_segments" not in response["data"][0] + document.get_data_source_detail_dict.assert_called_once_with(session=session) + document.get_dataset_process_rule.assert_called_once_with(session=session) + document.get_doc_metadata_details.assert_called_once_with(session=session) + document.get_hit_count.assert_called_once_with(session=session) def test_post_success(self, app: Flask, patch_tenant, patch_dataset, patch_permission): api = DatasetDocumentListApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - payload = {"indexing_technique": "economy"} created_dataset = make_dataset() created_document = make_document() - + session = MagicMock() + session.scalar.return_value = 0 with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -378,10 +408,8 @@ class TestDatasetDocumentListApi: "controllers.console.datasets.datasets_document.DocumentService.save_document_with_dataset_id", return_value=([created_document], "batch-1"), ), - patch("models.dataset.db.session.scalar", return_value=0), ): - response = method(api, user, "ds-1") - + response = method(api, session, user, "ds-1") assert "documents" in response assert response["dataset"]["id"] == "ds-1" assert response["documents"][0]["id"] == "doc-1" @@ -392,76 +420,59 @@ class TestDatasetDocumentListApi: def test_post_forbidden(self, app: Flask): api = DatasetDocumentListApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = MagicMock(is_dataset_editor=False) - with ( app.test_request_context("/", json={}), patch.object(type(console_ns), "payload", {}), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), ): with pytest.raises(Forbidden): - method(api, user, "ds-1") + method(api, MagicMock(), user, "ds-1") def test_get_with_fetch_true_and_invalid_fetch(self, app: Flask, patch_tenant, patch_dataset, patch_permission): api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - pagination = MagicMock(items=[make_serializable_document()], total=1) - with ( app.test_request_context("/?fetch=maybe"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.enrich_documents_with_summary_index_status", return_value=None, ), ): - response = method(api, tenant_id, user, "ds-1") - + response = method(api, MagicMock(), tenant_id, user, "ds-1") assert response["total"] == 1 def test_get_sort_hit_count(self, app: Flask, patch_tenant, patch_dataset, patch_permission): api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - pagination = MagicMock(items=[], total=0) - with ( app.test_request_context("/?sort=hit_count"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.enrich_documents_with_summary_index_status", return_value=None, ), ): - response = method(api, tenant_id, user, "ds-1") - + response = method(api, MagicMock(), tenant_id, user, "ds-1") assert response["total"] == 0 class TestDatasetInitApi: def test_post_success_serializes_created_dataset_and_documents(self, app: Flask, patch_tenant): api = DatasetInitApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, tenant_id = patch_tenant - payload = {"indexing_technique": "economy"} created_dataset = make_dataset() created_document = make_document(id="doc-init") - + session = MagicMock() + session.scalar.return_value = 0 with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -473,10 +484,8 @@ class TestDatasetInitApi: "controllers.console.datasets.datasets_document.DocumentService.save_document_without_dataset_id", return_value=(created_dataset, [created_document], "batch-init"), ), - patch("models.dataset.db.session.scalar", return_value=0), ): - response = method(api, tenant_id, user) - + response = method(api, session, tenant_id, user) assert response["dataset"]["id"] == "ds-1" assert response["documents"][0]["id"] == "doc-init" assert response["documents"][0]["data_source_info"] == {} @@ -487,37 +496,29 @@ class TestDatasetInitApi: class TestDocumentApi: def test_get_success(self, app: Flask, patch_tenant): api = DocumentApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = make_document_detail() - with ( app.test_request_context("/"), patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_process_rules", - return_value={}, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") assert status == 200 def test_get_invalid_metadata(self, app: Flask, patch_tenant): api = DocumentApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - with app.test_request_context("/?metadata=wrong"), patch.object(api, "get_document", return_value=MagicMock()): with pytest.raises(InvalidMetadataError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") def test_delete_success(self, app: Flask, patch_tenant, patch_dataset): api = DocumentApi() - method = inspect.unwrap(api.delete) + method = unwrap(api.delete) user, tenant_id = patch_tenant - with ( app.test_request_context("/"), patch( @@ -525,20 +526,15 @@ class TestDocumentApi: return_value=None, ), patch.object(api, "get_document", return_value=MagicMock()), - patch( - "controllers.console.datasets.datasets_document.DocumentService.delete_document", - return_value=None, - ), + patch("controllers.console.datasets.datasets_document.DocumentService.delete_document", return_value=None), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") assert status == 204 def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_dataset): api = DocumentApi() - method = inspect.unwrap(api.delete) + method = unwrap(api.delete) user, tenant_id = patch_tenant - with ( app.test_request_context("/"), patch( @@ -552,17 +548,15 @@ class TestDocumentApi: ), ): with pytest.raises(DocumentIndexingError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") class TestDocumentDownloadApi: def test_download_success(self, app: Flask, patch_tenant): api = DocumentDownloadApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock() - with ( app.test_request_context("/"), patch.object(api, "get_document", return_value=document), @@ -571,109 +565,68 @@ class TestDocumentDownloadApi: return_value="url", ), ): - response = method(api, tenant_id, user, "ds-1", "doc-1") - + response = method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") assert response["url"] == "url" class TestDocumentProcessingApi: def test_processing_forbidden_when_not_editor(self, app: Flask): api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) user = MagicMock(is_dataset_editor=False) - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=MagicMock()), - ): + with app.test_request_context("/"), patch.object(api, "get_document", return_value=MagicMock()): with pytest.raises(Forbidden): - method(api, "tenant-1", user, "ds-1", "doc-1", "pause") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "pause") def test_resume_from_error_state(self, app: Flask, patch_tenant): api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, tenant_id = patch_tenant - doc = MagicMock(indexing_status=IndexingStatus.ERROR, is_paused=True) - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=doc), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), - ): - _, status = method(api, tenant_id, user, "ds-1", "doc-1", "resume") - + session = MagicMock() + with app.test_request_context("/"), patch.object(api, "get_document", return_value=doc): + _, status = method(api, session, tenant_id, user, "ds-1", "doc-1", "resume") assert status == 200 def test_resume_success(self, app: Flask, patch_tenant): api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, tenant_id = patch_tenant - document = MagicMock(indexing_status=IndexingStatus.PAUSED, is_paused=True) - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), - ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1", "resume") - + session = MagicMock() + with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1", "resume") assert status == 200 def test_pause_success(self, app: Flask, patch_tenant): api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, tenant_id = patch_tenant - document = MagicMock(indexing_status="indexing") - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), - ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1", "pause") - + session = MagicMock() + with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1", "pause") assert status == 200 def test_pause_invalid(self, app: Flask, patch_tenant): api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, tenant_id = patch_tenant - document = MagicMock(indexing_status=IndexingStatus.COMPLETED) - with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): with pytest.raises(InvalidActionError): - method(api, tenant_id, user, "ds-1", "doc-1", "pause") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1", "pause") class TestDocumentMetadataApi: def test_put_metadata_schema_filtering(self, app: Flask, patch_tenant): api = DocumentMetadataApi() - method = inspect.unwrap(api.put) + method = unwrap(api.put) user, tenant_id = patch_tenant - doc = MagicMock() - - payload = { - "doc_type": "invoice", - "doc_metadata": {"amount": 10, "invalid": "x"}, - } - + payload = {"doc_type": "invoice", "doc_metadata": {"amount": 10, "invalid": "x"}} schema = {"amount": int} - + session = MagicMock() with ( app.test_request_context("/", json=payload), patch.object(api, "get_document", return_value=doc), @@ -681,24 +634,17 @@ class TestDocumentMetadataApi: "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"invoice": schema}, ), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), ): - method(api, tenant_id, user, "ds-1", "doc-1") - + method(api, session, tenant_id, user, "ds-1", "doc-1") assert doc.doc_metadata == {"amount": 10} def test_put_success(self, app: Flask, patch_tenant): api = DocumentMetadataApi() - method = inspect.unwrap(api.put) + method = unwrap(api.put) user, tenant_id = patch_tenant - document = MagicMock() - payload = {"doc_type": "others", "doc_metadata": {"a": 1}} - + session = MagicMock() with ( app.test_request_context("/", json=payload), patch.object(api, "get_document", return_value=document), @@ -706,31 +652,23 @@ class TestDocumentMetadataApi: "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"others": {}}, ), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 def test_put_invalid_payload(self, app: Flask, patch_tenant): api = DocumentMetadataApi() - method = inspect.unwrap(api.put) + method = unwrap(api.put) user, tenant_id = patch_tenant - with app.test_request_context("/", json={}), patch.object(api, "get_document", return_value=MagicMock()): with pytest.raises(ValueError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") def test_put_invalid_doc_type(self, app: Flask, patch_tenant): api = DocumentMetadataApi() - method = inspect.unwrap(api.put) + method = unwrap(api.put) user, tenant_id = patch_tenant - payload = {"doc_type": "invalid", "doc_metadata": {}} - with ( app.test_request_context("/", json=payload), patch.object(api, "get_document", return_value=MagicMock()), @@ -740,15 +678,14 @@ class TestDocumentMetadataApi: ), ): with pytest.raises(ValueError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") class TestDocumentStatusApi: def test_patch_success(self, app: Flask, patch_tenant, patch_dataset): api = DocumentStatusApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, _ = patch_tenant - with ( app.test_request_context("/?document_id=doc-1"), patch( @@ -764,15 +701,13 @@ class TestDocumentStatusApi: return_value=None, ), ): - response, status = method(api, user, "ds-1", "enable") - + response, status = method(api, MagicMock(), user, "ds-1", "enable") assert status == 200 def test_patch_invalid_action(self, app: Flask, patch_tenant, patch_dataset): api = DocumentStatusApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, _ = patch_tenant - with ( app.test_request_context("/?document_id=doc-1"), patch( @@ -789,89 +724,58 @@ class TestDocumentStatusApi: ), ): with pytest.raises(InvalidActionError): - method(api, user, "ds-1", "enable") + method(api, MagicMock(), user, "ds-1", "enable") class TestDocumentRetryApi: def test_retry_archived_document_skipped(self, app: Flask, patch_tenant, patch_dataset): api = DocumentRetryApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"document_ids": ["doc-1"]} - doc = MagicMock(indexing_status="indexing") - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", - return_value=doc, - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.check_archived", - return_value=True, - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.retry_document", - ) as retry_mock, + patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=doc), + patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=True), + patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_mock, ): - resp, status = method(api, "ds-1") - + resp, status = method(api, MagicMock(), "ds-1") assert status == 204 retry_mock.assert_called_once_with("ds-1", [], ANY) def test_retry_success(self, app: Flask, patch_tenant, patch_dataset): api = DocumentRetryApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"document_ids": ["doc-1"]} - document = MagicMock(indexing_status=IndexingStatus.INDEXING, archived=False) - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), + patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False), patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.check_archived", - return_value=False, - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.retry_document", - return_value=None, + "controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None ) as retry_mock, ): - response, status = method(api, "ds-1") - + response, status = method(api, MagicMock(), "ds-1") assert status == 204 retry_mock.assert_called_once_with("ds-1", [document], ANY) def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_dataset): api = DocumentRetryApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"document_ids": ["doc-1"]} - document = MagicMock(indexing_status=IndexingStatus.COMPLETED, archived=False) - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.retry_document", - return_value=None, + "controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None ) as retry_mock, ): - response, status = method(api, "ds-1") - + response, status = method(api, MagicMock(), "ds-1") assert status == 204 retry_mock.assert_called_once_with("ds-1", [], ANY) @@ -879,126 +783,86 @@ class TestDocumentRetryApi: class TestDocumentPipelineExecutionLogApi: def test_get_log_success(self, app: Flask, patch_tenant, patch_dataset): api = DocumentPipelineExecutionLogApi() - method = inspect.unwrap(api.get) - - log = MagicMock( - datasource_info="{}", - datasource_type="file", - input_data={}, - datasource_node_id="n1", - ) - + method = unwrap(api.get) + log = MagicMock(datasource_info="{}", datasource_type="file", input_data={}, datasource_node_id="n1") + session = MagicMock() + session.scalar.return_value = log with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=log, + "controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=MagicMock() ), ): - response, status = method(api, "ds-1", "doc-1") - + response, status = method(api, session, "ds-1", "doc-1") assert status == 200 class TestDocumentGenerateSummaryApi: def test_generate_summary_missing_documents(self, app: Flask, patch_tenant, patch_permission): api = DocumentGenerateSummaryApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - - dataset = MagicMock( - indexing_technique="high_quality", - summary_index_setting={"enable": True}, - ) - + dataset = MagicMock(indexing_technique="high_quality", summary_index_setting={"enable": True}) payload = {"document_list": ["doc-1", "doc-2"]} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids", return_value=[MagicMock(id="doc-1")], ), ): with pytest.raises(NotFound): - method(api, user, "ds-1") + method(api, MagicMock(), user, "ds-1") def test_generate_not_enabled(self, app: Flask, patch_tenant, patch_permission): api = DocumentGenerateSummaryApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - dataset = MagicMock(indexing_technique="high_quality", summary_index_setting={"enable": False}) - payload = {"document_list": ["doc-1"]} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), ): with pytest.raises(ValueError): - method(api, user, "ds-1") + method(api, MagicMock(), user, "ds-1") def test_generate_summary_success_with_qa_skip(self, app: Flask, patch_tenant, patch_permission): api = DocumentGenerateSummaryApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - - dataset = MagicMock( - indexing_technique="high_quality", - summary_index_setting={"enable": True}, - ) - + dataset = MagicMock(indexing_technique="high_quality", summary_index_setting={"enable": True}) doc1 = MagicMock(id="doc-1", doc_form=IndexStructureType.QA_INDEX) doc2 = MagicMock(id="doc-2", doc_form=IndexStructureType.PARAGRAPH_INDEX) - payload = {"document_list": ["doc-1", "doc-2"]} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids", return_value=[doc1, doc2], ), patch( - "controllers.console.datasets.datasets_document.generate_summary_index_task.delay", - return_value=None, + "controllers.console.datasets.datasets_document.generate_summary_index_task.delay", return_value=None ), ): - response, status = method(api, user, "ds-1") - + response, status = method(api, MagicMock(), user, "ds-1") assert status == 200 class TestDocumentSummaryStatusApi: def test_get_success(self, app: Flask, patch_tenant, patch_permission): api = DocumentSummaryStatusApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() ), patch( "services.summary_index_service.SummaryIndexService.get_document_summary_status_detail", @@ -1015,8 +879,7 @@ class TestDocumentSummaryStatusApi: }, ), ): - response, status = method(api, user, "ds-1", "doc-1") - + response, status = method(api, MagicMock(), user, "ds-1", "doc-1") assert status == 200 assert response["summary_status"]["timeout"] == 1 assert response["summaries"][0]["status"] == "timeout" @@ -1025,9 +888,8 @@ class TestDocumentSummaryStatusApi: class TestDocumentIndexingEstimateApi: def test_indexing_estimate_file_not_found(self, app: Flask, patch_tenant): api = DocumentIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type=DataSourceType.UPLOAD_FILE, @@ -1036,23 +898,16 @@ class TestDocumentIndexingEstimateApi: doc_form=IndexStructureType.PARAGRAPH_INDEX, dataset_process_rule=None, ) - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=None, - ), - ): + session = MagicMock() + session.scalar.return_value = None + with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): with pytest.raises(NotFound): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, session, tenant_id, user, "ds-1", "doc-1") def test_indexing_estimate_generic_exception(self, app: Flask, patch_tenant): api = DocumentIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type=DataSourceType.UPLOAD_FILE, @@ -1061,82 +916,61 @@ class TestDocumentIndexingEstimateApi: doc_form=IndexStructureType.PARAGRAPH_INDEX, dataset_process_rule=None, ) - upload_file = MagicMock() - mock_indexing_runner = MagicMock() mock_indexing_runner.indexing_estimate.side_effect = RuntimeError("Some indexing error") - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/"), patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=upload_file, - ), - patch( - "controllers.console.datasets.datasets_document.ExtractSetting", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_document.IndexingRunner", - return_value=mock_indexing_runner, - ), + patch("controllers.console.datasets.datasets_document.ExtractSetting", return_value=MagicMock()), + patch("controllers.console.datasets.datasets_document.IndexingRunner", return_value=mock_indexing_runner), ): with pytest.raises(IndexingEstimateError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, session, tenant_id, user, "ds-1", "doc-1") def test_get_finished(self, app: Flask, patch_tenant): api = DocumentIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock(indexing_status=IndexingStatus.COMPLETED) - with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): with pytest.raises(DocumentAlreadyFinishedError): - method(api, tenant_id, user, "ds-1", "doc-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") class TestDocumentBatchDownloadZipApi: def test_post_no_documents(self, app: Flask, patch_tenant): api = DocumentBatchDownloadZipApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, tenant_id = patch_tenant - payload: dict[str, list[str]] = {"document_ids": []} - with app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError): - method(api, tenant_id, user, "ds-1") + method(api, MagicMock(), tenant_id, user, "ds-1") class TestDatasetDocumentListApiDelete: def test_delete_success(self, app: Flask, patch_tenant, patch_dataset): """Test successful deletion of documents""" api = DatasetDocumentListApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) with ( app.test_request_context("/?document_id=doc-1&document_id=doc-2"), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_model_setting", return_value=None, ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.delete_documents", - return_value=None, - ), + patch("controllers.console.datasets.datasets_document.DocumentService.delete_documents", return_value=None), ): - response, status = method(api, "ds-1") - + response, status = method(api, MagicMock(), "ds-1") assert status == 204 def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_dataset): """Test deletion with indexing error""" api = DatasetDocumentListApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) with ( app.test_request_context("/?document_id=doc-1"), patch( @@ -1149,30 +983,25 @@ class TestDatasetDocumentListApiDelete: ), ): with pytest.raises(DocumentIndexingError): - method(api, "ds-1") + method(api, MagicMock(), "ds-1") def test_delete_dataset_not_found(self, app: Flask, patch_tenant): """Test deletion when dataset not found""" api = DatasetDocumentListApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) with ( app.test_request_context("/?document_id=doc-1"), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "ds-1") + method(api, MagicMock(), "ds-1") class TestDocumentBatchIndexingEstimateApi: def test_batch_indexing_estimate_website(self, app: Flask, patch_tenant): api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - doc = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type=DataSourceType.WEBSITE_CRAWL, @@ -1185,7 +1014,6 @@ class TestDocumentBatchIndexingEstimateApi: }, doc_form=IndexStructureType.PARAGRAPH_INDEX, ) - with ( app.test_request_context("/"), patch.object(api, "get_batch_documents", return_value=[doc]), @@ -1194,15 +1022,13 @@ class TestDocumentBatchIndexingEstimateApi: return_value=IndexingEstimate(total_segments=2, preview=[]), ), ): - resp, status = method(api, tenant_id, user, "ds-1", "batch-1") - + resp, status = method(api, MagicMock(), tenant_id, user, "ds-1", "batch-1") assert status == 200 def test_batch_indexing_estimate_notion(self, app: Flask, patch_tenant): api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - doc = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type=DataSourceType.NOTION_IMPORT, @@ -1214,7 +1040,6 @@ class TestDocumentBatchIndexingEstimateApi: }, doc_form=IndexStructureType.PARAGRAPH_INDEX, ) - with ( app.test_request_context("/"), patch.object(api, "get_batch_documents", return_value=[doc]), @@ -1223,43 +1048,38 @@ class TestDocumentBatchIndexingEstimateApi: return_value=IndexingEstimate(total_segments=1, preview=[]), ), ): - resp, status = method(api, tenant_id, user, "ds-1", "batch-1") - + resp, status = method(api, MagicMock(), tenant_id, user, "ds-1", "batch-1") assert status == 200 def test_batch_estimate_unsupported_datasource(self, app: Flask, patch_tenant): api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type="unknown", data_source_info_dict={}, doc_form=IndexStructureType.PARAGRAPH_INDEX, ) - with app.test_request_context("/"), patch.object(api, "get_batch_documents", return_value=[document]): with pytest.raises(ValueError): - method(api, tenant_id, user, "ds-1", "batch-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "batch-1") def test_get_batch_estimate_invalid_batch(self, app: Flask, patch_tenant): """Test batch estimation with invalid batch""" api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - with app.test_request_context("/"), patch.object(api, "get_batch_documents", side_effect=NotFound()): with pytest.raises(NotFound): - method(api, tenant_id, user, "ds-1", "invalid-batch") + method(api, MagicMock(), tenant_id, user, "ds-1", "invalid-batch") class TestDocumentBatchIndexingStatusApi: def test_get_batch_status_success_serializes_status_shape(self, app: Flask, patch_tenant): api = DocumentBatchIndexingStatusApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - document = MagicMock( id="doc-1", indexing_status=IndexingStatus.COMPLETED, @@ -1273,17 +1093,10 @@ class TestDocumentBatchIndexingStatusApi: error=None, stopped_at=None, ) - - with ( - app.test_request_context("/"), - patch.object(api, "get_batch_documents", return_value=[document]), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - side_effect=[2, 3], - ), - ): - response = method(api, user, "ds-1", "batch-1") - + session = MagicMock() + session.scalar.side_effect = [2, 3] + with app.test_request_context("/"), patch.object(api, "get_batch_documents", return_value=[document]): + response = method(api, session, user, "ds-1", "batch-1") assert response == { "data": [ { @@ -1306,20 +1119,18 @@ class TestDocumentBatchIndexingStatusApi: def test_get_batch_status_invalid_batch(self, app: Flask, patch_tenant): """Test batch status with invalid batch""" api = DocumentBatchIndexingStatusApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - with app.test_request_context("/"), patch.object(api, "get_batch_documents", side_effect=NotFound()): with pytest.raises(NotFound): - method(api, user, "ds-1", "invalid-batch") + method(api, MagicMock(), user, "ds-1", "invalid-batch") class TestDocumentIndexingStatusApi: def test_get_status_success_serializes_status_shape(self, app: Flask, patch_tenant): api = DocumentIndexingStatusApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock( id="doc-1", indexing_status=IndexingStatus.INDEXING, @@ -1333,17 +1144,10 @@ class TestDocumentIndexingStatusApi: error=None, stopped_at=None, ) - - with ( - app.test_request_context("/"), - patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - side_effect=[1, 4], - ), - ): - response = method(api, tenant_id, user, "ds-1", "doc-1") - + session = MagicMock() + session.scalar.side_effect = [1, 4] + with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): + response = method(api, session, tenant_id, user, "ds-1", "doc-1") assert response["id"] == "doc-1" assert response["indexing_status"] == "indexing" assert response["completed_segments"] == 1 @@ -1352,29 +1156,27 @@ class TestDocumentIndexingStatusApi: def test_get_status_document_not_found(self, app: Flask, patch_tenant): """Test getting status for non-existent document""" api = DocumentIndexingStatusApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - with app.test_request_context("/"), patch.object(api, "get_document", side_effect=NotFound()): with pytest.raises(NotFound): - method(api, tenant_id, user, "ds-1", "invalid-doc") + method(api, MagicMock(), tenant_id, user, "ds-1", "invalid-doc") class TestDocumentRenameApi: def test_post_success_serializes_document_shape(self, app: Flask, patch_tenant): api = DocumentRenameApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - payload = {"name": "Renamed Document"} renamed_document = make_document(id="doc-renamed", name="Renamed Document") - + session = MagicMock() + session.scalar.return_value = 0 with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=make_dataset(), + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=make_dataset() ), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_operator_permission", @@ -1384,10 +1186,8 @@ class TestDocumentRenameApi: "controllers.console.datasets.datasets_document.DocumentService.rename_document", return_value=renamed_document, ), - patch("models.dataset.db.session.scalar", return_value=0), ): - response = method(api, user, "ds-1", "doc-1") - + response = method(api, session, user, "ds-1", "doc-1") assert response["id"] == "doc-renamed" assert response["name"] == "Renamed Document" assert response["data_source_info"] == {} @@ -1399,41 +1199,29 @@ class TestDocumentApiMetadata: def test_get_with_only_option(self, app: Flask, patch_tenant): """Test get with 'only' metadata option""" api = DocumentApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = make_document_detail(doc_metadata_details=[]) - with ( app.test_request_context("/?metadata=only"), patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_process_rules", - return_value={}, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") assert status == 200 def test_get_with_without_option(self, app: Flask, patch_tenant): """Test get with 'without' metadata option""" api = DocumentApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = make_document_detail() - with ( app.test_request_context("/?metadata=without"), patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_process_rules", - return_value={}, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") assert status == 200 @@ -1441,50 +1229,40 @@ class TestDocumentGenerateSummaryApiSuccess: def test_generate_not_enabled_high_quality(self, app: Flask, patch_tenant, patch_permission): """Test summary generation on non-high-quality dataset""" api = DocumentGenerateSummaryApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user, _ = patch_tenant - dataset = MagicMock(indexing_technique="economy", summary_index_setting={"enable": True}) - payload = {"document_list": ["doc-1"]} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), ): with pytest.raises(ValueError): - method(api, user, "ds-1") + method(api, MagicMock(), user, "ds-1") class TestDocumentProcessingApiResume: def test_resume_invalid_status(self, app: Flask, patch_tenant): """Test resume on non-paused document""" api = DocumentProcessingApi() - method = inspect.unwrap(api.patch) + method = unwrap(api.patch) user, tenant_id = patch_tenant - document = MagicMock(indexing_status=IndexingStatus.COMPLETED, is_paused=False) - with app.test_request_context("/"), patch.object(api, "get_document", return_value=document): with pytest.raises(InvalidActionError): - method(api, tenant_id, user, "ds-1", "doc-1", "resume") + method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1", "resume") class TestDocumentPermissionCases: def test_document_batch_get_permission_denied(self, app: Flask, patch_tenant): api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() ), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", @@ -1492,18 +1270,16 @@ class TestDocumentPermissionCases: ), ): with pytest.raises(Forbidden): - method(api, tenant_id, user, "ds-1", "batch-1") + method(api, MagicMock(), tenant_id, user, "ds-1", "batch-1") def test_document_batch_get_documents_not_found(self, app: Flask, patch_tenant): api = DocumentBatchIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() ), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", @@ -1511,98 +1287,68 @@ class TestDocumentPermissionCases: ), patch.object(api, "get_batch_documents", return_value=None), ): - response, status = method(api, tenant_id, user, "ds-1", "batch-1") - + response, status = method(api, MagicMock(), tenant_id, user, "ds-1", "batch-1") assert status == 200 - assert response == { - "tokens": 0, - "total_price": 0, - "currency": "USD", - "total_segments": 0, - "preview": [], - } + assert response == {"tokens": 0, "total_price": 0, "currency": "USD", "total_segments": 0, "preview": []} def test_document_tenant_mismatch(self, app: Flask): api = DocumentApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) user = MagicMock(is_dataset_editor=True) - document = MagicMock( - tenant_id="other-tenant", - dataset_process_rule=None, - ) - + document = MagicMock(tenant_id="other-tenant", dataset_process_rule=None) with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_process_rules", - return_value={}, + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() ), + patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), + patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_process_rule_get_by_document_success(self, app: Flask, patch_tenant): api = GetProcessRuleApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, _ = patch_tenant - document = MagicMock(dataset_id="ds-1") - process_rule = MagicMock(mode="custom", rules_dict={"a": 1}) - + session = MagicMock() + dataset = MagicMock() + dataset.get_latest_process_rule.return_value = SimpleNamespace(mode="custom", rules_dict={"a": 1}) with ( app.test_request_context("/?document_id=doc-1"), patch( - "controllers.console.datasets.datasets_document.db.get_or_404", + "controllers.console.datasets.datasets_document.DocumentService.get_document_by_id", return_value=document, ), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), - ), + patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=process_rule, - ), ): - result = method(api, user) - + result = method(api, session, user) if isinstance(result, tuple): response, status = result else: - response, status = result, 200 - + response, status = (result, 200) assert status == 200 assert response["mode"] == "custom" def test_process_rule_permission_denied(self, app: Flask): api = GetProcessRuleApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) user = MagicMock(is_dataset_editor=True) document = MagicMock(dataset_id="ds-1") - + session = MagicMock() with ( app.test_request_context("/?document_id=doc-1"), patch( - "controllers.console.datasets.datasets_document.db.get_or_404", + "controllers.console.datasets.datasets_document.DocumentService.get_document_by_id", return_value=document, ), patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() ), patch( "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission", @@ -1610,47 +1356,36 @@ class TestDocumentPermissionCases: ), ): with pytest.raises(Forbidden): - method(api, user) + method(api, session, user) class TestDocumentListAdvancedCases: def test_document_list_with_multiple_sort_options(self, app: Flask, patch_tenant, patch_dataset, patch_permission): """Test document list with different sort options""" api = DatasetDocumentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - pagination = MagicMock(items=[make_serializable_document()], total=1) - with ( app.test_request_context("/?sort=updated_at"), - patch( - "controllers.console.datasets.datasets_document.paginate_query", - return_value=pagination, - ), + patch("controllers.console.datasets.datasets_document.paginate_query", return_value=pagination), patch( "controllers.console.datasets.datasets_document.DocumentService.enrich_documents_with_summary_index_status", return_value=None, ), ): - response = method(api, tenant_id, user, "ds-1") - + response = method(api, MagicMock(), tenant_id, user, "ds-1") assert response["total"] == 1 def test_document_metadata_with_schema_validation(self, app: Flask, patch_tenant): """Test document metadata update with schema validation""" api = DocumentMetadataApi() - method = inspect.unwrap(api.put) + method = unwrap(api.put) user, tenant_id = patch_tenant - doc = MagicMock() - payload = { - "doc_type": "contract", - "doc_metadata": {"amount": 5000, "currency": "USD", "invalid_field": "x"}, - } - + payload = {"doc_type": "contract", "doc_metadata": {"amount": 5000, "currency": "USD", "invalid_field": "x"}} schema = {"amount": int, "currency": str} - + session = MagicMock() with ( app.test_request_context("/", json=payload), patch.object(api, "get_document", return_value=doc), @@ -1658,13 +1393,8 @@ class TestDocumentListAdvancedCases: "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"contract": schema}, ), - patch( - "controllers.console.datasets.datasets_document.db.session.commit", - return_value=None, - ), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 assert doc.doc_metadata == {"amount": 5000, "currency": "USD"} @@ -1672,9 +1402,8 @@ class TestDocumentListAdvancedCases: class TestDocumentIndexingEdgeCases: def test_document_indexing_with_extraction_setting(self, app: Flask, patch_tenant): api = DocumentIndexingEstimateApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user, tenant_id = patch_tenant - document = MagicMock( indexing_status=IndexingStatus.INDEXING, data_source_type=DataSourceType.UPLOAD_FILE, @@ -1683,25 +1412,17 @@ class TestDocumentIndexingEdgeCases: doc_form=IndexStructureType.PARAGRAPH_INDEX, dataset_process_rule=None, ) - upload_file = MagicMock() - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/"), patch.object(api, "get_document", return_value=document), - patch( - "controllers.console.datasets.datasets_document.db.session.scalar", - return_value=upload_file, - ), - patch( - "controllers.console.datasets.datasets_document.ExtractSetting", - return_value=MagicMock(), - ), + patch("controllers.console.datasets.datasets_document.ExtractSetting", return_value=MagicMock()), patch( "controllers.console.datasets.datasets_document.IndexingRunner.indexing_estimate", return_value=IndexingEstimate(total_segments=5, preview=[]), ), ): - response, status = method(api, tenant_id, user, "ds-1", "doc-1") - + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py index 6288fe363f5..fa6d73078b7 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py @@ -8,11 +8,12 @@ upload-file documents, and rejects unsupported or missing file cases. from __future__ import annotations import importlib -import inspect import sys from collections import UserDict +from inspect import unwrap from io import BytesIO from types import SimpleNamespace +from unittest.mock import MagicMock from zipfile import ZipFile import pytest @@ -198,8 +199,8 @@ def test_batch_download_zip_returns_send_file( json={"document_ids": ["11111111-1111-1111-1111-111111111111", "22222222-2222-2222-2222-222222222222"]}, ): api = datasets_document_module.DocumentBatchDownloadZipApi() - method = inspect.unwrap(api.post) - result = method(api, "tenant-123", _mock_user(), dataset_id="ds-1") + method = unwrap(api.post) + result = method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1") # Assert: we returned via send_file with correct mime type and attachment. assert result["_send_file_kwargs"]["mimetype"] == "application/zip" @@ -264,8 +265,8 @@ def test_batch_download_zip_response_is_openable_zip( json={"document_ids": ["33333333-3333-3333-3333-333333333333", "44444444-4444-4444-4444-444444444444"]}, ): api = datasets_document_module.DocumentBatchDownloadZipApi() - method = inspect.unwrap(api.post) - response = method(api, "tenant-123", _mock_user(), dataset_id="ds-1") + method = unwrap(api.post) + response = method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1") # Assert: response body is a valid ZIP and contains the expected entries. response.direct_passthrough = False @@ -308,9 +309,9 @@ def test_batch_download_zip_rejects_non_upload_file_document( json={"document_ids": ["55555555-5555-5555-5555-555555555555"]}, ): api = datasets_document_module.DocumentBatchDownloadZipApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) with pytest.raises(NotFound): - method(api, "tenant-123", _mock_user(), dataset_id="ds-1") + method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1") def test_document_download_returns_url_for_upload_file_document( @@ -331,8 +332,8 @@ def test_document_download_returns_url_for_upload_file_document( # Build a request context then call the resource method directly. with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() - method = inspect.unwrap(api.get) - result = method(api, "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") + method = unwrap(api.get) + result = method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") assert result == {"url": "https://example.com/signed"} @@ -354,9 +355,9 @@ def test_document_download_rejects_non_upload_file_document( with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) with pytest.raises(NotFound): - method(api, "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") + method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") def test_document_download_rejects_missing_upload_file_id( @@ -376,9 +377,9 @@ def test_document_download_rejects_missing_upload_file_id( with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) with pytest.raises(NotFound): - method(api, "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") + method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") def test_document_download_rejects_when_upload_file_record_missing( @@ -398,9 +399,9 @@ def test_document_download_rejects_when_upload_file_record_missing( with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) with pytest.raises(NotFound): - method(api, "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") + method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") def test_document_download_rejects_tenant_mismatch( @@ -420,6 +421,6 @@ def test_document_download_rejects_tenant_mismatch( with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) with pytest.raises(Forbidden): - method(api, "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") + method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py index 93fd25a610e..73c028da03b 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py @@ -1,6 +1,7 @@ -import inspect +from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from typing import Any, cast +from unittest.mock import MagicMock, call, patch import pytest from flask import Flask @@ -19,14 +20,10 @@ from controllers.console.datasets.datasets_segments import ( DatasetDocumentSegmentListApi, DatasetDocumentSegmentUpdateApi, ) -from controllers.console.datasets.error import ( - ChildChunkDeleteIndexError, - ChildChunkIndexingError, - InvalidActionError, -) +from controllers.console.datasets.error import ChildChunkDeleteIndexError, ChildChunkIndexingError, InvalidActionError from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.rag.index_processor.constant.index_type import IndexStructureType -from fields.segment_fields import segment_response_with_summary +from fields.segment_fields import segment_response_with_summary, segment_responses_with_summaries from libs.datetime_utils import naive_utc_now from models.dataset import ChildChunk, DocumentSegment from models.enums import SegmentStatus, SegmentType @@ -119,164 +116,138 @@ def _bind_dataset_document(dataset, document, dataset_id: str = "ds-1", document def test_segment_response_with_summary(): segment = _segment() - + session = MagicMock() with ( - patch("models.dataset.db.session.scalar", return_value=None), - patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), + patch.object(DocumentSegment, "get_child_chunks", autospec=True, return_value=[]) as get_child_chunks, + patch.object(DocumentSegment, "get_attachments", autospec=True, return_value=[]) as get_attachments, ): - result = segment_response_with_summary(segment, "summary") - + result = segment_response_with_summary(segment, "summary", session=session) assert result.summary == "summary" assert result.id == segment.id + get_child_chunks.assert_called_once_with(segment, session=session, include_full_doc=False) + get_attachments.assert_called_once_with(segment, session=session) + + +def test_segment_responses_with_summaries_reuses_caller_session(): + segments = [_segment(), _segment()] + segments[1].id = "seg-2" + session = MagicMock() + expected_responses = [MagicMock(), MagicMock()] + + with patch( + "fields.segment_fields.segment_response_with_summary", side_effect=expected_responses + ) as serialize_segment: + responses = segment_responses_with_summaries( + segments, + {"seg-1": "summary-1", "seg-2": None}, + session=session, + ) + + assert responses == expected_responses + assert serialize_segment.call_args_list == [ + call(segments[0], "summary-1", session=session), + call(segments[1], None, session=session), + ] class TestDatasetDocumentSegmentListApi: def test_get_success(self, app: Flask): api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) dataset = MagicMock() document = MagicMock() user = MagicMock() - segment = _segment() - + session = MagicMock() + session.get.return_value = None + session.execute.return_value.all.return_value = [] pagination = MagicMock() pagination.items = [segment] pagination.total = 1 pagination.pages = 1 - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.paginate_query", - return_value=pagination, - ), - patch( - "services.summary_index_service.SummaryIndexService.get_segments_summaries", - return_value={}, - ), - patch("models.dataset.db.session.scalar", return_value=None), - patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), + patch("controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination), + patch("services.summary_index_service.SummaryIndexService.get_segments_summaries", return_value={}), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") assert status == 200 def test_get_dataset_not_found(self, app: Flask): api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user = MagicMock() - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_get_permission_denied(self, app: Flask): api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) dataset = MagicMock() user = MagicMock() - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", side_effect=services.errors.account.NoPermissionError("no access"), ), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") class TestDatasetDocumentSegmentApi: def test_patch_success(self, app: Flask): api = DatasetDocumentSegmentApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() dataset.indexing_technique = "economy" - document = MagicMock() document.id = "doc-1" - with ( app.test_request_context("/?segment_id=s1&segment_id=s2"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.redis_client.get", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.redis_client.get", return_value=None), patch( "controllers.console.datasets.datasets_segments.SegmentService.update_segments_status", return_value=None, ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1", "enable") - + response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") assert status == 200 assert response["result"] == "success" def test_patch_document_indexing_in_progress(self, app: Flask): api = DatasetDocumentSegmentApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() dataset.indexing_technique = "economy" - document = MagicMock() document.id = "doc-1" - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, @@ -285,38 +256,23 @@ class TestDatasetDocumentSegmentApi: "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.redis_client.get", - return_value=b"running", - ), + patch("controllers.console.datasets.datasets_segments.redis_client.get", return_value=b"running"), ): with pytest.raises(InvalidActionError): - method(api, "tenant-1", user, "ds-1", "doc-1", "disable") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "disable") def test_patch_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( - indexing_technique="high_quality", - embedding_model_provider="openai", - embedding_model="text-embed", + indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock(id="doc-1") - with ( app.test_request_context("/?segment_id=s1"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, @@ -331,32 +287,20 @@ class TestDatasetDocumentSegmentApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1", user, "ds-1", "doc-1", "enable") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") def test_patch_provider_token_not_init(self, app: Flask): api = DatasetDocumentSegmentApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( - indexing_technique="high_quality", - embedding_model_provider="openai", - embedding_model="text-embed", + indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock(id="doc-1") - with ( app.test_request_context("/?segment_id=s1"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, @@ -371,39 +315,30 @@ class TestDatasetDocumentSegmentApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1", user, "ds-1", "doc-1", "enable") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") class TestDatasetDocumentSegmentAddApi: def test_post_success(self, app: Flask): api = DatasetDocumentSegmentAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "hello"} - user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() dataset.indexing_technique = "economy" - document = MagicMock() document.doc_form = IndexStructureType.PARAGRAPH_INDEX _bind_dataset_document(dataset, document) - segment = _segment() - + session = MagicMock() + session.get.return_value = None + session.execute.return_value.all.return_value = [] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, @@ -412,126 +347,84 @@ class TestDatasetDocumentSegmentAddApi: "controllers.console.datasets.datasets_segments.SegmentService.segment_create_args_validate", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.create_segment", - return_value=segment, - ), + patch("controllers.console.datasets.datasets_segments.SegmentService.create_segment", return_value=segment), patch( "controllers.console.datasets.datasets_segments.SummaryIndexService.get_segment_summary", return_value=None, ), - patch("models.dataset.db.session.scalar", return_value=None), - patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") assert status == 200 assert response["data"]["id"] == "seg-1" def test_post_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( - indexing_technique="high_quality", - embedding_model_provider="openai", - embedding_model="text-embed", + indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.ModelManager.get_model_instance", side_effect=LLMBadRequestError(), ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_post_provider_token_not_init(self, app: Flask): api = DatasetDocumentSegmentAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( - indexing_technique="high_quality", - embedding_model_provider="openai", - embedding_model="text-embed", + indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.ModelManager.get_model_instance", side_effect=ProviderTokenNotInitError("token missing"), ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") class TestDatasetDocumentSegmentUpdateApi: def test_patch_success(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) payload = {"content": "updated"} - user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() dataset.indexing_technique = "economy" - document = MagicMock() document.doc_form = IndexStructureType.PARAGRAPH_INDEX _bind_dataset_document(dataset, document) - segment = _segment() - + session = MagicMock() + session.get.return_value = None + session.execute.return_value.all.return_value = [] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -541,118 +434,82 @@ class TestDatasetDocumentSegmentUpdateApi: "controllers.console.datasets.datasets_segments.SegmentService.segment_create_args_validate", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.update_segment", - return_value=segment, - ), + patch("controllers.console.datasets.datasets_segments.SegmentService.update_segment", return_value=segment), patch( "controllers.console.datasets.datasets_segments.SummaryIndexService.get_segment_summary", return_value=None, ), - patch("models.dataset.db.session.scalar", return_value=None), - patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1", "seg-1") assert status == 200 assert "data" in response def test_patch_document_outside_dataset_is_not_found(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) payload = {"content": "updated"} user = MagicMock(is_dataset_editor=True) dataset = MagicMock(id="ds-1", tenant_id="tenant-1", indexing_technique="economy") document = MagicMock(id="doc-1", dataset_id="other-dataset", tenant_id="tenant-1") - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") def test_patch_segment_not_found(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) payload = {"content": "updated"} user = MagicMock(is_dataset_editor=True) dataset = MagicMock(indexing_technique="economy") document = MagicMock() _bind_dataset_document(dataset, document) - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=None, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=None ), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") def test_patch_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( - indexing_technique="high_quality", - embedding_model_provider="openai", - embedding_model="text-embed", + indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, @@ -667,189 +524,145 @@ class TestDatasetDocumentSegmentUpdateApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") class TestDatasetDocumentSegmentBatchImportApi: def test_post_success(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock(spec=UploadFile) upload_file.name = "test.csv" user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=upload_file, - ), - patch( - "controllers.console.datasets.datasets_segments.redis_client.setnx", - return_value=True, + "controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=MagicMock() ), + patch("controllers.console.datasets.datasets_segments.redis_client.setnx", return_value=True), patch( "controllers.console.datasets.datasets_segments.batch_create_segment_to_index_task.delay", return_value=None, ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") assert status == 200 assert response["job_status"] == "waiting" def test_post_dataset_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = None with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_post_document_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = None with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=None, + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_post_upload_file_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = None with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=None, + "controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=MagicMock() ), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, session, "tenant-1", user, "ds-1", "doc-1") def test_post_invalid_file_type(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock() upload_file.name = "test.txt" user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=upload_file, + "controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=MagicMock() ), ): with pytest.raises(ValueError): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, session, "tenant-1", user, "ds-1", "doc-1") def test_post_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock() upload_file.name = "test.csv" user = MagicMock(id="u1") - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=MagicMock() ), patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=upload_file, - ), - patch( - "controllers.console.datasets.datasets_segments.redis_client.setnx", - side_effect=Exception("redis down"), + "controllers.console.datasets.datasets_segments.redis_client.setnx", side_effect=Exception("redis down") ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") assert status == 500 assert "error" in response def test_get_job_not_found_in_redis(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.redis_client.get", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.redis_client.get", return_value=None), ): with pytest.raises(ValueError): method(api, job_id="job-1") @@ -857,33 +670,25 @@ class TestDatasetDocumentSegmentBatchImportApi: class TestChildChunkAddApi: def test_patch_documents_batch_update_payload(self): - api_doc = getattr(ChildChunkAddApi.patch, "__apidoc__") # noqa: B009 + patch_method = cast(Any, ChildChunkAddApi.patch) + api_doc = cast(dict[str, Any], patch_method.__apidoc__) expected_model = ChildChunkBatchUpdatePayload.__name__ - assert [model.name for model in api_doc["expect"]] == [expected_model] def test_get_uses_default_pagination_for_malformed_ints(self, app: Flask): api = ChildChunkAddApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) dataset = MagicMock() document = _bind_dataset_document(dataset, MagicMock()) pagination = MagicMock(items=[], total=0, pages=0) - with ( app.test_request_context("/?page=bad&limit="), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_model_setting", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=MagicMock(), @@ -893,44 +698,33 @@ class TestChildChunkAddApi: return_value=pagination, ) as get_child_chunks, ): - response, status = method(api, "tenant-1", "ds-1", "doc-1", "seg-1") - + response, status = method(api, MagicMock(), "tenant-1", "ds-1", "doc-1", "seg-1") assert status == 200 assert response["page"] == 1 assert response["limit"] == 20 - get_child_chunks.assert_called_once_with("seg-1", "doc-1", "ds-1", 1, 20, None) + session = get_child_chunks.call_args.kwargs["session"] + assert isinstance(session, MagicMock) + assert get_child_chunks.call_args.args == ("seg-1", "doc-1", "ds-1", 1, 20, None) def test_post_success(self, app: Flask): api = ChildChunkAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "child"} - user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() dataset.indexing_technique = "economy" - document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = _child_chunk() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -941,38 +735,26 @@ class TestChildChunkAddApi: return_value=child_chunk, ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") - + response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") assert status == 200 assert response["data"]["id"] == "cc-1" def test_post_child_chunk_indexing_error(self, app: Flask): api = ChildChunkAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "child"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock(indexing_technique="economy") document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -984,65 +766,47 @@ class TestChildChunkAddApi: ), ): with pytest.raises(ChildChunkIndexingError): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") def test_post_permission_denied(self, app: Flask): api = ChildChunkAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"content": "child"} user = MagicMock(is_dataset_editor=True) dataset = MagicMock(indexing_technique="economy") document = MagicMock() _bind_dataset_document(dataset, document) - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", side_effect=services.errors.account.NoPermissionError("no access"), ), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") class TestChildChunkUpdateApi: def test_delete_success(self, app: Flask): api = ChildChunkUpdateApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) user = MagicMock() user.is_dataset_editor = True - dataset = MagicMock() document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = MagicMock() - with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", @@ -1053,40 +817,28 @@ class TestChildChunkUpdateApi: return_value=None, ), patch( - "controllers.console.datasets.datasets_segments.SegmentService.delete_child_chunk", - return_value=None, + "controllers.console.datasets.datasets_segments.SegmentService.delete_child_chunk", return_value=None ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") - + response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") assert status == 204 assert response == "" def test_delete_child_chunk_index_error(self, app: Flask): api = ChildChunkUpdateApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() child_chunk = MagicMock() - with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", @@ -1102,31 +854,22 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(ChildChunkDeleteIndexError): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") def test_delete_child_chunk_not_found(self, app: Flask): api = ChildChunkUpdateApi() - method = inspect.unwrap(api.delete) - + method = unwrap(api.delete) user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() - with ( app.test_request_context("/"), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", @@ -1138,33 +881,24 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") def test_patch_child_chunk_not_found(self, app: Flask): api = ChildChunkUpdateApi() - method = inspect.unwrap(api.patch) - + method = unwrap(api.patch) payload = {"content": "updated child"} user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() _bind_dataset_document(dataset, document) segment = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=segment, + "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", return_value=segment ), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunk_by_segment_ref", @@ -1176,97 +910,71 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") class TestSegmentListAdvancedCases: def test_segment_list_with_keyword_filter(self, app: Flask): api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) dataset = MagicMock() document = MagicMock() user = MagicMock() - segment = _segment() - + session = MagicMock() + session.get.return_value = None + session.execute.return_value.all.return_value = [] pagination = MagicMock(items=[segment], total=1, pages=1) - with ( app.test_request_context("/?keyword=test"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.paginate_query", - return_value=pagination, - ), - patch( - "services.summary_index_service.SummaryIndexService.get_segments_summaries", - return_value={}, - ), - patch("models.dataset.db.session.scalar", return_value=None), - patch("models.dataset.db.session.execute", return_value=MagicMock(all=MagicMock(return_value=[]))), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), + patch("controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination), + patch("services.summary_index_service.SummaryIndexService.get_segments_summaries", return_value={}), ): - result = method(api, "tenant-1", user, "ds-1", "doc-1") - + result = method(api, session, "tenant-1", user, "ds-1", "doc-1") if isinstance(result, tuple): response, status = result else: - response, status = result, 200 - + response, status = (result, 200) assert status == 200 assert response["total"] == 1 def test_segment_list_postgres_keyword_filter_handles_scalar_keywords(self, app: Flask): api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) dataset = MagicMock() document = MagicMock() user = MagicMock() pagination = MagicMock(items=[], total=0, pages=0) - with ( app.test_request_context("/?keyword=test"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.dify_config", SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME="postgresql"), ), patch( - "controllers.console.datasets.datasets_segments.paginate_query", - return_value=pagination, + "controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination ) as paginate_mock, ): method( api, + MagicMock(), "11111111-1111-1111-1111-111111111111", user, "22222222-2222-2222-2222-222222222222", "33333333-3333-3333-3333-333333333333", ) - query = paginate_mock.call_args.args[0] sql = str(query.compile(compile_kwargs={"literal_binds": True})) assert "jsonb_array_elements_text(CASE" in sql @@ -1275,14 +983,12 @@ class TestSegmentListAdvancedCases: def test_segment_list_permission_denied(self, app: Flask): """Test segment list with permission denied""" api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user = MagicMock() - with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=MagicMock(), + "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() ), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", @@ -1290,48 +996,35 @@ class TestSegmentListAdvancedCases: ), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_segment_list_dataset_not_found(self, app: Flask): """Test segment list with dataset not found""" api = DatasetDocumentSegmentListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user = MagicMock() - with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") class TestSegmentOperationCases: def test_segment_add_with_provider_token_error(self, app: Flask): """Test segment add with provider token not initialized""" api = DatasetDocumentSegmentAddApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() - payload = {"content": "new content", "answer": None} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, @@ -1342,91 +1035,60 @@ class TestSegmentOperationCases: ), ): with pytest.raises(ProviderTokenNotInitError): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_batch_import_with_document_not_found(self, app: Flask): """Test batch import with document not found""" api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = MagicMock(is_dataset_editor=True) dataset = MagicMock() - payload = {"upload_file_id": "file-1"} - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=None), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") def test_batch_import_with_invalid_file(self, app: Flask): """Test batch import with invalid file type""" api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() - upload_file = None # File not found - + upload_file = None payload = {"upload_file_id": "file-1"} - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=upload_file, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), ): with pytest.raises(NotFound): - method(api, "tenant-1", user, "ds-1", "doc-1") + method(api, session, "tenant-1", user, "ds-1", "doc-1") def test_batch_import_with_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() upload_file = MagicMock(spec=UploadFile, extension="csv", id="file-1") upload_file.name = "test.csv" - payload = {"upload_file_id": "file-1"} - + session = MagicMock() + session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "controllers.console.datasets.datasets_segments.DocumentService.get_document", - return_value=document, - ), - patch( - "controllers.console.datasets.datasets_segments.db.session.scalar", - return_value=upload_file, - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), + patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", return_value=None, @@ -1436,21 +1098,16 @@ class TestSegmentOperationCases: side_effect=Exception("Task failed"), ), ): - response, status = method(api, "tenant-1", user, "ds-1", "doc-1") - + response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") assert status == 500 assert "error" in response def test_batch_import_get_job_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) with ( app.test_request_context("/?job_id=invalid-job"), - patch( - "controllers.console.datasets.datasets_segments.redis_client.get", - return_value=None, - ), + patch("controllers.console.datasets.datasets_segments.redis_client.get", return_value=None), ): with pytest.raises(ValueError): method(api, "invalid-job") 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 6cdd8c84ddd..5addfd66e7f 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -62,10 +62,12 @@ def _external_api_dict(api_id: str = "api-1") -> dict: def _external_api_object(api_id: str = "api-1") -> SimpleNamespace: payload = _external_api_dict(api_id) + dataset_bindings = [SimpleNamespace(**binding) for binding in payload["dataset_bindings"]] return SimpleNamespace( **{ **payload, - "dataset_bindings": [SimpleNamespace(**binding) for binding in payload["dataset_bindings"]], + "dataset_bindings": dataset_bindings, + "get_dataset_bindings": MagicMock(return_value=dataset_bindings), } ) @@ -161,6 +163,7 @@ class TestExternalApiTemplateListApi: method = inspect.unwrap(api.get) api_item = _external_api_object("api-1") + session = MagicMock() with ( app.test_request_context("/?page=2&limit=1&keyword=vector"), @@ -170,7 +173,7 @@ class TestExternalApiTemplateListApi: return_value=([api_item], 3), ) as get_external_knowledge_apis, ): - resp, status = method(api, "tenant-1") + resp, status = method(api, session, "tenant-1") assert status == 200 assert resp == { @@ -180,7 +183,8 @@ class TestExternalApiTemplateListApi: "total": 3, "page": 2, } - get_external_knowledge_apis.assert_called_once_with(2, 1, "tenant-1", "vector") + api_item.get_dataset_bindings.assert_called_once_with(session=session) + get_external_knowledge_apis.assert_called_once_with(2, 1, "tenant-1", "vector", session=ANY) def test_post_success_uses_validated_payload_and_returns_template(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() @@ -212,6 +216,7 @@ class TestExternalApiTemplateListApi: assert status == 201 assert resp == _external_api_dict("api-created") + created.get_dataset_bindings.assert_called_once_with(session=session) validate_api_list.assert_called_once_with(payload["settings"]) create_external_knowledge_api.assert_called_once_with( tenant_id="tenant-1", @@ -274,6 +279,7 @@ class TestExternalApiTemplateApi: assert status == 200 assert resp == _external_api_dict("api-detail") + template.get_dataset_bindings.assert_called_once_with(session=session) get_external_knowledge_api.assert_called_once_with( external_knowledge_api_id="api-detail", tenant_id="tenant-1", session=session ) @@ -322,6 +328,7 @@ class TestExternalApiTemplateApi: assert status == 200 assert resp == _external_api_dict("api-updated") + updated.get_dataset_bindings.assert_called_once_with(session=session) validate_api_list.assert_called_once_with(payload["settings"]) update_external_knowledge_api.assert_called_once_with( tenant_id="tenant-1", @@ -391,6 +398,10 @@ class TestExternalDatasetCreateApi: "create_external_dataset", return_value=dataset, ) as create_external_dataset, + patch( + "controllers.console.datasets.external.dataset_detail_response_source", + return_value=dataset, + ) as dataset_response_source, ): session = MagicMock() resp, status = method(api, session, "tenant-1", current_user) @@ -403,6 +414,7 @@ class TestExternalDatasetCreateApi: args=payload, session=session, ) + dataset_response_source.assert_called_once_with(dataset, session=session) def test_create_forbidden(self, app: Flask, current_user: Account): current_user.role = TenantAccountRole.NORMAL @@ -578,7 +590,7 @@ class TestExternalApiTemplateListApiAdvanced: return_value=(templates, 25), ) as get_external_knowledge_apis, ): - resp, status = method(api, "tenant-1") + resp, status = method(api, MagicMock(), "tenant-1") assert status == 200 assert resp == { @@ -588,7 +600,7 @@ class TestExternalApiTemplateListApiAdvanced: "total": 25, "page": 2, } - get_external_knowledge_apis.assert_called_once_with(2, 3, "tenant-1", None) + get_external_knowledge_apis.assert_called_once_with(2, 3, "tenant-1", None, session=ANY) class TestExternalDatasetCreateApiAdvanced: diff --git a/api/tests/unit_tests/controllers/console/datasets/test_hit_testing_base.py b/api/tests/unit_tests/controllers/console/datasets/test_hit_testing_base.py index b635fcf3efd..39ddd2ec787 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_hit_testing_base.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_hit_testing_base.py @@ -94,7 +94,7 @@ class TestGetAndValidateDataset: "check_dataset_permission", ), ): - result = DatasetsHitTestingBase.get_and_validate_dataset("dataset-1", account, "tenant-1") + result = DatasetsHitTestingBase.get_and_validate_dataset(Mock(), "dataset-1", account, "tenant-1") assert result == dataset @@ -105,7 +105,7 @@ class TestGetAndValidateDataset: return_value=None, ): with pytest.raises(NotFound, match="Dataset not found"): - DatasetsHitTestingBase.get_and_validate_dataset("dataset-1", account, "tenant-1") + DatasetsHitTestingBase.get_and_validate_dataset(Mock(), "dataset-1", account, "tenant-1") def test_permission_denied(self, dataset, account): with ( @@ -121,7 +121,7 @@ class TestGetAndValidateDataset: ), ): with pytest.raises(Forbidden, match="no access"): - DatasetsHitTestingBase.get_and_validate_dataset("dataset-1", account, "tenant-1") + DatasetsHitTestingBase.get_and_validate_dataset(Mock(), "dataset-1", account, "tenant-1") class TestHitTestingArgsCheck: diff --git a/api/tests/unit_tests/controllers/console/datasets/test_metadata.py b/api/tests/unit_tests/controllers/console/datasets/test_metadata.py index 785c0ac09f2..00f45d488a4 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_metadata.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_metadata.py @@ -17,10 +17,7 @@ from controllers.console.datasets.metadata import ( ) from models.account import Account from services.dataset_service import DatasetService -from services.entities.knowledge_entities.knowledge_entities import ( - MetadataArgs, - MetadataOperationData, -) +from services.entities.knowledge_entities.knowledge_entities import MetadataArgs, MetadataOperationData from services.metadata_service import MetadataService @@ -58,61 +55,28 @@ def metadata_id(): @pytest.fixture(autouse=True) def bypass_decorators(mocker: MockerFixture): """Bypass setup/login/license decorators.""" - mocker.patch( - "controllers.console.datasets.metadata.setup_required", - lambda f: f, - ) - mocker.patch( - "controllers.console.datasets.metadata.login_required", - lambda f: f, - ) - mocker.patch( - "controllers.console.datasets.metadata.account_initialization_required", - lambda f: f, - ) - mocker.patch( - "controllers.console.datasets.metadata.enterprise_license_required", - lambda f: f, - ) + mocker.patch("controllers.console.datasets.metadata.setup_required", lambda f: f) + mocker.patch("controllers.console.datasets.metadata.login_required", lambda f: f) + mocker.patch("controllers.console.datasets.metadata.account_initialization_required", lambda f: f) + mocker.patch("controllers.console.datasets.metadata.enterprise_license_required", lambda f: f) class TestDatasetMetadataCreateApi: def test_create_metadata_success(self, app: Flask, current_user, dataset, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.post) - payload = {"name": "author"} - with ( app.test_request_context("/"), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch.object(MetadataArgs, "model_validate", return_value=MagicMock()), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch.object( - MetadataArgs, - "model_validate", - return_value=MagicMock(), - ), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - ), - patch.object( - MetadataService, - "create_metadata", - return_value={"id": "m1", "type": "string", "name": "author"}, + MetadataService, "create_metadata", return_value={"id": "m1", "type": "string", "name": "author"} ), ): - result, status = method(api, "tenant-1", current_user, dataset_id) - + result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id) assert status == 201 assert result["type"] == "string" assert result["name"] == "author" @@ -120,47 +84,24 @@ class TestDatasetMetadataCreateApi: def test_create_metadata_dataset_not_found(self, app: Flask, current_user, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.post) - - valid_payload = { - "type": "string", - "name": "author", - } - + valid_payload = {"type": "string", "name": "author"} with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=valid_payload, - ), - patch.object( - MetadataArgs, - "model_validate", - return_value=MagicMock(), - ), - patch.object( - DatasetService, - "get_dataset", - return_value=None, - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=valid_payload), + patch.object(MetadataArgs, "model_validate", return_value=MagicMock()), + patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, "tenant-1", current_user, dataset_id) + method(api, MagicMock(), "tenant-1", current_user, dataset_id) class TestDatasetMetadataGetApi: def test_get_metadata_success(self, app: Flask, dataset, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.get) - with ( app.test_request_context("/"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object( MetadataService, "get_dataset_metadatas", @@ -170,8 +111,7 @@ class TestDatasetMetadataGetApi: }, ), ): - result, status = method(api, dataset_id) - + result, status = method(api, MagicMock(), dataset_id) assert status == 200 assert result["doc_metadata"] == [{"id": "m1", "name": "author", "type": "string", "count": 0}] assert result["built_in_field_enabled"] is False @@ -179,51 +119,28 @@ class TestDatasetMetadataGetApi: def test_get_metadata_dataset_not_found(self, app: Flask, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.get) - - with ( - app.test_request_context("/"), - patch.object( - DatasetService, - "get_dataset", - return_value=None, - ), - ): + with app.test_request_context("/"), patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(NotFound): - method(api, dataset_id) + method(api, MagicMock(), dataset_id) class TestDatasetMetadataApi: def test_update_metadata_success(self, app: Flask, current_user, dataset, dataset_id, metadata_id): api = DatasetMetadataApi() method = unwrap(api.patch) - payload = {"name": "updated-name"} - with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), patch.object( MetadataService, "update_metadata_name", return_value={"id": "m1", "type": "string", "name": "updated-name"}, ), ): - result, status = method(api, "tenant-1", current_user, dataset_id, metadata_id) - + result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id, metadata_id) assert status == 200 assert result["type"] == "string" assert result["name"] == "updated-name" @@ -231,25 +148,13 @@ class TestDatasetMetadataApi: def test_delete_metadata_success(self, app: Flask, current_user, dataset, dataset_id, metadata_id): api = DatasetMetadataApi() method = unwrap(api.delete) - with ( app.test_request_context("/"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - ), - patch.object( - MetadataService, - "delete_metadata", - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + patch.object(MetadataService, "delete_metadata"), ): - result, status = method(api, current_user, dataset_id, metadata_id) - + result, status = method(api, MagicMock(), current_user, dataset_id, metadata_id) assert status == 204 assert result == "" @@ -258,50 +163,30 @@ class TestDatasetMetadataBuiltInFieldApi: def test_get_built_in_fields(self, app: Flask): api = DatasetMetadataBuiltInFieldApi() method = unwrap(api.get) - with ( app.test_request_context("/"), patch.object( MetadataService, "get_built_in_fields", - return_value=[ - {"name": "document_name", "type": "string"}, - {"name": "source", "type": "string"}, - ], + return_value=[{"name": "document_name", "type": "string"}, {"name": "source", "type": "string"}], ), ): result, status = method(api) - assert status == 200 - assert result["fields"] == [ - {"name": "document_name", "type": "string"}, - {"name": "source", "type": "string"}, - ] + assert result["fields"] == [{"name": "document_name", "type": "string"}, {"name": "source", "type": "string"}] class TestDatasetMetadataBuiltInFieldActionApi: def test_enable_built_in_field(self, app: Flask, current_user, dataset, dataset_id): api = DatasetMetadataBuiltInFieldActionApi() method = unwrap(api.post) - with ( app.test_request_context("/"), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - ), - patch.object( - MetadataService, - "enable_built_in_field", - ), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + patch.object(MetadataService, "enable_built_in_field"), ): - result, status = method(api, current_user, dataset_id, "enable") - + result, status = method(api, MagicMock(), current_user, dataset_id, "enable") assert status == 204 assert result == "" @@ -310,37 +195,15 @@ class TestDocumentMetadataEditApi: def test_update_document_metadata_success(self, app: Flask, current_user, dataset, dataset_id): api = DocumentMetadataEditApi() method = unwrap(api.post) - payload = {"operation": "add", "metadata": {}} - with ( app.test_request_context("/"), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - patch.object( - DatasetService, - "get_dataset", - return_value=dataset, - ), - patch.object( - DatasetService, - "check_dataset_permission", - ), - patch.object( - MetadataOperationData, - "model_validate", - return_value=MagicMock(), - ), - patch.object( - MetadataService, - "update_documents_metadata", - ), + patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), + patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + patch.object(MetadataOperationData, "model_validate", return_value=MagicMock()), + patch.object(MetadataService, "update_documents_metadata"), ): - result, status = method(api, current_user, dataset_id) - + result, status = method(api, MagicMock(), current_user, dataset_id) assert status == 204 assert result == "" diff --git a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py index 2cfa938af80..80ee65a7094 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py @@ -2,9 +2,10 @@ from unittest.mock import Mock import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session from controllers.console.datasets.error import PipelineNotFoundError -from controllers.console.datasets.wraps import get_rag_pipeline +from controllers.console.datasets.wraps import get_rag_pipeline, load_rag_pipeline from models.dataset import Pipeline @@ -27,13 +28,15 @@ class TestGetRagPipeline: return_value=(Mock(), "tenant-1"), ) - mocker.patch( - "controllers.console.datasets.wraps.db.session.scalar", + session_factory = mocker.patch("controllers.console.datasets.wraps.db.session") + get_pipeline_by_id = mocker.patch( + "controllers.console.datasets.wraps.RagPipelineService.get_pipeline_by_id", return_value=None, ) with pytest.raises(PipelineNotFoundError): dummy_view(pipeline_id="pipeline-1") + get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value) def test_pipeline_found_and_injected(self, mocker: MockerFixture): pipeline = Mock(spec=Pipeline) @@ -49,14 +52,34 @@ class TestGetRagPipeline: return_value=(Mock(), "tenant-1"), ) - mocker.patch( - "controllers.console.datasets.wraps.db.session.scalar", + session_factory = mocker.patch("controllers.console.datasets.wraps.db.session") + get_pipeline_by_id = mocker.patch( + "controllers.console.datasets.wraps.RagPipelineService.get_pipeline_by_id", return_value=pipeline, ) result = dummy_view(pipeline_id="pipeline-1") assert result is pipeline + get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value) + + def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture): + pipeline = Mock(spec=Pipeline) + session = Mock(spec=Session) + + mocker.patch( + "controllers.console.datasets.wraps.current_account_with_tenant", + return_value=(Mock(), "tenant-1"), + ) + get_pipeline_by_id = mocker.patch( + "controllers.console.datasets.wraps.RagPipelineService.get_pipeline_by_id", + return_value=pipeline, + ) + + result = load_rag_pipeline(session, "pipeline-1") + + assert result is pipeline + get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session) def test_pipeline_id_removed_from_kwargs(self, mocker: MockerFixture): pipeline = Mock(spec=Pipeline) @@ -71,8 +94,9 @@ class TestGetRagPipeline: return_value=(Mock(), "tenant-1"), ) + session_factory = mocker.patch("controllers.console.datasets.wraps.db.session") mocker.patch( - "controllers.console.datasets.wraps.db.session.scalar", + "controllers.console.datasets.wraps.RagPipelineService.get_pipeline_by_id", return_value=pipeline, ) @@ -92,15 +116,13 @@ class TestGetRagPipeline: return_value=(Mock(), "tenant-1"), ) - mock_scalar = mocker.patch( - "controllers.console.datasets.wraps.db.session.scalar", + session_factory = mocker.patch("controllers.console.datasets.wraps.db.session") + get_pipeline_by_id = mocker.patch( + "controllers.console.datasets.wraps.RagPipelineService.get_pipeline_by_id", return_value=pipeline, ) result = dummy_view(pipeline_id=123) assert result is pipeline - # Verify the pipeline_id was cast to string in the where clause - stmt = mock_scalar.call_args[0][0] - where_clauses = stmt.whereclause.clauses - assert where_clauses[0].right.value == "123" + get_pipeline_by_id.assert_called_once_with("123", "tenant-1", session=session_factory.return_value) diff --git a/api/tests/unit_tests/controllers/console/explore/test_audio.py b/api/tests/unit_tests/controllers/console/explore/test_audio.py index 387cb97f61b..e3c856926ca 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_audio.py +++ b/api/tests/unit_tests/controllers/console/explore/test_audio.py @@ -45,6 +45,7 @@ def installed_app(): app.app = MagicMock() app.app.id = "app-1" app.app.tenant_id = "tenant-1" + app.app_with_session.return_value = app.app return app diff --git a/api/tests/unit_tests/controllers/console/explore/test_completion.py b/api/tests/unit_tests/controllers/console/explore/test_completion.py index 64c87043889..2e2e0f5e7e8 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_completion.py +++ b/api/tests/unit_tests/controllers/console/explore/test_completion.py @@ -26,12 +26,19 @@ def user(): @pytest.fixture def completion_app(): - return MagicMock(app=MagicMock(mode=AppMode.COMPLETION)) + return _installed_app(AppMode.COMPLETION) @pytest.fixture def chat_app(): - return MagicMock(app=MagicMock(mode=AppMode.CHAT)) + return _installed_app(AppMode.CHAT) + + +def _installed_app(mode: AppMode): + app = MagicMock(mode=mode) + installed_app = MagicMock(app=app) + installed_app.app_with_session.return_value = app + return installed_app @pytest.fixture @@ -76,7 +83,7 @@ class TestCompletionApi: api = completion_module.CompletionApi() method = unwrap(api.post) - installed_app = MagicMock(app=MagicMock(mode=AppMode.CHAT)) + installed_app = _installed_app(AppMode.CHAT) with pytest.raises(NotCompletionAppError): method(api, MagicMock(), user, installed_app) @@ -216,7 +223,7 @@ class TestCompletionStopApi: method = unwrap(api.post) with patch.object(completion_module.AppTaskService, "stop_task"): - resp, status = method(api, "u1", completion_app, "task-1") + resp, status = method(api, MagicMock(), "u1", completion_app, "task-1") assert status == 200 assert resp == {"result": "success"} @@ -225,10 +232,10 @@ class TestCompletionStopApi: api = completion_module.CompletionStopApi() method = unwrap(api.post) - installed_app = MagicMock(app=MagicMock(mode=AppMode.CHAT)) + installed_app = _installed_app(AppMode.CHAT) with pytest.raises(NotCompletionAppError): - method(api, "u1", installed_app, "task") + method(api, MagicMock(), "u1", installed_app, "task") class TestChatApi: @@ -258,7 +265,7 @@ class TestChatApi: api = completion_module.ChatApi() method = unwrap(api.post) - installed_app = MagicMock(app=MagicMock(mode=AppMode.COMPLETION)) + installed_app = _installed_app(AppMode.COMPLETION) with pytest.raises(NotChatAppError): method(api, MagicMock(), user, installed_app) @@ -315,13 +322,18 @@ class TestChatApi: # A nonexistent conversation_id must fail fast as 404, before the streaming # generator is created. Previously the lookup only ran inside the generator, # so an invalid id surfaced as a hang instead of a clean error. + conversation_id = str(uuid.uuid4()) payload_patch = patch.object( type(completion_module.console_ns), "payload", new_callable=PropertyMock, - return_value={"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())}, + return_value={"inputs": {}, "query": "hi", "conversation_id": conversation_id}, ) generate_mock = MagicMock(return_value={"ok": True}) + get_conversation_mock = MagicMock( + side_effect=completion_module.services.errors.conversation.ConversationNotExistsError() + ) + session = MagicMock() api = completion_module.ChatApi() method = unwrap(api.post) @@ -332,15 +344,16 @@ class TestChatApi: patch.object( completion_module.ConversationService, "get_conversation", - side_effect=completion_module.services.errors.conversation.ConversationNotExistsError(), + get_conversation_mock, ), patch.object(completion_module.AppGenerateService, "generate", generate_mock), ): with pytest.raises(completion_module.NotFound): - method(api, MagicMock(), user, chat_app) + method(api, session, user, chat_app) # The lookup must run before generation, so the generator is never started. generate_mock.assert_not_called() + assert get_conversation_mock.call_args.kwargs["session"] is session def test_app_unavailable_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -444,7 +457,7 @@ class TestChatStopApi: api = completion_module.ChatStopApi() method = unwrap(api.post) with patch.object(completion_module.AppTaskService, "stop_task"): - resp, status = method(api, "u1", chat_app, "task-1") + resp, status = method(api, MagicMock(), "u1", chat_app, "task-1") assert status == 200 assert resp == {"result": "success"} @@ -453,7 +466,7 @@ class TestChatStopApi: api = completion_module.ChatStopApi() method = unwrap(api.post) - installed_app = MagicMock(app=MagicMock(mode=AppMode.COMPLETION)) + installed_app = _installed_app(AppMode.COMPLETION) with pytest.raises(NotChatAppError): - method(api, "u1", installed_app, "task") + method(api, MagicMock(), "u1", installed_app, "task") diff --git a/api/tests/unit_tests/controllers/console/explore/test_message.py b/api/tests/unit_tests/controllers/console/explore/test_message.py index 9a93fe5626f..a66869488b8 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_message.py @@ -49,6 +49,8 @@ def make_message(): msg.query = "hello" msg.re_sign_file_url_answer = "" msg.user_feedback = MagicMock(rating=None) + msg.inputs_with_session.return_value = msg.inputs + msg.user_feedback_with_session.return_value = msg.user_feedback msg.total_price = None msg.currency = None msg.status = "normal" @@ -56,13 +58,20 @@ def make_message(): return msg +def make_installed_app(mode: str | None = None): + app_model = MagicMock(mode=mode) + installed_app = MagicMock() + installed_app.app = app_model + installed_app.app_with_session.return_value = app_model + return installed_app + + class TestMessageListApi: def test_get_success(self, app: Flask): api = module.MessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") pagination = MagicMock( limit=20, @@ -91,8 +100,7 @@ class TestMessageListApi: api = module.MessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with pytest.raises(NotChatAppError): method(MagicMock(), installed_app) @@ -101,8 +109,7 @@ class TestMessageListApi: api = module.MessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( app.test_request_context( @@ -122,8 +129,7 @@ class TestMessageListApi: api = module.MessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( app.test_request_context( @@ -145,8 +151,7 @@ class TestMessageFeedbackApi: api = module.MessageFeedbackApi() method = unwrap(api.post) - installed_app = MagicMock() - installed_app.app = MagicMock() + installed_app = make_installed_app() with ( app.test_request_context("/", json={"rating": "like"}), @@ -163,8 +168,7 @@ class TestMessageFeedbackApi: api = module.MessageFeedbackApi() method = unwrap(api.post) - installed_app = MagicMock() - installed_app.app = MagicMock() + installed_app = make_installed_app() with ( app.test_request_context("/", json={}), @@ -183,8 +187,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -210,8 +213,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with pytest.raises(NotCompletionAppError): method(MagicMock(), MagicMock(), installed_app, "mid") @@ -220,8 +222,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -241,8 +242,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -262,8 +262,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -283,8 +282,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -304,8 +302,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -325,8 +322,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -346,8 +342,7 @@ class TestMessageMoreLikeThisApi: api = module.MessageMoreLikeThisApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with ( app.test_request_context( @@ -369,8 +364,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -387,8 +381,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") with pytest.raises(NotChatAppError): method(MagicMock(), installed_app, "mid") @@ -397,8 +390,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -414,8 +406,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -431,8 +422,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -448,8 +438,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -465,8 +454,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -482,8 +470,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -499,8 +486,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( @@ -516,8 +502,7 @@ class TestMessageSuggestedQuestionApi: api = module.MessageSuggestedQuestionApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with ( patch.object( diff --git a/api/tests/unit_tests/controllers/console/explore/test_parameter.py b/api/tests/unit_tests/controllers/console/explore/test_parameter.py index 9ee9403baaf..5a94027a3c9 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_parameter.py +++ b/api/tests/unit_tests/controllers/console/explore/test_parameter.py @@ -8,12 +8,18 @@ from controllers.console.app.error import AppUnavailableError from models.model import AppMode +def _installed_app(app): + installed_app = MagicMock(app=app) + installed_app.app_with_session.return_value = app + return installed_app + + class TestAppParameterApi: def test_get_app_none(self): api = module.AppParameterApi() method = unwrap(api.get) - installed_app = MagicMock(app=None) + installed_app = _installed_app(None) with pytest.raises(AppUnavailableError): method(installed_app) @@ -30,8 +36,9 @@ class TestAppParameterApi: mode=AppMode.ADVANCED_CHAT, workflow=workflow, ) + app.workflow_with_session.return_value = workflow - installed_app = MagicMock(app=app) + installed_app = _installed_app(app) with ( patch.object( @@ -57,8 +64,9 @@ class TestAppParameterApi: mode=AppMode.ADVANCED_CHAT, workflow=None, ) + app.workflow_with_session.return_value = None - installed_app = MagicMock(app=app) + installed_app = _installed_app(app) with pytest.raises(AppUnavailableError): method(installed_app) @@ -74,8 +82,10 @@ class TestAppParameterApi: mode=AppMode.CHAT, app_model_config=app_model_config, ) + app.id = "app-1" + app.app_model_config_with_session.return_value = app_model_config - installed_app = MagicMock(app=app) + installed_app = _installed_app(app) with ( patch.object( @@ -88,6 +98,7 @@ class TestAppParameterApi: "model_validate", return_value=MagicMock(model_dump=lambda **_: {"ok": True}), ), + patch.object(module, "load_annotation_reply_config", return_value=None), ): result = method(installed_app) @@ -101,8 +112,9 @@ class TestAppParameterApi: mode=AppMode.CHAT, app_model_config=None, ) + app.app_model_config_with_session.return_value = None - installed_app = MagicMock(app=app) + installed_app = _installed_app(app) with pytest.raises(AppUnavailableError): method(installed_app) @@ -114,7 +126,7 @@ class TestExploreAppMetaApi: method = unwrap(api.get) app = MagicMock() - installed_app = MagicMock(app=app) + installed_app = _installed_app(app) with patch.object( module.AppService, @@ -129,7 +141,7 @@ class TestExploreAppMetaApi: api = module.ExploreAppMetaApi() method = unwrap(api.get) - installed_app = MagicMock(app=None) + installed_app = _installed_app(None) with pytest.raises(ValueError): method(installed_app) 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 f210d0d5d04..ae36d69f7dd 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 @@ -20,10 +20,20 @@ def make_saved_message(): msg.query = "hello" msg.answer = "world" msg.user_feedback = MagicMock(rating="like") + msg.inputs_with_session.return_value = msg.inputs + msg.user_feedback_with_session.return_value = msg.user_feedback msg.created_at = None return msg +def make_installed_app(mode: str): + app_model = MagicMock(mode=mode) + installed_app = MagicMock() + installed_app.app = app_model + installed_app.app_with_session.return_value = app_model + return installed_app + + @pytest.fixture def payload_patch(): def _patch(payload): @@ -42,8 +52,7 @@ class TestSavedMessageListApi: api = module.SavedMessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") pagination = MagicMock( limit=20, @@ -72,8 +81,7 @@ class TestSavedMessageListApi: api = module.SavedMessageListApi() method = unwrap(api.get) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with pytest.raises(NotCompletionAppError): method(api, MagicMock(), installed_app) @@ -82,8 +90,7 @@ class TestSavedMessageListApi: api = module.SavedMessageListApi() method = unwrap(api.post) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") payload = {"message_id": str(uuid4())} current_user = MagicMock() @@ -103,8 +110,7 @@ class TestSavedMessageListApi: api = module.SavedMessageListApi() method = unwrap(api.post) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") payload = {"message_id": str(uuid4())} @@ -126,8 +132,7 @@ class TestSavedMessageApi: api = module.SavedMessageApi() method = unwrap(api.delete) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="completion") + installed_app = make_installed_app(mode="completion") current_user = MagicMock() with ( @@ -144,8 +149,7 @@ class TestSavedMessageApi: api = module.SavedMessageApi() method = unwrap(api.delete) - installed_app = MagicMock() - installed_app.app = MagicMock(mode="chat") + installed_app = make_installed_app(mode="chat") with pytest.raises(NotCompletionAppError): method(api, MagicMock(), installed_app, str(uuid4())) diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial.py b/api/tests/unit_tests/controllers/console/explore/test_trial.py index d810bf4da70..ed853952255 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -1,4 +1,5 @@ from datetime import UTC, datetime +from inspect import getsource, signature from inspect import unwrap as inspect_unwrap from io import BytesIO from types import SimpleNamespace @@ -121,6 +122,7 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): api = module.DatasetListApi() method = unwrap(api.get) app_model = SimpleNamespace(tenant_id="tenant-1") + session = MagicMock() with ( app.test_request_context("/?page=1&limit=20&ids=dataset-1"), @@ -130,9 +132,9 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): return_value=([DatasetListItem()], 1), ) as get_datasets, ): - result = method(api, app_model) + result = method(api, session, app_model) - get_datasets.assert_called_once_with(["dataset-1"], "tenant-1") + get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=session) assert result == { "data": [ { @@ -154,6 +156,38 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): } +@pytest.mark.parametrize( + "api_type", + [module.TrialSitApi, module.TrialAppParameterApi, module.AppApi, module.AppWorkflowApi, module.DatasetListApi], +) +def test_trial_app_handlers_use_explicit_read_session(api_type: type) -> None: + source = getsource(api_type.get) + + assert "@with_session(write=False)\n @get_app_model_with_trial(None)" in source + assert tuple(signature(api_type.get).parameters)[:3] == ("self", "session", "app_model") + + +def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + session = MagicMock() + app_model = MagicMock() + response_view = MagicMock() + get_app = MagicMock(return_value=app_model) + build_view = MagicMock(return_value=response_view) + validated = MagicMock() + validated.model_dump.return_value = {"id": "app-1"} + monkeypatch.setattr(module, "AppService", lambda: SimpleNamespace(get_app=get_app)) + monkeypatch.setattr(module, "AppResponseView", build_view) + monkeypatch.setattr(module.TrialAppDetailResponse, "model_validate", MagicMock(return_value=validated)) + + with app.test_request_context("/"): + result = unwrap(module.AppApi.get)(module.AppApi(), session, app_model) + + assert result == {"id": "app-1"} + get_app.assert_called_once_with(app_model, session=session) + build_view.assert_called_once_with(app_model, session=session) + module.TrialAppDetailResponse.model_validate.assert_called_once_with(response_view, from_attributes=True) + + class TestTrialAppWorkflowRunApi: def test_not_workflow_app(self, app: Flask, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -645,18 +679,25 @@ class TestTrialAppParameterApi: method = unwrap(api.get) with pytest.raises(AppUnavailableError): - method(api, None) + method(api, MagicMock(), None) def test_success_non_workflow(self, valid_parameters: dict[str, object]) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) - app_model = MagicMock( + app_model_config = MagicMock(app_id="app-1") + app_model_config.to_dict.return_value = {"user_input_form": []} + app_model = SimpleNamespace( mode=AppMode.CHAT, - app_model_config=MagicMock(to_dict=lambda: {"user_input_form": []}), + app_model_config_with_session=MagicMock(return_value=app_model_config), ) + session = MagicMock() + annotation_reply = {"enabled": False} with ( + patch.object( + module, "load_annotation_reply_config", return_value=annotation_reply + ) as load_annotation_reply, patch.object( module, "get_parameters_from_feature_dict", @@ -668,9 +709,38 @@ class TestTrialAppParameterApi: return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), ), ): - result = method(api, app_model) + result = method(api, session, app_model) assert result == {"ok": True} + app_model.app_model_config_with_session.assert_called_once_with(session=session) + load_annotation_reply.assert_called_once_with(session, "app-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + + def test_success_workflow(self, valid_parameters: dict[str, object]) -> None: + api = module.TrialAppParameterApi() + method = unwrap(api.get) + + workflow = MagicMock(features_dict={}) + workflow.user_input_form.return_value = [] + app_model = SimpleNamespace( + mode=AppMode.WORKFLOW, + workflow_with_session=MagicMock(return_value=workflow), + ) + session = MagicMock() + + with ( + patch.object(module, "get_parameters_from_feature_dict", return_value=valid_parameters), + patch.object( + module.ParametersResponse, + "model_validate", + return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), + ), + ): + result = method(api, session, app_model) + + assert result == {"ok": True} + app_model.workflow_with_session.assert_called_once_with(session=session) + workflow.user_input_form.assert_called_once_with(to_old_structure=True) class TestTrialChatAudioApi: @@ -1007,49 +1077,111 @@ class TestTrialSitApi: method = unwrap(api.get) app_model = MagicMock() app_model.id = "a1" + session = MagicMock() + session.scalar.return_value = None - with app.test_request_context("/"), patch.object(module.db.session, "scalar") as mock_scalar: - mock_scalar.return_value = None + with app.test_request_context("/"): with pytest.raises(Forbidden): - method(api, app_model) + method(api, session, app_model) + + session.scalar.assert_called_once() def test_archived_tenant(self, app: Flask) -> None: api = module.TrialSitApi() method = unwrap(api.get) site = MagicMock() - app_model = MagicMock() - app_model.id = "a1" - app_model.tenant = MagicMock() - app_model.tenant.status = TenantStatus.ARCHIVE + app_model = SimpleNamespace(id="a1", tenant_id="tenant-1") + tenant = SimpleNamespace(status=TenantStatus.ARCHIVE) + session = MagicMock() + session.scalar.return_value = site - with app.test_request_context("/"), patch.object(module.db.session, "scalar") as mock_scalar: - mock_scalar.return_value = site + with ( + app.test_request_context("/"), + patch.object(module.TenantService, "get_tenant_by_id", return_value=tenant) as get_tenant_by_id, + ): with pytest.raises(Forbidden): - method(api, app_model) + method(api, session, app_model) + + session.scalar.assert_called_once() + get_tenant_by_id.assert_called_once_with("tenant-1", session=session) def test_success(self, app: Flask) -> None: api = module.TrialSitApi() method = unwrap(api.get) site = MagicMock() - app_model = MagicMock() - app_model.id = "a1" - app_model.tenant = MagicMock() - app_model.tenant.status = TenantStatus.NORMAL + app_model = SimpleNamespace(id="a1", tenant_id="tenant-1") + tenant = SimpleNamespace(status=TenantStatus.NORMAL) + session = MagicMock() + session.scalar.return_value = site with ( app.test_request_context("/"), - patch.object(module.db.session, "scalar") as mock_scalar, + patch.object(module.TenantService, "get_tenant_by_id", return_value=tenant) as get_tenant_by_id, patch.object(module.SiteResponse, "model_validate") as mock_validate, ): - mock_scalar.return_value = site mock_validate_result = MagicMock() mock_validate_result.model_dump.return_value = {"name": "test", "icon": "icon"} mock_validate.return_value = mock_validate_result - result = method(api, app_model) + result = method(api, session, app_model) assert result == {"name": "test", "icon": "icon"} + session.scalar.assert_called_once() + get_tenant_by_id.assert_called_once_with("tenant-1", session=session) + + +class TestAppWorkflowApi: + def test_uses_injected_session(self) -> None: + api = module.AppWorkflowApi() + method = unwrap(api.get) + created_by = SimpleNamespace(id="account-1", name="Creator", email="creator@example.com") + workflow = SimpleNamespace( + id="workflow-1", + graph_dict={"nodes": []}, + features_dict={}, + unique_hash="workflow-hash", + version="draft", + marked_name="", + marked_comment="", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + updated_at=datetime(2024, 1, 2, tzinfo=UTC), + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + get_created_by_account=MagicMock(return_value=created_by), + get_updated_by_account=MagicMock(return_value=None), + get_tool_published=MagicMock(return_value=True), + ) + app_model = SimpleNamespace( + workflow_id="workflow-1", + workflow_with_session=MagicMock(return_value=workflow), + ) + session = MagicMock() + + result = method(api, session, app_model) + + assert result == { + "id": "workflow-1", + "graph": {"nodes": []}, + "features": {}, + "hash": "workflow-hash", + "version": "draft", + "marked_name": "", + "marked_comment": "", + "created_by": {"id": "account-1", "name": "Creator", "email": "creator@example.com"}, + "created_at": 1704067200, + "updated_by": None, + "updated_at": 1704153600, + "tool_published": True, + "environment_variables": [], + "conversation_variables": [], + "rag_pipeline_variables": [], + } + app_model.workflow_with_session.assert_called_once_with(session=session) + workflow.get_created_by_account.assert_called_once_with(session=session) + workflow.get_updated_by_account.assert_called_once_with(session=session) + workflow.get_tool_published.assert_called_once_with(session=session) class TestTrialChatAudioApiExceptionHandlers: diff --git a/api/tests/unit_tests/controllers/console/explore/test_workflow.py b/api/tests/unit_tests/controllers/console/explore/test_workflow.py index 70cd9fd5cc8..8bcce0c22b2 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/explore/test_workflow.py @@ -36,14 +36,18 @@ def workflow_app(): @pytest.fixture def installed_workflow_app(workflow_app): - return MagicMock(app=workflow_app) + installed_app = MagicMock(app=workflow_app) + installed_app.app_with_session.return_value = workflow_app + return installed_app @pytest.fixture def non_workflow_installed_app(): app = MagicMock() app.mode = AppMode.CHAT - return MagicMock(app=app) + installed_app = MagicMock(app=app) + installed_app.app_with_session.return_value = app + return installed_app @pytest.fixture @@ -112,7 +116,7 @@ class TestInstalledAppWorkflowTaskStopApi: method = unwrap(api.post) with pytest.raises(NotWorkflowAppError): - method(non_workflow_installed_app, "task-1") + method(api, MagicMock(), non_workflow_installed_app, "task-1") def test_success(self, installed_workflow_app): api = InstalledAppWorkflowTaskStopApi() @@ -122,7 +126,7 @@ class TestInstalledAppWorkflowTaskStopApi: patch("controllers.console.explore.workflow.AppQueueManager.set_stop_flag_no_user_check") as stop_flag, patch("controllers.console.explore.workflow.GraphEngineManager.send_stop_command") as send_stop, ): - result = method(installed_workflow_app, "task-1") + result = method(api, MagicMock(), installed_workflow_app, "task-1") stop_flag.assert_called_once_with("task-1") send_stop.assert_called_once_with("task-1") diff --git a/api/tests/unit_tests/controllers/console/explore/test_wraps.py b/api/tests/unit_tests/controllers/console/explore/test_wraps.py index ab361da93a9..a60c13315b8 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/explore/test_wraps.py @@ -24,8 +24,10 @@ from models import AccountTrialAppRecord, App, AppMode, InstalledApp, TrialApp def _bind_database(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: - monkeypatch.setattr(wraps_module.db, "session", sqlite_session) - monkeypatch.setattr(model_module.db, "session", sqlite_session) + session_proxy = MagicMock(wraps=sqlite_session) + session_proxy.return_value = sqlite_session + monkeypatch.setattr(wraps_module.db, "session", session_proxy) + monkeypatch.setattr(model_module.db, "session", session_proxy) def _app() -> App: diff --git a/api/tests/unit_tests/controllers/console/test_apikey.py b/api/tests/unit_tests/controllers/console/test_apikey.py index 1517ff5ed8e..0435ad0996a 100644 --- a/api/tests/unit_tests/controllers/console/test_apikey.py +++ b/api/tests/unit_tests/controllers/console/test_apikey.py @@ -4,7 +4,7 @@ import inspect from collections.abc import Callable from types import SimpleNamespace from typing import cast -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest from werkzeug.exceptions import Forbidden @@ -44,8 +44,14 @@ def _make_account(role: TenantAccountRole) -> Account: return account -def test_list_api_keys_uses_injected_tenant_id() -> None: +def test_list_api_keys_uses_injected_session_and_tenant_id() -> None: resource = _make_list_resource() + raw_get = cast( + Callable[[BaseApiKeyListResource, object, str, str], dict[str, object]], + inspect.unwrap(BaseApiKeyListResource.get), + ) + session = MagicMock() + session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() api_key = SimpleNamespace( id="key-1", type=ApiTokenType.APP, @@ -53,16 +59,12 @@ def test_list_api_keys_uses_injected_tenant_id() -> None: last_used_at=None, created_at=None, ) + session.scalars.return_value.all.return_value = [api_key] - with ( - patch("controllers.console.apikey._get_resource") as get_resource, - patch("controllers.console.apikey.db") as db_mock, - ): - db_mock.session.scalars.return_value.all.return_value = [api_key] + result = raw_get(resource, session, "app-1", "tenant-1") - result = resource.get("app-1", "tenant-1") - - get_resource.assert_called_once_with("app-1", "tenant-1", App) + session.execute.assert_called_once() + session.scalars.assert_called_once() assert result == { "data": [ { @@ -76,66 +78,85 @@ def test_list_api_keys_uses_injected_tenant_id() -> None: } -def test_create_api_key_uses_injected_tenant_id() -> None: +def test_create_api_key_uses_injected_session_and_tenant_id() -> None: resource = _make_list_resource() raw_post = cast( - Callable[[BaseApiKeyListResource, str, str], tuple[dict[str, object], int]], + Callable[[BaseApiKeyListResource, object, str, str], tuple[dict[str, object], int]], inspect.unwrap(BaseApiKeyListResource.post), ) + session = MagicMock() + session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() + session.scalar.return_value = 0 def add_api_token(api_token: ApiToken) -> None: api_token.id = "key-1" - with ( - patch("controllers.console.apikey._get_resource") as get_resource, - patch("controllers.console.apikey.db") as db_mock, - patch("controllers.console.apikey.ApiToken.generate_api_key", return_value="app-generated-token"), - ): - db_mock.session.scalar.return_value = 0 - db_mock.session.add.side_effect = add_api_token + with patch( + "controllers.console.apikey.ApiToken.generate_api_key", return_value="app-generated-token" + ) as generate_api_key: + session.add.side_effect = add_api_token - result, status = raw_post(resource, "app-1", "tenant-1") + result, status = raw_post(resource, session, "app-1", "tenant-1") - get_resource.assert_called_once_with("app-1", "tenant-1", App) assert status == 201 assert result["token"] == "app-generated-token" - api_token = db_mock.session.add.call_args.args[0] + api_token = session.add.call_args.args[0] assert api_token.app_id == "app-1" assert api_token.tenant_id == "tenant-1" assert api_token.type == ApiTokenType.APP - db_mock.session.commit.assert_called_once() + generate_api_key.assert_called_once_with("app-", 24, session=session) + session.execute.assert_called_once() + session.scalar.assert_called_once() + session.commit.assert_called_once() def test_delete_api_key_rejects_non_admin_account() -> None: resource = _make_key_resource() + raw_delete = cast( + Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], + inspect.unwrap(BaseApiKeyResource.delete), + ) + session = MagicMock() + session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - with ( - patch("controllers.console.apikey._get_resource") as get_resource, - patch("controllers.console.apikey.db") as db_mock, - ): - with pytest.raises(Forbidden): - resource.delete("app-1", "key-1", "tenant-1", _make_account(TenantAccountRole.NORMAL)) + with pytest.raises(Forbidden): + raw_delete( + resource, + session, + "app-1", + "key-1", + "tenant-1", + _make_account(TenantAccountRole.NORMAL), + ) - get_resource.assert_called_once_with("app-1", "tenant-1", App) - db_mock.session.scalar.assert_not_called() + session.execute.assert_called_once() + session.scalar.assert_not_called() -def test_delete_api_key_uses_injected_user_and_tenant() -> None: +def test_delete_api_key_uses_injected_session_user_and_tenant() -> None: resource = _make_key_resource() + raw_delete = cast( + Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], + inspect.unwrap(BaseApiKeyResource.delete), + ) api_key = SimpleNamespace(token="app-token", type=ApiTokenType.APP) + session = MagicMock() + session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() + session.scalar.return_value = api_key - with ( - patch("controllers.console.apikey._get_resource") as get_resource, - patch("controllers.console.apikey.db") as db_mock, - patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache, - ): - db_mock.session.scalar.return_value = api_key + with patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache: + result, status = raw_delete( + resource, + session, + "app-1", + "key-1", + "tenant-1", + _make_account(TenantAccountRole.OWNER), + ) - result, status = resource.delete("app-1", "key-1", "tenant-1", _make_account(TenantAccountRole.OWNER)) - - get_resource.assert_called_once_with("app-1", "tenant-1", App) delete_cache.assert_called_once_with("app-token", ApiTokenType.APP) - db_mock.session.execute.assert_called_once() - db_mock.session.commit.assert_called_once() + assert session.execute.call_count == 2 + session.scalar.assert_called_once() + session.commit.assert_called_once() assert result == "" assert status == 204 diff --git a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py index 28a94387198..14cf4584851 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py @@ -1,8 +1,8 @@ -import inspect import logging from http import HTTPStatus +from inspect import unwrap from io import BytesIO -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -72,21 +72,16 @@ def make_account_with_tenant(tenant: Tenant) -> Account: class TestTenantListApi: def test_get_success_saas_path(self, app: Flask): api = TenantListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant1 = make_tenant("t1", name="Tenant 1") tenant2 = make_tenant("t2", name="Tenant 2") last_opened_at = naive_utc_now() user = make_account() - with ( app.test_request_context("/workspaces"), patch( "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[ - (tenant1, make_membership(last_opened_at=last_opened_at)), - (tenant2, make_membership()), - ], + return_value=[(tenant1, make_membership(last_opened_at=last_opened_at)), (tenant2, make_membership())], ), patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", False), patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", True), @@ -100,8 +95,7 @@ class TestTenantListApi: ) as get_plan_bulk_mock, patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, ): - result, status = method(api, "t1", user) - + result, status = method(api, MagicMock(), "t1", user) assert status == HTTPStatus.OK assert len(result["workspaces"]) == 2 assert result["workspaces"][0]["current"] is True @@ -119,16 +113,13 @@ class TestTenantListApi: (SaaS contract treats enabled as on; display follows subscription.plan). """ api = TenantListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant1 = make_tenant("t1", name="Tenant 1") tenant2 = make_tenant("t2", name="Tenant 2") - features_t2 = MagicMock() features_t2.billing.enabled = False features_t2.billing.subscription.plan = CloudPlan.PROFESSIONAL user = make_account() - with ( app.test_request_context("/workspaces"), patch( @@ -143,12 +134,10 @@ class TestTenantListApi: return_value={"t1": {"plan": CloudPlan.TEAM, "expiration_date": 0}}, ) as get_plan_bulk_mock, patch( - "controllers.console.workspace.workspace.FeatureService.get_features", - return_value=features_t2, + "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features_t2 ) as get_features_mock, ): - result, status = method(api, "t1", user) - + result, status = method(api, MagicMock(), "t1", user) assert status == HTTPStatus.OK assert result["workspaces"][0]["plan"] == CloudPlan.TEAM assert result["workspaces"][1]["plan"] == CloudPlan.PROFESSIONAL @@ -164,16 +153,13 @@ class TestTenantListApi: so we simulate the real failure mode by returning empty dict for non-empty input. """ api = TenantListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant1 = make_tenant("t1", name="Tenant 1") tenant2 = make_tenant("t2", name="Tenant 2") - features = MagicMock() features.billing.enabled = False features.billing.subscription.plan = CloudPlan.TEAM user = make_account() - with ( app.test_request_context("/workspaces"), caplog.at_level(logging.WARNING, logger="controllers.console.workspace.workspace"), @@ -185,16 +171,13 @@ class TestTenantListApi: patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", True), patch("controllers.console.workspace.workspace.dify_config.EDITION", "CLOUD"), patch( - "controllers.console.workspace.workspace.BillingService.get_plan_bulk", - return_value={}, # Simulates real failure: empty result for non-empty input + "controllers.console.workspace.workspace.BillingService.get_plan_bulk", return_value={} ) as get_plan_bulk_mock, patch( - "controllers.console.workspace.workspace.FeatureService.get_features", - return_value=features, + "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features ) as get_features_mock, ): - result, status = method(api, "t2", user) - + result, status = method(api, MagicMock(), "t2", user) assert status == HTTPStatus.OK assert result["workspaces"][0]["plan"] == CloudPlan.TEAM assert result["workspaces"][1]["plan"] == CloudPlan.TEAM @@ -204,15 +187,12 @@ class TestTenantListApi: def test_get_billing_disabled_community_path(self, app: Flask): api = TenantListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant = make_tenant("t1", name="Tenant") - features = MagicMock() features.billing.enabled = False features.billing.subscription.plan = CloudPlan.SANDBOX user = make_account() - with ( app.test_request_context("/workspaces"), patch( @@ -223,24 +203,20 @@ class TestTenantListApi: patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", False), patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), patch( - "controllers.console.workspace.workspace.FeatureService.get_features", - return_value=features, + "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features ) as get_features_mock, ): - result, status = method(api, "t1", user) - + result, status = method(api, MagicMock(), "t1", user) assert status == HTTPStatus.OK assert result["workspaces"][0]["plan"] == CloudPlan.SANDBOX get_features_mock.assert_called_once_with("t1", exclude_vector_space=True) def test_get_enterprise_only_skips_feature_service(self, app: Flask): api = TenantListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant1 = make_tenant("t1", name="Tenant 1") tenant2 = make_tenant("t2", name="Tenant 2") user = make_account() - with ( app.test_request_context("/workspaces"), patch( @@ -252,8 +228,7 @@ class TestTenantListApi: patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, ): - result, status = method(api, "t2", user) - + result, status = method(api, MagicMock(), "t2", user) assert status == HTTPStatus.OK assert result["workspaces"][0]["plan"] == CloudPlan.SANDBOX assert result["workspaces"][1]["plan"] == CloudPlan.SANDBOX @@ -263,22 +238,17 @@ class TestTenantListApi: def test_get_enterprise_only_with_empty_tenants(self, app: Flask): api = TenantListApi() - method = inspect.unwrap(api.get) + method = unwrap(api.get) user = make_account() - with ( app.test_request_context("/workspaces"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[], - ), + patch("controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", return_value=[]), patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", True), patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", False), patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, ): - result, status = method(api, None, user) - + result, status = method(api, MagicMock(), None, user) assert status == HTTPStatus.OK assert result["workspaces"] == [] get_features_mock.assert_not_called() @@ -287,34 +257,28 @@ class TestTenantListApi: class TestWorkspaceListApi: def test_get_success(self, app: Flask): api = WorkspaceListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant = make_tenant("t1", name="T") paginate_result = MagicMock(items=[tenant], has_next=False, total=1) - with ( app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 20}), patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result), ): - result, status = method(api) - + result, status = method(api, MagicMock()) assert status == HTTPStatus.OK assert result["total"] == 1 assert result["has_more"] is False def test_get_has_next_true(self, app: Flask): api = WorkspaceListApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) tenant = make_tenant("t1", name="T") paginate_result = MagicMock(items=[tenant], has_next=True, total=10) - with ( app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 1}), patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result), ): - result, status = method(api) - + result, status = method(api, MagicMock()) assert status == HTTPStatus.OK assert result["has_more"] is True @@ -322,72 +286,61 @@ class TestWorkspaceListApi: class TestTenantApi: def test_post_active_tenant(self, app: Flask): api = TenantApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) tenant = make_tenant() user = make_account_with_tenant(tenant) - with ( app.test_request_context("/workspaces/current"), patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} ), ): - result, status = method(api, user) - + result, status = method(api, MagicMock(), user) assert status == HTTPStatus.OK assert result["id"] == "t1" def test_post_archived_with_switch(self, app: Flask): api = TenantApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) archived = make_tenant(status=TenantStatus.ARCHIVE) new_tenant = make_tenant("new") user = make_account_with_tenant(archived) - with ( app.test_request_context("/workspaces/current"), patch("controllers.console.workspace.workspace.TenantService.get_join_tenants", return_value=[new_tenant]), - patch("controllers.console.workspace.workspace.TenantService.switch_tenant"), + patch("controllers.console.workspace.workspace.TenantService.switch_tenant") as switch_tenant, patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "new"} ), ): - result, status = method(api, user) - + result, status = method(api, MagicMock(), user) assert result["id"] == "new" + switch_tenant.assert_called_once_with(user, new_tenant.id, session=ANY) def test_post_archived_no_tenant(self, app: Flask): api = TenantApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) user = make_account_with_tenant(make_tenant(status=TenantStatus.ARCHIVE)) - with ( app.test_request_context("/workspaces/current"), patch("controllers.console.workspace.workspace.TenantService.get_join_tenants", return_value=[]), ): with pytest.raises(Unauthorized): - method(api, user) + method(api, MagicMock(), user) def test_post_info_path(self, app: Flask, caplog: pytest.LogCaptureFixture): api = TenantApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) tenant = make_tenant() user = make_account_with_tenant(tenant) - with ( app.test_request_context("/info"), caplog.at_level(logging.WARNING, logger="controllers.console.workspace.workspace"), patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", - return_value={"id": "t1"}, + "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} ), ): - result, status = method(api, user) - + result, status = method(api, MagicMock(), user) assert "Deprecated URL /info was used." in caplog.messages assert status == HTTPStatus.OK @@ -396,14 +349,8 @@ class TestTenantInfoResponse: def test_tenant_info_response_normalizes_enum_and_datetime(self): created_at = naive_utc_now() payload = TenantInfoResponse.model_validate( - { - "id": "t1", - "status": TenantStatus.NORMAL, - "plan": CloudPlan.TEAM, - "created_at": created_at, - } + {"id": "t1", "status": TenantStatus.NORMAL, "plan": CloudPlan.TEAM, "created_at": created_at} ).model_dump(mode="json") - assert payload["status"] == "normal" assert payload["plan"] == "team" assert payload["created_at"] == int(created_at.timestamp()) @@ -419,110 +366,91 @@ class TestTenantInfoResponse: }, } ).model_dump(mode="json") - - assert payload["custom_config"] == { - "remove_webapp_brand": True, - "replace_webapp_logo": "logo-file-id", - } + assert payload["custom_config"] == {"remove_webapp_brand": True, "replace_webapp_logo": "logo-file-id"} class TestSwitchWorkspaceApi: def test_switch_success(self, app: Flask): api = SwitchWorkspaceApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"tenant_id": "t2"} tenant = make_tenant("t2") user = make_account() - with ( app.test_request_context("/workspaces/switch", json=payload), - patch("controllers.console.workspace.workspace.TenantService.switch_tenant"), - patch("controllers.console.workspace.workspace.db.session.get") as get_mock, + patch("controllers.console.workspace.workspace.TenantService.switch_tenant") as switch_tenant, patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t2"} ), ): - get_mock.return_value = tenant - result = method(api, user) - + session = MagicMock() + session.get.return_value = tenant + result = method(api, session, user) assert result["result"] == "success" + switch_tenant.assert_called_once_with(user, "t2", session=session) def test_switch_not_linked(self, app: Flask): api = SwitchWorkspaceApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"tenant_id": "bad"} user = make_account() - with ( app.test_request_context("/workspaces/switch", json=payload), patch("controllers.console.workspace.workspace.TenantService.switch_tenant", side_effect=Exception), ): with pytest.raises(AccountNotLinkTenantError): - method(api, user) + method(api, MagicMock(), user) def test_switch_tenant_not_found(self, app: Flask): api = SwitchWorkspaceApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"tenant_id": "missing"} user = make_account() - with ( app.test_request_context("/workspaces/switch", json=payload), patch("controllers.console.workspace.workspace.TenantService.switch_tenant"), - patch("controllers.console.workspace.workspace.db.session.get") as get_mock, ): - get_mock.return_value = None - + session = MagicMock() + session.get.return_value = None with pytest.raises(ValueError): - method(api, user) + method(api, session, user) class TestCustomConfigWorkspaceApi: def test_post_success(self, app: Flask): api = CustomConfigWorkspaceApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) tenant = make_tenant(custom_config={}) - payload = {"remove_webapp_brand": True} - + events = [] + with ( + app.test_request_context("/workspaces/custom-config", json=payload), + patch( + "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", + side_effect=lambda *args, **kwargs: events.append("get_tenant_info") or {"id": "t1"}, + ), + ): + session = MagicMock() + session.get.return_value = tenant + session.commit.side_effect = lambda: events.append("commit") + result = method(api, session, "t1") + assert result["result"] == "success" + assert events == ["commit", "get_tenant_info"] + + def test_logo_fallback(self, app: Flask): + api = CustomConfigWorkspaceApi() + method = unwrap(api.post) + tenant = make_tenant(custom_config={"replace_webapp_logo": "old-logo"}) + payload = {"remove_webapp_brand": False} with ( app.test_request_context("/workspaces/custom-config", json=payload), - patch("controllers.console.workspace.workspace.db.get_or_404", return_value=tenant), - patch("controllers.console.workspace.workspace.db.session.commit"), patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} ), ): - result = method(api, "t1") - - assert result["result"] == "success" - - def test_logo_fallback(self, app: Flask): - api = CustomConfigWorkspaceApi() - method = inspect.unwrap(api.post) - - tenant = make_tenant(custom_config={"replace_webapp_logo": "old-logo"}) - - payload = {"remove_webapp_brand": False} - - with ( - app.test_request_context("/workspaces/custom-config", json=payload), - patch( - "controllers.console.workspace.workspace.db.get_or_404", - return_value=tenant, - ), - patch("controllers.console.workspace.workspace.db.session.commit"), - patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", - return_value={"id": "t1"}, - ), - ): - result = method(api, "t1") - + session = MagicMock() + session.get.return_value = tenant + result = method(api, session, "t1") assert tenant.custom_config_dict["replace_webapp_logo"] == "old-logo" assert result["result"] == "success" @@ -530,137 +458,84 @@ class TestCustomConfigWorkspaceApi: class TestWebappLogoWorkspaceApi: def test_no_file(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) + method = unwrap(api.post) user = make_account() - with app.test_request_context("/upload", data={}): with pytest.raises(NoFileUploadedError): method(api, user) def test_too_many_files(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - - data = { - "file": MagicMock(), - "extra": MagicMock(), - } + method = unwrap(api.post) + data = {"file": MagicMock(), "extra": MagicMock()} user = make_account() - with app.test_request_context("/upload", data=data): with pytest.raises(TooManyFilesError): method(api, user) def test_invalid_extension(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) file = MagicMock(filename="test.txt") user = make_account() - with app.test_request_context("/upload", data={"file": file}): with pytest.raises(UnsupportedFileTypeError): method(api, user) def test_upload_success(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - - file = FileStorage( - stream=BytesIO(b"data"), - filename="logo.png", - content_type="image/png", - ) - + method = unwrap(api.post) + file = FileStorage(stream=BytesIO(b"data"), filename="logo.png", content_type="image/png") upload = MagicMock(id="file1") user = make_account() - with ( - app.test_request_context( - "/upload", - data={"file": file}, - content_type="multipart/form-data", - ), + app.test_request_context("/upload", data={"file": file}, content_type="multipart/form-data"), patch("controllers.console.workspace.workspace.FileService") as fs, patch("controllers.console.workspace.workspace.db") as mock_db, ): mock_db.engine = MagicMock() fs.return_value.upload_file.return_value = upload - result, status = method(api, user) - assert status == HTTPStatus.CREATED assert result == {"id": "file1"} assert WorkspaceLogoUploadResponse.model_validate(result).model_dump(mode="json") == {"id": "file1"} def test_filename_missing(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - - file = FileStorage( - stream=BytesIO(b"data"), - filename="", - content_type="image/png", - ) + method = unwrap(api.post) + file = FileStorage(stream=BytesIO(b"data"), filename="", content_type="image/png") user = make_account() - - with app.test_request_context( - "/upload", - data={"file": file}, - content_type="multipart/form-data", - ): + with app.test_request_context("/upload", data={"file": file}, content_type="multipart/form-data"): with pytest.raises(FilenameNotExistsError): method(api, user) def test_file_too_large(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - - file = FileStorage( - stream=BytesIO(b"x"), - filename="logo.png", - content_type="image/png", - ) + method = unwrap(api.post) + file = FileStorage(stream=BytesIO(b"x"), filename="logo.png", content_type="image/png") user = make_account() - with ( - app.test_request_context( - "/upload", - data={"file": file}, - content_type="multipart/form-data", - ), + app.test_request_context("/upload", data={"file": file}, content_type="multipart/form-data"), patch("controllers.console.workspace.workspace.FileService") as fs, patch("controllers.console.workspace.workspace.db") as mock_db, ): mock_db.engine = MagicMock() fs.return_value.upload_file.side_effect = services.errors.file.FileTooLargeError("too big") - with pytest.raises(FileTooLargeError): method(api, user) def test_service_unsupported_file(self, app: Flask): api = WebappLogoWorkspaceApi() - method = inspect.unwrap(api.post) - - file = FileStorage( - stream=BytesIO(b"x"), - filename="logo.png", - content_type="image/png", - ) + method = unwrap(api.post) + file = FileStorage(stream=BytesIO(b"x"), filename="logo.png", content_type="image/png") user = make_account() - with ( - app.test_request_context( - "/upload", - data={"file": file}, - content_type="multipart/form-data", - ), + app.test_request_context("/upload", data={"file": file}, content_type="multipart/form-data"), patch("controllers.console.workspace.workspace.FileService") as fs, patch("controllers.console.workspace.workspace.db") as mock_db, ): mock_db.engine = MagicMock() fs.return_value.upload_file.side_effect = services.errors.file.UnsupportedFileTypeError() - with pytest.raises(UnsupportedFileTypeError): method(api, user) @@ -668,49 +543,40 @@ class TestWebappLogoWorkspaceApi: class TestWorkspaceInfoApi: def test_post_success(self, app: Flask): api = WorkspaceInfoApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) tenant = make_tenant() - payload = {"name": "New Name"} - + events = [] with ( app.test_request_context("/workspaces/info", json=payload), - patch("controllers.console.workspace.workspace.db.get_or_404", return_value=tenant), - patch("controllers.console.workspace.workspace.db.session.commit"), patch( "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", - return_value={"id": "t1", "name": "New Name"}, + side_effect=lambda *args, **kwargs: ( + events.append("get_tenant_info") or {"id": "t1", "name": "New Name"} + ), ), ): - result = method(api, "t1") - + session = MagicMock() + session.get.return_value = tenant + session.commit.side_effect = lambda: events.append("commit") + result = method(api, session, "t1") assert result["result"] == "success" + assert events == ["commit", "get_tenant_info"] def test_no_current_tenant(self, app: Flask): api = WorkspaceInfoApi() - method = inspect.unwrap(api.post) - + method = unwrap(api.post) payload = {"name": "X"} - - with ( - app.test_request_context("/workspaces/info", json=payload), - ): + with app.test_request_context("/workspaces/info", json=payload): with pytest.raises(ValueError): - method(api, None) + method(api, MagicMock(), None) class TestWorkspacePermissionApi: def test_get_success(self, app: Flask): api = WorkspacePermissionApi() - method = inspect.unwrap(api.get) - - permission = MagicMock( - workspace_id="t1", - allow_member_invite=True, - allow_owner_transfer=False, - ) - + method = unwrap(api.get) + permission = MagicMock(workspace_id="t1", allow_member_invite=True, allow_owner_transfer=False) with ( app.test_request_context("/permission"), patch( @@ -719,20 +585,14 @@ class TestWorkspacePermissionApi: ), ): result, status = method(api, "t1") - assert status == HTTPStatus.OK - expected = { - "workspace_id": "t1", - "allow_member_invite": True, - "allow_owner_transfer": False, - } + expected = {"workspace_id": "t1", "allow_member_invite": True, "allow_owner_transfer": False} assert result == expected assert WorkspacePermissionResponse.model_validate(result).model_dump(mode="json") == expected def test_no_current_tenant(self, app: Flask): api = WorkspacePermissionApi() - method = inspect.unwrap(api.get) - + method = unwrap(api.get) with app.test_request_context("/permission"): with pytest.raises(ValueError): method(api, None) 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 ad84eed1f5e..c7f788dcb55 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 @@ -145,7 +145,7 @@ class TestEnterpriseAppDSLImport: body, status_code = result assert status_code == 200 assert body["status"] == "completed" - mock_account.set_tenant_id.assert_called_once_with("ws-123") + mock_account.set_tenant_id_with_session.assert_called_once_with("ws-123", session=self._mock_session) self._mock_session.commit.assert_called_once_with() self._mock_session.rollback.assert_not_called() diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py index 093a53ba0d2..3b714a84da1 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py @@ -1,4 +1,5 @@ import uuid +from contextlib import nullcontext from unittest.mock import MagicMock, patch import pytest @@ -152,10 +153,14 @@ def test_load_account_skips_when_already_set(): def test_load_account_sets_current_tenant_when_tenant_present(): account = MagicMock() tenant = MagicMock() + session = MagicMock() data = _make_auth_data(account_id=uuid.uuid4(), tenant=tenant) - with patch("controllers.openapi.auth.prepare.AccountService.get_account_by_id", return_value=account): + with ( + patch("controllers.openapi.auth.prepare.AccountService.get_account_by_id", return_value=account), + patch("controllers.openapi.auth.prepare.session_factory.create_session", return_value=nullcontext(session)), + ): load_account(data) - assert account.current_tenant is tenant + account.set_current_tenant_with_session.assert_called_once_with(tenant, session=session) def test_load_account_raises_unauthorized_when_not_found(): diff --git a/api/tests/unit_tests/controllers/openapi/test_account.py b/api/tests/unit_tests/controllers/openapi/test_account.py index b4af7dda4a6..d4b28fa554e 100644 --- a/api/tests/unit_tests/controllers/openapi/test_account.py +++ b/api/tests/unit_tests/controllers/openapi/test_account.py @@ -4,7 +4,6 @@ import builtins import sys import uuid from types import SimpleNamespace -from unittest.mock import MagicMock import pytest from flask import Flask @@ -171,7 +170,6 @@ def _stub_session_deps(monkeypatch: pytest.MonkeyPatch, rows): mod = sys.modules[_ACCOUNT_MOD] monkeypatch.setattr(mod, "get_auth_ctx", lambda: SimpleNamespace()) monkeypatch.setattr(mod, "list_active_sessions", lambda *args, **kwargs: rows) - monkeypatch.setattr(mod, "db", MagicMock()) def test_sessions_list_valid_query_parses_page_and_limit(app: Flask, monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py index 708a0e59865..f5e9b16d29c 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py @@ -1,4 +1,5 @@ from types import SimpleNamespace +from unittest.mock import MagicMock from controllers.openapi._input_schema import EMPTY_INPUT_SCHEMA from controllers.openapi.apps import _EMPTY_PARAMETERS, build_app_describe_response @@ -25,28 +26,36 @@ def _app() -> _FakeApp: def test_fields_none_returns_all_blocks(monkeypatch): - monkeypatch.setattr("controllers.openapi.apps.parameters_payload", lambda app: {"k": "v"}) - monkeypatch.setattr("controllers.openapi.apps.build_input_schema", lambda app: {"s": 1}) - resp = build_app_describe_response(_app(), None) + app = _app() + session = MagicMock() + parameters_payload = MagicMock(return_value={"k": "v"}) + input_schema = MagicMock(return_value={"s": 1}) + monkeypatch.setattr("controllers.openapi.apps.parameters_payload", parameters_payload) + monkeypatch.setattr("controllers.openapi.apps.build_input_schema", input_schema) + resp = build_app_describe_response(app, None, session=session) assert resp.info is not None assert resp.info.name == "Demo" assert resp.parameters == {"k": "v"} assert resp.input_schema == {"s": 1} + parameters_payload.assert_called_once_with(app, session=session) + input_schema.assert_called_once_with(app, session=session) def test_fields_subset_limits_blocks(monkeypatch): - monkeypatch.setattr("controllers.openapi.apps.parameters_payload", lambda app: {"k": "v"}) - monkeypatch.setattr("controllers.openapi.apps.build_input_schema", lambda app: {"s": 1}) - resp = build_app_describe_response(_app(), ["info"]) + session = MagicMock() + monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) + monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) + resp = build_app_describe_response(_app(), ["info"], session=session) assert resp.info is not None assert resp.parameters is None assert resp.input_schema is None def test_info_omits_author_and_tags(monkeypatch): - monkeypatch.setattr("controllers.openapi.apps.parameters_payload", lambda app: {}) - monkeypatch.setattr("controllers.openapi.apps.build_input_schema", lambda app: {}) - resp = build_app_describe_response(_app(), ["info"]) + session = MagicMock() + monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={})) + monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={})) + resp = build_app_describe_response(_app(), ["info"], session=session) assert resp.info is not None # Usage-face describe must not expose creator identity or tags (cross-tenant leak). assert not hasattr(resp.info, "author") @@ -54,20 +63,20 @@ def test_info_omits_author_and_tags(monkeypatch): def test_parameters_fallback_on_app_unavailable(monkeypatch): - def _raise(app): + def _raise(app, *, session): raise AppUnavailableError() monkeypatch.setattr("controllers.openapi.apps.parameters_payload", _raise) - monkeypatch.setattr("controllers.openapi.apps.build_input_schema", lambda app: {"s": 1}) - resp = build_app_describe_response(_app(), ["parameters"]) + monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) + resp = build_app_describe_response(_app(), ["parameters"], session=MagicMock()) assert resp.parameters == dict(_EMPTY_PARAMETERS) def test_input_schema_fallback_on_app_unavailable(monkeypatch): - def _raise(app): + def _raise(app, *, session): raise AppUnavailableError() - monkeypatch.setattr("controllers.openapi.apps.parameters_payload", lambda app: {"k": "v"}) + monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", _raise) - resp = build_app_describe_response(_app(), ["input_schema"]) + resp = build_app_describe_response(_app(), ["input_schema"], session=MagicMock()) assert resp.input_schema == dict(EMPTY_INPUT_SCHEMA) diff --git a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py index 12bf4c696f2..2e9e7bc06a8 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py @@ -5,6 +5,7 @@ HTTP plumbing or DB. Pin the response shapes that are CLI contracts. from __future__ import annotations from types import SimpleNamespace +from unittest.mock import MagicMock import pytest @@ -35,8 +36,14 @@ def _fake_app(**overrides): def test_parameters_payload_raises_app_unavailable_when_no_config(): + app = _fake_app(mode="chat") + app.app_model_config_with_session = MagicMock(return_value=None) + session = MagicMock() + with pytest.raises(AppUnavailableError): - parameters_payload(_fake_app(mode="chat", app_model_config=None)) + parameters_payload(app, session=session) + + app.app_model_config_with_session.assert_called_once_with(session=session) def test_empty_parameters_constant_matches_describe_fallback_shape(): diff --git a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py index 0f530e3c3c7..49f7bea5cd3 100644 --- a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py +++ b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py @@ -8,10 +8,17 @@ dropping them. Mode/name/page/limit have the same shape as AppListQuery. from __future__ import annotations +import inspect +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + import pytest from pydantic import ValidationError -from controllers.openapi.apps_permitted_external import PermittedExternalAppsListQuery +from controllers.openapi.apps_permitted_external import ( + PermittedExternalAppDescribeApi, + PermittedExternalAppsListQuery, +) from ._mode_constants import NON_LISTABLE_MODES @@ -60,3 +67,22 @@ def test_query_accepts_valid_mode(): q = PermittedExternalAppsListQuery.model_validate({"mode": "chat"}) assert q.mode is not None assert q.mode.value == "chat" + + +def test_describe_forwards_request_session_to_response_builder(): + api = PermittedExternalAppDescribeApi() + method = inspect.unwrap(api.get) + session = MagicMock() + app = MagicMock() + auth_data = SimpleNamespace(app=app) + query = SimpleNamespace(fields={"info"}) + response = object() + + with patch( + "controllers.openapi.apps_permitted_external.build_app_describe_response", + return_value=response, + ) as build_response: + result = method(api, session, "app-id", auth_data=auth_data, query=query) + + assert result is response + build_response.assert_called_once_with(app, query.fields, session=session) diff --git a/api/tests/unit_tests/controllers/openapi/test_input_schema.py b/api/tests/unit_tests/controllers/openapi/test_input_schema.py index 73cb978ac1e..133072ad33e 100644 --- a/api/tests/unit_tests/controllers/openapi/test_input_schema.py +++ b/api/tests/unit_tests/controllers/openapi/test_input_schema.py @@ -95,28 +95,37 @@ from models.model import AppMode def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bool | None = None): - """Returns a MagicMock whose .mode + workflow / app_model_config branch is wired up.""" + """Returns a MagicMock whose explicit config getters are wired up.""" app = MagicMock() app.mode = mode if mode in (AppMode.WORKFLOW, AppMode.ADVANCED_CHAT): if has_workflow is False: - app.workflow = None + app.workflow_with_session.return_value = None else: - app.workflow = MagicMock() - app.workflow.user_input_form.return_value = form or [] - app.workflow.features_dict = {} + workflow = MagicMock() + workflow.user_input_form.return_value = form or [] + workflow.features_dict = {} + app.workflow_with_session.return_value = workflow else: if has_workflow is False: - app.app_model_config = None + app.app_model_config_with_session.return_value = None else: - app.app_model_config = MagicMock() - app.app_model_config.to_dict.return_value = {"user_input_form": form or []} + app_model_config = MagicMock() + app_model_config.to_dict.return_value = {"user_input_form": form or []} + app.app_model_config_with_session.return_value = app_model_config return app +def _session() -> MagicMock: + session = MagicMock() + session.scalar.return_value = None + return session + + def test_chat_mode_includes_query() -> None: app = _stub_app(AppMode.CHAT, form=[{"text-input": {"variable": "x", "label": "X", "required": True}}]) - schema = build_input_schema(app) + session = _session() + schema = build_input_schema(app, session=session) assert schema["$schema"] == "https://json-schema.org/draft/2020-12/schema" assert "query" in schema["properties"] assert schema["properties"]["query"]["type"] == "string" @@ -124,30 +133,32 @@ def test_chat_mode_includes_query() -> None: assert "query" in schema["required"] assert "inputs" in schema["required"] assert schema["properties"]["inputs"]["additionalProperties"] is False + app.app_model_config_with_session.assert_called_once_with(session=session) + app.app_model_config_with_session.return_value.to_dict.assert_called_once_with(annotation_reply={"enabled": False}) def test_agent_chat_mode_includes_query() -> None: app = _stub_app(AppMode.AGENT_CHAT, form=[]) - schema = build_input_schema(app) + schema = build_input_schema(app, session=_session()) assert "query" in schema["properties"] def test_advanced_chat_mode_includes_query() -> None: app = _stub_app(AppMode.ADVANCED_CHAT, form=[]) - schema = build_input_schema(app) + schema = build_input_schema(app, session=_session()) assert "query" in schema["properties"] def test_workflow_mode_omits_query() -> None: app = _stub_app(AppMode.WORKFLOW, form=[]) - schema = build_input_schema(app) + schema = build_input_schema(app, session=_session()) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] def test_completion_mode_omits_query() -> None: app = _stub_app(AppMode.COMPLETION, form=[]) - schema = build_input_schema(app) + schema = build_input_schema(app, session=_session()) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] @@ -160,20 +171,20 @@ def test_inputs_required_driven_by_form() -> None: {"text-input": {"variable": "context", "label": "Context", "required": False}}, ], ) - schema = build_input_schema(app) + schema = build_input_schema(app, session=_session()) assert schema["properties"]["inputs"]["required"] == ["industry"] def test_misconfigured_chat_raises_app_unavailable() -> None: app = _stub_app(AppMode.CHAT, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app) + build_input_schema(app, session=_session()) def test_misconfigured_workflow_raises_app_unavailable() -> None: app = _stub_app(AppMode.WORKFLOW, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app) + build_input_schema(app, session=_session()) def test_empty_input_schema_sentinel_shape() -> None: 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 c3afdf95ea5..13176cc3b27 100644 --- a/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py +++ b/api/tests/unit_tests/controllers/openapi/test_workspaces_members.py @@ -29,7 +29,7 @@ from flask import Flask from flask.views import MethodView from pydantic import ValidationError from sqlalchemy import Engine, select -from sqlalchemy.orm import Session, scoped_session, sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, NotFound, UnprocessableEntity from controllers.openapi import bp as openapi_bp @@ -96,15 +96,10 @@ def database_session(sqlite_engine: Engine): models = (Account, Tenant, TenantAccountJoin) tables = [model.metadata.tables[model.__tablename__] for model in models] TypeBase.metadata.create_all(sqlite_engine, tables=tables) - database_session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) - try: - with ( - patch.object(workspaces_module.db, "session", database_session), - patch("models.account.db", SimpleNamespace(engine=sqlite_engine)), - ): - yield database_session() - finally: - database_session.remove() + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + factory = SimpleNamespace(get_session_maker=lambda: session_maker, create_session=session_maker) + with patch("controllers.common.session.session_factory", factory), session_maker() as session: + yield session def _rule(app: Flask, path: str): 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 52695a26683..6d60af54e2e 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 ANY, Mock +from unittest.mock import MagicMock, Mock import pytest from flask import Flask @@ -35,36 +35,25 @@ from extensions.ext_redis import redis_client from models.model import App from services.annotation_service import AppAnnotationService -# --------------------------------------------------------------------------- -# Pydantic Model Tests -# --------------------------------------------------------------------------- - class TestAnnotationCreatePayload: """Test suite for AnnotationCreatePayload Pydantic model.""" def test_payload_with_question_and_answer(self): """Test payload with required fields.""" - payload = AnnotationCreatePayload( - question="What is AI?", - answer="AI is artificial intelligence.", - ) + payload = AnnotationCreatePayload(question="What is AI?", answer="AI is artificial intelligence.") assert payload.question == "What is AI?" assert payload.answer == "AI is artificial intelligence." def test_payload_with_unicode_content(self): """Test payload with unicode content.""" - payload = AnnotationCreatePayload( - question="什么是人工智能?", - answer="人工智能是模拟人类智能的技术。", - ) + payload = AnnotationCreatePayload(question="什么是人工智能?", answer="人工智能是模拟人类智能的技术。") assert payload.question == "什么是人工智能?" def test_payload_with_special_characters(self): """Test payload with special characters.""" payload = AnnotationCreatePayload( - question="What is AI?", - answer="AI & ML are related fields with 100% growth!", + question="What is AI?", answer="AI & ML are related fields with 100% growth!" ) assert "" in payload.question @@ -75,9 +64,7 @@ class TestAnnotationReplyActionPayload: def test_payload_with_all_fields(self): """Test payload with all fields.""" payload = AnnotationReplyActionPayload( - score_threshold=0.8, - embedding_provider_name="openai", - embedding_model_name="text-embedding-ada-002", + score_threshold=0.8, embedding_provider_name="openai", embedding_model_name="text-embedding-ada-002" ) assert payload.score_threshold == 0.8 assert payload.embedding_provider_name == "openai" @@ -86,18 +73,14 @@ class TestAnnotationReplyActionPayload: def test_payload_with_different_provider(self): """Test payload with different embedding provider.""" payload = AnnotationReplyActionPayload( - score_threshold=0.75, - embedding_provider_name="azure_openai", - embedding_model_name="text-embedding-3-small", + score_threshold=0.75, embedding_provider_name="azure_openai", embedding_model_name="text-embedding-3-small" ) assert payload.embedding_provider_name == "azure_openai" def test_payload_with_zero_threshold(self): """Test payload with zero score threshold.""" payload = AnnotationReplyActionPayload( - score_threshold=0.0, - embedding_provider_name="local", - embedding_model_name="default", + score_threshold=0.0, embedding_provider_name="local", embedding_model_name="default" ) assert payload.score_threshold == 0.0 @@ -105,14 +88,12 @@ class TestAnnotationReplyActionPayload: class TestAnnotationListQuery: def test_defaults(self) -> None: query = AnnotationListQuery.model_validate({}) - assert query.page == 1 assert query.limit == 20 assert query.keyword == "" def test_valid_numeric_strings(self) -> None: query = AnnotationListQuery.model_validate({"page": "2", "limit": "5", "keyword": "refund"}) - assert query.page == 2 assert query.limit == 5 assert query.keyword == "refund" @@ -124,11 +105,6 @@ class TestAnnotationListQuery: AnnotationListQuery.model_validate({field: value}) -# --------------------------------------------------------------------------- -# Model and Error Pattern Tests -# --------------------------------------------------------------------------- - - class TestAppModelPatterns: """Test App model patterns used by annotation controller.""" @@ -138,7 +114,6 @@ class TestAppModelPatterns: app.id = str(uuid.uuid4()) app.status = "normal" app.enable_api = True - assert app.id is not None assert app.status == "normal" assert app.enable_api @@ -154,7 +129,6 @@ class TestAppModelPatterns: """Test app with archived status.""" app = Mock(spec=App) app.status = "archived" - assert app.status == "archived" @@ -185,18 +159,15 @@ class TestAnnotationReplyActionApi: def test_enable(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: enable_mock = Mock(return_value={"job_id": "job-1", "job_status": "waiting"}) monkeypatch.setattr(AppAnnotationService, "enable_app_annotation", enable_mock) - api = AnnotationReplyActionApi() handler = unwrap(api.post) app_model = SimpleNamespace(id="app", tenant_id="tenant") - with app.test_request_context( "/apps/annotation-reply/enable", method="POST", json={"score_threshold": 0.5, "embedding_provider_name": "p", "embedding_model_name": "m"}, ): response, status = handler(api, app_model=app_model, action="enable") - assert status == 200 assert response == {"job_id": "job-1", "job_status": "waiting"} enable_mock.assert_called_once() @@ -204,18 +175,15 @@ class TestAnnotationReplyActionApi: def test_disable(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: disable_mock = Mock(return_value={"job_id": "job-1", "job_status": "waiting"}) monkeypatch.setattr(AppAnnotationService, "disable_app_annotation", disable_mock) - api = AnnotationReplyActionApi() handler = unwrap(api.post) app_model = SimpleNamespace(id="app", tenant_id="tenant") - with app.test_request_context( "/apps/annotation-reply/disable", method="POST", json={"score_threshold": 0.5, "embedding_provider_name": "p", "embedding_model_name": "m"}, ): response, status = handler(api, app_model=app_model, action="disable") - assert status == 200 assert response == {"job_id": "job-1", "job_status": "waiting"} disable_mock.assert_called_once() @@ -224,28 +192,24 @@ class TestAnnotationReplyActionApi: class TestAnnotationReplyActionStatusApi: def test_missing_job(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(redis_client, "get", lambda *_args, **_kwargs: None) - api = AnnotationReplyActionStatusApi() handler = unwrap(api.get) app_model = SimpleNamespace(id="app") - with pytest.raises(ValueError): handler(api, app_model=app_model, job_id="j1", action="enable") def test_error(self, monkeypatch: pytest.MonkeyPatch) -> None: + def _get(key): if "error" in key: return b"oops" return b"error" monkeypatch.setattr(redis_client, "get", _get) - api = AnnotationReplyActionStatusApi() handler = unwrap(api.get) app_model = SimpleNamespace(id="app") - response, status = handler(api, app_model=app_model, job_id="j1", action="enable") - assert status == 200 assert response["job_status"] == "error" assert response["error_msg"] == "oops" @@ -256,34 +220,32 @@ class TestAnnotationListApi: annotation = SimpleNamespace(id="a1", question="q", content="a", created_at=0) get_mock = Mock(return_value=([annotation], 1)) monkeypatch.setattr(AppAnnotationService, "get_annotation_list_by_app_id", get_mock) - api = AnnotationListApi() handler = unwrap(api.get) app_model = SimpleNamespace(id="app") - with app.test_request_context("/apps/annotations", method="GET"): - response = handler(api, app_model=app_model) - + response = handler(api, MagicMock(), app_model=app_model) assert response["page"] == 1 assert response["limit"] == 20 - get_mock.assert_called_once_with("app", 1, 20, "", session=ANY) + session = get_mock.call_args.args[-1] + assert isinstance(session, MagicMock) + assert get_mock.call_args.args == ("app", 1, 20, "", session) 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) get_mock = Mock(return_value=([annotation], 1)) monkeypatch.setattr(AppAnnotationService, "get_annotation_list_by_app_id", get_mock) - api = AnnotationListApi() handler = unwrap(api.get) app_model = SimpleNamespace(id="app") - with app.test_request_context("/apps/annotations?page=2&limit=5&keyword=refund", method="GET"): - response = handler(api, app_model=app_model) - + response = handler(api, MagicMock(), app_model=app_model) assert response["total"] == 1 assert response["page"] == 2 assert response["limit"] == 5 - get_mock.assert_called_once_with("app", 2, 5, "refund", session=ANY) + session = get_mock.call_args.args[-1] + assert isinstance(session, MagicMock) + assert get_mock.call_args.args == ("app", 2, 5, "refund", session) @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( @@ -291,32 +253,24 @@ class TestAnnotationListApi: ) -> None: get_mock = Mock(return_value=([], 0)) monkeypatch.setattr(AppAnnotationService, "get_annotation_list_by_app_id", get_mock) - api = AnnotationListApi() handler = unwrap(api.get) app_model = SimpleNamespace(id="app") - with app.test_request_context(f"/apps/annotations?{query_string}", method="GET"): with pytest.raises(ValidationError): - handler(api, app_model=app_model) - + handler(api, MagicMock(), app_model=app_model) get_mock.assert_not_called() def test_create(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: annotation = SimpleNamespace(id="a1", question="q", content="a", created_at=0) monkeypatch.setattr( - AppAnnotationService, - "insert_app_annotation_directly", - lambda *_args, **_kwargs: annotation, + AppAnnotationService, "insert_app_annotation_directly", lambda *_args, **_kwargs: annotation ) - api = AnnotationListApi() handler = unwrap(api.post) app_model = SimpleNamespace(id="app") - with app.test_request_context("/apps/annotations", method="POST", json={"question": "q", "answer": "a"}): - response, status = handler(api, app_model=app_model) - + response, status = handler(api, MagicMock(), app_model=app_model) assert status == HTTPStatus.CREATED assert response["question"] == "q" @@ -325,25 +279,18 @@ class TestAnnotationUpdateDeleteApi: def test_update_delete(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: annotation = SimpleNamespace(id="a1", question="q", content="a", created_at=0) monkeypatch.setattr( - AppAnnotationService, - "update_app_annotation_directly", - lambda *_args, **_kwargs: annotation, + AppAnnotationService, "update_app_annotation_directly", lambda *_args, **_kwargs: annotation ) delete_mock = Mock() monkeypatch.setattr(AppAnnotationService, "delete_app_annotation", delete_mock) - api = AnnotationUpdateDeleteApi() put_handler = unwrap(api.put) delete_handler = unwrap(api.delete) app_model = SimpleNamespace(id="app", tenant_id="tenant") - with app.test_request_context("/apps/annotations/1", method="PUT", json={"question": "q", "answer": "a"}): - response = put_handler(api, app_model=app_model, annotation_id="1") - + response = put_handler(api, MagicMock(), app_model=app_model, annotation_id="1") assert response["answer"] == "a" - with app.test_request_context("/apps/annotations/1", method="DELETE"): - response, status = delete_handler(api, app_model=app_model, annotation_id="1") - + response, status = delete_handler(api, MagicMock(), app_model=app_model, annotation_id="1") assert status == 204 delete_mock.assert_called_once() 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 04e9220ad55..c48cb343950 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 @@ -40,6 +40,8 @@ class TestAppParameterApi: app.mode = AppMode.CHAT app.status = "normal" app.enable_api = True + app.app_model_config_with_session.return_value = None + app.workflow_with_session.return_value = None return app @patch("controllers.service_api.wraps.user_logged_in") @@ -60,6 +62,7 @@ class TestAppParameterApi: "suggested_questions": [], } mock_app_model.app_model_config = mock_config + mock_app_model.app_model_config_with_session.return_value = mock_config mock_app_model.workflow = None # Mock authentication @@ -83,7 +86,13 @@ class TestAppParameterApi: setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) # Act - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): + with ( + app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}), + patch( + "controllers.service_api.app.app.load_annotation_reply_config", + return_value={"enabled": False}, + ), + ): api = AppParameterApi() response = api.get() @@ -108,6 +117,7 @@ class TestAppParameterApi: mock_workflow.features_dict = {"suggested_questions": []} mock_workflow.user_input_form.return_value = [{"type": "text", "label": "Input", "variable": "input"}] mock_app_model.workflow = mock_workflow + mock_app_model.workflow_with_session.return_value = mock_workflow mock_app_model.app_model_config = None # Mock authentication @@ -184,7 +194,7 @@ class TestAppParameterApi: assert response["user_input_form"] == [ {"text-input": {"label": "topic", "variable": "topic", "required": True}} ] - mock_get_agent_parameters.assert_called_once_with(mock_app_model) + mock_get_agent_parameters.assert_called_once_with(mock_app_model, session=ANY) @patch("controllers.service_api.wraps.user_logged_in") @patch("controllers.service_api.wraps.current_app") diff --git a/api/tests/unit_tests/controllers/service_api/app/test_audio.py b/api/tests/unit_tests/controllers/service_api/app/test_audio.py index ee4c7f8828f..28ebbb6f190 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_audio.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_audio.py @@ -167,6 +167,7 @@ class TestAudioServiceMockedBehavior: result = AudioService.transcript_asr( app_model=mock_app, file=mock_file, + session=Mock(), end_user="user_123", ) 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 0ac0dde0f5a..de3ffa82018 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 @@ -688,11 +688,8 @@ class TestChatApiController: # A well-formed but nonexistent conversation_id must fail fast as 404, before the # streaming generator is created. Previously the lookup only ran inside the generator, # so an invalid id surfaced as a hang instead of a clean error. - monkeypatch.setattr( - ConversationService, - "get_conversation", - lambda *_args, **_kwargs: (_ for _ in ()).throw(ConversationNotExistsError()), - ) + get_conversation_mock = Mock(side_effect=ConversationNotExistsError()) + monkeypatch.setattr(ConversationService, "get_conversation", get_conversation_mock) generate_mock = Mock(return_value={"text": "unused"}) monkeypatch.setattr(AppGenerateService, "generate", generate_mock) @@ -711,6 +708,7 @@ class TestChatApiController: # The lookup must run before generation, so the generator is never started. generate_mock.assert_not_called() + assert get_conversation_mock.call_args.kwargs["session"] is orm_session class TestChatStopApiController: diff --git a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py index 5ab1bd4ff5b..129220cbc9c 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py @@ -353,6 +353,7 @@ class TestHitlServiceApi: app_model.tenant_id = "tenant-id" app_model.max_active_requests = 0 app_model.is_agent = False + app_model.is_agent_with_session.return_value = False user = MagicMock() user.id = "user-id" 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 0b1ca8741a9..c7bac28b694 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 @@ -14,8 +14,9 @@ Focus on: - API endpoint business logic and error handling """ +import inspect import uuid -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -40,6 +41,23 @@ from models.enums import IndexingStatus, SegmentType from services.dataset_service import DocumentService, SegmentService +def _session_factory_mock(): + mock_factory = MagicMock() + mock_session = MagicMock() + mock_factory.session = mock_session + + transaction = MagicMock() + transaction.__enter__.return_value = mock_session + transaction.__exit__.return_value = None + mock_factory.get_session_maker.return_value.begin.return_value = transaction + + read_session = MagicMock() + read_session.__enter__.return_value = mock_session + read_session.__exit__.return_value = None + mock_factory.create_session.return_value = read_session + return mock_factory + + def _segment_response_dict(summary: str | None = None): return { "id": "seg-1", @@ -393,6 +411,7 @@ class TestSegmentServiceMockedBehavior: tenant_id=mock_document.tenant_id, page=1, limit=20, + session=MagicMock(), ) assert len(segments) == 2 @@ -483,6 +502,7 @@ class TestChildChunkServiceMockedBehavior: dataset_id=str(uuid.uuid4()), page=1, limit=20, + session=MagicMock(), ) assert len(result.items) == 2 @@ -885,10 +905,10 @@ class TestSegmentApiGet: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_segments_success( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -902,7 +922,7 @@ class TestSegmentApiGet: """Test successful segment list retrieval.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset( mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX ) @@ -923,14 +943,15 @@ class TestSegmentApiGet: assert "data" in response assert "total" in response assert response["page"] == 1 + mock_dump_segments.assert_called_once_with([mock_segment], {}, session=session_factory.session) @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") - def test_list_segments_dataset_not_found(self, mock_db, mock_account_fn, app, mock_tenant, mock_dataset): + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) + def test_list_segments_dataset_not_found(self, session_factory, mock_account_fn, app, mock_tenant, mock_dataset): """Test 404 when dataset not found.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None # Act & Assert with app.test_request_context( @@ -943,14 +964,14 @@ class TestSegmentApiGet: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_segments_document_not_found( - self, mock_db, mock_account_fn, mock_doc_svc, app, mock_tenant, mock_dataset + self, session_factory, mock_account_fn, mock_doc_svc, app, mock_tenant, mock_dataset ): """Test 404 when document not found.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = None # Act & Assert @@ -999,14 +1020,14 @@ class TestSegmentApiPost: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_segments_success( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1023,7 +1044,7 @@ class TestSegmentApiPost: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" @@ -1046,23 +1067,28 @@ class TestSegmentApiPost: headers={"Authorization": "Bearer test_token"}, ): api = SegmentApi() - response, status = api.post(tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id") + response, status = api.post( + tenant_id=mock_tenant.id, + dataset_id=mock_dataset.id, + document_id="doc-id", + ) # Assert assert status == 200 assert "data" in response assert "doc_form" in response + mock_dump_segments.assert_called_once_with([mock_segment], {}, session=session_factory.session) @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_segments_missing_segments( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, app: Flask, @@ -1075,7 +1101,7 @@ class TestSegmentApiPost: mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" @@ -1090,7 +1116,11 @@ class TestSegmentApiPost: headers={"Authorization": "Bearer test_token"}, ): api = SegmentApi() - response, status = api.post(tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id") + response, status = api.post( + tenant_id=mock_tenant.id, + dataset_id=mock_dataset.id, + document_id="doc-id", + ) # Assert assert status == 400 @@ -1098,14 +1128,14 @@ class TestSegmentApiPost: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_segments_document_not_completed( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, app: Flask, @@ -1117,7 +1147,7 @@ class TestSegmentApiPost: self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "indexing" # Not completed @@ -1132,7 +1162,11 @@ class TestSegmentApiPost: ): api = SegmentApi() with pytest.raises(NotFound): - api.post(tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id") + api.post( + tenant_id=mock_tenant.id, + dataset_id=mock_dataset.id, + document_id="doc-id", + ) class TestDatasetSegmentApiDelete: @@ -1143,19 +1177,14 @@ class TestDatasetSegmentApiDelete: unwrapped method directly to bypass the billing decorator. """ - @staticmethod - def _call_delete(api: DatasetSegmentApi, **kwargs): - """Call the unwrapped delete to skip billing decorators.""" - return api.delete.__wrapped__(api, **kwargs) - @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_segment_success( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_dataset_svc, @@ -1168,7 +1197,7 @@ class TestDatasetSegmentApiDelete: """Test successful segment deletion.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc = _document_for_dataset(mock_dataset) @@ -1183,8 +1212,10 @@ class TestDatasetSegmentApiDelete: method="DELETE", ): api = DatasetSegmentApi() - response = self._call_delete( + delete = inspect.unwrap(api.delete) + response = delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -1193,15 +1224,17 @@ 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, session_factory.session + ) @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_segment_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1212,7 +1245,7 @@ class TestDatasetSegmentApiDelete: """Test 404 when segment not found.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc = _document_for_dataset(mock_dataset) mock_doc.indexing_status = "completed" @@ -1228,9 +1261,11 @@ class TestDatasetSegmentApiDelete: method="DELETE", ): api = DatasetSegmentApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -1240,10 +1275,10 @@ class TestDatasetSegmentApiDelete: @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_segment_dataset_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_dataset_svc, @@ -1254,7 +1289,7 @@ class TestDatasetSegmentApiDelete: """Test 404 when dataset not found for delete.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None # Act & Assert with app.test_request_context( @@ -1262,9 +1297,11 @@ class TestDatasetSegmentApiDelete: method="DELETE", ): api = DatasetSegmentApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -1274,10 +1311,10 @@ class TestDatasetSegmentApiDelete: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_segment_document_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1288,7 +1325,7 @@ class TestDatasetSegmentApiDelete: """Test 404 when document not found for delete.""" # Arrange mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc_svc.get_document.return_value = None @@ -1298,9 +1335,11 @@ class TestDatasetSegmentApiDelete: method="DELETE", ): api = DatasetSegmentApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -1341,14 +1380,14 @@ class TestDatasetSegmentApiUpdate: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_update_segment_success( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1364,7 +1403,7 @@ class TestDatasetSegmentApiUpdate: self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc_svc.get_document.return_value = _document_for_dataset( mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX @@ -1393,18 +1432,19 @@ class TestDatasetSegmentApiUpdate: assert status == 200 assert "data" in response mock_seg_svc.update_segment.assert_called_once() + mock_dump_segment.assert_called_once_with(updated, None, session=session_factory.session) @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_update_segment_dataset_not_found( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1415,7 +1455,7 @@ class TestDatasetSegmentApiUpdate: """Test 404 when dataset not found for update.""" self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id", @@ -1436,14 +1476,14 @@ class TestDatasetSegmentApiUpdate: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_update_segment_not_found( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1456,7 +1496,7 @@ class TestDatasetSegmentApiUpdate: self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = None @@ -1490,10 +1530,10 @@ class TestDatasetSegmentApiGetSingle: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_get_single_segment_success( self, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1507,7 +1547,7 @@ class TestDatasetSegmentApiGetSingle: ): """Test successful single segment retrieval.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc = _document_for_dataset(mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX) mock_doc_svc.get_document.return_value = mock_doc @@ -1530,6 +1570,7 @@ class TestDatasetSegmentApiGetSingle: assert status == 200 assert "data" in response assert response["doc_form"] == IndexStructureType.PARAGRAPH_INDEX + mock_dump_segment.assert_called_once_with(mock_segment, None, session=session_factory.session) @patch("controllers.service_api.dataset.segment.segment_response_with_summary") @patch("controllers.service_api.dataset.segment.SummaryIndexService.get_segment_summary") @@ -1537,10 +1578,10 @@ class TestDatasetSegmentApiGetSingle: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_get_single_segment_includes_summary( self, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1554,7 +1595,7 @@ class TestDatasetSegmentApiGetSingle: ): """Test that single segment response includes summary content from SummaryIndexService.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc = _document_for_dataset(mock_dataset, doc_form=IndexStructureType.PARAGRAPH_INDEX) mock_doc_svc.get_document.return_value = mock_doc @@ -1577,12 +1618,17 @@ class TestDatasetSegmentApiGetSingle: assert status == 200 assert response["data"]["summary"] == "This is the segment summary" + mock_dump_segment.assert_called_once_with( + mock_segment, + "This is the segment summary", + session=session_factory.session, + ) @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_get_single_segment_dataset_not_found( self, - mock_db, + session_factory, mock_account_fn, app: Flask, mock_tenant, @@ -1590,7 +1636,7 @@ class TestDatasetSegmentApiGetSingle: ): """Test 404 when dataset not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id", @@ -1608,10 +1654,10 @@ class TestDatasetSegmentApiGetSingle: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_get_single_segment_document_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1621,7 +1667,7 @@ class TestDatasetSegmentApiGetSingle: ): """Test 404 when document not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc_svc.get_document.return_value = None @@ -1642,10 +1688,10 @@ class TestDatasetSegmentApiGetSingle: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.DatasetService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_get_single_segment_segment_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_dataset_svc, mock_doc_svc, @@ -1656,7 +1702,7 @@ class TestDatasetSegmentApiGetSingle: ): """Test 404 when segment not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_dataset_svc.check_dataset_model_setting.return_value = None mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = None @@ -1685,10 +1731,10 @@ class TestChildChunkApiGet: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_child_chunks_success( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1698,7 +1744,7 @@ class TestChildChunkApiGet: ): """Test successful child chunk list retrieval.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = Mock() @@ -1725,10 +1771,10 @@ class TestChildChunkApiGet: assert response["page"] == 1 @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_child_chunks_dataset_not_found( self, - mock_db, + session_factory, mock_account_fn, app: Flask, mock_tenant, @@ -1736,7 +1782,7 @@ class TestChildChunkApiGet: ): """Test 404 when dataset not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id/child_chunks", @@ -1753,10 +1799,10 @@ class TestChildChunkApiGet: @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_child_chunks_document_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, app: Flask, @@ -1765,7 +1811,7 @@ class TestChildChunkApiGet: ): """Test 404 when document not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = None with app.test_request_context( @@ -1784,10 +1830,10 @@ class TestChildChunkApiGet: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_list_child_chunks_segment_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1797,7 +1843,7 @@ class TestChildChunkApiGet: ): """Test 404 when segment not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = None @@ -1841,14 +1887,14 @@ class TestChildChunkApiPost: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_child_chunk_success( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1860,7 +1906,7 @@ class TestChildChunkApiPost: self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = Mock() mock_child = _child_chunk() @@ -1884,14 +1930,14 @@ class TestChildChunkApiPost: assert "data" in response @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_child_chunk_dataset_not_found( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, app: Flask, mock_tenant, @@ -1900,7 +1946,7 @@ class TestChildChunkApiPost: """Test 404 when dataset not found.""" self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = None + session_factory.session.scalar.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/doc-id/segments/seg-id/child_chunks", @@ -1920,14 +1966,14 @@ class TestChildChunkApiPost: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") def test_create_child_chunk_segment_not_found( self, mock_validate_token, mock_feature_svc, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1938,7 +1984,7 @@ class TestChildChunkApiPost: """Test 404 when segment not found.""" self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant.id) mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) mock_seg_svc.get_segment_by_ref.return_value = None @@ -1967,21 +2013,13 @@ class TestDatasetChildChunkApiDelete: through both layers. """ - @staticmethod - def _call_delete(api: DatasetChildChunkApi, **kwargs): - """Unwrap through both decorator layers.""" - fn = api.delete - while hasattr(fn, "__wrapped__"): - fn = fn.__wrapped__ - return fn(api, **kwargs) - @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_child_chunk_success( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -1991,7 +2029,7 @@ class TestDatasetChildChunkApiDelete: ): """Test successful child chunk deletion.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc = _document_for_dataset(mock_dataset) mock_doc_svc.get_document.return_value = mock_doc @@ -2013,8 +2051,10 @@ class TestDatasetChildChunkApiDelete: method="DELETE", ): api = DatasetChildChunkApi() - response = self._call_delete( + delete = inspect.unwrap(api.delete) + response = delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -2028,10 +2068,10 @@ class TestDatasetChildChunkApiDelete: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_child_chunk_not_found( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -2041,7 +2081,7 @@ class TestDatasetChildChunkApiDelete: ): """Test 404 when child chunk not found.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) @@ -2056,9 +2096,11 @@ class TestDatasetChildChunkApiDelete: method="DELETE", ): api = DatasetChildChunkApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -2069,10 +2111,10 @@ class TestDatasetChildChunkApiDelete: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_child_chunk_segment_document_mismatch( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -2082,7 +2124,7 @@ class TestDatasetChildChunkApiDelete: ): """Test 404 when segment does not belong to the document.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) @@ -2093,9 +2135,11 @@ class TestDatasetChildChunkApiDelete: method="DELETE", ): api = DatasetChildChunkApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", @@ -2106,10 +2150,10 @@ class TestDatasetChildChunkApiDelete: @patch("controllers.service_api.dataset.segment.SegmentService") @patch("controllers.service_api.dataset.segment.DocumentService") @patch("controllers.service_api.dataset.segment.current_account_with_tenant") - @patch("controllers.service_api.dataset.segment.db") + @patch("controllers.common.session.session_factory", new_callable=_session_factory_mock) def test_delete_child_chunk_wrong_segment( self, - mock_db, + session_factory, mock_account_fn, mock_doc_svc, mock_seg_svc, @@ -2119,7 +2163,7 @@ class TestDatasetChildChunkApiDelete: ): """Test 404 when child chunk does not belong to the segment.""" mock_account_fn.return_value = (Mock(), mock_tenant.id) - mock_db.session.scalar.return_value = mock_dataset + session_factory.session.scalar.return_value = mock_dataset mock_doc_svc.get_document.return_value = _document_for_dataset(mock_dataset) segment_id = str(uuid.uuid4()) @@ -2135,9 +2179,11 @@ class TestDatasetChildChunkApiDelete: method="DELETE", ): api = DatasetChildChunkApi() + delete = inspect.unwrap(api.delete) with pytest.raises(NotFound): - self._call_delete( + delete( api, + session_factory.session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, document_id="doc-id", 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 ef379afe5c6..9713caabeb5 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 @@ -15,11 +15,12 @@ Focus on: - API endpoint business logic and error handling """ +import inspect import json import uuid -from dataclasses import dataclass, field +from dataclasses import dataclass from datetime import UTC, datetime -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -52,27 +53,13 @@ def _document_data_source_info() -> dict[str, str]: return {"type": "website_crawl", "url": "https://example.com/docs", "title": "Docs"} -@dataclass -class _DocumentModelSessionStub: - scalar_values: list[object] = field(default_factory=list) - - def scalar(self, *args: object, **kwargs: object) -> object: - if self.scalar_values: - return self.scalar_values.pop(0) - return None - - def scalars(self, *args: object, **kwargs: object) -> object: - result = Mock() - result.all.return_value = [] - return result - - def get(self, *args: object, **kwargs: object) -> None: - return None - - -@dataclass -class _DocumentModelDbStub: - session: _DocumentModelSessionStub = field(default_factory=_DocumentModelSessionStub) +def _unwrap_non_wrapped_controller(view): + while view.__closure__: + inner_functions = [cell.cell_contents for cell in view.__closure__ if inspect.isfunction(cell.cell_contents)] + if not inner_functions: + break + view = inner_functions[-1] + return inspect.unwrap(view) @dataclass @@ -164,7 +151,7 @@ def _expected_document_response(document: Document) -> dict[str, object]: "position": document.position, "data_source_type": document.data_source_type, "data_source_info": document.data_source_info_dict, - "data_source_detail_dict": document.data_source_detail_dict, + "data_source_detail_dict": document.data_source_info_dict, "dataset_process_rule_id": document.dataset_process_rule_id, "name": document.name, "created_from": document.created_from, @@ -621,17 +608,18 @@ class TestDocumentServiceSaveValidation: pass mock_check_form.side_effect = TestStopError() + session = Mock() # Skip actual logic by mocking dependent calls or raising error to stop early with pytest.raises(TestStopError): # We just want to check check_doc_form is called early - DocumentService.save_document_with_dataset_id(dataset, config, Mock(), session=Mock()) + DocumentService.save_document_with_dataset_id(dataset, config, Mock(), session=session) # This will fail if we raise exception before check_doc_form, # but check_doc_form is the first thing called. # Ideally we'd mock everything to completion, but for unit validation: # We can just verify check_doc_form was called if we mock it to not raise. - mock_check_form.assert_called_once() + mock_check_form.assert_called_once_with(dataset, config.doc_form, session=session) # ============================================================================= @@ -661,7 +649,6 @@ class TestDocumentApiGet: id=str(uuid.uuid4()), tenant_id=mock_tenant, name="test_document.txt", - dataset_process_rule_id=str(uuid.uuid4()), word_count=100, ) @@ -682,6 +669,8 @@ class TestDocumentApiGet: mock_doc_svc.get_document.return_value = mock_doc_detail mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} + session = MagicMock() + session.scalar.side_effect = [5, 0] # Act with app.test_request_context( @@ -689,11 +678,14 @@ class TestDocumentApiGet: method="GET", ): api = DocumentApi() - with ( - patch.object(api, "get_dataset", return_value=mock_dataset), - patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([5, 5, 5, 5, 0]))), - ): - response = api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_doc_detail.id) + with patch.object(api, "get_dataset", return_value=mock_dataset): + response = inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_doc_detail.id, + ) # Assert assert response == { @@ -739,6 +731,8 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant) mock_doc_svc.get_document.return_value = None + session = MagicMock() + session.scalar.return_value = mock_dataset # Act & Assert with app.test_request_context( @@ -748,7 +742,13 @@ class TestDocumentApiGet: api = DocumentApi() with patch.object(api, "get_dataset", return_value=mock_dataset): with pytest.raises(NotFound): - api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id="nonexistent") + inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id="nonexistent", + ) @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_forbidden_wrong_tenant( @@ -761,6 +761,8 @@ class TestDocumentApiGet: mock_doc_detail.tenant_id = "different-tenant-id" mock_doc_svc.get_document.return_value = mock_doc_detail + session = MagicMock() + session.scalar.return_value = mock_dataset # Act & Assert with app.test_request_context( @@ -770,7 +772,13 @@ class TestDocumentApiGet: api = DocumentApi() with patch.object(api, "get_dataset", return_value=mock_dataset): with pytest.raises(Forbidden): - api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_doc_detail.id) + inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_doc_detail.id, + ) @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_metadata_only( @@ -782,6 +790,8 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + session = MagicMock() + session.scalar.return_value = mock_dataset # Act with app.test_request_context( @@ -790,7 +800,13 @@ class TestDocumentApiGet: ): api = DocumentApi() with patch.object(api, "get_dataset", return_value=mock_dataset): - response = api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_doc_detail.id) + response = inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_doc_detail.id, + ) # Assert — metadata='only' returns only id, doc_type, doc_metadata assert response["id"] == mock_doc_detail.id @@ -817,6 +833,8 @@ class TestDocumentApiGet: mock_doc_svc.get_document.return_value = mock_doc_detail mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} + session = MagicMock() + session.scalar.side_effect = [5, 0] # Act with app.test_request_context( @@ -824,11 +842,14 @@ class TestDocumentApiGet: method="GET", ): api = DocumentApi() - with ( - patch.object(api, "get_dataset", return_value=mock_dataset), - patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([5, 5, 5, 5, 0]))), - ): - response = api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_doc_detail.id) + with patch.object(api, "get_dataset", return_value=mock_dataset): + response = inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_doc_detail.id, + ) # Assert — metadata='without' omits doc_type / doc_metadata assert response["id"] == mock_doc_detail.id @@ -880,6 +901,8 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + session = MagicMock() + session.scalar.return_value = mock_dataset # Act & Assert with app.test_request_context( @@ -889,7 +912,13 @@ class TestDocumentApiGet: api = DocumentApi() with patch.object(api, "get_dataset", return_value=mock_dataset): with pytest.raises(InvalidMetadataError): - api.get(tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_doc_detail.id) + inspect.unwrap(type(api).get)( + api, + session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_doc_detail.id, + ) class TestDocumentApiDelete: @@ -897,17 +926,10 @@ class TestDocumentApiDelete: ``delete`` is wrapped by ``@cloud_edition_billing_rate_limit_check`` which internally calls ``validate_and_get_api_token``. To bypass the decorator - we call the original function via ``__wrapped__`` (preserved by - ``functools.wraps``). ``delete`` loads the dataset via - ``db.session.scalar(select(Dataset)...)``, so we patch ``db`` at the - controller module. + we call the original function via ``inspect.unwrap`` (preserved by + ``functools.wraps``) and pass the test session explicitly. """ - @staticmethod - def _call_delete(api: DocumentApi, **kwargs): - """Call the unwrapped delete to skip billing decorators.""" - return api.delete.__wrapped__(api, **kwargs) - @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") def test_delete_document_success(self, mock_db, mock_doc_svc, app: Flask, mock_tenant, mock_document): @@ -927,13 +949,18 @@ class TestDocumentApiDelete: method="DELETE", ): api = DocumentApi() - response = self._call_delete( - api, tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_document.id + delete = inspect.unwrap(type(api).delete) + response = delete( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_document.id, ) # 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") @@ -953,8 +980,15 @@ class TestDocumentApiDelete: method="DELETE", ): api = DocumentApi() + delete = inspect.unwrap(type(api).delete) with pytest.raises(NotFound): - self._call_delete(api, tenant_id=mock_tenant, dataset_id=dataset_id, document_id=document_id) + delete( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=document_id, + ) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") @@ -974,8 +1008,15 @@ class TestDocumentApiDelete: method="DELETE", ): api = DocumentApi() + delete = inspect.unwrap(type(api).delete) with pytest.raises(ArchivedDocumentImmutableError): - self._call_delete(api, tenant_id=mock_tenant, dataset_id=dataset_id, document_id=mock_document.id) + delete( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=mock_document.id, + ) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") @@ -992,8 +1033,15 @@ class TestDocumentApiDelete: method="DELETE", ): api = DocumentApi() + delete = inspect.unwrap(type(api).delete) with pytest.raises(ValueError, match="Dataset does not exist."): - self._call_delete(api, tenant_id=mock_tenant, dataset_id=dataset_id, document_id=document_id) + delete( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=dataset_id, + document_id=document_id, + ) class TestDocumentListApi: @@ -1005,7 +1053,7 @@ class TestDocumentListApi: def test_list_documents_success(self, mock_db, mock_doc_svc, mock_paginate, app: Flask, mock_tenant, mock_dataset): """Test successful document list retrieval.""" # Arrange - mock_db.session.scalar.return_value = mock_dataset + mock_db.session.scalar.side_effect = [mock_dataset, 0, 0] documents = [ make_serializable_document( @@ -1028,8 +1076,9 @@ class TestDocumentListApi: method="GET", ): api = DocumentListApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([0, 0]))): - response = api.get(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + response = inspect.unwrap(type(api).get)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) # Assert assert response == { @@ -1055,7 +1104,7 @@ class TestDocumentListApi: ): api = DocumentListApi() with pytest.raises(NotFound): - api.get(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + inspect.unwrap(type(api).get)(api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id) class TestDocumentIndexingStatusApi: @@ -1085,7 +1134,13 @@ class TestDocumentIndexingStatusApi: method="GET", ): api = DocumentIndexingStatusApi() - response = api.get(tenant_id=mock_tenant, dataset_id=mock_dataset.id, batch=batch_id) + response = inspect.unwrap(type(api).get)( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + batch=batch_id, + ) # Assert assert response == { @@ -1121,7 +1176,13 @@ class TestDocumentIndexingStatusApi: ): api = DocumentIndexingStatusApi() with pytest.raises(NotFound): - api.get(tenant_id=mock_tenant, dataset_id=mock_dataset.id, batch=batch_id) + inspect.unwrap(type(api).get)( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + batch=batch_id, + ) @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") @@ -1141,7 +1202,13 @@ class TestDocumentIndexingStatusApi: ): api = DocumentIndexingStatusApi() with pytest.raises(NotFound): - api.get(tenant_id=mock_tenant, dataset_id=mock_dataset.id, batch=batch_id) + inspect.unwrap(type(api).get)( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + batch=batch_id, + ) class TestDocumentAddByTextApi: @@ -1206,7 +1273,7 @@ class TestDocumentAddByTextApi: # Arrange — neutralise billing decorators self._setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) - mock_db.session.scalar.return_value = mock_dataset + mock_db.session.scalar.side_effect = [mock_dataset, None, 0] mock_dataset.indexing_technique = "economy" mock_current_user.id = str(uuid.uuid4()) @@ -1235,8 +1302,9 @@ class TestDocumentAddByTextApi: headers={"Authorization": "Bearer test_token"}, ): api = DocumentAddByTextApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([object(), 0]))): - response, status = api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + response, status = _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) # Assert assert (response, status) == ( @@ -1244,6 +1312,7 @@ class TestDocumentAddByTextApi: 200, ) assert "data_source_info_dict" not in response["document"] + assert mock_doc_svc.save_document_with_dataset_id.call_args.kwargs["session"] is mock_db.session @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") @@ -1266,7 +1335,9 @@ class TestDocumentAddByTextApi: ): api = DocumentAddByTextApi() with pytest.raises(ValueError, match="Dataset does not exist."): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") @@ -1295,7 +1366,9 @@ class TestDocumentAddByTextApi: ): api = DocumentAddByTextApi() with pytest.raises(ValueError, match="indexing_technique is required."): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) class TestArchivedDocumentImmutableError: @@ -1390,7 +1463,7 @@ class TestDocumentUpdateByTextApiPost: """Test successful document update by text.""" _setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) mock_dataset.indexing_technique = "economy" - mock_db.session.scalar.return_value = mock_dataset + mock_db.session.scalar.side_effect = [mock_dataset, None, 0] mock_current_user.id = "user-1" mock_upload = Mock() @@ -1409,12 +1482,13 @@ class TestDocumentUpdateByTextApiPost: headers={"Authorization": "Bearer test_token"}, ): api = DocumentUpdateByTextApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([object(), 0]))): - response, status = api.post( - tenant_id=mock_tenant, - dataset_id=mock_dataset.id, - document_id=doc_id, - ) + response, status = _unwrap_non_wrapped_controller(type(api).post)( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + document_id=doc_id, + ) assert (response, status) == ( {"document": _expected_document_response(mock_document), "batch": "batch-1"}, @@ -1446,7 +1520,9 @@ class TestDocumentUpdateByTextApiPost: ): api = DocumentUpdateByTextApi() with pytest.raises(ValueError, match="Dataset does not exist"): - api.post( + _unwrap_non_wrapped_controller(type(api).post)( + api, + mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id, document_id=doc_id, @@ -1483,7 +1559,7 @@ class TestDocumentAddByFileApiPost: mock_dataset.provider = "vendor" mock_dataset.indexing_technique = "economy" mock_dataset.chunk_structure = None - mock_db.session.scalar.return_value = mock_dataset + mock_db.session.scalar.side_effect = [mock_dataset, 0] mock_current_user.id = "user-1" mock_upload = Mock() @@ -1496,7 +1572,10 @@ class TestDocumentAddByFileApiPost: from io import BytesIO - data = {"file": (BytesIO(b"content"), "test.pdf", "application/pdf")} + data = { + "file": (BytesIO(b"content"), "test.pdf", "application/pdf"), + "data": json.dumps({"process_rule": {"mode": "automatic", "rules": None}}), + } with app.test_request_context( f"/datasets/{mock_dataset.id}/document/create-by-file", method="POST", @@ -1505,8 +1584,9 @@ class TestDocumentAddByFileApiPost: headers={"Authorization": "Bearer test_token"}, ): api = DocumentAddByFileApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([object(), 0]))): - response, status = api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + response, status = _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) assert (response, status) == ( {"document": _expected_document_response(mock_document), "batch": "batch-file"}, @@ -1541,7 +1621,9 @@ class TestDocumentAddByFileApiPost: ): api = DocumentAddByFileApi() with pytest.raises(ValueError, match="Dataset does not exist"): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) @patch("controllers.service_api.dataset.document.db") @patch("controllers.service_api.wraps.FeatureService") @@ -1572,7 +1654,9 @@ class TestDocumentAddByFileApiPost: ): api = DocumentAddByFileApi() with pytest.raises(ValueError, match="External datasets"): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) @patch("controllers.service_api.dataset.document.db") @patch("controllers.service_api.wraps.FeatureService") @@ -1604,7 +1688,9 @@ class TestDocumentAddByFileApiPost: ): api = DocumentAddByFileApi() with pytest.raises(NoFileUploadedError): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) @patch("controllers.service_api.dataset.document.db") @patch("controllers.service_api.wraps.FeatureService") @@ -1637,7 +1723,9 @@ class TestDocumentAddByFileApiPost: ): api = DocumentAddByFileApi() with pytest.raises(ValueError, match="indexing_technique is required"): - api.post(tenant_id=mock_tenant, dataset_id=mock_dataset.id) + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) class TestDocumentUpdateByFileApiPatch: @@ -1669,24 +1757,28 @@ class TestDocumentUpdateByFileApiPatch: ) doc_id = str(uuid.uuid4()) + session = MagicMock() + session.scalar.return_value = 0 with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/{doc_id}/{route_name}", method="POST", headers={"Authorization": "Bearer test_token"}, ): api = DeprecatedDocumentUpdateByFileApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([0]))): - response, status = api.post( - tenant_id=mock_tenant, - dataset_id=mock_dataset.id, - document_id=doc_id, - ) + response, status = _unwrap_non_wrapped_controller(type(api).post)( + api, + session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + document_id=doc_id, + ) assert (response, status) == ( {"document": _expected_document_response(mock_update_document_by_file.return_value[0]), "batch": "batch-1"}, 200, ) mock_update_document_by_file.assert_called_once_with( + session=session, tenant_id=mock_tenant, dataset_id=mock_dataset.id, document_id=doc_id, @@ -1721,7 +1813,9 @@ class TestDocumentUpdateByFileApiPatch: ): api = DocumentApi() with pytest.raises(ValueError, match="Dataset does not exist"): - api.patch( + _unwrap_non_wrapped_controller(type(api).patch)( + api, + mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id, document_id=doc_id, @@ -1757,7 +1851,9 @@ class TestDocumentUpdateByFileApiPatch: ): api = DocumentApi() with pytest.raises(ValueError, match="External datasets"): - api.patch( + _unwrap_non_wrapped_controller(type(api).patch)( + api, + mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id, document_id=doc_id, @@ -1786,7 +1882,7 @@ class TestDocumentUpdateByFileApiPatch: mock_dataset.indexing_technique = "economy" mock_dataset.provider = "vendor" mock_dataset.chunk_structure = None - mock_db.session.scalar.return_value = mock_dataset + mock_db.session.scalar.side_effect = [mock_dataset, None, 0] mock_current_user.id = "user-1" mock_upload = Mock() @@ -1809,12 +1905,13 @@ class TestDocumentUpdateByFileApiPatch: headers={"Authorization": "Bearer test_token"}, ): api = DocumentApi() - with patch("models.dataset.db", _DocumentModelDbStub(_DocumentModelSessionStub([object(), 0]))): - response, status = api.patch( - tenant_id=mock_tenant, - dataset_id=mock_dataset.id, - document_id=doc_id, - ) + response, status = _unwrap_non_wrapped_controller(type(api).patch)( + api, + mock_db.session, + tenant_id=mock_tenant, + dataset_id=mock_dataset.id, + document_id=doc_id, + ) assert (response, status) == ( {"document": _expected_document_response(mock_document), "batch": "batch-1"}, 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 dd1322a6344..c30762911e2 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 @@ -17,7 +17,7 @@ Decorator strategy: import uuid from inspect import unwrap -from unittest.mock import ANY, Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -64,8 +64,8 @@ class TestDatasetMetadataCreatePost: """ @staticmethod - def _call_post(api, **kwargs): - return unwrap(api.post)(api, **kwargs) + def _call_post(api, session: MagicMock, **kwargs): + return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -91,8 +91,10 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() + session = MagicMock() response, status = self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, ) @@ -118,9 +120,11 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() + session = MagicMock() with pytest.raises(NotFound): self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, ) @@ -194,8 +198,8 @@ class TestDatasetMetadataServiceApiPatch: """ @staticmethod - def _call_patch(api, **kwargs): - return unwrap(api.patch)(api, **kwargs) + def _call_patch(api, session: MagicMock, **kwargs): + return unwrap(api.patch)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -221,8 +225,10 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "New Name"}, ): api = DatasetMetadataServiceApi() + session = MagicMock() response, status = self._call_patch( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, metadata_id=metadata_id, @@ -250,9 +256,11 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "x"}, ): api = DatasetMetadataServiceApi() + session = MagicMock() with pytest.raises(NotFound): self._call_patch( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, metadata_id=metadata_id, @@ -266,8 +274,8 @@ class TestDatasetMetadataServiceApiDelete: """ @staticmethod - def _call_delete(api, **kwargs): - return unwrap(api.delete)(api, **kwargs) + def _call_delete(api, session: MagicMock, **kwargs): + return unwrap(api.delete)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -292,8 +300,10 @@ class TestDatasetMetadataServiceApiDelete: method="DELETE", ): api = DatasetMetadataServiceApi() + session = MagicMock() response = self._call_delete( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, metadata_id=metadata_id, @@ -319,9 +329,11 @@ class TestDatasetMetadataServiceApiDelete: method="DELETE", ): api = DatasetMetadataServiceApi() + session = MagicMock() with pytest.raises(NotFound): self._call_delete( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, metadata_id=metadata_id, @@ -375,8 +387,8 @@ class TestDatasetMetadataBuiltInFieldAction: """ @staticmethod - def _call_post(api, **kwargs): - return unwrap(api.post)(api, **kwargs) + def _call_post(api, session: MagicMock, **kwargs): + return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -399,8 +411,10 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() + session = MagicMock() response, status = self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, action="enable", @@ -408,7 +422,7 @@ class TestDatasetMetadataBuiltInFieldAction: assert status == 200 assert response["result"] == "success" - mock_meta_svc.enable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) + mock_meta_svc.enable_built_in_field.assert_called_once_with(mock_dataset, session) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -431,15 +445,17 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() + session = MagicMock() response, status = self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, action="disable", ) assert status == 200 - mock_meta_svc.disable_built_in_field.assert_called_once_with(mock_dataset, session=ANY) + mock_meta_svc.disable_built_in_field.assert_called_once_with(mock_dataset, session) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_action_dataset_not_found( @@ -457,9 +473,11 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() + session = MagicMock() with pytest.raises(NotFound): self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, action="enable", @@ -478,8 +496,8 @@ class TestDocumentMetadataEditPost: """ @staticmethod - def _call_post(api, **kwargs): - return unwrap(api.post)(api, **kwargs) + def _call_post(api, session: MagicMock, **kwargs): + return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @patch("controllers.service_api.dataset.metadata.DatasetService") @@ -504,8 +522,10 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() + session = MagicMock() response, status = self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, ) @@ -530,9 +550,11 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() + session = MagicMock() with pytest.raises(NotFound): self._call_post( api, + session, tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, ) diff --git a/api/tests/unit_tests/controllers/web/test_app.py b/api/tests/unit_tests/controllers/web/test_app.py index 73f308dc749..5494e4bbc03 100644 --- a/api/tests/unit_tests/controllers/web/test_app.py +++ b/api/tests/unit_tests/controllers/web/test_app.py @@ -23,7 +23,10 @@ class TestAppParameterApi: features_dict=features_dict, user_input_form=lambda to_old_structure=False: [], ) - app_model = SimpleNamespace(mode="advanced-chat", workflow=workflow) + app_model = SimpleNamespace( + mode="advanced-chat", + workflow_with_session=lambda *, session: workflow, + ) with ( app.test_request_context("/parameters"), @@ -42,7 +45,10 @@ class TestAppParameterApi: features_dict=features_dict, user_input_form=lambda to_old_structure=False: [{"var": "x"}], ) - app_model = SimpleNamespace(mode="workflow", workflow=workflow) + app_model = SimpleNamespace( + mode="workflow", + workflow_with_session=lambda *, session: workflow, + ) with ( app.test_request_context("/parameters"), @@ -55,19 +61,27 @@ class TestAppParameterApi: mock_params.assert_called_once_with(features_dict=features_dict, user_input_form=[{"var": "x"}]) def test_advanced_chat_mode_no_workflow_raises(self, app: Flask) -> None: - app_model = SimpleNamespace(mode="advanced-chat", workflow=None) + app_model = SimpleNamespace( + mode="advanced-chat", + workflow_with_session=lambda *, session: None, + ) with app.test_request_context("/parameters"): with pytest.raises(AppUnavailableError): AppParameterApi().get(app_model, SimpleNamespace()) def test_standard_mode_uses_app_model_config(self, app: Flask) -> None: - config = SimpleNamespace(to_dict=lambda: {"user_input_form": [{"var": "y"}], "key": "val"}) - app_model = SimpleNamespace(mode="chat", app_model_config=config) + config = SimpleNamespace(to_dict=lambda **_kwargs: {"user_input_form": [{"var": "y"}], "key": "val"}) + app_model = SimpleNamespace( + id="app-1", + mode="chat", + app_model_config_with_session=lambda *, session: config, + ) with ( app.test_request_context("/parameters"), patch("controllers.web.app.get_parameters_from_feature_dict", return_value={}) as mock_params, patch("controllers.web.app.fields.Parameters") as mock_fields, + patch("controllers.web.app.load_annotation_reply_config", return_value={"enabled": False}), ): mock_fields.model_validate.return_value.model_dump.return_value = {} AppParameterApi().get(app_model, SimpleNamespace()) @@ -76,7 +90,10 @@ class TestAppParameterApi: assert call_kwargs.kwargs["user_input_form"] == [{"var": "y"}] def test_standard_mode_no_config_raises(self, app: Flask) -> None: - app_model = SimpleNamespace(mode="chat", app_model_config=None) + app_model = SimpleNamespace( + mode="chat", + app_model_config_with_session=lambda *, session: None, + ) with app.test_request_context("/parameters"): with pytest.raises(AppUnavailableError): AppParameterApi().get(app_model, SimpleNamespace()) diff --git a/api/tests/unit_tests/controllers/web/test_completion.py b/api/tests/unit_tests/controllers/web/test_completion.py index 49a88802470..e3bbbe2c87c 100644 --- a/api/tests/unit_tests/controllers/web/test_completion.py +++ b/api/tests/unit_tests/controllers/web/test_completion.py @@ -2,6 +2,8 @@ from __future__ import annotations +import uuid +from inspect import unwrap from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -157,6 +159,24 @@ class TestChatApi: with pytest.raises(AgentNotPublishedError): ChatApi().post(app_model, _end_user()) + @patch("controllers.web.completion.AppGenerateService.generate", return_value="response") + @patch("controllers.web.completion.ConversationService.get_conversation") + @patch("controllers.web.completion.web_ns") + def test_conversation_validation_uses_request_session( + self, + mock_ns: MagicMock, + mock_get_conversation: MagicMock, + mock_generate: MagicMock, + app: Flask, + ) -> None: + mock_ns.payload = {"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())} + session = MagicMock() + + with app.test_request_context("/chat-messages", method="POST"): + unwrap(ChatApi.post)(ChatApi(), session, _chat_app(), _end_user()) + + assert mock_get_conversation.call_args.kwargs["session"] is session + # --------------------------------------------------------------------------- # ChatStopApi diff --git a/api/tests/unit_tests/core/agent/test_base_agent_runner.py b/api/tests/unit_tests/core/agent/test_base_agent_runner.py index 735cc437bd1..7411164512c 100644 --- a/api/tests/unit_tests/core/agent/test_base_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_base_agent_runner.py @@ -6,6 +6,7 @@ import pytest from pytest_mock import MockerFixture import core.agent.base_agent_runner as module +import models.model as model_module from core.agent.base_agent_runner import BaseAgentRunner # ========================================================== @@ -21,7 +22,7 @@ def mock_db_session(mocker: MockerFixture): @pytest.fixture -def runner(mocker: MockerFixture, mock_db_session): +def runner(mocker: MockerFixture): r = BaseAgentRunner.__new__(BaseAgentRunner) r.tenant_id = "tenant" r.user_id = "user" @@ -90,6 +91,9 @@ class TestCreateAgentThought: result = runner.create_agent_thought("m", "msg", "tool", "input", ["f1"]) assert result == "10" assert runner.agent_thought_count == 1 + mock_db_session.add.assert_called_once_with(mock_thought) + mock_db_session.commit.assert_called_once_with() + mock_db_session.close.assert_called_once_with() def test_without_files(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): mock_thought = mocker.MagicMock(id=11) @@ -151,6 +155,8 @@ class TestSaveAgentThought: assert agent.answer == "answer" assert agent.tokens == 3 assert "tool1" in json.loads(agent.tool_labels_str) + mock_db_session.commit.assert_called_once_with() + mock_db_session.close.assert_called_once_with() def test_label_fallback_when_none(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): agent = self.setup_agent(mocker) @@ -216,15 +222,18 @@ class TestSaveAgentThought: class TestOrganizeUserPrompt: def test_no_files(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): - mock_db_session.scalars.return_value.all.return_value = [] + caller_session = mocker.MagicMock() + caller_session.scalars.return_value.all.return_value = [] msg = mocker.MagicMock(id="1", query="hello", app_model_config=None) - result = runner.organize_agent_user_prompt(msg) + result = runner.organize_agent_user_prompt(msg, session=caller_session) assert result.content == "hello" + assert mock_db_session.mock_calls == [] def test_with_files_no_config(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): mock_db_session.scalars.return_value.all.return_value = [mocker.MagicMock()] msg = mocker.MagicMock(id="1", query="hello", app_model_config=None) - result = runner.organize_agent_user_prompt(msg) + msg.app_model_config_with_session.return_value = None + result = runner.organize_agent_user_prompt(msg, session=mock_db_session) assert result.content == "hello" def test_image_detail_low_fallback(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): @@ -235,10 +244,18 @@ class TestOrganizeUserPrompt: mocker.patch.object(module.file_factory, "build_from_message_files", return_value=[]) msg = mocker.MagicMock(id="1", query="hello") - msg.app_model_config.to_dict.return_value = {} + app_model_config = mocker.MagicMock() + app_model_config.app_id = "app1" + app_model_config.to_dict.return_value = {} + msg.app_model_config_with_session.return_value = app_model_config + load_annotation_reply_config = mocker.patch.object( + module, "load_annotation_reply_config", return_value={"enabled": False} + ) - result = runner.organize_agent_user_prompt(msg) + result = runner.organize_agent_user_prompt(msg, session=mock_db_session) assert result.content == "hello" + load_annotation_reply_config.assert_called_once_with(mock_db_session, "app1") + app_model_config.to_dict.assert_called_once_with(annotation_reply={"enabled": False}) # ========================================================== @@ -248,23 +265,27 @@ class TestOrganizeUserPrompt: class TestOrganizeHistory: def test_empty(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): - mock_db_session.execute.return_value.scalars.return_value.all.return_value = [] + caller_session = mocker.MagicMock() + caller_session.execute.return_value.scalars.return_value.all.return_value = [] mocker.patch.object(module, "extract_thread_messages", return_value=[]) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=caller_session) assert result == [] + assert mock_db_session.mock_calls == [] def test_with_answer_only(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): msg = mocker.MagicMock(id="m1", answer="ans", agent_thoughts=[], app_model_config=None) + msg.agent_thoughts_with_session.return_value = [] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert any(isinstance(x, module.AssistantPromptMessage) for x in result) def test_skip_current_message(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): msg = mocker.MagicMock(id="msg_current", agent_thoughts=[], answer="ans", app_model_config=None) mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert result == [] def test_with_tool_calls_invalid_json(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): @@ -275,21 +296,25 @@ class TestOrganizeHistory: thought="thinking", ) msg = mocker.MagicMock(id="m2", agent_thoughts=[thought], answer=None, app_model_config=None) + msg.agent_thoughts_with_session.return_value = [thought] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) mocker.patch("uuid.uuid4", return_value="uuid") - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert isinstance(result, list) def test_empty_tool_name_split(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): thought = mocker.MagicMock(tool=";", thought="thinking") msg = mocker.MagicMock(id="m5", agent_thoughts=[thought], answer=None, app_model_config=None) + msg.agent_thoughts_with_session.return_value = [thought] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert isinstance(result, list) def test_valid_json_tool_flow(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): @@ -306,12 +331,14 @@ class TestOrganizeHistory: answer=None, app_model_config=None, ) + msg.agent_thoughts_with_session.return_value = [thought] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) mocker.patch("uuid.uuid4", return_value="uuid") - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert isinstance(result, list) @@ -424,19 +451,25 @@ class TestAdditionalCoverage: mocker.patch.object(module, "TextPromptMessageContent", side_effect=lambda **kw: MagicMock(**kw)) msg = mocker.MagicMock(id="1", query="hello") - msg.app_model_config.to_dict.return_value = {} + app_model_config = mocker.MagicMock() + app_model_config.app_id = "app1" + app_model_config.to_dict.return_value = {} + msg.app_model_config_with_session.return_value = app_model_config + mocker.patch.object(module, "load_annotation_reply_config", return_value={"enabled": False}) - result = runner.organize_agent_user_prompt(msg) + result = runner.organize_agent_user_prompt(msg, session=mock_db_session) assert result is not None def test_organize_history_without_tool_names(self, runner: BaseAgentRunner, mock_db_session, mocker: MockerFixture): thought = mocker.MagicMock(tool=None, thought="thinking") msg = mocker.MagicMock(id="m3", agent_thoughts=[thought], answer=None, app_model_config=None) + msg.agent_thoughts_with_session.return_value = [thought] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert isinstance(result, list) def test_organize_history_multiple_tools_split( @@ -449,12 +482,14 @@ class TestAdditionalCoverage: thought="thinking", ) msg = mocker.MagicMock(id="m4", agent_thoughts=[thought], answer=None, app_model_config=None) + msg.agent_thoughts_with_session.return_value = [thought] + msg.app_model_config_with_session.return_value = None mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) mocker.patch("uuid.uuid4", return_value="uuid") - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert isinstance(result, list) @@ -480,12 +515,14 @@ class TestConvertDatasetRetrieverTool: class TestBaseAgentRunnerInit: def test_init_sets_stream_tool_call_and_files(self, mocker: MockerFixture): - session = mocker.MagicMock() - session.scalar.return_value = 2 - mocker.patch.object(module.db, "session", session) - - mocker.patch.object(BaseAgentRunner, "organize_agent_history", return_value=[]) - mocker.patch.object(module.DatasetRetrieverTool, "get_dataset_tools", return_value=["ds_tool"]) + caller_session = mocker.MagicMock() + caller_session.scalar.return_value = 2 + global_session = mocker.MagicMock() + mocker.patch.object(model_module.db, "session", global_session) + organize_agent_history = mocker.patch.object(BaseAgentRunner, "organize_agent_history", return_value=[]) + get_dataset_tools = mocker.patch.object( + module.DatasetRetrieverTool, "get_dataset_tools", return_value=["ds_tool"] + ) llm = mocker.MagicMock() llm.get_model_schema.return_value = mocker.MagicMock( @@ -503,7 +540,7 @@ class TestBaseAgentRunnerInit: message = mocker.MagicMock(id="msg1", conversation_id="conv1") runner = BaseAgentRunner( - session=session, + session=caller_session, tenant_id="tenant", application_generate_entity=app_generate, conversation=mocker.MagicMock(), @@ -520,6 +557,9 @@ class TestBaseAgentRunnerInit: assert runner.files == ["file1"] assert runner.dataset_tools == ["ds_tool"] assert runner.agent_thought_count == 2 + organize_agent_history.assert_called_once_with(session=caller_session, prompt_messages=[]) + assert get_dataset_tools.call_args.kwargs["session"] is caller_session + assert global_session.mock_calls == [] class TestBaseAgentRunnerCoverage: @@ -599,7 +639,7 @@ class TestBaseAgentRunnerCoverage: system_message = module.SystemPromptMessage(content="sys") - result = runner.organize_agent_history([system_message]) + result = runner.organize_agent_history([system_message], session=mock_db_session) assert system_message in result @@ -613,6 +653,7 @@ class TestBaseAgentRunnerCoverage: thought="thinking", ) msg = mocker.MagicMock(id="m6", agent_thoughts=[thought], answer=None, app_model_config=None) + msg.agent_thoughts_with_session.return_value = [thought] mock_db_session.execute.return_value.scalars.return_value.all.return_value = [msg] mocker.patch.object(module, "extract_thread_messages", return_value=[msg]) @@ -624,6 +665,6 @@ class TestBaseAgentRunnerCoverage: return_value=module.UserPromptMessage(content="user"), ) - result = runner.organize_agent_history([]) + result = runner.organize_agent_history([], session=mock_db_session) assert any(isinstance(item, module.ToolPromptMessage) for item in result) diff --git a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py index b0886ea8d19..43987bd65d0 100644 --- a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py @@ -204,6 +204,8 @@ class TestRun: results = list(runner.run(runner.session, message, "query", {})) assert isinstance(results, list) + assert "session" not in runner.create_agent_thought.call_args.kwargs + assert all("session" not in call.kwargs for call in runner.save_agent_thought.call_args_list) def test_run_with_action_and_tool_invocation(self, runner: DummyRunner, mocker: MockerFixture): message = MagicMock() @@ -335,13 +337,23 @@ class TestRun: def test_run_when_no_action_branch(self, runner: DummyRunner, mocker: MockerFixture): message = MagicMock() message.id = "msg-id" + events: list[str] = [] + session = MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session.close.side_effect = lambda: events.append("close") + def provider_chunks(): + events.append("first-chunk") + yield "chunk" + + runner.model_instance.invoke_llm.return_value = provider_chunks() mocker.patch( "core.agent.cot_agent_runner.CotAgentOutputParser.handle_react_stream_output", - return_value=[], + side_effect=lambda chunks, _usage: list(chunks), ) - results = list(runner.run(runner.session, message, "query", {})) + results = list(runner.run(session, message, "query", {})) + assert events == ["commit", "close", "first-chunk"] assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"} assert results[-1].delta.message.content == "" diff --git a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py index f32ded9edff..d70605c96ef 100644 --- a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py @@ -300,6 +300,8 @@ class TestRunMethod: outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) == 1 + assert "session" not in runner.create_agent_thought.call_args.kwargs + assert "session" not in runner.save_agent_thought.call_args.kwargs assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"} runner.queue_manager.publish.assert_called() @@ -309,16 +311,22 @@ class TestRunMethod: def test_run_streaming_branch(self, runner: FunctionCallAgentRunner): message = MagicMock(id="m1") runner.stream_tool_call = True + events: list[str] = [] + session = MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session.close.side_effect = lambda: events.append("close") content = [TextPromptMessageContent(data="hi")] chunk = DummyChunk(message=DummyMessage(content=content), usage=build_usage()) def generator(): + events.append("first-chunk") yield chunk runner.model_instance.invoke_llm.return_value = generator() - outputs = list(runner.run(runner.session, message, "query")) + outputs = list(runner.run(session, message, "query")) + assert events == ["commit", "close", "first-chunk"] assert len(outputs) == 1 def test_run_streaming_tool_calls_list_content(self, runner: FunctionCallAgentRunner): diff --git a/api/tests/unit_tests/core/app/app_config/easy_ui_based_app/test_dataset_manager.py b/api/tests/unit_tests/core/app/app_config/easy_ui_based_app/test_dataset_manager.py index e4e4f99c6d4..ccfa51280e1 100644 --- a/api/tests/unit_tests/core/app/app_config/easy_ui_based_app/test_dataset_manager.py +++ b/api/tests/unit_tests/core/app/app_config/easy_ui_based_app/test_dataset_manager.py @@ -208,7 +208,7 @@ class TestDatasetConfigManagerConvert: class TestValidateAndSetDefaults: def test_validate_sets_defaults(self): config = {} - updated, fields = DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.CHAT, config) + updated, fields = DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.CHAT, config, MagicMock()) assert "dataset_configs" in updated assert updated["dataset_configs"]["retrieval_model"] == "single" assert isinstance(fields, list) @@ -216,7 +216,7 @@ class TestValidateAndSetDefaults: def test_validate_raises_when_dataset_configs_not_dict(self): config = {"dataset_configs": "invalid"} with pytest.raises(AttributeError): - DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.CHAT, config) + DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.CHAT, config, MagicMock()) def test_validate_requires_query_variable_in_completion_mode(self, valid_uuid): config = { @@ -228,7 +228,7 @@ class TestValidateAndSetDefaults: } } with pytest.raises(ValueError): - DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.COMPLETION, config) + DatasetConfigManager.validate_and_set_defaults("tenant1", AppMode.COMPLETION, config, MagicMock()) # ============================== @@ -239,7 +239,9 @@ class TestValidateAndSetDefaults: class TestExtractDatasetConfig: def test_extract_sets_defaults(self): config = {} - result = DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + result = DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) assert "agent_mode" in result assert result["agent_mode"]["enabled"] is False assert result["agent_mode"]["tools"] == [] @@ -247,17 +249,23 @@ class TestExtractDatasetConfig: def test_extract_invalid_agent_mode_type(self): config = {"agent_mode": "invalid"} with pytest.raises(ValueError): - DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) def test_extract_invalid_enabled_type(self): config = {"agent_mode": {"enabled": "yes"}} with pytest.raises(ValueError): - DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) def test_extract_invalid_tools_type(self): config = {"agent_mode": {"enabled": True, "tools": "invalid"}} with pytest.raises(ValueError): - DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) def test_extract_invalid_uuid(self, mocker: MockerFixture): invalid_uuid = "not-a-uuid" @@ -269,7 +277,9 @@ class TestExtractDatasetConfig: } } with pytest.raises(ValueError): - DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) def test_extract_dataset_not_exists(self, valid_uuid, mocker: MockerFixture): mocker.patch( @@ -284,7 +294,9 @@ class TestExtractDatasetConfig: } } with pytest.raises(ValueError): - DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) # ============================== @@ -301,14 +313,14 @@ class TestIsDatasetExists: return_value=mock_dataset, ) - assert DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid) + assert DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid, MagicMock()) def test_dataset_exists_false_when_not_found(self, mocker: MockerFixture, valid_uuid): mocker.patch( "core.app.app_config.easy_ui_based_app.dataset.manager.DatasetService.get_dataset", return_value=None, ) - assert not DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid) + assert not DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid, MagicMock()) def test_dataset_exists_false_when_tenant_mismatch(self, mocker: MockerFixture, valid_uuid): mock_dataset = MagicMock() @@ -317,7 +329,7 @@ class TestIsDatasetExists: "core.app.app_config.easy_ui_based_app.dataset.manager.DatasetService.get_dataset", return_value=mock_dataset, ) - assert not DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid) + assert not DatasetConfigManager.is_dataset_exists("tenant1", valid_uuid, MagicMock()) # ============================== @@ -338,6 +350,8 @@ class TestExtractDatasetConfigForLegacyCompatibility: } } - result = DatasetConfigManager.extract_dataset_config_for_legacy_compatibility("tenant1", AppMode.CHAT, config) + result = DatasetConfigManager.extract_dataset_config_for_legacy_compatibility( + "tenant1", AppMode.CHAT, config, MagicMock() + ) assert result["agent_mode"]["tools"] == [{}] 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 41e14af72de..f347a5fae7e 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 @@ -37,6 +37,7 @@ class TestAdvancedChatAppGeneratorValidation: invoke_from=InvokeFrom.WEB_APP, workflow_run_id="run-id", streaming=False, + session=MagicMock(), ) def test_generate_requires_string_query(self): @@ -51,6 +52,7 @@ class TestAdvancedChatAppGeneratorValidation: invoke_from=InvokeFrom.WEB_APP, workflow_run_id="run-id", streaming=False, + session=MagicMock(), ) def test_single_iteration_generate_validates_args(self): @@ -64,6 +66,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args={"inputs": {}}, streaming=False, + session=MagicMock(), ) with pytest.raises(ValueError, match="inputs is required"): @@ -74,6 +77,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args={}, streaming=False, + session=MagicMock(), ) def test_single_loop_generate_validates_args(self): @@ -87,6 +91,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs={}), streaming=False, + session=MagicMock(), ) with pytest.raises(ValueError, match="inputs is required"): @@ -97,6 +102,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs=None), streaming=False, + session=MagicMock(), ) @@ -120,10 +126,12 @@ class TestAdvancedChatAppGeneratorInternals: built_files: list[object] = [] build_files_called = {"called": False} captured: dict[str, object] = {} + session = MagicMock() + get_conversation = MagicMock(return_value=conversation) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.ConversationService.get_conversation", - lambda **kwargs: conversation, + get_conversation, ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.FileUploadConfigManager.convert", @@ -189,12 +197,15 @@ class TestAdvancedChatAppGeneratorInternals: invoke_from=InvokeFrom.WEB_APP, workflow_run_id="run-id", streaming=False, + session=session, ) assert result == {"ok": True} assert captured["conversation"] is conversation assert captured["application_generate_entity"].files == built_files + assert captured["session"] is session assert build_files_called["called"] is True + assert get_conversation.call_args.kwargs["session"] is session def test_resume_delegates_to_generate(self, monkeypatch: pytest.MonkeyPatch): generator = AdvancedChatAppGenerator() @@ -230,6 +241,7 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), + session=MagicMock(), application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), @@ -247,8 +259,10 @@ class TestAdvancedChatAppGeneratorInternals: app_config = self._build_app_config() captured: dict[str, object] = {} prefill_calls: list[object] = [] + draft_sessions: list[object] = [] var_loader = SimpleNamespace(loader="draft") workflow = SimpleNamespace(id="workflow-id") + session = MagicMock() monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.AdvancedChatAppConfigManager.get_app_config", @@ -273,7 +287,7 @@ class TestAdvancedChatAppGeneratorInternals: class _DraftVarService: def __init__(self, session): - _ = session + draft_sessions.append(session) def prefill_conversation_variable_default_values(self, workflow, user_id): prefill_calls.append((workflow, user_id)) @@ -293,11 +307,14 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="user-id"), args={"inputs": {"foo": "bar"}, "trace_session_id": "session-1"}, streaming=False, + session=session, ) assert result == {"ok": True} assert prefill_calls == [(workflow, "user-id")] + assert draft_sessions == [session] assert captured["variable_loader"] is var_loader + assert captured["session"] is session assert captured["application_generate_entity"].single_iteration_run.node_id == "node-1" assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" @@ -306,8 +323,10 @@ class TestAdvancedChatAppGeneratorInternals: app_config = self._build_app_config() captured: dict[str, object] = {} prefill_calls: list[object] = [] + draft_sessions: list[object] = [] var_loader = SimpleNamespace(loader="draft") workflow = SimpleNamespace(id="workflow-id") + session = MagicMock() monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.AdvancedChatAppConfigManager.get_app_config", @@ -332,7 +351,7 @@ class TestAdvancedChatAppGeneratorInternals: class _DraftVarService: def __init__(self, session): - _ = session + draft_sessions.append(session) def prefill_conversation_variable_default_values(self, workflow, user_id): prefill_calls.append((workflow, user_id)) @@ -352,11 +371,14 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="user-id"), args=SimpleNamespace(inputs={"foo": "bar"}, trace_session_id="session-1"), streaming=False, + session=session, ) assert result == {"ok": True} assert prefill_calls == [(workflow, "user-id")] + assert draft_sessions == [session] assert captured["variable_loader"] is var_loader + assert captured["session"] is session assert captured["application_generate_entity"].single_loop_run.node_id == "node-2" assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" @@ -391,9 +413,13 @@ class TestAdvancedChatAppGeneratorInternals: db_session = SimpleNamespace(commit=MagicMock(), refresh=MagicMock(), close=MagicMock()) captured: dict[str, object] = {} thread_data: dict[str, object] = {} + init_records = MagicMock(return_value=(conversation, message)) + get_thread_messages_length = MagicMock(return_value=2) - monkeypatch.setattr(generator, "_init_generate_records", lambda *args: (conversation, message)) - monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.get_thread_messages_length", lambda _: 2) + monkeypatch.setattr(generator, "_init_generate_records", init_records) + monkeypatch.setattr( + "core.app.apps.advanced_chat.app_generator.get_thread_messages_length", get_thread_messages_length + ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.MessageBasedAppQueueManager", lambda **kwargs: SimpleNamespace(**kwargs), @@ -438,6 +464,7 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="user"), invoke_from=InvokeFrom.WEB_APP, application_generate_entity=application_generate_entity, + session=db_session, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), conversation=None, @@ -450,6 +477,8 @@ class TestAdvancedChatAppGeneratorInternals: assert thread_data["started"] is True assert "pause-layer" in thread_data["kwargs"]["graph_engine_layers"] assert generator._dialogue_count == 3 + assert init_records.call_args.kwargs["session"] is db_session + get_thread_messages_length.assert_called_once_with(conversation.id, session=db_session) db_session.commit.assert_called_once() db_session.refresh.assert_called_once_with(conversation) db_session.close.assert_called_once() @@ -488,10 +517,13 @@ class TestAdvancedChatAppGeneratorInternals: ) db_session = SimpleNamespace(close=MagicMock(), commit=MagicMock(), refresh=MagicMock()) init_records = MagicMock() + get_thread_messages_length = MagicMock(return_value=0) thread_data: dict[str, object] = {} monkeypatch.setattr(generator, "_init_generate_records", init_records) - monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.get_thread_messages_length", lambda _: 0) + monkeypatch.setattr( + "core.app.apps.advanced_chat.app_generator.get_thread_messages_length", get_thread_messages_length + ) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.MessageBasedAppQueueManager", lambda **kwargs: SimpleNamespace(**kwargs), @@ -530,6 +562,7 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="user"), invoke_from=InvokeFrom.WEB_APP, application_generate_entity=application_generate_entity, + session=db_session, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), conversation=conversation, @@ -539,6 +572,7 @@ class TestAdvancedChatAppGeneratorInternals: assert response == {"raw": True} init_records.assert_not_called() + get_thread_messages_length.assert_called_once_with(conversation.id, session=db_session) assert thread_data["started"] is True db_session.commit.assert_not_called() db_session.refresh.assert_not_called() @@ -1171,6 +1205,7 @@ class TestAdvancedChatAppGeneratorInternals: invoke_from=InvokeFrom.DEBUGGER, workflow_run_id="run-id", streaming=False, + session=MagicMock(), ) assert result == {"ok": True} @@ -1250,6 +1285,7 @@ class TestAdvancedChatAppGeneratorInternals: invoke_from=InvokeFrom.SERVICE_API, workflow_run_id="run-id", streaming=False, + session=MagicMock(), ) assert captured["application_generate_entity"].parent_message_id == UUID_NIL @@ -1313,6 +1349,7 @@ class TestAdvancedChatAppGeneratorResume: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), + session=MagicMock(), application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), @@ -1360,6 +1397,7 @@ class TestAdvancedChatAppGeneratorResume: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), + session=MagicMock(), application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py index 2e3f7645c73..5a7b58e581a 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py @@ -6,7 +6,7 @@ import pytest import core.app.apps.advanced_chat.app_runner as module from core.app.apps.advanced_chat.app_runner import AdvancedChatAppRunner from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom -from core.app.entities.queue_entities import QueueStopEvent +from core.app.entities.queue_entities import QueueAnnotationReplyEvent, QueueStopEvent from core.moderation.base import ModerationError MINIMAL_GRAPH = { @@ -192,6 +192,38 @@ def test_run_returns_early_when_direct_output_via_handle_input_moderation(build_ mock_init_graph.assert_not_called() +def test_run_publishes_annotation_after_commit(build_runner): + runner = build_runner + events: list[str] = [] + session = MagicMock() + session.scalar.return_value = MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session_context = MagicMock() + session_context.__enter__.return_value = session + session_context.__exit__.return_value = False + annotation_reply = MagicMock(id="annotation-1", content="annotated answer") + + def publish(event): + if isinstance(event, QueueAnnotationReplyEvent): + events.append("publish") + + with ( + _patch_common_run_deps(runner), + patch.object(module, "create_session", return_value=session_context), + patch.object( + runner, + "handle_input_moderation", + return_value=(False, runner.application_generate_entity.inputs, runner.application_generate_entity.query), + ), + patch.object(runner, "handle_annotation_reply", return_value=annotation_reply), + patch.object(runner, "_publish_event", side_effect=publish), + patch.object(runner, "_complete_with_stream_output"), + ): + runner.run() + + assert events == ["commit", "publish"] + + def test_run_closes_scoped_session_before_workflow_run(build_runner): runner = build_runner events = [] @@ -200,7 +232,7 @@ def test_run_closes_scoped_session_before_workflow_run(build_runner): mock_session.scalar.return_value = MagicMock() session_context = MagicMock() session_context.__enter__.return_value = mock_session - session_context.__exit__.return_value = False + session_context.__exit__.side_effect = lambda exc_type, exc, tb: events.append("close") or False workflow_entry = MagicMock() @@ -216,7 +248,6 @@ def test_run_closes_scoped_session_before_workflow_run(build_runner): patch.object(module, "RedisChannel"), patch.object(module, "redis_client"), patch.object(module, "WorkflowEntry", return_value=workflow_entry), - patch.object(module.db.session, "close", side_effect=lambda: events.append("close")), patch.object( runner, "handle_input_moderation", @@ -228,4 +259,4 @@ def test_run_closes_scoped_session_before_workflow_run(build_runner): ): runner.run() - assert events == ["close", "run"] + assert events[-2:] == ["close", "run"] diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py index 7ab15d40e30..1f975c74fb1 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py @@ -580,22 +580,28 @@ class TestAdvancedChatGenerateTaskPipeline: assert result is False assert seen == ["token"] - def test_handle_retriever_and_annotation_events(self): + def test_handle_retriever_and_annotation_events(self, monkeypatch: pytest.MonkeyPatch): pipeline = _make_pipeline() calls = {"retriever": 0, "annotation": 0} def _hit_retriever(event): calls["retriever"] += 1 - def _hit_annotation(event): + def _hit_annotation(_manager, event, session): calls["annotation"] += 1 pipeline._message_cycle_manager.handle_retriever_resources = _hit_retriever - pipeline._message_cycle_manager.handle_annotation_reply = _hit_annotation + monkeypatch.setattr(type(pipeline._message_cycle_manager), "handle_annotation_reply", _hit_annotation) retriever_event = QueueRetrieverResourcesEvent(retriever_resources=[]) annotation_event = QueueAnnotationReplyEvent(message_annotation_id="ann") + @contextmanager + def _fake_session(): + yield SimpleNamespace() + + monkeypatch.setattr(pipeline, "_database_session", _fake_session) + assert list(pipeline._handle_retriever_resources_event(retriever_event)) == [] assert list(pipeline._handle_annotation_reply_event(annotation_event)) == [] assert calls == {"retriever": 1, "annotation": 1} diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_app_config_manager.py b/api/tests/unit_tests/core/app/apps/agent_app/test_app_config_manager.py index 73c53ae98f5..d714e20ffcf 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_app_config_manager.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_app_config_manager.py @@ -28,7 +28,7 @@ def _soul() -> AgentSoulConfig: def test_model_and_prompt_come_from_soul(): - d = AgentAppConfigManager._synthesize_config_dict(_soul(), None) + d = AgentAppConfigManager._synthesize_config_dict(_soul(), None, annotation_reply=None) assert d["model"] == { "provider": "langgenius/openai/openai", "name": "gpt-4o-mini", @@ -45,14 +45,18 @@ def test_model_and_prompt_come_from_soul(): def test_feature_flags_come_from_app_model_config_when_present(): # Q3: opener/follow-up/etc. live on app_model_config; model/prompt stay from Soul. fake_amc = SimpleNamespace( - to_dict=lambda: { + to_dict=lambda **_: { "opening_statement": "Hi, I'm Iris.", "suggested_questions_after_answer": {"enabled": True}, "model": {"provider": "should-be-overridden", "name": "old"}, "pre_prompt": "old prompt", } ) - d = AgentAppConfigManager._synthesize_config_dict(_soul(), fake_amc) # type: ignore[arg-type] + d = AgentAppConfigManager._synthesize_config_dict( + _soul(), + fake_amc, + annotation_reply={"enabled": False}, # type: ignore[arg-type] + ) # feature flags preserved assert d["opening_statement"] == "Hi, I'm Iris." assert d["suggested_questions_after_answer"] == {"enabled": True} @@ -62,7 +66,7 @@ def test_feature_flags_come_from_app_model_config_when_present(): def test_missing_soul_model_leaves_no_model_key(): - d = AgentAppConfigManager._synthesize_config_dict(AgentSoulConfig(), None) + d = AgentAppConfigManager._synthesize_config_dict(AgentSoulConfig(), None, annotation_reply=None) assert "model" not in d assert d["pre_prompt"] == "" assert d["file_upload"] == { @@ -77,7 +81,7 @@ def test_missing_soul_model_leaves_no_model_key(): def test_soul_file_upload_overrides_legacy_app_model_config(): fake_amc = SimpleNamespace( - to_dict=lambda: { + to_dict=lambda **_: { "file_upload": { "enabled": False, "image": {"enabled": False}, @@ -85,7 +89,11 @@ def test_soul_file_upload_overrides_legacy_app_model_config(): } ) - d = AgentAppConfigManager._synthesize_config_dict(AgentSoulConfig(), fake_amc) # type: ignore[arg-type] + d = AgentAppConfigManager._synthesize_config_dict( + AgentSoulConfig(), + fake_amc, + annotation_reply={"enabled": False}, # type: ignore[arg-type] + ) assert d["file_upload"] == { "allowed_file_extensions": ["JPG", "JPEG", "PNG", "GIF", "WEBP", "SVG"], @@ -100,7 +108,7 @@ def test_soul_file_upload_overrides_legacy_app_model_config(): def test_prompt_type_defaults_to_simple(): # PromptTemplateConfigManager.convert requires prompt_type; an Agent App with # no legacy app_model_config must still get the "simple" slot synthesized. - d = AgentAppConfigManager._synthesize_config_dict(_soul(), None) + d = AgentAppConfigManager._synthesize_config_dict(_soul(), None, annotation_reply=None) assert d["prompt_type"] == "simple" @@ -115,6 +123,7 @@ def test_get_app_config_has_null_model_config_id_without_legacy_row(): app_config = AgentAppConfigManager.get_app_config( app_model=app_model, # type: ignore[arg-type] agent_soul=_soul(), + annotation_reply=None, app_model_config=None, conversation=None, ) diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py b/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py index ef030e7e170..12c429e85b9 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py @@ -10,19 +10,22 @@ manager is replaced with a no-op so the thread body can run inline. from __future__ import annotations import contextlib +import inspect import json import pytest from pytest_mock import MockerFixture +import core.app.apps.agent_app.app_generator as module from core.app.apps.agent_app.app_generator import ( AgentAppGenerator, AgentAppGeneratorError, ) from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import AGENT_RUNTIME_EXIT_INTENT_ARG, InvokeFrom, UserFrom +from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.workflow.file_reference import build_file_reference -from models import Account +from models import Account, AppModelConfig from models.agent import AgentConfigDraftType MODULE = "core.app.apps.agent_app.app_generator" @@ -50,6 +53,7 @@ class TestGenerateGuards: user=DummyAccount("u"), args={}, invoke_from=InvokeFrom.WEB_APP, + session=mocker.MagicMock(), streaming=False, ) @@ -60,6 +64,7 @@ class TestGenerateGuards: user=DummyAccount("u"), args={"inputs": {}}, invoke_from=InvokeFrom.WEB_APP, + session=mocker.MagicMock(), ) def test_rejects_blank_query(self, generator: AgentAppGenerator, mocker: MockerFixture): @@ -69,6 +74,7 @@ class TestGenerateGuards: user=DummyAccount("u"), args={"query": " ", "inputs": {}}, invoke_from=InvokeFrom.WEB_APP, + session=mocker.MagicMock(), ) @@ -85,7 +91,9 @@ class TestGenerateSuccess: def test_generate_orchestrates_and_starts_worker(self, generator, mocker: MockerFixture): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") + app_model.app_model_config_id = "config-1" user = DummyAccount("user") + session = mocker.MagicMock() generator._resolve_agent = mocker.MagicMock( return_value=(mocker.MagicMock(id="agent1"), "snap1", "snapshot", mocker.MagicMock()) @@ -107,7 +115,7 @@ class TestGenerateSuccess: ) mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) thread_obj = mocker.MagicMock() - mocker.patch(f"{MODULE}.threading.Thread", return_value=thread_obj) + thread_constructor = mocker.patch(f"{MODULE}.threading.Thread", return_value=thread_obj) mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"}) file_mappings = [ { @@ -123,17 +131,22 @@ class TestGenerateSuccess: user=user, args={"query": "hello", "inputs": {"name": "world"}, "files": file_mappings}, invoke_from=InvokeFrom.WEB_APP, + session=session, streaming=True, ) assert result == {"result": "ok"} thread_obj.start.assert_called_once() + worker_call = thread_constructor.call_args + inspect.signature(worker_call.kwargs["target"]).bind(**worker_call.kwargs["kwargs"]) generator._resolve_agent.assert_called_once_with( app_model, invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=user, + session=session, ) + session.get.assert_called_once_with(AppModelConfig, "config-1") assert generate_entity.call_args.kwargs["prompt_file_mappings"] == file_mappings assert generate_entity.call_args.kwargs["agent_runtime_exit_intent"] == "suspend" @@ -168,6 +181,7 @@ class TestGenerateSuccess: user=user, args={"query": "hello", "inputs": {}, AGENT_RUNTIME_EXIT_INTENT_ARG: "delete"}, invoke_from=InvokeFrom.DEBUGGER, + session=mocker.MagicMock(), streaming=True, ) @@ -204,6 +218,7 @@ class TestGenerateSuccess: user=user, args={"query": "hello", "inputs": {}, AGENT_RUNTIME_EXIT_INTENT_ARG: "bogus"}, invoke_from=InvokeFrom.DEBUGGER, + session=mocker.MagicMock(), streaming=True, ) @@ -223,22 +238,32 @@ class TestGenerateSuccess: f"{MODULE}.ConversationService.get_conversation", return_value=mocker.MagicMock(id="conv") ) mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[])) + mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False}) mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"}) + session = mocker.MagicMock() + user = DummyAccount("user") generator.generate( app_model=app_model, - user=DummyAccount("user"), + user=user, args={"query": "hi", "inputs": {}, "conversation_id": "conv"}, invoke_from=InvokeFrom.WEB_APP, + session=session, streaming=True, ) - get_conv.assert_called_once() + get_conv.assert_called_once_with( + app_model=app_model, + conversation_id="conv", + user=user, + session=session, + ) + assert generator._init_generate_records.call_args.kwargs["session"] is session def test_generate_does_not_include_trace_session_id_in_extras( self, generator: AgentAppGenerator, mocker: MockerFixture @@ -273,6 +298,7 @@ class TestGenerateSuccess: user=user, args={"query": "hello", "inputs": {}, "trace_session_id": "session-1"}, invoke_from=InvokeFrom.WEB_APP, + session=mocker.MagicMock(), streaming=True, ) @@ -299,11 +325,20 @@ class TestGenerateWorker: ): generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock(id="conv")) generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock(id="msg")) - generator._run_input_guards = mocker.MagicMock(return_value=(handled, guard_query)) + generator._run_input_guards = mocker.MagicMock(return_value=(handled, guard_query, None)) generator._resolve_agent_by_id = mocker.MagicMock( return_value=(mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) ) - mocker.patch(f"{MODULE}.db.session.get", return_value=mocker.MagicMock(id="app1")) + session = mocker.MagicMock() + session.get.return_value = mocker.MagicMock(id="app1") + session_context = mocker.MagicMock() + session_context.__enter__.return_value = session + session_maker = mocker.patch(f"{MODULE}.session_factory.get_session_maker").return_value + session_maker.begin.return_value = session_context + resolver_session = mocker.MagicMock() + resolver_context = mocker.MagicMock() + resolver_context.__enter__.return_value = resolver_session + mocker.patch(f"{MODULE}.session_factory.create_session", return_value=resolver_context) mocker.patch(f"{MODULE}.db.session.close") mocker.patch(f"{MODULE}.DifyRunContext", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.build_dify_model_access", return_value=(mocker.MagicMock(), None)) @@ -315,7 +350,7 @@ class TestGenerateWorker: if run_side_effect is not None: runner.run.side_effect = run_side_effect mocker.patch(f"{MODULE}.AgentAppRunner", return_value=runner) - return runner + return runner, resolver_session def _call( self, @@ -349,14 +384,15 @@ class TestGenerateWorker: ) def test_happy_path_runs_backend(self, generator: AgentAppGenerator, mocker: MockerFixture): - runner = self._wire(generator, mocker) + runner, resolver_session = self._wire(generator, mocker) queue_manager = mocker.MagicMock() self._call(generator, mocker, queue_manager) runner.run.assert_called_once() + assert generator._resolve_agent_by_id.call_args.kwargs["session"] is resolver_session queue_manager.publish_error.assert_not_called() def test_worker_passes_runtime_session_scope_to_runner(self, generator, mocker: MockerFixture): - runner = self._wire(generator, mocker) + runner, _ = self._wire(generator, mocker) queue_manager = mocker.MagicMock() self._call(generator, mocker, queue_manager, runtime_session_snapshot_id=None) @@ -365,7 +401,7 @@ class TestGenerateWorker: assert runner.run.call_args.kwargs["session_scope_snapshot_id"] is None def test_worker_forwards_runtime_exit_intent_to_runner(self, generator, mocker: MockerFixture): - runner = self._wire(generator, mocker) + runner, _ = self._wire(generator, mocker) queue_manager = mocker.MagicMock() self._call(generator, mocker, queue_manager, agent_runtime_exit_intent="delete") @@ -373,7 +409,7 @@ class TestGenerateWorker: assert runner.run.call_args.kwargs["agent_runtime_exit_intent"] == "delete" def test_worker_appends_prompt_files_to_backend_query(self, generator, mocker: MockerFixture): - runner = self._wire(generator, mocker, guard_query="你看得见这张图片吗") + runner, _ = self._wire(generator, mocker, guard_query="你看得见这张图片吗") queue_manager = mocker.MagicMock() file_mappings = [ { @@ -416,16 +452,36 @@ class TestGenerateWorker: ) def test_input_guard_short_circuit_skips_backend(self, generator, mocker: MockerFixture): - runner = self._wire(generator, mocker, handled=True) + runner, _ = self._wire(generator, mocker, handled=True) queue_manager = mocker.MagicMock() self._call(generator, mocker, queue_manager) runner.run.assert_not_called() + def test_annotation_reply_publishes_after_guard_transaction_commits(self, generator, mocker: MockerFixture): + runner, _ = self._wire(generator, mocker, handled=True) + annotation_reply = mocker.MagicMock(id="annotation-1", content="annotated answer") + generator._run_input_guards.return_value = (True, "query", annotation_reply) + events: list[str] = [] + guard_context = module.session_factory.get_session_maker.return_value.begin.return_value + guard_context.__exit__.side_effect = lambda *args: events.append("commit") or False + queue_manager = mocker.MagicMock() + + def publish(event, *_args): + if isinstance(event, QueueAnnotationReplyEvent): + events.append("publish") + + queue_manager.publish.side_effect = publish + + self._call(generator, mocker, queue_manager) + + assert events == ["commit", "publish"] + runner.run.assert_not_called() + def test_resume_skips_input_guards_and_consumes_reply(self, generator, mocker: MockerFixture): # ENG-638 (review): on resume the replayed query is NOT new end-user input. # Input guards must be skipped, even if moderation/annotation would match, # so the run continues and the human reply (deferred_tool_results) is used. - runner = self._wire(generator, mocker, handled=True) # guards WOULD short-circuit + runner, _ = self._wire(generator, mocker, handled=True) # guards WOULD short-circuit queue_manager = mocker.MagicMock() self._call(generator, mocker, queue_manager, is_resume=True, query="the approved reply") @@ -460,46 +516,65 @@ class TestResumeAfterFormSubmission: return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg")) ) generator._handle_response = mocker.MagicMock(return_value=None) - mocker.patch( + get_conversation = mocker.patch( f"{MODULE}.ConversationService.get_conversation", return_value=mocker.MagicMock(id="conv", invoke_from=InvokeFrom.WEB_APP), ) mocker.patch(f"{MODULE}.AgentAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables=[])) + mocker.patch(f"{MODULE}.load_annotation_reply_config", return_value={"enabled": False}) mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.AgentAppRuntimeSessionStore") - return mocker.patch( - f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user") + return ( + mocker.patch( + f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user") + ), + get_conversation, ) def test_resume_resends_paused_turn_query(self, generator, mocker: MockerFixture): - entity = self._wire(generator, mocker) - db_mock = mocker.patch(f"{MODULE}.db") - db_mock.session.scalar.return_value = mocker.MagicMock(query="original question") + entity, get_conversation = self._wire(generator, mocker) + app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") + app_model.app_model_config_id = "config-1" + user = DummyAccount("user") + session = mocker.MagicMock() + session.get.return_value = mocker.MagicMock() + session.scalar.return_value = mocker.MagicMock(query="original question") generator.resume_after_form_submission( - app_model=mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent"), - user=DummyAccount("user"), + app_model=app_model, + user=user, conversation_id="conv", invoke_from=InvokeFrom.WEB_APP, + session=session, ) # The paused turn's query is re-sent verbatim — never blank. assert entity.call_args.kwargs["query"] == "original question" assert "agent_runtime_exit_intent" not in entity.call_args.kwargs + get_conversation.assert_called_once_with( + app_model=app_model, + conversation_id="conv", + user=user, + session=session, + ) + assert generator._init_generate_records.call_args.kwargs["session"] is session + session.get.assert_called_once_with(AppModelConfig, "config-1") + assert generator._resolve_agent.call_args.kwargs["session"] is session def test_resume_falls_back_to_placeholder_when_no_paused_message(self, generator, mocker: MockerFixture): - entity = self._wire(generator, mocker) - db_mock = mocker.patch(f"{MODULE}.db") - db_mock.session.scalar.return_value = None + entity, _ = self._wire(generator, mocker) + session = mocker.MagicMock() + session.scalar.return_value = None generator.resume_after_form_submission( app_model=mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent"), user=DummyAccount("user"), conversation_id="conv", invoke_from=InvokeFrom.WEB_APP, + session=session, ) # No prior user message -> a non-blank placeholder, still never blank. @@ -514,16 +589,20 @@ class TestResumeAfterFormSubmission: scope=mocker.MagicMock(agent_config_snapshot_id="draft-build-1") ) draft_row = mocker.MagicMock(draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id="user") - db_mock = mocker.patch(f"{MODULE}.db") - db_mock.session.scalar.side_effect = [draft_row, mocker.MagicMock(query="original question")] account_user = mocker.MagicMock(spec=Account) account_user.id = "user" + app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") + app_model.app_model_config_id = "config-1" + session = mocker.MagicMock() + session.scalar.side_effect = [draft_row, mocker.MagicMock(query="original question")] generator.resume_after_form_submission( - app_model=mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent"), + app_model=app_model, user=account_user, conversation_id="conv", invoke_from=InvokeFrom.DEBUGGER, + session=session, ) assert generator._resolve_agent.call_args.kwargs["draft_type"] == "debug_build" + assert generator._resolve_agent.call_args.kwargs["session"] is session diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py b/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py index 24dd61cb951..ec2706a2068 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py @@ -10,6 +10,7 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any +from unittest.mock import MagicMock import pytest @@ -17,7 +18,6 @@ import core.app.features.annotation_reply.annotation_reply as annotation_mod import core.moderation.input_moderation as input_moderation_mod from core.app.apps.agent_app.app_generator import AgentAppGenerator from core.app.entities.queue_entities import ( - QueueAnnotationReplyEvent, QueueLLMChunkEvent, QueueMessageEndEvent, ) @@ -80,7 +80,8 @@ class TestRunInputGuards: _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() - handled, query = AgentAppGenerator()._run_input_guards( + handled, query, annotation_reply = AgentAppGenerator()._run_input_guards( + session=MagicMock(), application_generate_entity=_make_entity("hello"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -89,6 +90,7 @@ class TestRunInputGuards: assert handled is False assert query == "hello" + assert annotation_reply is None assert qm.events == [] def test_moderation_override_sanitizes_query(self, monkeypatch: pytest.MonkeyPatch): @@ -96,7 +98,8 @@ class TestRunInputGuards: _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() - handled, query = AgentAppGenerator()._run_input_guards( + handled, query, annotation_reply = AgentAppGenerator()._run_input_guards( + session=MagicMock(), application_generate_entity=_make_entity("leak my secret"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -105,6 +108,7 @@ class TestRunInputGuards: assert handled is False assert query == "[redacted]" + assert annotation_reply is None assert qm.events == [] def test_moderation_block_short_circuits(self, monkeypatch: pytest.MonkeyPatch): @@ -112,7 +116,8 @@ class TestRunInputGuards: _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() - handled, _ = AgentAppGenerator()._run_input_guards( + handled, _, annotation_reply = AgentAppGenerator()._run_input_guards( + session=MagicMock(), application_generate_entity=_make_entity("forbidden"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -120,6 +125,7 @@ class TestRunInputGuards: ) assert handled is True + assert annotation_reply is None assert any(isinstance(e, QueueLLMChunkEvent) for e in qm.events) assert _answer_text(qm.events) == "blocked preset answer" assert _saved_user_query(qm.events) == "forbidden" @@ -129,7 +135,8 @@ class TestRunInputGuards: _patch_annotation(monkeypatch, reply=SimpleNamespace(id="anno-1", content="I am the annotated Iris.")) qm = _FakeQueueManager() - handled, _ = AgentAppGenerator()._run_input_guards( + handled, _, annotation_reply = AgentAppGenerator()._run_input_guards( + session=MagicMock(), application_generate_entity=_make_entity("what is your name"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -137,8 +144,5 @@ class TestRunInputGuards: ) assert handled is True - annotation_events = [e for e in qm.events if isinstance(e, QueueAnnotationReplyEvent)] - assert len(annotation_events) == 1 - assert annotation_events[0].message_annotation_id == "anno-1" - assert _answer_text(qm.events) == "I am the annotated Iris." - assert _saved_user_query(qm.events) == "what is your name" + assert annotation_reply.id == "anno-1" + assert qm.events == [] diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py index c00ca9ac0b1..3fd3eee1b87 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py @@ -12,10 +12,9 @@ from typing import Any import pytest -from core.app.apps.agent_app import app_generator as gen_mod from core.app.apps.agent_app.app_generator import AgentAppGenerator, AgentAppGeneratorError, AgentAppNotPublishedError from core.app.entities.app_invoke_entities import InvokeFrom -from models.agent import AgentSource +from models.agent import AgentConfigDraftType, AgentSource _SOUL_DICT = { "model": { @@ -28,17 +27,21 @@ _SOUL_DICT = { class _FakeScalarSession: - """db.session stub: scalar() pops the next queued row (ignores the stmt).""" + """Session stub whose scalar() pops the next queued row.""" def __init__(self, values: list[Any]) -> None: self._values = list(values) + self.added: list[Any] = [] + self.flush_count = 0 def scalar(self, _stmt: Any) -> Any: return self._values.pop(0) if self._values else None + def add(self, value: Any) -> None: + self.added.append(value) -def _patch_session(monkeypatch, values: list[Any]) -> None: - monkeypatch.setattr(gen_mod, "db", SimpleNamespace(session=_FakeScalarSession(values))) + def flush(self) -> None: + self.flush_count += 1 def _snapshot() -> SimpleNamespace: @@ -46,13 +49,13 @@ def _snapshot() -> SimpleNamespace: class TestResolveAgentById: - def test_success_returns_agent_snapshot_soul(self, monkeypatch: pytest.MonkeyPatch): + def test_success_returns_agent_snapshot_soul(self): agent = SimpleNamespace(id="agent-1") snapshot = _snapshot() - _patch_session(monkeypatch, [agent, snapshot]) + session = _FakeScalarSession([agent, snapshot]) resolved_agent, resolved_snapshot, soul = AgentAppGenerator._resolve_agent_by_id( - tenant_id="t1", agent_id="agent-1", snapshot_id="snap-1" + tenant_id="t1", agent_id="agent-1", snapshot_id="snap-1", session=session ) assert resolved_agent is agent @@ -61,24 +64,55 @@ class TestResolveAgentById: assert soul.model is not None assert soul.model.model == "gpt-4o-mini" - def test_agent_missing_raises(self, monkeypatch: pytest.MonkeyPatch): - _patch_session(monkeypatch, [None]) + def test_agent_missing_raises(self): + session = _FakeScalarSession([None]) with pytest.raises(AgentAppGeneratorError, match="Agent not found"): - AgentAppGenerator._resolve_agent_by_id(tenant_id="t1", agent_id="x", snapshot_id="snap-1") + AgentAppGenerator._resolve_agent_by_id(tenant_id="t1", agent_id="x", snapshot_id="snap-1", session=session) - def test_no_published_version_raises(self, monkeypatch: pytest.MonkeyPatch): - _patch_session(monkeypatch, [SimpleNamespace(id="agent-1")]) + def test_no_published_version_raises(self): + session = _FakeScalarSession([SimpleNamespace(id="agent-1")]) with pytest.raises(AgentAppGeneratorError, match="no published version"): - AgentAppGenerator._resolve_agent_by_id(tenant_id="t1", agent_id="agent-1", snapshot_id=None) + AgentAppGenerator._resolve_agent_by_id( + tenant_id="t1", agent_id="agent-1", snapshot_id=None, session=session + ) - def test_snapshot_missing_raises(self, monkeypatch: pytest.MonkeyPatch): - _patch_session(monkeypatch, [SimpleNamespace(id="agent-1"), None]) + def test_snapshot_missing_raises(self): + session = _FakeScalarSession([SimpleNamespace(id="agent-1"), None]) with pytest.raises(AgentAppGeneratorError, match="published version not found"): - AgentAppGenerator._resolve_agent_by_id(tenant_id="t1", agent_id="agent-1", snapshot_id="snap-1") + AgentAppGenerator._resolve_agent_by_id( + tenant_id="t1", + agent_id="agent-1", + snapshot_id="snap-1", + session=session, + ) + + +class TestResolveDebugDraft: + def test_missing_shared_draft_is_created_with_supplied_session(self): + agent = SimpleNamespace( + id="agent-1", + active_config_snapshot_id="snap-1", + created_by="creator-1", + updated_by="updater-1", + ) + session = _FakeScalarSession([None, SimpleNamespace(id="agent-1"), _snapshot()]) + + draft = AgentAppGenerator._resolve_debug_draft( + tenant_id="t1", + agent=agent, + draft_type=None, + account_id=None, + session=session, + ) + + assert draft.draft_type == AgentConfigDraftType.DRAFT + assert draft.base_snapshot_id == "snap-1" + assert session.added == [draft] + assert session.flush_count == 1 class TestResolveAgent: - def test_success_chains_to_resolve_by_id(self, monkeypatch: pytest.MonkeyPatch): + def test_success_chains_to_resolve_by_id(self): bound_agent = SimpleNamespace( id="agent-1", source=AgentSource.AGENT_APP, @@ -88,7 +122,7 @@ class TestResolveAgent: inner_agent = SimpleNamespace(id="agent-1") snapshot = _snapshot() # scalar order: bound agent (in _resolve_agent), then agent + snapshot (in _resolve_agent_by_id) - _patch_session(monkeypatch, [bound_agent, inner_agent, snapshot]) + session = _FakeScalarSession([bound_agent, inner_agent, snapshot]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") agent, config_id, config_version_kind, soul = AgentAppGenerator()._resolve_agent( @@ -96,6 +130,7 @@ class TestResolveAgent: invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] assert agent is bound_agent @@ -103,7 +138,7 @@ class TestResolveAgent: assert config_version_kind == "snapshot" assert soul.model is not None - def test_unpublished_draft_still_resolves_active_snapshot(self, monkeypatch: pytest.MonkeyPatch): + def test_unpublished_draft_still_resolves_active_snapshot(self): bound_agent = SimpleNamespace( id="agent-1", source=AgentSource.AGENT_APP, @@ -112,7 +147,7 @@ class TestResolveAgent: ) inner_agent = SimpleNamespace(id="agent-1") snapshot = _snapshot() - _patch_session(monkeypatch, [bound_agent, inner_agent, snapshot]) + session = _FakeScalarSession([bound_agent, inner_agent, snapshot]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") agent, config_id, config_version_kind, soul = AgentAppGenerator()._resolve_agent( @@ -120,6 +155,7 @@ class TestResolveAgent: invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] assert agent is bound_agent @@ -127,14 +163,14 @@ class TestResolveAgent: assert config_version_kind == "snapshot" assert soul.prompt.system_prompt == "You are Iris." - def test_unpublished_imported_agent_is_not_available_to_public_runtime(self, monkeypatch: pytest.MonkeyPatch): + def test_unpublished_imported_agent_is_not_available_to_public_runtime(self): bound_agent = SimpleNamespace( id="agent-1", source=AgentSource.IMPORTED, active_config_snapshot_id="snap-1", active_config_is_published=False, ) - _patch_session(monkeypatch, [bound_agent]) + session = _FakeScalarSession([bound_agent]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") with pytest.raises(AgentAppNotPublishedError, match="not been published"): @@ -143,9 +179,10 @@ class TestResolveAgent: invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] - def test_unpublished_imported_agent_remains_available_to_debugger(self, monkeypatch: pytest.MonkeyPatch): + def test_unpublished_imported_agent_remains_available_to_debugger(self): bound_agent = SimpleNamespace( id="agent-1", source=AgentSource.IMPORTED, @@ -153,7 +190,7 @@ class TestResolveAgent: active_config_is_published=False, ) draft = SimpleNamespace(id="draft-1", draft_type="draft", config_snapshot_dict=_SOUL_DICT) - _patch_session(monkeypatch, [bound_agent, draft]) + session = _FakeScalarSession([bound_agent, draft]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") agent, config_id, config_version_kind, soul = AgentAppGenerator()._resolve_agent( @@ -161,6 +198,7 @@ class TestResolveAgent: invoke_from=InvokeFrom.DEBUGGER, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] assert agent is bound_agent @@ -168,14 +206,14 @@ class TestResolveAgent: assert config_version_kind == "draft" assert soul.prompt.system_prompt == "You are Iris." - def test_agent_without_active_snapshot_raises_before_model_resolution(self, monkeypatch: pytest.MonkeyPatch): + def test_agent_without_active_snapshot_raises_before_model_resolution(self): bound_agent = SimpleNamespace( id="agent-1", source=AgentSource.AGENT_APP, active_config_snapshot_id=None, active_config_is_published=False, ) - _patch_session(monkeypatch, [bound_agent]) + session = _FakeScalarSession([bound_agent]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") with pytest.raises(AgentAppNotPublishedError, match="not been published"): @@ -184,10 +222,11 @@ class TestResolveAgent: invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] - def test_unbound_app_raises(self, monkeypatch: pytest.MonkeyPatch): - _patch_session(monkeypatch, [None]) + def test_unbound_app_raises(self): + session = _FakeScalarSession([None]) app_model = SimpleNamespace(id="app-1", tenant_id="t1") with pytest.raises(AgentAppGeneratorError, match="has no bound Agent"): AgentAppGenerator()._resolve_agent( @@ -195,4 +234,5 @@ class TestResolveAgent: invoke_from=InvokeFrom.WEB_APP, draft_type=None, user=SimpleNamespace(id="user-1"), + session=session, ) # type: ignore[arg-type] diff --git a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_config_manager.py b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_config_manager.py index d47b70e9502..8dd79e347cf 100644 --- a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_config_manager.py +++ b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_config_manager.py @@ -1,5 +1,6 @@ import uuid from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from pytest_mock import MockerFixture @@ -39,6 +40,7 @@ class TestAgentChatAppConfigManagerGetAppConfig: app_model_config=app_model_config, conversation=None, override_config_dict=override_config, + annotation_reply=None, ) assert result.app_model_config_dict == override_config @@ -50,6 +52,7 @@ class TestAgentChatAppConfigManagerGetAppConfig: app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat") app_model_config = mocker.MagicMock(id="cfg1") app_model_config.to_dict.return_value = {"model": {"provider": "p"}} + annotation_reply = {"enabled": False} conversation = mocker.MagicMock() mocker.patch("core.app.apps.agent_chat.app_config_manager.ModelConfigManager.convert") @@ -72,15 +75,18 @@ class TestAgentChatAppConfigManagerGetAppConfig: app_model_config=app_model_config, conversation=conversation, override_config_dict=None, + annotation_reply=annotation_reply, ) assert result.app_model_config_dict == app_model_config.to_dict.return_value assert result.app_model_config_from.value == "conversation-specific-config" + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) def test_get_app_config_latest_config(self, mocker: MockerFixture): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat") app_model_config = mocker.MagicMock(id="cfg1") app_model_config.to_dict.return_value = {"model": {"provider": "p"}} + annotation_reply = {"enabled": False} mocker.patch("core.app.apps.agent_chat.app_config_manager.ModelConfigManager.convert") mocker.patch("core.app.apps.agent_chat.app_config_manager.PromptTemplateConfigManager.convert") @@ -102,9 +108,11 @@ class TestAgentChatAppConfigManagerGetAppConfig: app_model_config=app_model_config, conversation=None, override_config_dict=None, + annotation_reply=annotation_reply, ) assert result.app_model_config_from.value == "app-latest-config" + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) class TestAgentChatAppConfigManagerConfigValidate: @@ -147,7 +155,7 @@ class TestAgentChatAppConfigManagerConfigValidate: mocker.patch.object( AgentChatAppConfigManager, "validate_agent_mode_and_set_defaults", - side_effect=lambda tenant_id, cfg: return_with_key("agent_mode"), + side_effect=lambda tenant_id, cfg, session: return_with_key("agent_mode"), ) mocker.patch( "core.app.apps.agent_chat.app_config_manager.OpeningStatementConfigManager.validate_and_set_defaults", @@ -171,14 +179,14 @@ class TestAgentChatAppConfigManagerConfigValidate: ) mocker.patch( "core.app.apps.agent_chat.app_config_manager.DatasetConfigManager.validate_and_set_defaults", - side_effect=lambda tenant_id, app_mode, cfg: return_with_key("dataset"), + side_effect=lambda tenant_id, app_mode, cfg, session: return_with_key("dataset"), ) mocker.patch( "core.app.apps.agent_chat.app_config_manager.SensitiveWordAvoidanceConfigManager.validate_and_set_defaults", side_effect=lambda tenant_id, cfg: return_with_key("moderation"), ) - filtered = AgentChatAppConfigManager.config_validate("tenant", config) + filtered = AgentChatAppConfigManager.config_validate("tenant", config, mocker.MagicMock()) assert set(filtered.keys()) == { "model", "user_input_form", @@ -199,7 +207,7 @@ class TestAgentChatAppConfigManagerConfigValidate: class TestValidateAgentModeAndSetDefaults: def test_defaults_when_missing(self): config = {} - updated, keys = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config) + updated, keys = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config, MagicMock()) assert "agent_mode" in updated assert updated["agent_mode"]["enabled"] is False assert updated["agent_mode"]["tools"] == [] @@ -211,34 +219,38 @@ class TestValidateAgentModeAndSetDefaults: ) def test_agent_mode_type_validation(self, agent_mode): with pytest.raises(ValueError): - AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", {"agent_mode": agent_mode}) + AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( + "tenant", {"agent_mode": agent_mode}, MagicMock() + ) def test_agent_mode_empty_list_defaults(self): config = {"agent_mode": []} - updated, _ = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config) + updated, _ = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config, MagicMock()) assert updated["agent_mode"]["enabled"] is False assert updated["agent_mode"]["tools"] == [] def test_enabled_must_be_bool(self): with pytest.raises(ValueError): - AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", {"agent_mode": {"enabled": "yes"}}) + AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( + "tenant", {"agent_mode": {"enabled": "yes"}}, MagicMock() + ) def test_strategy_must_be_valid(self): with pytest.raises(ValueError): AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( - "tenant", {"agent_mode": {"enabled": True, "strategy": "invalid"}} + "tenant", {"agent_mode": {"enabled": True, "strategy": "invalid"}}, MagicMock() ) def test_tools_must_be_list(self): with pytest.raises(ValueError): AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( - "tenant", {"agent_mode": {"enabled": True, "tools": "not-list"}} + "tenant", {"agent_mode": {"enabled": True, "tools": "not-list"}}, MagicMock() ) def test_old_tool_dataset_requires_id(self): with pytest.raises(ValueError): AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( - "tenant", {"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": True}}]}} + "tenant", {"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": True}}]}}, MagicMock() ) def test_old_tool_dataset_id_must_be_uuid(self): @@ -246,6 +258,7 @@ class TestValidateAgentModeAndSetDefaults: AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( "tenant", {"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": True, "id": "bad"}}]}}, + MagicMock(), ) def test_old_tool_dataset_id_not_exists(self, mocker: MockerFixture): @@ -258,6 +271,7 @@ class TestValidateAgentModeAndSetDefaults: AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( "tenant", {"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": True, "id": dataset_id}}]}}, + MagicMock(), ) def test_old_tool_enabled_must_be_bool(self): @@ -265,6 +279,7 @@ class TestValidateAgentModeAndSetDefaults: AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( "tenant", {"agent_mode": {"enabled": True, "tools": [{"dataset": {"enabled": "yes", "id": str(uuid.uuid4())}}]}}, + MagicMock(), ) @pytest.mark.parametrize("missing_key", ["provider_type", "provider_id", "tool_name", "tool_parameters"]) @@ -273,7 +288,7 @@ class TestValidateAgentModeAndSetDefaults: tool.pop(missing_key, None) with pytest.raises(ValueError): AgentChatAppConfigManager.validate_agent_mode_and_set_defaults( - "tenant", {"agent_mode": {"enabled": True, "tools": [tool]}} + "tenant", {"agent_mode": {"enabled": True, "tools": [tool]}}, MagicMock() ) def test_valid_old_and_new_style_tools(self, mocker: MockerFixture): @@ -298,6 +313,6 @@ class TestValidateAgentModeAndSetDefaults: } } - updated, _ = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config) + updated, _ = AgentChatAppConfigManager.validate_agent_mode_and_set_defaults("tenant", config, MagicMock()) assert updated["agent_mode"]["tools"][0]["dataset"]["enabled"] is False assert updated["agent_mode"]["tools"][1]["enabled"] is False diff --git a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py index 60297b45181..0dadae0064b 100644 --- a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py @@ -1,4 +1,5 @@ import contextlib +import inspect import logging import pytest @@ -33,19 +34,33 @@ class TestAgentChatAppGeneratorGenerate: app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): - generator.generate(app_model=app_model, user=user, args={}, invoke_from=mocker.MagicMock(), streaming=False) + generator.generate( + session=mocker.MagicMock(), + app_model=app_model, + user=user, + args={}, + invoke_from=mocker.MagicMock(), + streaming=False, + ) def test_generate_requires_query(self, generator, mocker: MockerFixture): app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): - generator.generate(app_model=app_model, user=user, args={"inputs": {}}, invoke_from=mocker.MagicMock()) + generator.generate( + session=mocker.MagicMock(), + app_model=app_model, + user=user, + args={"inputs": {}}, + invoke_from=mocker.MagicMock(), + ) def test_generate_rejects_non_string_query(self, generator, mocker: MockerFixture): app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): generator.generate( + session=mocker.MagicMock(), app_model=app_model, user=user, args={"query": 123, "inputs": {}}, @@ -58,6 +73,7 @@ class TestAgentChatAppGeneratorGenerate: with pytest.raises(ValueError): generator.generate( + session=mocker.MagicMock(), app_model=app_model, user=user, args={"query": "hi", "inputs": {}, "model_config": {"model": {"provider": "p"}}}, @@ -120,9 +136,6 @@ class TestAgentChatAppGeneratorGenerate: "core.app.apps.agent_chat.app_generator.threading.Thread", return_value=thread_obj, ) - session = mocker.MagicMock() - mocker.patch("core.app.apps.agent_chat.app_generator.db.session", return_value=session) - mocker.patch( "core.app.apps.agent_chat.app_generator.AgentChatAppGenerateResponseConverter.convert", return_value={"result": "ok"}, @@ -141,18 +154,30 @@ class TestAgentChatAppGeneratorGenerate: "files": [{"id": "f1"}], "trace_session_id": "session-1", } + session = mocker.MagicMock() - result = generator.generate(app_model=app_model, user=user, args=args, invoke_from=invoke_from, streaming=True) + result = generator.generate( + session=session, + app_model=app_model, + user=user, + args=args, + invoke_from=invoke_from, + streaming=True, + ) assert result == {"result": "ok"} + assert generator._get_app_model_config.call_args.kwargs["session"] is session + assert generator._init_generate_records.call_args.kwargs["session"] is session assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" - assert thread_constructor.call_args.kwargs["kwargs"]["session"] is session + worker_call = thread_constructor.call_args + inspect.signature(worker_call.kwargs["target"]).bind(**worker_call.kwargs["kwargs"]) thread_obj.start.assert_called_once() def test_generate_without_file_config(self, generator, mocker: MockerFixture): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat") - app_model_config = mocker.MagicMock(id="cfg1") + app_model_config = mocker.MagicMock(id="cfg1", app_id="app1") app_model_config.to_dict.return_value = {"model": {"provider": "p"}} + annotation_reply = {"enabled": False} user = DummyAccount("user") @@ -163,7 +188,11 @@ class TestAgentChatAppGeneratorGenerate: ) generator._handle_response = mocker.MagicMock(return_value="response") - mocker.patch( + load_annotation_reply_config = mocker.patch( + "core.app.apps.agent_chat.app_generator.load_annotation_reply_config", + return_value=annotation_reply, + ) + get_app_config = mocker.patch( "core.app.apps.agent_chat.app_generator.AgentChatAppConfigManager.get_app_config", return_value=mocker.MagicMock(variables={}, prompt_template=mocker.MagicMock(), external_data_variables=[]), ) @@ -206,8 +235,10 @@ class TestAgentChatAppGeneratorGenerate: ) args = {"query": "hello", "inputs": {"name": "world"}} + session = mocker.MagicMock() result = generator.generate( + session=session, app_model=app_model, user=user, args=args, @@ -216,6 +247,9 @@ class TestAgentChatAppGeneratorGenerate: ) assert result == {"result": "ok"} + load_annotation_reply_config.assert_called_once_with(session, "app1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply class TestAgentChatAppGeneratorWorker: @@ -235,11 +269,13 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = GenerateTaskStoppedError() mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - mocker.patch("core.app.apps.agent_chat.app_generator.db.session.close") + session_cm = mocker.MagicMock() + session_cm.__enter__.return_value = mocker.MagicMock() + create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") + create_session.return_value = session_cm generator._generate_worker( flask_app=mocker.MagicMock(), - session=mocker.MagicMock(), context=mocker.MagicMock(), application_generate_entity=mocker.MagicMock(), queue_manager=queue_manager, @@ -266,11 +302,13 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = error mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - mocker.patch("core.app.apps.agent_chat.app_generator.db.session.close") + session_cm = mocker.MagicMock() + session_cm.__enter__.return_value = mocker.MagicMock() + create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") + create_session.return_value = session_cm generator._generate_worker( flask_app=mocker.MagicMock(), - session=mocker.MagicMock(), context=mocker.MagicMock(), application_generate_entity=mocker.MagicMock(), queue_manager=queue_manager, @@ -290,14 +328,16 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = ValueError("bad") mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - mocker.patch("core.app.apps.agent_chat.app_generator.db.session.close") + session_cm = mocker.MagicMock() + session_cm.__enter__.return_value = mocker.MagicMock() + create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") + create_session.return_value = session_cm mocker.patch("core.app.apps.agent_chat.app_generator.dify_config", new=mocker.MagicMock(DEBUG=True)) with caplog.at_level(logging.ERROR, logger="core.app.apps.agent_chat.app_generator"): generator._generate_worker( flask_app=mocker.MagicMock(), - session=mocker.MagicMock(), context=mocker.MagicMock(), application_generate_entity=mocker.MagicMock(), queue_manager=queue_manager, diff --git a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py index a3a879bee3d..f4caca4e1d7 100644 --- a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py @@ -33,7 +33,7 @@ class TestAgentChatAppRunnerRun: patch_create_session(mocker, return_value=None) with pytest.raises(ValueError): - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) def test_run_moderation_error_direct_output(self, runner: AgentChatAppRunner, mocker: MockerFixture): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") @@ -54,7 +54,7 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "moderation_for_inputs", side_effect=ModerationError("bad")) mocker.patch.object(runner, "direct_output") - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) runner.direct_output.assert_called_once() @@ -78,13 +78,15 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) annotation = mocker.MagicMock(id="anno", content="answer") - mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=annotation) + annotation_query = mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=annotation) mocker.patch.object(runner, "direct_output") queue_manager = mocker.MagicMock() - runner.run(mocker.MagicMock(), generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock()) + write_session = mocker.MagicMock() + runner.run(generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock(), write_session) queue_manager.publish.assert_called_once() + assert annotation_query.call_args.kwargs["session"] is write_session runner.direct_output.assert_called_once() def test_run_hosting_moderation_short_circuits(self, runner: AgentChatAppRunner, mocker: MockerFixture): @@ -109,7 +111,7 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) mocker.patch.object(runner, "check_hosting_moderation", return_value=True) - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) def test_run_model_schema_missing(self, runner: AgentChatAppRunner, mocker: MockerFixture): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") @@ -144,7 +146,7 @@ class TestAgentChatAppRunnerRun: mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) with pytest.raises(ValueError): - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) @pytest.mark.parametrize( ("mode", "expected_runner"), @@ -198,11 +200,16 @@ class TestAgentChatAppRunnerRun: runner_instance = mocker.MagicMock() runner_cls.return_value = runner_instance - runner_instance.run.return_value = [] + events: list[str] = [] + runner_instance.run.side_effect = lambda **_kwargs: events.append("agent-run") or [] mocker.patch.object(runner, "_handle_invoke_result") + session = mocker.MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session.close.side_effect = lambda: events.append("close") - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, session) + assert events == ["commit", "close", "commit", "close", "agent-run"] runner_instance.run.assert_called_once() runner._handle_invoke_result.assert_called_once() @@ -247,7 +254,7 @@ class TestAgentChatAppRunnerRun: patch_create_session(mocker, side_effect=[app_record, conversation, message]) with pytest.raises(ValueError): - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) def test_run_function_calling_strategy_selected_by_features( self, runner: AgentChatAppRunner, mocker: MockerFixture @@ -299,7 +306,7 @@ class TestAgentChatAppRunnerRun: runner_instance.run.return_value = [] mocker.patch.object(runner, "_handle_invoke_result") - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) assert app_config.agent.strategy == AgentEntity.Strategy.FUNCTION_CALLING runner_instance.run.assert_called_once() @@ -334,11 +341,11 @@ class TestAgentChatAppRunnerRun: with pytest.raises(ValueError): runner.run( - mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"), + mocker.MagicMock(), ) def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture): @@ -371,11 +378,11 @@ class TestAgentChatAppRunnerRun: with pytest.raises(ValueError): runner.run( - mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"), + mocker.MagicMock(), ) def test_run_invalid_agent_strategy_raises(self, runner: AgentChatAppRunner, mocker: MockerFixture): @@ -419,4 +426,4 @@ class TestAgentChatAppRunnerRun: patch_create_session(mocker, side_effect=[app_record, conversation, message]) with pytest.raises(ValueError): - runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) diff --git a/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py b/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py index 271d007be6e..24ec3116b32 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py @@ -1,5 +1,5 @@ from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch from core.app.app_config.entities import EasyUIBasedAppModelConfigFrom, ModelConfigEntity, PromptTemplateEntity from core.app.apps.chat.app_config_manager import ChatAppConfigManager @@ -35,12 +35,47 @@ class TestChatAppConfigManager: app_model_config=app_model_config, conversation=None, override_config_dict=override, + annotation_reply=None, ) assert app_config.app_model_config_from == EasyUIBasedAppModelConfigFrom.ARGS assert app_config.app_model_config_dict == override assert app_config.app_mode == AppMode.CHAT + def test_get_app_config_uses_injected_annotation_reply(self): + app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.CHAT.value) + app_model_config = SimpleNamespace( + id="config-1", + to_dict=MagicMock(return_value={"model": "m"}), + ) + annotation_reply = {"enabled": False} + + model_entity = ModelConfigEntity(provider="p", model="m") + prompt_entity = PromptTemplateEntity( + prompt_type=PromptTemplateEntity.PromptType.SIMPLE, + simple_prompt_template="hi", + ) + + with ( + patch("core.app.apps.chat.app_config_manager.ModelConfigManager.convert", return_value=model_entity), + patch( + "core.app.apps.chat.app_config_manager.PromptTemplateConfigManager.convert", return_value=prompt_entity + ), + patch( + "core.app.apps.chat.app_config_manager.SensitiveWordAvoidanceConfigManager.convert", + return_value=None, + ), + patch("core.app.apps.chat.app_config_manager.DatasetConfigManager.convert", return_value=None), + patch("core.app.apps.chat.app_config_manager.BasicVariablesConfigManager.convert", return_value=([], [])), + ): + ChatAppConfigManager.get_app_config( + app_model=app_model, + app_model_config=app_model_config, + annotation_reply=annotation_reply, + ) + + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + def test_config_validate_filters_related_keys(self): config = {"extra": 1} @@ -98,7 +133,7 @@ class TestChatAppConfigManager: side_effect=_add_key("sensitive_word_avoidance", 11), ), ): - filtered = ChatAppConfigManager.config_validate(tenant_id="t1", config=config) + filtered = ChatAppConfigManager.config_validate(session=MagicMock(), tenant_id="t1", config=config) assert filtered["model"] == 1 assert filtered["inputs"] == 2 diff --git a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py index 5b4e7bacc68..bdb706bde09 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py @@ -72,10 +72,19 @@ class TestChatAppGenerator: generator = ChatAppGenerator() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") user = SimpleNamespace(id="user-1", session_id="session-1") - args = {"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}, "trace_session_id": "session-1"} + session = MagicMock() + args = { + "query": "hi", + "inputs": {}, + "conversation_id": "conversation-1", + "model_config": {"foo": "bar"}, + "trace_session_id": "session-1", + } with ( - patch("core.app.apps.chat.app_generator.ConversationService.get_conversation", return_value=None), + patch( + "core.app.apps.chat.app_generator.ConversationService.get_conversation", return_value=None + ) as get_conversation, patch("core.app.apps.chat.app_generator.ChatAppConfigManager.config_validate", return_value={"x": 1}), patch( "core.app.apps.chat.app_generator.ChatAppConfigManager.get_app_config", @@ -107,11 +116,45 @@ class TestChatAppGenerator: patch("core.app.apps.chat.app_generator.threading.Thread") as mock_thread, ): mock_thread.return_value.start.return_value = None - result = generator.generate(MagicMock(), app_model, user, args, InvokeFrom.DEBUGGER, streaming=False) + result = generator.generate(app_model, user, args, InvokeFrom.DEBUGGER, streaming=False, session=session) assert result == {"ok": True} + assert get_conversation.call_args.kwargs["session"] is session assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" + def test_generate_uses_session_for_annotation_reply(self): + generator = ChatAppGenerator() + app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") + app_model_config = MagicMock(id="config-1", app_id="app-1") + annotation_reply = {"enabled": False} + user = SimpleNamespace(id="user-1", session_id="session-1") + session = MagicMock() + + with ( + patch.object(ChatAppGenerator, "_get_app_model_config", return_value=app_model_config), + patch( + "core.app.apps.chat.app_generator.load_annotation_reply_config", + return_value=annotation_reply, + ) as load_annotation_reply_config, + patch("core.app.apps.chat.app_generator.FileUploadConfigManager.convert", return_value=None), + patch( + "core.app.apps.chat.app_generator.ChatAppConfigManager.get_app_config", + side_effect=RuntimeError("stop after app config"), + ) as get_app_config, + ): + with pytest.raises(RuntimeError, match="stop after app config"): + generator.generate( + app_model, + user, + {"query": "hi", "inputs": {}}, + InvokeFrom.WEB_APP, + session=session, + ) + + load_annotation_reply_config.assert_called_once_with(session, "app-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply + def test_generate_rejects_model_config_override_for_non_debugger(self): generator = ChatAppGenerator() with pytest.raises(ValueError): @@ -138,11 +181,12 @@ class TestChatAppGenerator: patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=InvokeAuthorizationError()), + patch("core.app.apps.chat.app_generator.session_factory.create_session") as create_session, patch("core.app.apps.chat.app_generator.db.session.close"), ): + create_session.return_value.__enter__.return_value = MagicMock() generator._generate_worker( flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), - session=MagicMock(), application_generate_entity=entity, queue_manager=queue_manager, conversation_id="c1", @@ -155,11 +199,12 @@ class TestChatAppGenerator: patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=GenerateTaskStoppedError()), + patch("core.app.apps.chat.app_generator.session_factory.create_session") as create_session, patch("core.app.apps.chat.app_generator.db.session.close"), ): + create_session.return_value.__enter__.return_value = MagicMock() generator._generate_worker( flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), - session=MagicMock(), application_generate_entity=entity, queue_manager=queue_manager, conversation_id="c1", @@ -189,7 +234,7 @@ class TestChatAppRunner: with patched_create_session(return_value=None): with pytest.raises(ValueError): runner.run( - MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() ) def test_run_moderation_error_direct_output(self): @@ -222,7 +267,7 @@ class TestChatAppRunner: patch.object(ChatAppRunner, "direct_output") as mock_direct, ): runner.run( - MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() ) mock_direct.assert_called_once() @@ -256,13 +301,15 @@ class TestChatAppRunner: patched_create_session(return_value=SimpleNamespace(id="app-1", tenant_id="tenant-1")), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), - patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=annotation), + patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=annotation) as annotation_query, patch.object(ChatAppRunner, "direct_output") as mock_direct, ): queue_manager = DummyQueueManager() - runner.run(MagicMock(), app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) + write_session = MagicMock() + runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), write_session) assert any(isinstance(item[0], QueueAnnotationReplyEvent) for item in queue_manager.published) + assert annotation_query.call_args.kwargs["session"] is write_session mock_direct.assert_called_once() def test_run_returns_when_hosting_moderation_blocks(self): @@ -296,10 +343,10 @@ class TestChatAppRunner: patch.object(ChatAppRunner, "check_hosting_moderation", return_value=True), ): runner.run( - MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() ) - def test_run_closes_scoped_session_before_stream_consumption(self): + def test_run_closes_explicit_session_before_stream_consumption(self): runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", @@ -325,6 +372,9 @@ class TestChatAppRunner: events = [] queue_manager = DummyQueueManager() model_instance = MagicMock() + session = MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session.close.side_effect = lambda: events.append("close") def invoke_stream(): events.append("first-chunk") @@ -347,12 +397,11 @@ class TestChatAppRunner: side_effect=lambda invoke_result, **kwargs: list(invoke_result), ) as mock_handle, patch("core.app.apps.chat.app_runner.ModelInstance", return_value=model_instance), - patch("core.app.apps.chat.app_runner.db.session.close", side_effect=lambda: events.append("close")), ): model_instance.invoke_llm.side_effect = invoke_llm - runner.run(MagicMock(), app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), session) - assert events == ["close", "invoke", "first-chunk"] + assert events == ["commit", "close", "commit", "close", "invoke", "first-chunk"] mock_handle.assert_called_once_with( invoke_result=ANY, queue_manager=queue_manager, diff --git a/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py b/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py index 130264972a3..6d90aa7e53b 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_base_app_runner_multimodal.py @@ -7,7 +7,6 @@ import pytest from core.app.apps.base_app_runner import AppRunner from core.app.entities.app_invoke_entities import InvokeFrom -from core.app.entities.queue_entities import QueueMessageFileEvent from graphon.file import FileTransferMethod, FileType from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent from models.enums import CreatorUserRole @@ -81,54 +80,40 @@ class TestBaseAppRunnerMultimodal: mock_msg_file_class.return_value = mock_message_file file_session = MagicMock() - mock_session_factory = MagicMock() - mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) - mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) + # Act + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory) as mock_sm: - with patch("core.app.apps.base_app_runner.db") as mock_db: - # Act - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method( - runner, *args, **kwargs - ) + message_file_id = runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + # Assert + mock_mgr.create_file_by_url.assert_called_once_with( + user_id=mock_user_id, + tenant_id=mock_tenant_id, + file_url=image_url, + conversation_id=None, + ) - # Assert - mock_mgr.create_file_by_url.assert_called_once_with( - user_id=mock_user_id, - tenant_id=mock_tenant_id, - file_url=image_url, - conversation_id=None, - ) + mock_msg_file_class.assert_called_once() + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["message_id"] == mock_message_id + assert call_kwargs["type"] == FileType.IMAGE + assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE + assert call_kwargs["belongs_to"] == "assistant" + assert call_kwargs["created_by"] == mock_user_id - mock_msg_file_class.assert_called_once() - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["message_id"] == mock_message_id - assert call_kwargs["type"] == FileType.IMAGE - assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE - assert call_kwargs["belongs_to"] == "assistant" - assert call_kwargs["created_by"] == mock_user_id - - # Verify independent session was used (not db.session) - mock_sm.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - file_session.add.assert_called_once_with(mock_message_file) - mock_db.session.commit.assert_not_called() - mock_db.session.close.assert_not_called() - - # Verify event was published - mock_queue_manager.publish.assert_called_once() - publish_call = mock_queue_manager.publish.call_args - assert isinstance(publish_call[0][0], QueueMessageFileEvent) - assert publish_call[0][0].message_file_id == mock_message_file.id + file_session.add.assert_called_once_with(mock_message_file) + file_session.flush.assert_called_once() + assert message_file_id == mock_message_file.id + mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_with_base64( self, @@ -163,41 +148,34 @@ class TestBaseAppRunnerMultimodal: mock_msg_file_class.return_value = mock_message_file file_session = MagicMock() - mock_session_factory = MagicMock() - mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) - mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): - with patch("core.app.apps.base_app_runner.db") as mock_db: - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method( - runner, *args, **kwargs - ) + message_file_id = runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + mock_mgr.create_file_by_raw.assert_called_once() + call_kwargs = mock_mgr.create_file_by_raw.call_args[1] + assert call_kwargs["user_id"] == mock_user_id + assert call_kwargs["tenant_id"] == mock_tenant_id + assert call_kwargs["conversation_id"] is None + assert "file_binary" in call_kwargs + assert call_kwargs["mimetype"] == "image/png" + assert call_kwargs["filename"].startswith("generated_image") + assert call_kwargs["filename"].endswith(".png") - mock_mgr.create_file_by_raw.assert_called_once() - call_kwargs = mock_mgr.create_file_by_raw.call_args[1] - assert call_kwargs["user_id"] == mock_user_id - assert call_kwargs["tenant_id"] == mock_tenant_id - assert call_kwargs["conversation_id"] is None - assert "file_binary" in call_kwargs - assert call_kwargs["mimetype"] == "image/png" - assert call_kwargs["filename"].startswith("generated_image") - assert call_kwargs["filename"].endswith(".png") - - mock_msg_file_class.assert_called_once() - file_session.add.assert_called_once() - mock_db.session.commit.assert_not_called() - - mock_queue_manager.publish.assert_called_once() + mock_msg_file_class.assert_called_once() + file_session.add.assert_called_once() + file_session.flush.assert_called_once() + assert message_file_id == mock_message_file.id + mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_with_base64_data_uri( self, @@ -230,29 +208,22 @@ class TestBaseAppRunnerMultimodal: mock_msg_file_class.return_value = mock_message_file file_session = MagicMock() - mock_session_factory = MagicMock() - mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) - mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): - with patch("core.app.apps.base_app_runner.db"): - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method( - runner, *args, **kwargs - ) + runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) - - mock_mgr.create_file_by_raw.assert_called_once() - call_kwargs = mock_mgr.create_file_by_raw.call_args[1] - assert "file_binary" in call_kwargs + mock_mgr.create_file_by_raw.assert_called_once() + call_kwargs = mock_mgr.create_file_by_raw.call_args[1] + assert "file_binary" in call_kwargs def test_handle_multimodal_image_content_without_url_or_base64( self, @@ -272,22 +243,23 @@ class TestBaseAppRunnerMultimodal: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - with patch("core.app.apps.base_app_runner.db"): - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + file_session = MagicMock() + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - mock_mgr_class.assert_not_called() - mock_msg_file_class.assert_not_called() - mock_queue_manager.publish.assert_not_called() + mock_mgr_class.assert_not_called() + mock_msg_file_class.assert_not_called() + mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_with_error( self, @@ -311,21 +283,22 @@ class TestBaseAppRunnerMultimodal: mock_mgr_class.return_value = mock_mgr with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: - with patch("core.app.apps.base_app_runner.db"): - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) + file_session = MagicMock() + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) + runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - mock_msg_file_class.assert_not_called() - mock_queue_manager.publish.assert_not_called() + mock_msg_file_class.assert_not_called() + mock_queue_manager.publish.assert_not_called() def test_handle_multimodal_image_content_debugger_mode( self, @@ -355,28 +328,21 @@ class TestBaseAppRunnerMultimodal: mock_msg_file_class.return_value = mock_message_file file_session = MagicMock() - mock_session_factory = MagicMock() - mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) - mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): - with patch("core.app.apps.base_app_runner.db"): - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method( - runner, *args, **kwargs - ) + runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) - - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT def test_handle_multimodal_image_content_service_api_mode( self, @@ -406,25 +372,18 @@ class TestBaseAppRunnerMultimodal: mock_msg_file_class.return_value = mock_message_file file_session = MagicMock() - mock_session_factory = MagicMock() - mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session) - mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False) + runner = MagicMock() + method = AppRunner._handle_multimodal_image_content + runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) - with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory): - with patch("core.app.apps.base_app_runner.db"): - runner = MagicMock() - method = AppRunner._handle_multimodal_image_content - runner._handle_multimodal_image_content = lambda *args, **kwargs: method( - runner, *args, **kwargs - ) + runner._handle_multimodal_image_content( + session=file_session, + content=content, + message_id=mock_message_id, + user_id=mock_user_id, + tenant_id=mock_tenant_id, + queue_manager=mock_queue_manager, + ) - runner._handle_multimodal_image_content( - content=content, - message_id=mock_message_id, - user_id=mock_user_id, - tenant_id=mock_tenant_id, - queue_manager=mock_queue_manager, - ) - - call_kwargs = mock_msg_file_class.call_args[1] - assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER + call_kwargs = mock_msg_file_class.call_args[1] + assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER diff --git a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py index 69522b193bf..6de22abbe1b 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py @@ -65,7 +65,7 @@ class TestCompletionAppRunner: with patched_create_session(return_value=None): with pytest.raises(ValueError): - runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock()) + runner.run(app_generate_entity, MagicMock(), MagicMock(), MagicMock()) def test_run_moderation_error_outputs_direct(self, runner, mocker: MockerFixture): app_record = MagicMock(id="app1", tenant_id="tenant") @@ -79,7 +79,7 @@ class TestCompletionAppRunner: runner._handle_invoke_result = MagicMock() with patched_create_session(return_value=app_record): - runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) runner.direct_output.assert_called_once() runner._handle_invoke_result.assert_not_called() @@ -96,7 +96,7 @@ class TestCompletionAppRunner: runner._handle_invoke_result = MagicMock() with patched_create_session(return_value=app_record): - runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) runner._handle_invoke_result.assert_not_called() @@ -133,19 +133,22 @@ class TestCompletionAppRunner: mocker.patch.object(module, "ModelInstance", return_value=model_instance) with patched_create_session(return_value=app_record): - runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant")) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant"), MagicMock()) dataset_retrieval.retrieve.assert_called_once() assert dataset_retrieval.retrieve.call_args.kwargs["query"] == "query_from_input" runner._handle_invoke_result.assert_called_once() - def test_run_closes_scoped_session_before_stream_consumption(self, runner, mocker: MockerFixture): + def test_run_closes_explicit_session_before_stream_consumption(self, runner, mocker: MockerFixture): app_record = MagicMock(id="app1", tenant_id="tenant") app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config) queue_manager = MagicMock() events = [] + session = MagicMock() + session.commit.side_effect = lambda: events.append("commit") + session.close.side_effect = lambda: events.append("close") runner.organize_prompt_messages = MagicMock(return_value=([], None)) runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query")) runner.check_hosting_moderation = MagicMock(return_value=False) @@ -164,12 +167,11 @@ class TestCompletionAppRunner: model_instance.invoke_llm.side_effect = invoke_llm mocker.patch.object(module, "ModelInstance", return_value=model_instance) - mocker.patch.object(module.db.session, "close", side_effect=lambda: events.append("close")) with patched_create_session(return_value=app_record): - runner.run(MagicMock(), app_generate_entity, queue_manager, MagicMock(id="msg")) + runner.run(app_generate_entity, queue_manager, MagicMock(id="msg"), session) - assert events == ["close", "invoke", "first-chunk"] + assert events == ["commit", "close", "invoke", "first-chunk"] runner._handle_invoke_result.assert_called_once_with( invoke_result=ANY, queue_manager=queue_manager, @@ -190,7 +192,7 @@ class TestCompletionAppRunner: runner.check_hosting_moderation = MagicMock(return_value=True) with patched_create_session(return_value=app_record): - runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) assert ( runner.organize_prompt_messages.call_args.kwargs["image_detail_config"] diff --git a/api/tests/unit_tests/core/app/apps/completion/test_completion_app_config_manager.py b/api/tests/unit_tests/core/app/apps/completion/test_completion_app_config_manager.py index 353162be8c5..bb2ae23ef96 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_completion_app_config_manager.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_completion_app_config_manager.py @@ -29,6 +29,7 @@ class TestCompletionAppConfigManager: app_model=app_model, app_model_config=app_model_config, override_config_dict=override_config, + annotation_reply=None, ) assert result.app_model_config_from == EasyUIBasedAppModelConfigFrom.ARGS @@ -41,6 +42,7 @@ class TestCompletionAppConfigManager: app_model = MagicMock(tenant_id="tenant", id="app1", mode=AppMode.COMPLETION) app_model_config = MagicMock(id="cfg1") app_model_config.to_dict.return_value = {"model": {"provider": "x"}} + annotation_reply = {"enabled": False} mocker.patch.object(module.ModelConfigManager, "convert", return_value="model") mocker.patch.object(module.PromptTemplateConfigManager, "convert", return_value="prompt") @@ -50,10 +52,15 @@ class TestCompletionAppConfigManager: mocker.patch.object(module.BasicVariablesConfigManager, "convert", return_value=([], [])) mocker.patch.object(module, "CompletionAppConfig", side_effect=lambda **kwargs: SimpleNamespace(**kwargs)) - result = CompletionAppConfigManager.get_app_config(app_model=app_model, app_model_config=app_model_config) + result = CompletionAppConfigManager.get_app_config( + app_model=app_model, + app_model_config=app_model_config, + annotation_reply=annotation_reply, + ) assert result.app_model_config_from == EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG assert result.app_model_config_dict == {"model": {"provider": "x"}} + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) def test_config_validate_filters_related_keys(self, mocker: MockerFixture): config = { @@ -109,7 +116,7 @@ class TestCompletionAppConfigManager: return_value=(config, ["moderation"]), ) - filtered = CompletionAppConfigManager.config_validate("tenant", config) + filtered = CompletionAppConfigManager.config_validate("tenant", config, MagicMock()) assert "extra" not in filtered assert set(filtered.keys()) == { diff --git a/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py b/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py index de0456851bd..22a4030183c 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py @@ -1,6 +1,6 @@ import contextlib from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, call import pytest from pydantic import ValidationError @@ -10,6 +10,7 @@ import core.app.apps.completion.app_generator as module from core.app.apps.completion.app_generator import CompletionAppGenerator from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom +from graphon.file import FILE_MODEL_IDENTITY from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from services.errors.app import MoreLikeThisDisabledError from services.errors.message import MessageNotExistsError @@ -39,7 +40,7 @@ def generator(mocker: MockerFixture): def _build_app_model(): - return MagicMock(tenant_id="tenant", id="app1", mode="completion") + return MagicMock(tenant_id="tenant", id="app1", mode="completion", app_model_config_id="cfg-current") def _build_user(): @@ -47,7 +48,7 @@ def _build_user(): def _build_app_model_config(): - config = MagicMock(id="cfg") + config = MagicMock(id="cfg", app_id="app1") config.to_dict.return_value = {"model": {"provider": "x"}} return config @@ -78,11 +79,21 @@ class TestCompletionAppGenerator: def test_generate_success_no_file_config(self, generator, mocker: MockerFixture): app_model_config = _build_app_model_config() mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config) + annotation_reply = {"enabled": False} + load_annotation_reply_config = mocker.patch.object( + module, + "load_annotation_reply_config", + return_value=annotation_reply, + ) mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=None) mocker.patch.object(module.file_factory, "build_from_mappings") app_config = MagicMock(variables=["v"], to_dict=MagicMock(return_value={})) - mocker.patch.object(module.CompletionAppConfigManager, "get_app_config", return_value=app_config) + get_app_config = mocker.patch.object( + module.CompletionAppConfigManager, + "get_app_config", + return_value=app_config, + ) mocker.patch.object(module.ModelConfigConverter, "convert", return_value=MagicMock()) mocker.patch.object(generator, "_prepare_user_inputs", return_value={"k": "v"}) @@ -94,8 +105,9 @@ class TestCompletionAppGenerator: mocker.patch.object(generator, "_handle_response", return_value="response") mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") + session = MagicMock() result = generator.generate( - session=MagicMock(), + session=session, app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {"a": 1}, "files": [], "trace_session_id": "session-1"}, @@ -106,6 +118,9 @@ class TestCompletionAppGenerator: assert result == "converted" assert generator.generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" module.file_factory.build_from_mappings.assert_not_called() + load_annotation_reply_config.assert_called_once_with(session, "app1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply def test_generate_success_with_files(self, generator, mocker: MockerFixture): app_model_config = _build_app_model_config() @@ -178,7 +193,6 @@ class TestCompletionAppGenerator: def test_generate_more_like_this_message_not_found(self, generator, mocker: MockerFixture): session = mocker.MagicMock() session.scalar.return_value = None - mocker.patch.object(module.db, "session", session) with pytest.raises(MessageNotExistsError): generator.generate_more_like_this( @@ -191,12 +205,12 @@ class TestCompletionAppGenerator: def test_generate_more_like_this_disabled(self, generator, mocker: MockerFixture): app_model = _build_app_model() - app_model.app_model_config = MagicMock(more_like_this=False, more_like_this_dict={"enabled": False}) + current_config = MagicMock(more_like_this=False, more_like_this_dict={"enabled": False}) message = MagicMock() session = mocker.MagicMock() session.scalar.return_value = message - mocker.patch.object(module.db, "session", session) + session.get.return_value = current_config with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( @@ -209,12 +223,11 @@ class TestCompletionAppGenerator: def test_generate_more_like_this_app_model_config_missing(self, generator, mocker: MockerFixture): app_model = _build_app_model() - app_model.app_model_config = None + app_model.app_model_config_id = None message = MagicMock() session = mocker.MagicMock() session.scalar.return_value = message - mocker.patch.object(module.db, "session", session) with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( @@ -227,12 +240,13 @@ class TestCompletionAppGenerator: def test_generate_more_like_this_message_config_none(self, generator, mocker: MockerFixture): app_model = _build_app_model() - app_model.app_model_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) + current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) - message = MagicMock(app_model_config=None) + message = MagicMock(conversation_id="conv-1") + conversation = MagicMock(app_model_config_id=None) session = mocker.MagicMock() session.scalar.return_value = message - mocker.patch.object(module.db, "session", session) + session.get.side_effect = [current_config, conversation] with pytest.raises(ValueError): generator.generate_more_like_this( @@ -245,27 +259,49 @@ class TestCompletionAppGenerator: def test_generate_more_like_this_success(self, generator, mocker: MockerFixture): app_model = _build_app_model() - app_model.app_model_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) + current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) - message = MagicMock() - message.message_files = [{"id": "f"}] - message.inputs = {"a": 1} - message.query = "q" + message = module.Message(id="msg", app_id="app1", conversation_id="conv-1", query="q") + message.inputs = {"attachment": {"dify_model_identity": FILE_MODEL_IDENTITY}} + message_files = [{"id": "f"}] + message_files_with_session = mocker.patch.object( + module.Message, + "message_files_with_session", + return_value=message_files, + ) - app_model_config = MagicMock() + app_model_config = MagicMock(app_id="app1") app_model_config.to_dict.return_value = { "model": {"completion_params": {"temperature": 0.1}}, "file_upload": {"enabled": True}, } - message.app_model_config = app_model_config + annotation_reply = {"enabled": False} + load_annotation_reply_config = mocker.patch.object( + module, + "load_annotation_reply_config", + return_value=annotation_reply, + ) + conversation = MagicMock(app_model_config_id="cfg-message") session = mocker.MagicMock() - session.scalar.return_value = message - mocker.patch.object(module.db, "session", session) + session.scalar.side_effect = [message, "tenant"] + session.get.side_effect = [current_config, conversation, app_model_config] + + global_session = MagicMock() + global_session.scalar.side_effect = AssertionError("global session must not be used") + global_session.scalars.side_effect = AssertionError("global session must not be used") + mocker.patch.object(module.db, "session", global_session) + + def restore_input_file(*, file_mapping, tenant_resolver): + assert file_mapping["dify_model_identity"] == FILE_MODEL_IDENTITY + assert tenant_resolver() == "tenant" + return "input-file" + + mocker.patch("models.model.build_file_from_input_mapping", side_effect=restore_input_file) file_extra_config = MagicMock() mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=file_extra_config) - mocker.patch.object(module.file_factory, "build_from_mappings", return_value=["file1"]) + build_from_mappings = mocker.patch.object(module.file_factory, "build_from_mappings", return_value=["file1"]) app_config = MagicMock(variables=["v"], to_dict=MagicMock(return_value={})) get_app_config = mocker.patch.object( @@ -293,6 +329,23 @@ class TestCompletionAppGenerator: ) assert result == "converted" + assert session.get.call_args_list == [ + call(module.AppModelConfig, "cfg-current"), + call(module.Conversation, "conv-1"), + call(module.AppModelConfig, "cfg-message"), + ] + load_annotation_reply_config.assert_called_once_with(session, "app1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert session.scalar.call_count == 2 + message_files_with_session.assert_called_once_with(session=session) + build_from_mappings.assert_called_once_with( + mappings=message_files, + tenant_id="tenant", + config=file_extra_config, + access_controller=generator._file_access_controller, + ) + assert global_session.mock_calls == [] + assert generator.generate_entity.call_args.kwargs["inputs"] == {"attachment": "input-file"} override_dict = get_app_config.call_args.kwargs["override_config_dict"] assert override_dict["model"]["completion_params"]["temperature"] == 0.9 @@ -317,7 +370,11 @@ class TestCompletionAppGenerator: flask_app.app_context.return_value = contextlib.nullcontext() session = mocker.MagicMock() - mocker.patch.object(module.db, "session", session) + session_context = mocker.MagicMock() + session_context.__enter__.return_value = session + create_session = mocker.patch.object(module.session_factory, "create_session") + create_session.return_value = session_context + mocker.patch.object(module.db, "session") mocker.patch.object(generator, "_get_message", return_value=MagicMock()) @@ -328,7 +385,6 @@ class TestCompletionAppGenerator: queue_manager = MagicMock() generator._generate_worker( flask_app=flask_app, - session=session, application_generate_entity=MagicMock(), queue_manager=queue_manager, message_id="msg", diff --git a/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py b/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py index 67cea557118..f2b8179160b 100644 --- a/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py +++ b/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py @@ -81,6 +81,9 @@ def _dummy_preserve(*args, **kwargs): class DummySession: def __init__(self): self.scalar = MagicMock() + self.add = MagicMock() + self.flush = MagicMock() + self.commit = MagicMock() def __enter__(self): return self @@ -98,6 +101,7 @@ def test_generate_dataset_missing(generator, mocker: MockerFixture): with pytest.raises(ValueError): generator.generate( + session=session, pipeline=pipeline, workflow=_build_workflow(), user=_build_user(), @@ -140,6 +144,7 @@ def test_generate_debugger_calls_generate(generator, mocker: MockerFixture): mocker.patch.object(generator, "_generate", return_value={"result": "ok"}) result = generator.generate( + session=session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -201,9 +206,6 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock mocker.patch.object(module, "DocumentPipelineExecutionLog", return_value=MagicMock()) - db_session = MagicMock() - mocker.patch.object(module.db, "session", db_session) - mocker.patch.object( module.DifyCoreRepositoryFactory, "create_workflow_execution_repository", @@ -219,6 +221,7 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock mocker.patch.object(module, "RagPipelineTaskProxy", return_value=task_proxy) result = generator.generate( + session=session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -230,6 +233,8 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock assert result["batch"] assert len(result["documents"]) == 2 check_limits.assert_called_once_with(len(datasource_info_list), features) + session.flush.assert_called_once_with() + session.commit.assert_called_once_with() task_proxy.delay.assert_called_once() @@ -259,11 +264,9 @@ def test_generate_published_pipeline_rejects_when_document_creation_limits_excee side_effect=ValueError("document limit exceeded"), ) - db_session = MagicMock() - mocker.patch.object(module.db, "session", db_session) - with pytest.raises(ValueError, match="document limit exceeded"): generator.generate( + session=session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -273,7 +276,7 @@ def test_generate_published_pipeline_rejects_when_document_creation_limits_excee ) check_limits.assert_called_once_with(len(datasource_info_list), features) - db_session.add.assert_not_called() + session.add.assert_not_called() def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture): @@ -309,6 +312,7 @@ def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture): mocker.patch.object(generator, "_generate", return_value={"result": "ok"}) result = generator.generate( + session=session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -399,6 +403,7 @@ def test_generate_raises_when_workflow_not_found(generator, mocker: MockerFixtur with pytest.raises(ValueError): generator._generate( + session=session, flask_app=flask_app, context=contextlib.nullcontext(), pipeline=_build_pipeline(), @@ -437,6 +442,7 @@ def test_generate_success_returns_converted(generator, mocker: MockerFixture): mocker.patch.object(module.WorkflowAppGenerateResponseConverter, "convert", return_value="converted") result = generator._generate( + session=session, flask_app=flask_app, context=contextlib.nullcontext(), pipeline=_build_pipeline(), @@ -459,11 +465,18 @@ def test_generate_success_returns_converted(generator, mocker: MockerFixture): def test_single_iteration_generate_validates_inputs(generator, mocker: MockerFixture): with pytest.raises(ValueError): - generator.single_iteration_generate(_build_pipeline(), _build_workflow(), "", _build_user(), {}) + generator.single_iteration_generate( + _build_pipeline(), _build_workflow(), "", _build_user(), {}, session=DummySession() + ) with pytest.raises(ValueError): generator.single_iteration_generate( - _build_pipeline(), _build_workflow(), "node", _build_user(), {"inputs": None} + _build_pipeline(), + _build_workflow(), + "node", + _build_user(), + {"inputs": None}, + session=DummySession(), ) @@ -481,6 +494,7 @@ def test_single_iteration_generate_dataset_required(generator, mocker: MockerFix "node", _build_user(), {"inputs": {"a": 1}}, + session=session, ) @@ -519,6 +533,7 @@ def test_single_iteration_generate_success(generator, mocker: MockerFixture): _build_user(), {"inputs": {"a": 1}}, streaming=False, + session=session, ) assert result == {"ok": True} @@ -559,6 +574,7 @@ def test_single_loop_generate_success(generator, mocker: MockerFixture): _build_user(), {"inputs": {"a": 1}}, streaming=False, + session=session, ) assert result == {"ok": True} 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 ef12f0be965..d9e66cd26a6 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 @@ -47,7 +47,7 @@ def _make_generate_entity(app_config: WorkflowUIBasedAppConfig) -> AdvancedChatA @pytest.fixture(autouse=True) -def _mock_db_session(monkeypatch: pytest.MonkeyPatch): +def mock_db_session(monkeypatch: pytest.MonkeyPatch): session = MagicMock() def refresh_side_effect(obj): @@ -64,20 +64,24 @@ def _mock_db_session(monkeypatch: pytest.MonkeyPatch): return session -def test_init_generate_records_sets_conversation_metadata(): +def test_init_generate_records_sets_conversation_metadata(mock_db_session): app_config = _make_app_config() entity = _make_generate_entity(app_config) generator = AdvancedChatAppGenerator() - conversation, _ = generator._init_generate_records(entity, conversation=None) + conversation, _ = generator._init_generate_records( + entity, + conversation=None, + session=mock_db_session, + ) assert entity.conversation_id == "generated-conversation-id" assert conversation.id == "generated-conversation-id" assert entity.is_new_conversation is True -def test_init_generate_records_marks_existing_conversation(): +def test_init_generate_records_marks_existing_conversation(mock_db_session): app_config = _make_app_config() entity = _make_generate_entity(app_config) @@ -103,7 +107,11 @@ def test_init_generate_records_marks_existing_conversation(): generator = AdvancedChatAppGenerator() - conversation, _ = generator._init_generate_records(entity, conversation=existing_conversation) + conversation, _ = generator._init_generate_records( + entity, + conversation=existing_conversation, + session=mock_db_session, + ) assert entity.conversation_id == "existing-conversation-id" assert conversation is existing_conversation @@ -155,6 +163,7 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke ) captured: dict[str, object] = {} + session = MagicMock() def fake_generate(self, **kwargs): captured.update(kwargs) @@ -170,10 +179,12 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke invoke_from=InvokeFrom.SERVICE_API, workflow_run_id="workflow-run-id", streaming=False, + session=session, ) assert result == {"status": "ok"} assert captured["conversation"] is None + assert captured["session"] is session application_generate_entity = captured["application_generate_entity"] assert isinstance(application_generate_entity, AdvancedChatAppGenerateEntity) assert application_generate_entity.conversation_id is None diff --git a/api/tests/unit_tests/core/app/apps/test_base_app_runner.py b/api/tests/unit_tests/core/app/apps/test_base_app_runner.py index 1b22412ba6b..deb9ab4d2af 100644 --- a/api/tests/unit_tests/core/app/apps/test_base_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/test_base_app_runner.py @@ -1,6 +1,7 @@ from __future__ import annotations import logging +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import MagicMock @@ -15,7 +16,12 @@ from core.app.app_config.entities import ( from core.app.apps.base_app_runner import AppRunner from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom -from core.app.entities.queue_entities import QueueAgentMessageEvent, QueueLLMChunkEvent, QueueMessageEndEvent +from core.app.entities.queue_entities import ( + QueueAgentMessageEvent, + QueueLLMChunkEvent, + QueueMessageEndEvent, + QueueMessageFileEvent, +) from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.message_entities import ( AssistantPromptMessage, @@ -352,6 +358,55 @@ class TestAppRunner: assert queue.events[-1].llm_result.usage == usage assert "Failed to handle multimodal image output" in caplog.messages + def test_handle_invoke_result_stream_commits_message_file_before_publish(self, monkeypatch: pytest.MonkeyPatch): + runner = AppRunner() + runner._handle_multimodal_image_content = MagicMock(return_value="message-file-1") + session = MagicMock() + events: list[str] = [] + session.commit.side_effect = lambda: events.append("commit") + monkeypatch.setattr( + "core.app.apps.base_app_runner.session_factory.create_session", + lambda: nullcontext(session), + ) + queue = _QueueRecorder() + original_publish = queue.publish + + def publish(event, pub_from): + if isinstance(event, QueueMessageFileEvent): + events.append("publish") + original_publish(event, pub_from) + + queue.publish = publish + + def stream(): + yield LLMResultChunk( + model="model", + prompt_messages=[AssistantPromptMessage(content="prompt")], + delta=LLMResultChunkDelta( + index=0, + message=AssistantPromptMessage( + content=[ + ImagePromptMessageContent( + url="https://example.com/image.png", + format="png", + mime_type="image/png", + ) + ] + ), + ), + ) + + runner._handle_invoke_result_stream( + invoke_result=stream(), + queue_manager=queue, + agent=False, + message_id="message-1", + user_id="user-1", + tenant_id="tenant-1", + ) + + assert events == ["commit", "publish"] + def test_handle_invoke_result_stream_closes_generator_when_stopped(self): runner = AppRunner() chunk = LLMResultChunk( @@ -393,13 +448,13 @@ class TestAppRunner: mime_type="image/png", ) - db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock(), refresh=MagicMock()) + db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), refresh=MagicMock()) monkeypatch.setattr("core.app.apps.base_app_runner.ToolFileManager", lambda: MagicMock()) - monkeypatch.setattr("core.app.apps.base_app_runner.db", SimpleNamespace(session=db_session)) queue_manager = SimpleNamespace(invoke_from=InvokeFrom.SERVICE_API, publish=MagicMock()) runner._handle_multimodal_image_content( + session=db_session, content=content, message_id="message-id", user_id="user-id", @@ -471,7 +526,7 @@ class TestAppRunner: runner = AppRunner() monkeypatch.setattr( "core.app.apps.base_app_runner.AnnotationReplyFeature.query", - lambda self, app_record, message, query, user_id, invoke_from: "reply", + lambda self, app_record, message, query, user_id, invoke_from, session: "reply", ) response = runner.query_app_annotations_to_reply( @@ -480,6 +535,7 @@ class TestAppRunner: query="hello", user_id="user", invoke_from=InvokeFrom.WEB_APP, + session=MagicMock(), ) assert response == "reply" diff --git a/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py b/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py index 6a9b5e76194..368be54dbf0 100644 --- a/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py @@ -85,7 +85,7 @@ def _make_chat_generate_entity(app_config: EasyUIBasedAppConfig) -> ChatAppGener @pytest.fixture(autouse=True) -def _mock_db_session(monkeypatch: pytest.MonkeyPatch): +def mock_db_session(monkeypatch: pytest.MonkeyPatch): session = MagicMock() def refresh_side_effect(obj): @@ -102,13 +102,17 @@ def _mock_db_session(monkeypatch: pytest.MonkeyPatch): return session -def test_init_generate_records_skips_conversation_fields_for_non_conversation_entity(): +def test_init_generate_records_skips_conversation_fields_for_non_conversation_entity(mock_db_session): app_config = _make_app_config(AppMode.COMPLETION) entity = DummyCompletionGenerateEntity(app_config=app_config) generator = MessageBasedAppGenerator() - conversation, message = generator._init_generate_records(entity, conversation=None) + conversation, message = generator._init_generate_records( + entity, + conversation=None, + session=mock_db_session, + ) assert conversation.id == "generated-conversation-id" assert message.id == "generated-message-id" @@ -116,13 +120,17 @@ def test_init_generate_records_skips_conversation_fields_for_non_conversation_en assert hasattr(entity, "is_new_conversation") is False -def test_init_generate_records_sets_conversation_fields_for_chat_entity(): +def test_init_generate_records_sets_conversation_fields_for_chat_entity(mock_db_session): app_config = _make_app_config(AppMode.CHAT) entity = _make_chat_generate_entity(app_config) generator = MessageBasedAppGenerator() - conversation, _ = generator._init_generate_records(entity, conversation=None) + conversation, _ = generator._init_generate_records( + entity, + conversation=None, + session=mock_db_session, + ) assert entity.conversation_id == "generated-conversation-id" assert entity.is_new_conversation is True @@ -155,20 +163,23 @@ class TestMessageBasedAppGeneratorExtras: stream=False, ) - def test_get_app_model_config_requires_valid_config(self, monkeypatch: pytest.MonkeyPatch): + def test_get_app_model_config_requires_valid_config(self): generator = MessageBasedAppGenerator() app_model = SimpleNamespace(id="app", app_model_config_id=None, app_model_config=None) + session = MagicMock() with pytest.raises(AppModelConfigBrokenError): - generator._get_app_model_config(app_model, conversation=None) + generator._get_app_model_config(app_model, conversation=None, session=session) conversation = SimpleNamespace(app_model_config_id="missing-id") - monkeypatch.setattr( - message_based_app_generator, "db", SimpleNamespace(session=SimpleNamespace(scalar=lambda _: None)) - ) + session.scalar.return_value = None with pytest.raises(AppModelConfigBrokenError): - generator._get_app_model_config(app_model=SimpleNamespace(id="app"), conversation=conversation) + generator._get_app_model_config( + app_model=SimpleNamespace(id="app"), + conversation=conversation, + session=session, + ) def test_get_conversation_introduction_handles_missing_inputs(self): app_config = _make_app_config(AppMode.CHAT) diff --git a/api/tests/unit_tests/core/app/apps/test_pause_resume.py b/api/tests/unit_tests/core/app/apps/test_pause_resume.py index 835a8c22dd0..6bfda02eb4c 100644 --- a/api/tests/unit_tests/core/app/apps/test_pause_resume.py +++ b/api/tests/unit_tests/core/app/apps/test_pause_resume.py @@ -274,6 +274,7 @@ def test_advanced_chat_pause_resume_matches_baseline(mocker: MockerFixture): user=SimpleNamespace(), conversation=SimpleNamespace(id="conv"), message=SimpleNamespace(id="msg"), + session=SimpleNamespace(), application_generate_entity=SimpleNamespace( stream=False, invoke_from=InvokeFrom.SERVICE_API, diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py b/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py index 228ca2024e2..3509c349aef 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py @@ -66,6 +66,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args={"inputs": {}}, streaming=False, + session=Mock(), ) with pytest.raises(ValueError, match="inputs is required"): @@ -76,6 +77,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args={}, streaming=False, + session=Mock(), ) def test_single_loop_generate_validates_args(self): @@ -89,6 +91,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs={}), streaming=False, + session=Mock(), ) def test_single_iteration_generate_includes_trace_session_id_in_extras(self, monkeypatch: pytest.MonkeyPatch): @@ -134,6 +137,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(id="user-id"), args={"inputs": {"foo": "bar"}, "trace_session_id": "session-1"}, streaming=False, + session=Mock(), ) assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" @@ -181,6 +185,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(id="user-id"), args=SimpleNamespace(inputs={"foo": "bar"}, trace_session_id="session-1"), streaming=False, + session=Mock(), ) assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" @@ -193,6 +198,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs=None), streaming=False, + session=Mock(), ) diff --git a/api/tests/unit_tests/core/app/features/test_annotation_reply.py b/api/tests/unit_tests/core/app/features/test_annotation_reply.py index 2c9204e64fa..242364b301e 100644 --- a/api/tests/unit_tests/core/app/features/test_annotation_reply.py +++ b/api/tests/unit_tests/core/app/features/test_annotation_reply.py @@ -97,7 +97,9 @@ class TestAnnotationReplyFeature: vector_instance = Mock() vector_instance.search_by_vector.return_value = [document] - with patch("core.app.features.annotation_reply.annotation_reply.Vector", return_value=vector_instance): + with patch( + "core.app.features.annotation_reply.annotation_reply.Vector", return_value=vector_instance + ) as vector_cls: result = AnnotationReplyFeature().query( app_record=SimpleNamespace(id="app-1", tenant_id="tenant-1"), message=SimpleNamespace(id="msg-1"), @@ -111,6 +113,7 @@ class TestAnnotationReplyFeature: vector_instance.search_by_vector.assert_called_once_with( query="hi", top_k=1, score_threshold=1, filter={"group_id": ["app-1"]} ) + assert vector_cls.call_args.kwargs["session"] is sqlite_session sqlite_session.refresh(annotation) assert annotation.hit_count == 1 history = sqlite_session.scalar(select(AppAnnotationHitHistory)) diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py index 1c1bf391d3e..53a70553e8d 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py @@ -1,5 +1,5 @@ from types import SimpleNamespace -from unittest.mock import ANY, Mock, patch +from unittest.mock import Mock, patch import pytest @@ -31,6 +31,17 @@ from graphon.model_runtime.entities.message_entities import TextPromptMessageCon from models.model import AppMode +def _patch_stream_session(): + session = Mock() + session_cm = Mock() + session_cm.__enter__ = Mock(return_value=session) + session_cm.__exit__ = Mock(return_value=False) + return session, patch( + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + return_value=session_cm, + ) + + class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: """Test cases for EasyUIBasedGenerateTaskPipeline._process_stream_response method.""" @@ -251,17 +262,16 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline._save_message = Mock() pipeline._message_end_to_stream_response = Mock(return_value=Mock(spec=MessageEndStreamResponse)) - # Patch db.engine used inside pipeline for session creation - with patch( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", new=SimpleNamespace(engine=Mock()) - ): + session, patch_session = _patch_stream_session() + with patch_session: # Execute responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) # Assert assert len(responses) == 1 assert mock_task_state.llm_result == llm_result - pipeline._save_message.assert_called_once() + pipeline._save_message.assert_called_once_with(session=session, trace_manager=None) + session.commit.assert_called_once() pipeline._message_end_to_stream_response.assert_called_once() def test_error_event(self, pipeline): @@ -277,16 +287,15 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline.handle_error = Mock(return_value=Exception("Test error")) pipeline.error_to_stream_response = Mock(return_value=Mock(spec=ErrorStreamResponse)) - # Patch db.engine used inside pipeline for session creation - with patch( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", new=SimpleNamespace(engine=Mock()) - ): + session, patch_session = _patch_stream_session() + with patch_session: # Execute responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) # Assert assert len(responses) == 1 - pipeline.handle_error.assert_called_once() + pipeline.handle_error.assert_called_once_with(event=error_event, session=session, message_id="test-message-id") + session.commit.assert_called_once() pipeline.error_to_stream_response.assert_called_once() def test_ping_event(self, pipeline): @@ -364,15 +373,14 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline._save_message = Mock() pipeline._message_end_to_stream_response = Mock(return_value=Mock(spec=MessageEndStreamResponse)) - # Patch db.engine used inside pipeline for session creation - with patch( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", new=SimpleNamespace(engine=Mock()) - ): + session, patch_session = _patch_stream_session() + with patch_session: # Execute list(pipeline._process_stream_response(publisher=None, trace_manager=trace_manager)) # Assert - pipeline._save_message.assert_called_once_with(session=ANY, trace_manager=trace_manager) + pipeline._save_message.assert_called_once_with(session=session, trace_manager=trace_manager) + session.commit.assert_called_once() def test_multiple_events_sequence(self, pipeline, mock_message_cycle_manager, mock_task_state): """Test handling multiple events in sequence.""" diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py index 153157337e1..49ecc9358c8 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py @@ -65,10 +65,6 @@ class _DummyModelConf: self.model = "mock" -class _FakeDb: - engine: object = object() - - class _UnknownQueueEvent: pass @@ -395,12 +391,8 @@ class TestEasyUiBasedGenerateTaskPipeline: return None monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", - _Session, - ) - monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) responses = list(pipeline._process_stream_response(publisher=None)) @@ -574,12 +566,8 @@ class TestEasyUiBasedGenerateTaskPipeline: return _Result(message_files if self.calls == 1 else upload_files) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", - _Session, - ) - monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) monkeypatch.setattr( "core.app.task_pipeline.message_file_utils.file_helpers.get_signed_file_url", @@ -641,7 +629,7 @@ class TestEasyUiBasedGenerateTaskPipeline: _set_method( pipeline._message_cycle_manager, "handle_annotation_reply", - lambda event: _AnnotationReply(content="annotated"), + lambda event, session: _AnnotationReply(content="annotated"), ) _set_method(pipeline, "_agent_thought_to_stream_response", _agent_thought_response) _set_method(pipeline._message_cycle_manager, "message_file_to_stream_response", _file_response) @@ -663,12 +651,8 @@ class TestEasyUiBasedGenerateTaskPipeline: return None monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", - _Session, - ) - monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) responses = list(pipeline._process_stream_response(publisher=None)) @@ -709,12 +693,8 @@ class TestEasyUiBasedGenerateTaskPipeline: return agent_thought monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", - _Session, - ) - monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) response = pipeline._agent_thought_to_stream_response(QueueAgentThoughtEvent(agent_thought_id="thought")) @@ -757,14 +737,9 @@ class TestEasyUiBasedGenerateTaskPipeline: return agent_thought monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", - _Session, + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) - monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), - ) - response = pipeline._agent_thought_to_stream_response(QueueAgentThoughtEvent(agent_thought_id="thought")) assert response is not None @@ -1063,10 +1038,9 @@ class TestEasyUiBasedGenerateTaskPipeline: def commit(self): return None - monkeypatch.setattr("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", _Session) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) responses = list(pipeline._process_stream_response(publisher=None)) @@ -1321,10 +1295,9 @@ class TestEasyUiBasedGenerateTaskPipeline: def scalars(self, *args, **kwargs): return _Result() - monkeypatch.setattr("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", _Session) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) response = pipeline._message_end_to_stream_response() @@ -1361,10 +1334,9 @@ class TestEasyUiBasedGenerateTaskPipeline: def scalars(self, *args, **kwargs): return _Result() - monkeypatch.setattr("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", _Session) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) response = pipeline._message_end_to_stream_response() @@ -1426,10 +1398,9 @@ class TestEasyUiBasedGenerateTaskPipeline: self.calls += 1 return _Result(message_files if self.calls == 1 else []) - monkeypatch.setattr("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", _Session) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) monkeypatch.setattr( "core.app.task_pipeline.message_file_utils.file_helpers.get_signed_file_url", @@ -1488,10 +1459,9 @@ class TestEasyUiBasedGenerateTaskPipeline: def scalar(self, *args, **kwargs): return None - monkeypatch.setattr("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session", _Session) monkeypatch.setattr( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db", - _FakeDb(), + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + lambda: _Session(), ) response = pipeline._agent_thought_to_stream_response(QueueAgentThoughtEvent(agent_thought_id="missing")) diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py index b1c06e237a8..5d5807e1215 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py @@ -25,6 +25,16 @@ from graphon.file import FileTransferMethod, FileType from models.model import MessageFile, UploadFile +def _patch_create_session(mock_session): + session_cm = MagicMock() + session_cm.__enter__.return_value = mock_session + session_cm.__exit__.return_value = False + return patch( + "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", + return_value=session_cm, + ) + + class TestMessageEndStreamResponseFiles: """Test suite for files array population in message_end SSE event.""" @@ -92,15 +102,8 @@ class TestMessageEndStreamResponseFiles: def test_message_end_with_no_files(self, mock_pipeline): """Test that files array is empty when no MessageFile records exist.""" # Arrange - with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, - ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session + mock_session = MagicMock(spec=Session) + with _patch_create_session(mock_session): mock_session.scalars.return_value.all.return_value = [] # Act @@ -118,17 +121,11 @@ class TestMessageEndStreamResponseFiles: # Arrange mock_message_file_local.message_id = mock_pipeline._message_id + mock_session = MagicMock(spec=Session) with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, + _patch_create_session(mock_session), patch("core.app.task_pipeline.message_file_utils.file_helpers.get_signed_file_url") as mock_get_url, ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - # Mock database queries # First query: MessageFile mock_message_files_result = Mock() @@ -182,15 +179,8 @@ class TestMessageEndStreamResponseFiles: # Arrange mock_message_file_remote.message_id = mock_pipeline._message_id - with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, - ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - + mock_session = MagicMock(spec=Session) + with _patch_create_session(mock_session): # Mock database queries mock_scalars_result = Mock() mock_scalars_result.all.return_value = [mock_message_file_remote] @@ -223,15 +213,8 @@ class TestMessageEndStreamResponseFiles: mock_message_file_tool.message_id = mock_pipeline._message_id mock_message_file_tool.url = "https://example.com/tool_file.png" - with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, - ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - + mock_session = MagicMock(spec=Session) + with _patch_create_session(mock_session): # Mock database queries mock_scalars_result = Mock() mock_scalars_result.all.return_value = [mock_message_file_tool] @@ -257,17 +240,11 @@ class TestMessageEndStreamResponseFiles: mock_message_file_tool.message_id = mock_pipeline._message_id mock_message_file_tool.url = "tool_file_123.png" + mock_session = MagicMock(spec=Session) with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, + _patch_create_session(mock_session), patch("core.app.task_pipeline.message_file_utils.sign_tool_file") as mock_sign_tool, ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - # Mock database queries mock_scalars_result = Mock() mock_scalars_result.all.return_value = [mock_message_file_tool] @@ -297,15 +274,11 @@ class TestMessageEndStreamResponseFiles: mock_message_file_tool.message_id = mock_pipeline._message_id mock_message_file_tool.url = "tool_file_abc.verylongextension" + mock_session = MagicMock(spec=Session) with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, + _patch_create_session(mock_session), patch("core.app.task_pipeline.message_file_utils.sign_tool_file") as mock_sign_tool, ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session mock_scalars_result = Mock() mock_scalars_result.all.return_value = [mock_message_file_tool] mock_session.scalars.return_value = mock_scalars_result @@ -326,17 +299,11 @@ class TestMessageEndStreamResponseFiles: mock_message_file_local.message_id = mock_pipeline._message_id mock_message_file_remote.message_id = mock_pipeline._message_id + mock_session = MagicMock(spec=Session) with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, + _patch_create_session(mock_session), patch("core.app.task_pipeline.message_file_utils.file_helpers.get_signed_file_url") as mock_get_url, ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - # Mock database queries # First query: MessageFile mock_message_files_result = Mock() @@ -378,17 +345,11 @@ class TestMessageEndStreamResponseFiles: # Arrange mock_message_file_local.message_id = mock_pipeline._message_id + mock_session = MagicMock(spec=Session) with ( - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.db") as mock_db, - patch("core.app.task_pipeline.easy_ui_based_generate_task_pipeline.Session") as mock_session_class, + _patch_create_session(mock_session), patch("core.app.task_pipeline.message_file_utils.file_helpers.get_signed_file_url") as mock_get_url, ): - mock_engine = MagicMock() - mock_db.engine = mock_engine - - mock_session = MagicMock(spec=Session) - mock_session_class.return_value.__enter__.return_value = mock_session - # Mock database queries # First query: MessageFile mock_message_files_result = Mock() diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py b/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py index 0863d5a8f73..c0cc2d75386 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py @@ -11,7 +11,14 @@ from core.app.entities.queue_entities import QueueAnnotationReplyEvent, QueueRet from core.app.entities.task_entities import MessageStreamResponse, StreamEvent, TaskStateMetadata from core.app.task_pipeline.message_cycle_manager import MessageCycleManager from core.rag.entities import RetrievalSourceMetadata -from models.model import AppMode +from models.model import App, AppMode + + +def _patch_create_session(mock_session): + session_cm = Mock() + session_cm.__enter__ = Mock(return_value=mock_session) + session_cm.__exit__ = Mock(return_value=False) + return patch("core.app.task_pipeline.message_cycle_manager.session_factory.create_session", return_value=session_cm) class TestMessageCycleManagerOptimization: @@ -268,12 +275,10 @@ class TestMessageCycleManagerOptimization: db_session = Mock() db_session.scalar.return_value = None - with patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db: - mock_db.session = db_session + with _patch_create_session(db_session): message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-missing", "hello") db_session.commit.assert_not_called() - db_session.close.assert_not_called() def test_generate_conversation_name_worker_returns_when_app_missing(self, message_cycle_manager): """Return early when non-completion conversation has no app relation.""" @@ -281,39 +286,45 @@ class TestMessageCycleManagerOptimization: conversation = SimpleNamespace(mode=AppMode.CHAT, app=None, app_id="app-id") db_session = Mock() db_session.scalar.return_value = conversation + db_session.get.return_value = None - with patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db: - mock_db.session = db_session + with _patch_create_session(db_session): message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") db_session.commit.assert_not_called() - db_session.close.assert_not_called() def test_generate_conversation_name_worker_uses_cached_name(self, message_cycle_manager): """Use cached conversation name when present and avoid LLM call.""" flask_app = Flask(__name__) - conversation = SimpleNamespace( - mode=AppMode.CHAT, - app=SimpleNamespace(tenant_id="tenant-1"), - app_id="app-id", - name="", - ) + + class ConversationWithPoisonedApp: + mode = AppMode.CHAT + app_id = "app-id" + name = "" + + @property + def app(self): + raise AssertionError("conversation.app must not open an implicit session") + + conversation = ConversationWithPoisonedApp() + app_model = SimpleNamespace(tenant_id="tenant-1") db_session = Mock() db_session.scalar.return_value = conversation + db_session.get.return_value = app_model with ( - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, + _patch_create_session(db_session) as create_session, patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, ): - mock_db.session = db_session mock_redis.get.return_value = b"cached-title" message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") assert conversation.name == "cached-title" + create_session.assert_called_once_with() + db_session.get.assert_called_once_with(App, "app-id") db_session.commit.assert_called_once() - db_session.close.assert_called_once() mock_llm_generator.generate_conversation_name.assert_not_called() mock_redis.setex.assert_not_called() @@ -330,11 +341,10 @@ class TestMessageCycleManagerOptimization: db_session.scalar.return_value = conversation with ( - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, + _patch_create_session(db_session), patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, ): - mock_db.session = db_session mock_redis.get.return_value = None mock_llm_generator.generate_conversation_name.return_value = "generated-title" @@ -342,7 +352,6 @@ class TestMessageCycleManagerOptimization: assert conversation.name == "generated-title" db_session.commit.assert_called_once() - db_session.close.assert_called_once() mock_redis.setex.assert_called_once() def test_generate_conversation_name_worker_falls_back_when_generation_fails( @@ -361,12 +370,11 @@ class TestMessageCycleManagerOptimization: long_query = "q" * 60 with ( - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, + _patch_create_session(db_session), patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, patch("core.app.task_pipeline.message_cycle_manager.dify_config") as mock_dify_config, ): - mock_db.session = db_session mock_redis.get.return_value = None mock_llm_generator.generate_conversation_name.side_effect = RuntimeError("generation failed") mock_dify_config.DEBUG = True @@ -376,7 +384,6 @@ class TestMessageCycleManagerOptimization: assert conversation.name == (long_query[:47] + "...") db_session.commit.assert_called_once() - db_session.close.assert_called_once() assert any(record.levelno == logging.ERROR for record in caplog.records) def test_handle_annotation_reply_sets_metadata(self, message_cycle_manager): @@ -391,19 +398,24 @@ class TestMessageCycleManagerOptimization: annotation = SimpleNamespace( id="ann-1", account_id="acct-1", - account=SimpleNamespace(name="Alice"), ) + session = Mock() - with patch("core.app.task_pipeline.message_cycle_manager.AppAnnotationService") as mock_service: + with ( + patch("core.app.task_pipeline.message_cycle_manager.AppAnnotationService") as mock_service, + patch("core.app.task_pipeline.message_cycle_manager.AccountService") as account_service, + ): mock_service.get_annotation_by_id.return_value = annotation + account_service.get_account_by_id.return_value = SimpleNamespace(name="Alice") result = message_cycle_manager.handle_annotation_reply( - QueueAnnotationReplyEvent(message_annotation_id="ann-1") + QueueAnnotationReplyEvent(message_annotation_id="ann-1"), session ) assert result == annotation assert message_cycle_manager._task_state.metadata.annotation_reply.id == "ann-1" assert message_cycle_manager._task_state.metadata.annotation_reply.account.name == "Alice" + account_service.get_account_by_id.assert_called_once_with("acct-1", session=session) def test_handle_annotation_reply_returns_none_when_missing(self, message_cycle_manager): """Return None and keep metadata unchanged when annotation is not found.""" @@ -413,7 +425,7 @@ class TestMessageCycleManagerOptimization: mock_service.get_annotation_by_id.return_value = None result = message_cycle_manager.handle_annotation_reply( - QueueAnnotationReplyEvent(message_annotation_id="missing") + QueueAnnotationReplyEvent(message_annotation_id="missing"), Mock() ) assert result is None diff --git a/api/tests/unit_tests/core/ops/test_base_trace_instance.py b/api/tests/unit_tests/core/ops/test_base_trace_instance.py index 15a2af17ca4..76ee0336990 100644 --- a/api/tests/unit_tests/core/ops/test_base_trace_instance.py +++ b/api/tests/unit_tests/core/ops/test_base_trace_instance.py @@ -92,7 +92,6 @@ def test_get_service_account_with_tenant_success(mock_db_session): mock_account = MagicMock(spec=Account) mock_account.id = "creator_id" - mock_account.set_tenant_id = MagicMock() mock_tenant_join = MagicMock(spec=TenantAccountJoin) mock_tenant_join.tenant_id = "tenant_id" @@ -105,4 +104,4 @@ def test_get_service_account_with_tenant_success(mock_db_session): result = instance.get_service_account_with_tenant("some_app_id") assert result == mock_account - mock_account.set_tenant_id.assert_called_once_with("tenant_id") + mock_account.set_tenant_id_with_session.assert_called_once_with("tenant_id", session=mock_db_session) diff --git a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py index baedd873005..6fa94fbbca5 100644 --- a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py +++ b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py @@ -116,7 +116,6 @@ class TestPluginAppBackwardsInvocation: route = mocker.patch.object(PluginAppBackwardsInvocation, route_method, return_value={"routed": True}) result = PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -125,6 +124,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={"x": 1}, files=[], + session=MagicMock(), ) assert result == {"routed": True} @@ -143,7 +143,6 @@ class TestPluginAppBackwardsInvocation: route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True}) result = PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="", tenant_id="tenant", @@ -152,6 +151,7 @@ class TestPluginAppBackwardsInvocation: stream=True, inputs={}, files=[], + session=MagicMock(), ) assert result == {"ok": True} @@ -165,7 +165,6 @@ class TestPluginAppBackwardsInvocation: with pytest.raises(ValueError, match="missing query"): PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -174,6 +173,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], + session=MagicMock(), ) def test_invoke_app_unexpected_mode_raises(self, mocker: MockerFixture): @@ -182,7 +182,6 @@ class TestPluginAppBackwardsInvocation: with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -191,6 +190,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], + session=MagicMock(), ) @pytest.mark.parametrize( @@ -205,7 +205,6 @@ class TestPluginAppBackwardsInvocation: spy = mocker.patch(generator_path, return_value={"result": "ok"}) result = PluginAppBackwardsInvocation.invoke_chat_app( - MagicMock(), app=app, user=MagicMock(), conversation_id="conv-1", @@ -213,6 +212,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={"k": "v"}, files=[], + session=MagicMock(), ) assert result == {"result": "ok"} @@ -234,9 +234,9 @@ class TestPluginAppBackwardsInvocation: "core.plugin.backwards_invocation.app.AdvancedChatAppGenerator.generate", return_value={"result": "ok"}, ) + session = MagicMock() result = PluginAppBackwardsInvocation.invoke_chat_app( - MagicMock(), app=app, user=MagicMock(), conversation_id="conv-1", @@ -244,10 +244,12 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={"k": "v"}, files=[], + session=session, ) assert result == {"result": "ok"} call_kwargs = generator_spy.call_args.kwargs + assert call_kwargs["session"] is session pause_state_config = call_kwargs.get("pause_state_config") assert isinstance(pause_state_config, PauseStateLayerConfig) assert pause_state_config.state_owner_user_id == "owner-id" @@ -257,7 +259,6 @@ class TestPluginAppBackwardsInvocation: mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_chat_app( - MagicMock(), app=app, user=MagicMock(), conversation_id="conv-1", @@ -265,13 +266,13 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], + session=MagicMock(), ) def test_invoke_chat_app_unexpected_mode_raises(self): app = MagicMock(mode="invalid") with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_chat_app( - MagicMock(), app=app, user=MagicMock(), conversation_id="conv-1", @@ -279,6 +280,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], + session=MagicMock(), ) def test_invoke_workflow_app_injects_pause_state_config(self, mocker: MockerFixture): @@ -319,7 +321,6 @@ class TestPluginAppBackwardsInvocation: mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -328,6 +329,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], + session=MagicMock(), ) def test_invoke_completion_app(self, mocker: MockerFixture): @@ -336,7 +338,7 @@ class TestPluginAppBackwardsInvocation: ) app = MagicMock(mode=AppMode.COMPLETION) - result = PluginAppBackwardsInvocation.invoke_completion_app(MagicMock(), app, MagicMock(), False, {"x": 1}, []) + result = PluginAppBackwardsInvocation.invoke_completion_app(app, MagicMock(), False, {"x": 1}, [], MagicMock()) assert result == {"ok": 1} assert spy.call_count == 1 @@ -408,7 +410,6 @@ class TestPluginAppBackwardsInvocation: route = mocker.patch.object(PluginAppBackwardsInvocation, "invoke_workflow_app", return_value={"ok": True}) result = PluginAppBackwardsInvocation.invoke_app( - MagicMock(), app_id="app", user_id="wecom-sender-1", tenant_id="tenant", @@ -417,6 +418,7 @@ class TestPluginAppBackwardsInvocation: stream=True, inputs={}, files=[], + session=MagicMock(), ) assert result == {"ok": True} diff --git a/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py b/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py index 1f46634b892..3b38a9af40a 100644 --- a/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py +++ b/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py @@ -1,7 +1,6 @@ +from unittest.mock import MagicMock from uuid import uuid4 -from pytest_mock import MockerFixture - from constants import UUID_NIL from core.prompt.utils.extract_thread_messages import extract_thread_messages from core.prompt.utils.get_thread_messages_length import get_thread_messages_length @@ -105,32 +104,33 @@ def test_extract_thread_messages_breaks_when_parent_is_none(): assert result[0].id == id2 -def test_get_thread_messages_length_excludes_newly_created_empty_answer(mocker: MockerFixture): +def test_get_thread_messages_length_excludes_newly_created_empty_answer(): id1, id2 = str(uuid4()), str(uuid4()) messages = [ MockMessage(id2, id1, answer=""), # newest generated message should be excluded MockMessage(id1, UUID_NIL, answer="ok"), ] - mock_scalars = mocker.patch("core.prompt.utils.get_thread_messages_length.db.session.scalars") - mock_scalars.return_value.all.return_value = messages + session = MagicMock() + session.scalars.return_value.all.return_value = messages - length = get_thread_messages_length("conversation-1") + length = get_thread_messages_length("conversation-1", session=session) assert length == 1 - mock_scalars.assert_called_once() + session.scalars.assert_called_once() -def test_get_thread_messages_length_keeps_non_empty_latest_answer(mocker: MockerFixture): +def test_get_thread_messages_length_keeps_non_empty_latest_answer(): id1, id2 = str(uuid4()), str(uuid4()) messages = [ MockMessage(id2, id1, answer="latest-answer"), MockMessage(id1, UUID_NIL, answer="older-answer"), ] - mock_scalars = mocker.patch("core.prompt.utils.get_thread_messages_length.db.session.scalars") - mock_scalars.return_value.all.return_value = messages + session = MagicMock() + session.scalars.return_value.all.return_value = messages - length = get_thread_messages_length("conversation-2") + length = get_thread_messages_length("conversation-2", session=session) assert length == 2 + session.scalars.assert_called_once() diff --git a/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py b/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py index 1f3247590c4..96bb8e8bfc7 100644 --- a/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py +++ b/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py @@ -17,6 +17,7 @@ class TestDataPostProcessor: def test_init_sets_rerank_and_reorder_runners(self): rerank_runner = object() reorder_runner = object() + session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_runner", return_value=rerank_runner) as rerank_mock: with patch.object(DataPostProcessor, "_get_reorder_runner", return_value=reorder_runner) as reorder_mock: @@ -26,6 +27,7 @@ class TestDataPostProcessor: reranking_model={"config": "value"}, weights={"weight": "value"}, reorder_enabled=True, + session=session, ) assert processor.rerank_runner is rerank_runner @@ -35,6 +37,7 @@ class TestDataPostProcessor: "tenant-1", {"config": "value"}, {"weight": "value"}, + session=session, ) reorder_mock.assert_called_once_with(True) @@ -87,6 +90,7 @@ class TestDataPostProcessor: } expected_runner = object() processor = DataPostProcessor.__new__(DataPostProcessor) + session = MagicMock() with patch( "core.rag.data_post_processor.data_post_processor.RerankRunnerFactory.create_rerank_runner", @@ -97,6 +101,7 @@ class TestDataPostProcessor: tenant_id="tenant-1", reranking_model=None, weights=weights_config, + session=session, ) assert result is expected_runner @@ -114,6 +119,7 @@ class TestDataPostProcessor: "reranking_provider_name": "provider-x", "reranking_model_name": "model-y", } + session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=None) as model_mock: with patch( @@ -124,6 +130,7 @@ class TestDataPostProcessor: tenant_id="tenant-1", reranking_model=reranking_model, weights=None, + session=session, ) assert result is None @@ -134,6 +141,7 @@ class TestDataPostProcessor: processor = DataPostProcessor.__new__(DataPostProcessor) model_instance = object() expected_runner = object() + session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=model_instance): with patch( @@ -148,19 +156,22 @@ class TestDataPostProcessor: "reranking_model_name": "model-y", }, weights=None, + session=session, ) assert result is expected_runner factory_mock.assert_called_once_with( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=model_instance, + session=session, ) def test_get_rerank_runner_returns_none_for_unsupported_mode(self): processor = DataPostProcessor.__new__(DataPostProcessor) + session = MagicMock() - assert processor._get_rerank_runner("unsupported", "tenant-1", None, None) is None - assert processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None) is None + assert processor._get_rerank_runner("unsupported", "tenant-1", None, None, session=session) is None + assert processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None, session=session) is None def test_get_reorder_runner_by_flag(self): processor = DataPostProcessor.__new__(DataPostProcessor) diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py b/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py index 9224d6c88e8..8a576bc40a2 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py @@ -52,7 +52,7 @@ class _FakeSelect: def _dataset_keyword_table(data_source_type: str = "database", keyword_table_dict: dict[str, Any] | None = None): return SimpleNamespace( data_source_type=data_source_type, - keyword_table_dict=keyword_table_dict, + get_keyword_table_dict=MagicMock(return_value=keyword_table_dict), keyword_table="", ) @@ -62,19 +62,17 @@ def _dataset(dataset_keyword_table=None, keyword_number=None): id="dataset-1", tenant_id="tenant-1", keyword_number=keyword_number, - dataset_keyword_table=dataset_keyword_table, + get_dataset_keyword_table=MagicMock(return_value=dataset_keyword_table), ) @pytest.fixture def patched_runtime(monkeypatch: pytest.MonkeyPatch): session = MagicMock() - db = SimpleNamespace(session=session) storage = MagicMock() lock = MagicMock(return_value=_DummyLock()) redis_client = SimpleNamespace(lock=lock) - monkeypatch.setattr(jieba_module, "db", db) monkeypatch.setattr(jieba_module, "storage", storage) monkeypatch.setattr(jieba_module, "redis_client", redis_client) @@ -96,7 +94,8 @@ def test_create_indexes_documents_and_returns_self(monkeypatch: pytest.MonkeyPat [ Document(page_content="alpha", metadata={"doc_id": "node-1"}), SimpleNamespace(page_content="ignored", metadata=None), - ] + ], + patched_runtime.session, ) assert result is keyword @@ -105,6 +104,7 @@ def test_create_indexes_documents_and_returns_self(monkeypatch: pytest.MonkeyPat assert call_args[0] == "dataset-1" assert call_args[1] == "node-1" assert set(call_args[2]) == {"kw1", "kw2"} + assert call_args[3] is patched_runtime.session saved_table = keyword._save_dataset_keyword_table.call_args.args[0] assert saved_table["kw1"] == {"node-1"} assert saved_table["kw2"] == {"node-1"} @@ -125,14 +125,18 @@ def test_add_texts_supports_keywords_list_and_extract_fallback(monkeypatch: pyte Document(page_content="extract-this", metadata={"doc_id": "node-1"}), Document(page_content="use-manual", metadata={"doc_id": "node-2"}), ] - keyword.add_texts(texts, keywords_list=[[], ["manual"]]) + keyword.add_texts(texts, patched_runtime.session, keywords_list=[[], ["manual"]]) assert keyword._update_segment_keywords.call_count == 2 first_call = keyword._update_segment_keywords.call_args_list[0].args second_call = keyword._update_segment_keywords.call_args_list[1].args assert set(first_call[2]) == {"auto"} assert second_call[2] == ["manual"] - keyword._save_dataset_keyword_table.assert_called_once() + assert first_call[3] is patched_runtime.session + assert second_call[3] is patched_runtime.session + keyword._save_dataset_keyword_table.assert_called_once_with( + {"auto": {"node-1"}, "manual": {"node-2"}}, patched_runtime.session + ) def test_add_texts_without_keywords_list_always_uses_extractor(monkeypatch: pytest.MonkeyPatch, patched_runtime): @@ -145,33 +149,46 @@ def test_add_texts_without_keywords_list_always_uses_extractor(monkeypatch: pyte monkeypatch.setattr(keyword, "_update_segment_keywords", MagicMock()) monkeypatch.setattr(keyword, "_save_dataset_keyword_table", MagicMock()) - keyword.add_texts([Document(page_content="content", metadata={"doc_id": "node-1"})]) + keyword.add_texts([Document(page_content="content", metadata={"doc_id": "node-1"})], patched_runtime.session) handler.extract_keywords.assert_called_once_with("content", 1) assert set(keyword._update_segment_keywords.call_args.args[2]) == {"from-extractor"} + assert keyword._update_segment_keywords.call_args.args[3] is patched_runtime.session def test_text_exists_handles_missing_and_existing_keyword_table(monkeypatch: pytest.MonkeyPatch): - keyword = Jieba(_dataset(_dataset_keyword_table())) + keyword = Jieba(_dataset(_dataset_keyword_table(keyword_table_dict=None))) + session = MagicMock() + assert keyword.text_exists("node-1", session=session) is False - monkeypatch.setattr(keyword, "_get_dataset_keyword_table", MagicMock(return_value=None)) - assert keyword.text_exists("node-1") is False - - monkeypatch.setattr(keyword, "_get_dataset_keyword_table", MagicMock(return_value={"k": {"node-1", "node-2"}})) - assert keyword.text_exists("node-2") is True - assert keyword.text_exists("node-x") is False + keyword = Jieba( + _dataset( + _dataset_keyword_table( + keyword_table_dict={"__type__": "keyword_table", "__data__": {"table": {"k": {"node-1", "node-2"}}}} + ) + ) + ) + assert keyword.text_exists("node-2", session=session) is True + assert keyword.text_exists("node-x", session=session) is False def test_delete_by_ids_updates_table_when_present(monkeypatch: pytest.MonkeyPatch, patched_runtime): - keyword = Jieba(_dataset(_dataset_keyword_table())) + keyword = Jieba( + _dataset( + _dataset_keyword_table( + keyword_table_dict={"__type__": "keyword_table", "__data__": {"table": {"k": {"node-1", "node-2"}}}} + ) + ) + ) monkeypatch.setattr(keyword, "_get_dataset_keyword_table", MagicMock(return_value={"k": {"node-1", "node-2"}})) monkeypatch.setattr(keyword, "_delete_ids_from_keyword_table", MagicMock(return_value={"k": {"node-2"}})) monkeypatch.setattr(keyword, "_save_dataset_keyword_table", MagicMock()) - keyword.delete_by_ids(["node-1"]) + keyword.delete_by_ids(["node-1"], patched_runtime.session) + keyword._get_dataset_keyword_table.assert_called_once_with(patched_runtime.session) keyword._delete_ids_from_keyword_table.assert_called_once_with({"k": {"node-1", "node-2"}}, ["node-1"]) - keyword._save_dataset_keyword_table.assert_called_once_with({"k": {"node-2"}}) + keyword._save_dataset_keyword_table.assert_called_once_with({"k": {"node-2"}}, patched_runtime.session) def test_delete_by_ids_saves_none_when_keyword_table_is_missing(monkeypatch: pytest.MonkeyPatch, patched_runtime): @@ -180,10 +197,11 @@ def test_delete_by_ids_saves_none_when_keyword_table_is_missing(monkeypatch: pyt monkeypatch.setattr(keyword, "_delete_ids_from_keyword_table", MagicMock()) monkeypatch.setattr(keyword, "_save_dataset_keyword_table", MagicMock()) - keyword.delete_by_ids(["node-1"]) + keyword.delete_by_ids(["node-1"], patched_runtime.session) + keyword._get_dataset_keyword_table.assert_called_once_with(patched_runtime.session) keyword._delete_ids_from_keyword_table.assert_not_called() - keyword._save_dataset_keyword_table.assert_called_once_with(None) + keyword._save_dataset_keyword_table.assert_called_once_with(None, patched_runtime.session) def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch: pytest.MonkeyPatch, patched_runtime): @@ -205,10 +223,9 @@ def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch: monkeypatch.setattr(jieba_module, "DocumentSegment", _FakeDocumentSegment) monkeypatch.setattr(jieba_module, "select", lambda *_: _FakeSelect()) - monkeypatch.setattr(keyword, "_get_dataset_keyword_table", MagicMock(return_value={"k": {"node-1", "node-2"}})) monkeypatch.setattr(keyword, "_retrieve_ids_by_query", MagicMock(return_value=["node-1", "node-2"])) - documents = keyword.search("query", top_k=2, document_ids_filter=["doc-2"]) + documents = keyword.search("query", session=patched_runtime.session, top_k=2, document_ids_filter=["doc-2"]) assert len(documents) == 1 assert documents[0].page_content == "segment-content" @@ -221,11 +238,11 @@ def test_delete_removes_keyword_table_and_optional_file(monkeypatch: pytest.Monk file_keyword = _dataset_keyword_table(data_source_type="object_storage") keyword_db = Jieba(_dataset(db_keyword)) - keyword_db.delete() + keyword_db.delete(session=patched_runtime.session) patched_runtime.storage.delete.assert_not_called() keyword_file = Jieba(_dataset(file_keyword)) - keyword_file.delete() + keyword_file.delete(session=patched_runtime.session) patched_runtime.storage.delete.assert_called_once_with("keyword_files/tenant-1/dataset-1.txt") assert patched_runtime.session.delete.call_count == 2 @@ -235,20 +252,22 @@ def test_delete_removes_keyword_table_and_optional_file(monkeypatch: pytest.Monk def test_save_dataset_keyword_table_to_database(monkeypatch: pytest.MonkeyPatch, patched_runtime): dataset_keyword_table = _dataset_keyword_table(data_source_type="database") keyword = Jieba(_dataset(dataset_keyword_table)) + patched_runtime.session.scalar.return_value = dataset_keyword_table - keyword._save_dataset_keyword_table({"kw": {"node-1"}}) + keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session) assert '"__type__":"keyword_table"' in dataset_keyword_table.keyword_table assert '"index_id":"dataset-1"' in dataset_keyword_table.keyword_table - patched_runtime.session.commit.assert_called_once() + patched_runtime.session.flush.assert_called_once() def test_save_dataset_keyword_table_to_file_storage(monkeypatch: pytest.MonkeyPatch, patched_runtime): dataset_keyword_table = _dataset_keyword_table(data_source_type="file") keyword = Jieba(_dataset(dataset_keyword_table)) patched_runtime.storage.exists.return_value = True + patched_runtime.session.scalar.return_value = dataset_keyword_table - keyword._save_dataset_keyword_table({"kw": {"node-1"}}) + keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session) patched_runtime.storage.delete.assert_called_once_with("keyword_files/tenant-1/dataset-1.txt") patched_runtime.storage.save.assert_called_once() @@ -262,36 +281,28 @@ def test_get_dataset_keyword_table_returns_existing_table_data(monkeypatch: pyte keyword_table_dict={"__type__": "keyword_table", "__data__": {"table": {"kw": ["node-1"]}}} ) keyword = Jieba(_dataset(existing)) - assert keyword._get_dataset_keyword_table() == {"kw": ["node-1"]} + patched_runtime.session.scalar.return_value = existing + assert keyword._get_dataset_keyword_table(patched_runtime.session) == {"kw": ["node-1"]} missing_payload = _dataset_keyword_table(keyword_table_dict=None) keyword_with_missing_payload = Jieba(_dataset(missing_payload)) - assert keyword_with_missing_payload._get_dataset_keyword_table() == {} + patched_runtime.session.scalar.return_value = missing_payload + assert keyword_with_missing_payload._get_dataset_keyword_table(patched_runtime.session) == {} def test_get_dataset_keyword_table_creates_table_when_missing(monkeypatch: pytest.MonkeyPatch, patched_runtime): - created_tables: list[SimpleNamespace] = [] - - def _fake_dataset_keyword_table(**kwargs): - kwargs.setdefault("keyword_table", "") - kwargs.setdefault("keyword_table_dict", None) - table = SimpleNamespace(**kwargs) - created_tables.append(table) - return table - keyword = Jieba(_dataset(dataset_keyword_table=None)) - monkeypatch.setattr(jieba_module, "DatasetKeywordTable", _fake_dataset_keyword_table) monkeypatch.setattr(jieba_module.dify_config, "KEYWORD_DATA_SOURCE_TYPE", "database") + patched_runtime.session.scalar.return_value = None - result = keyword._get_dataset_keyword_table() + result = keyword._get_dataset_keyword_table(patched_runtime.session) assert result == {} - assert len(created_tables) == 1 - assert created_tables[0].dataset_id == "dataset-1" - assert created_tables[0].data_source_type == "database" - assert '"index_id":"dataset-1"' in created_tables[0].keyword_table - patched_runtime.session.add.assert_called_once_with(created_tables[0]) - patched_runtime.session.commit.assert_called_once() + created_table = patched_runtime.session.add.call_args.args[0] + assert created_table.dataset_id == "dataset-1" + assert created_table.data_source_type == "database" + assert '"index_id":"dataset-1"' in created_table.keyword_table + patched_runtime.session.flush.assert_called_once() def test_add_and_delete_ids_from_keyword_table_helpers(): @@ -335,37 +346,42 @@ def test_update_segment_keywords_updates_when_segment_exists(monkeypatch: pytest segment = SimpleNamespace(keywords=[]) patched_runtime.session.scalar.return_value = segment - keyword._update_segment_keywords("dataset-1", "node-1", ["kw1", "kw2"]) + keyword._update_segment_keywords("dataset-1", "node-1", ["kw1", "kw2"], patched_runtime.session) assert segment.keywords == ["kw1", "kw2"] patched_runtime.session.add.assert_called_once_with(segment) - patched_runtime.session.commit.assert_called_once() + patched_runtime.session.flush.assert_called_once() patched_runtime.session.reset_mock() patched_runtime.session.scalar.return_value = None - keyword._update_segment_keywords("dataset-1", "node-missing", ["kw3"]) + keyword._update_segment_keywords("dataset-1", "node-missing", ["kw3"], patched_runtime.session) patched_runtime.session.add.assert_not_called() - patched_runtime.session.commit.assert_not_called() + patched_runtime.session.flush.assert_not_called() -def test_create_segment_keywords_and_update_segment_keywords_index(monkeypatch: pytest.MonkeyPatch): +def test_create_segment_keywords_and_update_segment_keywords_index(monkeypatch: pytest.MonkeyPatch, patched_runtime): keyword = Jieba(_dataset(_dataset_keyword_table())) monkeypatch.setattr(keyword, "_get_dataset_keyword_table", MagicMock(return_value={})) monkeypatch.setattr(keyword, "_update_segment_keywords", MagicMock()) monkeypatch.setattr(keyword, "_save_dataset_keyword_table", MagicMock()) - keyword.create_segment_keywords("node-1", ["kw"]) - keyword._update_segment_keywords.assert_called_once_with("dataset-1", "node-1", ["kw"]) - keyword._save_dataset_keyword_table.assert_called_once() + keyword.create_segment_keywords("node-1", ["kw"], patched_runtime.session) + keyword._get_dataset_keyword_table.assert_called_once_with(patched_runtime.session) + keyword._update_segment_keywords.assert_called_once_with("dataset-1", "node-1", ["kw"], patched_runtime.session) + keyword._save_dataset_keyword_table.assert_called_once_with({"kw": {"node-1"}}, patched_runtime.session) + keyword._get_dataset_keyword_table.reset_mock() keyword._save_dataset_keyword_table.reset_mock() - keyword.update_segment_keywords_index("node-2", ["kw2"]) - keyword._save_dataset_keyword_table.assert_called_once() + keyword.update_segment_keywords_index("node-2", ["kw2"], patched_runtime.session) + keyword._get_dataset_keyword_table.assert_called_once_with(patched_runtime.session) + keyword._save_dataset_keyword_table.assert_called_once_with({"kw2": {"node-2"}}, patched_runtime.session) -def test_multi_create_segment_keywords_uses_provided_and_extracted_keywords(monkeypatch: pytest.MonkeyPatch): +def test_multi_create_segment_keywords_uses_provided_and_extracted_keywords( + monkeypatch: pytest.MonkeyPatch, patched_runtime +): keyword = Jieba(_dataset(_dataset_keyword_table(), keyword_number=2)) handler = MagicMock() handler.extract_keywords.return_value = {"auto"} @@ -380,7 +396,8 @@ def test_multi_create_segment_keywords_uses_provided_and_extracted_keywords(monk [ {"segment": first_segment, "keywords": ["manual"]}, {"segment": second_segment, "keywords": []}, - ] + ], + patched_runtime.session, ) assert first_segment.keywords == ["manual"] @@ -388,6 +405,7 @@ def test_multi_create_segment_keywords_uses_provided_and_extracted_keywords(monk saved_table = keyword._save_dataset_keyword_table.call_args.args[0] assert saved_table["manual"] == {"node-1"} assert saved_table["auto"] == {"node-2"} + assert keyword._save_dataset_keyword_table.call_args.args[1] is patched_runtime.session def test_set_orjson_default_and_dumps_with_sets(): diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py index 12855ed5647..5e60c7a1979 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py @@ -1,5 +1,6 @@ from types import SimpleNamespace from typing import override +from unittest.mock import MagicMock import pytest @@ -9,28 +10,28 @@ from core.rag.models.document import Document class _KeywordThatRaises(BaseKeyword): @override - def create(self, texts: list[Document], **kwargs): - return super().create(texts, **kwargs) + def create(self, texts: list[Document], session, **kwargs): + return super().create(texts, session, **kwargs) @override - def add_texts(self, texts: list[Document], **kwargs): - return super().add_texts(texts, **kwargs) + def add_texts(self, texts: list[Document], session, **kwargs): + return super().add_texts(texts, session, **kwargs) @override - def text_exists(self, id: str) -> bool: - return super().text_exists(id) + def text_exists(self, id: str, *, session) -> bool: + return super().text_exists(id, session=session) @override - def delete_by_ids(self, ids: list[str]): - return super().delete_by_ids(ids) + def delete_by_ids(self, ids: list[str], session, **kwargs): + return super().delete_by_ids(ids, session, **kwargs) @override - def delete(self): - return super().delete() + def delete(self, *, session): + return super().delete(session=session) @override - def search(self, query: str, **kwargs): - return super().search(query, **kwargs) + def search(self, query: str, *, session, **kwargs): + return super().search(query, session=session, **kwargs) class _KeywordForHelpers(BaseKeyword): @@ -39,50 +40,51 @@ class _KeywordForHelpers(BaseKeyword): self._existing_ids = existing_ids or set() @override - def create(self, texts: list[Document], **kwargs): + def create(self, texts: list[Document], session, **kwargs): return self @override - def add_texts(self, texts: list[Document], **kwargs): + def add_texts(self, texts: list[Document], session, **kwargs): return None @override - def text_exists(self, id: str) -> bool: + def text_exists(self, id: str, *, session) -> bool: return id in self._existing_ids @override - def delete_by_ids(self, ids: list[str]): + def delete_by_ids(self, ids: list[str], session, **kwargs): return None @override - def delete(self): + def delete(self, *, session): return None @override - def search(self, query: str, **kwargs): + def search(self, query: str, *, session, **kwargs): return [] def test_abstract_methods_raise_not_implemented(): keyword = _KeywordThatRaises(SimpleNamespace(id="dataset-1")) + session = MagicMock() with pytest.raises(NotImplementedError): - keyword.create([]) + keyword.create([], session) with pytest.raises(NotImplementedError): - keyword.add_texts([]) + keyword.add_texts([], session) with pytest.raises(NotImplementedError): - keyword.text_exists("doc-1") + keyword.text_exists("doc-1", session=session) with pytest.raises(NotImplementedError): - keyword.delete_by_ids(["doc-1"]) + keyword.delete_by_ids(["doc-1"], session) with pytest.raises(NotImplementedError): - keyword.delete() + keyword.delete(session=session) with pytest.raises(NotImplementedError): - keyword.search("query") + keyword.search("query", session=session) def test_filter_duplicate_texts_removes_existing_doc_ids(): @@ -93,7 +95,7 @@ def test_filter_duplicate_texts_removes_existing_doc_ids(): SimpleNamespace(page_content="without-metadata", metadata=None), ] - filtered = keyword._filter_duplicate_texts(texts) + filtered = keyword._filter_duplicate_texts(texts, session=MagicMock()) assert [text.metadata["doc_id"] for text in filtered if text.metadata] == ["keep"] assert any(text.metadata is None for text in filtered) diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py index e1765b17cb4..57df6bfa18a 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py @@ -48,19 +48,20 @@ def test_keyword_methods_forward_to_processor(): keyword._keyword_processor = processor docs = [Document(page_content="doc", metadata={"doc_id": "doc-1"})] - keyword.create(docs, foo="bar") - keyword.add_texts(docs, batch=True) - assert keyword.text_exists("doc-1") is True - keyword.delete_by_ids(["doc-1"]) - keyword.delete() - assert keyword.search("query", top_k=1) == processor.search.return_value + session = MagicMock() + keyword.create(docs, session, foo="bar") + keyword.add_texts(docs, session, batch=True, keywords_list=[["kw"]]) + assert keyword.text_exists("doc-1", session=session) is True + keyword.delete_by_ids(["doc-1"], session) + keyword.delete(session=session) + assert keyword.search("query", session=session, top_k=1) == processor.search.return_value - processor.create.assert_called_once_with(docs, foo="bar") - processor.add_texts.assert_called_once_with(docs, batch=True) - processor.text_exists.assert_called_once_with("doc-1") - processor.delete_by_ids.assert_called_once_with(["doc-1"]) - processor.delete.assert_called_once() - processor.search.assert_called_once_with("query", top_k=1) + processor.create.assert_called_once_with(docs, session, foo="bar") + processor.add_texts.assert_called_once_with(docs, session, batch=True, keywords_list=[["kw"]]) + processor.text_exists.assert_called_once_with("doc-1", session=session) + processor.delete_by_ids.assert_called_once_with(["doc-1"], session) + processor.delete.assert_called_once_with(session=session) + processor.search.assert_called_once_with("query", session=session, top_k=1) def test_keyword_getattr_returns_callable_and_raises_for_invalid_attributes(): 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 d8452d91e2c..6a89e6ec5c8 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 @@ -118,27 +118,17 @@ class _FakeScalarsResult: class _FakeSession: - def __init__(self, execute_payloads: list[list], summaries: list) -> None: - self._payloads = list(execute_payloads) - self._summaries = summaries + def __init__(self, execute_payloads: list[list], scalars_payloads: list[list]) -> None: + self._execute_payloads = list(execute_payloads) + self._scalars_payloads = list(scalars_payloads) def execute(self, stmt): - data = self._payloads.pop(0) if self._payloads else [] + data = self._execute_payloads.pop(0) if self._execute_payloads else [] return _FakeExecuteResult(data) def scalars(self, stmt): - return _FakeScalarsResult(self._summaries) - - -class _FakeSessionContext: - def __init__(self, session: _FakeSession) -> None: - self._session = session - - def __enter__(self) -> _FakeSession: - return self._session - - def __exit__(self, exc_type, exc, tb) -> bool: - return False + data = self._scalars_payloads.pop(0) if self._scalars_payloads else [] + return _FakeScalarsResult(data) class _SimpleRetrievalChildChunk: @@ -182,6 +172,17 @@ class TestRetrievalServiceInternals: app.app_context.return_value.__exit__.return_value = False return app + @pytest.fixture + def vector_session(self, monkeypatch: pytest.MonkeyPatch) -> MagicMock: + session = MagicMock() + session_context = MagicMock() + session_context.__enter__.return_value = session + session_class = MagicMock(return_value=session_context) + session.context = session_context + monkeypatch.setattr(retrieval_service_module, "Session", session_class) + monkeypatch.setattr(retrieval_service_module, "db", SimpleNamespace(engine=Mock())) + return session + def test_retrieve_with_attachment_ids_only(self, monkeypatch: pytest.MonkeyPatch, internal_dataset): with ( patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset", return_value=internal_dataset), @@ -283,19 +284,28 @@ class TestRetrievalServiceInternals: mock_keyword_class.return_value = keyword_instance all_documents: list[Document] = [] exceptions: list[str] = [] + session = MagicMock() + engine = Mock() - RetrievalService.keyword_search( - flask_app=internal_flask_app, - dataset_id=internal_dataset.id, - query='query "with quotes"', - top_k=5, - all_documents=all_documents, - exceptions=exceptions, - ) + with ( + patch.object(retrieval_service_module, "db", SimpleNamespace(engine=engine)), + patch("core.rag.datasource.retrieval_service.Session") as session_class, + ): + session_class.return_value.__enter__.return_value = session + RetrievalService.keyword_search( + flask_app=internal_flask_app, + dataset_id=internal_dataset.id, + query='query "with quotes"', + top_k=5, + all_documents=all_documents, + exceptions=exceptions, + ) assert len(all_documents) == 1 assert exceptions == [] keyword_instance.search.assert_called_once() + assert keyword_instance.search.call_args.kwargs["session"] is session + session_class.assert_called_once_with(engine) @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_keyword_search_appends_exception_when_dataset_missing(self, mock_get_dataset, internal_flask_app): @@ -326,23 +336,30 @@ class TestRetrievalServiceInternals: mock_keyword_class.return_value = keyword_instance all_documents: list[Document] = [] exceptions: list[str] = [] + session = MagicMock() - RetrievalService.keyword_search( - flask_app=internal_flask_app, - dataset_id=internal_dataset.id, - query="query", - top_k=2, - all_documents=all_documents, - exceptions=exceptions, - ) + with ( + patch.object(retrieval_service_module, "db", SimpleNamespace(engine=Mock())), + patch("core.rag.datasource.retrieval_service.Session") as session_class, + ): + session_class.return_value.__enter__.return_value = session + RetrievalService.keyword_search( + flask_app=internal_flask_app, + dataset_id=internal_dataset.id, + query="query", + top_k=2, + all_documents=all_documents, + exceptions=exceptions, + ) assert all_documents == [] assert exceptions == ["keyword failed"] + assert keyword_instance.search.call_args.kwargs["session"] is session @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_embedding_search_text_without_reranking( - self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app + self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app, vector_session ): internal_dataset.is_multimodal = False mock_get_dataset.return_value = internal_dataset @@ -368,12 +385,13 @@ class TestRetrievalServiceInternals: assert len(all_documents) == 1 assert exceptions == [] + mock_vector_class.assert_called_once_with(dataset=internal_dataset, session=vector_session) vector_instance.search_by_vector.assert_called_once() @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_embedding_search_image_non_multimodal_returns_early( - self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app + self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app, vector_session ): internal_dataset.is_multimodal = False mock_get_dataset.return_value = internal_dataset @@ -411,6 +429,7 @@ class TestRetrievalServiceInternals: mock_model_manager_class, internal_dataset, internal_flask_app, + vector_session, ): internal_dataset.is_multimodal = True mock_get_dataset.return_value = internal_dataset @@ -418,7 +437,12 @@ class TestRetrievalServiceInternals: reranked_docs = [create_mock_document("image-content-reranked", "img-doc", 0.97)] vector_instance = Mock() - vector_instance.search_by_file.return_value = original_docs + + def search_by_file(**_kwargs): + assert vector_session.context.__exit__.call_count == 0 + return original_docs + + vector_instance.search_by_file.side_effect = search_by_file mock_vector_class.return_value = vector_instance processor_instance = Mock() @@ -466,6 +490,7 @@ class TestRetrievalServiceInternals: mock_model_manager_class, internal_dataset, internal_flask_app, + vector_session, ): internal_dataset.is_multimodal = True mock_get_dataset.return_value = internal_dataset @@ -511,7 +536,13 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_embedding_search_text_with_reranking_non_multimodal( - self, mock_get_dataset, mock_vector_class, mock_processor_class, internal_dataset, internal_flask_app + self, + mock_get_dataset, + mock_vector_class, + mock_processor_class, + internal_dataset, + internal_flask_app, + vector_session, ): internal_dataset.is_multimodal = False mock_get_dataset.return_value = internal_dataset @@ -552,7 +583,7 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_embedding_search_appends_exception_when_vector_fails( - self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app + self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app, vector_session ): mock_get_dataset.return_value = internal_dataset vector_instance = Mock() @@ -580,7 +611,7 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_full_text_index_search_without_reranking( - self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app + self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app, vector_session ): mock_get_dataset.return_value = internal_dataset vector_instance = Mock() @@ -609,7 +640,13 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_full_text_index_search_with_reranking( - self, mock_get_dataset, mock_vector_class, mock_processor_class, internal_dataset, internal_flask_app + self, + mock_get_dataset, + mock_vector_class, + mock_processor_class, + internal_dataset, + internal_flask_app, + vector_session, ): mock_get_dataset.return_value = internal_dataset original_docs = [create_mock_document("fulltext", "ft-1", 0.68)] @@ -669,7 +706,7 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Vector") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_full_text_index_search_appends_exception_when_search_fails( - self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app + self, mock_get_dataset, mock_vector_class, internal_dataset, internal_flask_app, vector_session ): mock_get_dataset.return_value = internal_dataset vector_instance = Mock() @@ -694,12 +731,12 @@ class TestRetrievalServiceInternals: assert exceptions == ["fulltext failed"] def test_format_retrieval_documents_with_empty_input_returns_empty_list(self): - assert RetrievalService.format_retrieval_documents([]) == [] + assert RetrievalService.format_retrieval_documents(MagicMock(), []) == [] def test_format_retrieval_documents_without_document_id_returns_empty_list(self): documents = [Document(page_content="content", metadata={"doc_id": "doc-1", "score": 0.4}, provider="dify")] - assert RetrievalService.format_retrieval_documents(documents) == [] + assert RetrievalService.format_retrieval_documents(MagicMock(), documents) == [] def test_format_retrieval_documents_with_parent_child_summary_and_attachments( self, monkeypatch: pytest.MonkeyPatch @@ -716,13 +753,6 @@ class TestRetrievalServiceInternals: dataset_id="dataset-id", ) - scalars_result = Mock() - scalars_result.all.return_value = [ - dataset_doc_parent, - dataset_doc_text, - dataset_doc_parent_summary, - ] - monkeypatch.setattr(retrieval_service_module.db.session, "scalars", Mock(return_value=scalars_result)) monkeypatch.setattr(retrieval_service_module, "RetrievalChildChunk", _SimpleRetrievalChildChunk) monkeypatch.setattr(retrieval_service_module, "RetrievalSegments", _SimpleRetrievalSegment) @@ -826,16 +856,14 @@ class TestRetrievalServiceInternals: [segment_parent, segment_text], [segment_summary, segment_parent_summary], ], - summaries=[ - SimpleNamespace(chunk_id="segment-summary", summary_content="summary for text"), - SimpleNamespace(chunk_id="segment-parent-summary", summary_content="summary for parent"), + scalars_payloads=[ + [dataset_doc_parent, dataset_doc_text, dataset_doc_parent_summary], + [ + SimpleNamespace(chunk_id="segment-summary", summary_content="summary for text"), + SimpleNamespace(chunk_id="segment-parent-summary", summary_content="summary for parent"), + ], ], ) - monkeypatch.setattr( - retrieval_service_module.session_factory, - "create_session", - lambda: _FakeSessionContext(fake_session), - ) monkeypatch.setattr( RetrievalService, "get_segment_attachment_infos", @@ -867,7 +895,7 @@ class TestRetrievalServiceInternals: ], ) - result = RetrievalService.format_retrieval_documents(input_documents) + result = RetrievalService.format_retrieval_documents(fake_session, input_documents) assert len(result) == 4 result_by_segment_id = {item.segment.id: item for item in result} @@ -881,17 +909,16 @@ class TestRetrievalServiceInternals: assert result_by_segment_id["segment-parent-summary"].summary == "summary for parent" assert result_by_segment_id["segment-parent-summary"].child_chunks == [] - def test_format_retrieval_documents_rolls_back_and_raises_when_db_fails(self, monkeypatch: pytest.MonkeyPatch): - rollback = Mock() - monkeypatch.setattr(retrieval_service_module.db.session, "rollback", rollback) - monkeypatch.setattr(retrieval_service_module.db.session, "scalars", Mock(side_effect=RuntimeError("db error"))) + def test_format_retrieval_documents_rolls_back_and_raises_when_db_fails(self): + session = MagicMock() + session.scalars.side_effect = RuntimeError("db error") documents = [Document(page_content="content", metadata={"document_id": "doc-1"}, provider="dify")] with pytest.raises(RuntimeError, match="db error"): - RetrievalService.format_retrieval_documents(documents) + RetrievalService.format_retrieval_documents(session, documents) - rollback.assert_called_once() + session.rollback.assert_called_once() def test_retrieve_internal_returns_early_without_query_or_attachment(self, internal_dataset, internal_flask_app): all_documents: list[Document] = [] @@ -963,7 +990,7 @@ class TestRetrievalServiceInternals: ) def test_retrieve_internal_hybrid_weighted_attachment_flow( - self, monkeypatch: pytest.MonkeyPatch, internal_dataset, internal_flask_app + self, monkeypatch: pytest.MonkeyPatch, internal_dataset, internal_flask_app, vector_session ): executor = _ImmediateExecutor() monkeypatch.setattr(retrieval_service_module, "ThreadPoolExecutor", lambda *args, **kwargs: executor) diff --git a/api/tests/unit_tests/core/rag/datasource/test_retrieval_attachment_access.py b/api/tests/unit_tests/core/rag/datasource/test_retrieval_attachment_access.py index 8790452f27c..aa6dd66975c 100644 --- a/api/tests/unit_tests/core/rag/datasource/test_retrieval_attachment_access.py +++ b/api/tests/unit_tests/core/rag/datasource/test_retrieval_attachment_access.py @@ -152,7 +152,7 @@ def test_knowledge_retrieval_grants_returned_segments_to_current_scope( "multiple_retrieve", lambda **kwargs: [RagDocument(page_content="segment content", provider="dify")], ) - monkeypatch.setattr(RetrievalService, "format_retrieval_documents", lambda documents: [record]) + monkeypatch.setattr(RetrievalService, "format_retrieval_documents", lambda _session, documents: [record]) factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) monkeypatch.setattr(dataset_retrieval_module.session_factory, "create_session", factory) scope = FileAccessScope( diff --git a/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py b/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py index d86ffd26478..673832fa2d4 100644 --- a/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py +++ b/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py @@ -147,10 +147,11 @@ def test_get_vector_factory_entry_point_overrides_builtin(vector_factory_module, def test_vector_init_uses_default_and_custom_attributes(vector_factory_module): dataset = SimpleNamespace(id="dataset-1") + session = MagicMock() - with patch.object(vector_factory_module.Vector, "_init_vector", return_value="processor"): - default_vector = vector_factory_module.Vector(dataset) - custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"]) + with patch.object(vector_factory_module.Vector, "_init_vector", return_value="processor") as init_vector: + default_vector = vector_factory_module.Vector(dataset, session=session) + custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"], session=session) # `is_summary` and `original_chunk_id` must be in the default return-properties # projection so summary index retrieval works on backends that honor the list @@ -167,14 +168,17 @@ def test_vector_init_uses_default_and_custom_attributes(vector_factory_module): assert custom_vector._attributes == ["doc_id"] # ``_embeddings`` is now a lazy proxy that defers materializing the real # embedding model until ``embed_*`` is invoked, so cleanup paths never - # trigger billing/feature-service calls during ``Vector(dataset)`` + # trigger billing/feature-service calls during ``Vector(dataset, session=...)`` # construction. See ``_LazyEmbeddings``. assert isinstance(default_vector._embeddings, vector_factory_module._LazyEmbeddings) + assert default_vector._session is session + assert custom_vector._session is session assert default_vector._vector_processor == "processor" + assert [call.kwargs["session"] for call in init_vector.call_args_list] == [session, session] def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_module, monkeypatch: pytest.MonkeyPatch): - """``Vector(dataset)`` must not transitively call ``ModelManager`` during + """``Vector(dataset, session=...)`` must not transitively call ``ModelManager`` during construction. The real embedding model should only be materialized on the first ``embed_*`` call (i.e. create / search paths) so cleanup paths (``delete_by_ids`` / ``delete``) remain resilient to billing-API failures. @@ -237,7 +241,7 @@ def test_init_vector_prefers_dataset_index_struct(vector_factory_module, monkeyp vector._attributes = ["doc_id"] vector._embeddings = "embeddings" - result = vector._init_vector() + result = vector._init_vector(session=MagicMock()) assert result == "vector-processor" assert calls["vector_type"] == vector_factory_module.VectorType.UPSTASH @@ -257,11 +261,8 @@ def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch: monkeypatch.setattr(vector_factory_module, "Whitelist", SimpleNamespace(tenant_id=_Expr(), category=_Expr())) monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt")) - monkeypatch.setattr( - vector_factory_module, - "db", - SimpleNamespace(session=SimpleNamespace(scalars=lambda _stmt: SimpleNamespace(one_or_none=lambda: object()))), - ) + session = MagicMock() + session.scalars.return_value.one_or_none.return_value = object() monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE", vector_factory_module.VectorType.CHROMA) monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE_WHITELIST_ENABLE", True) monkeypatch.setattr( @@ -275,10 +276,11 @@ def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch: vector._attributes = ["doc_id"] vector._embeddings = "embeddings" - result = vector._init_vector() + result = vector._init_vector(session=session) assert result == "vector-processor" assert calls["vector_type"] == vector_factory_module.VectorType.TIDB_ON_QDRANT + session.scalars.assert_called_once_with("stmt") def test_init_vector_raises_when_vector_store_missing(vector_factory_module, monkeypatch: pytest.MonkeyPatch): @@ -291,7 +293,7 @@ def test_init_vector_raises_when_vector_store_missing(vector_factory_module, mon vector._embeddings = "embeddings" with pytest.raises(ValueError, match="Vector store must be specified"): - vector._init_vector() + vector._init_vector(session=MagicMock()) def test_create_batches_texts_and_skips_empty_input(vector_factory_module): @@ -357,18 +359,12 @@ def test_create_multimodal_filters_missing_uploads(vector_factory_module, monkey vector._embeddings = MagicMock() vector._embeddings.embed_multimodal_documents.return_value = [[0.1, 0.2]] vector._vector_processor = MagicMock() + session = MagicMock() + vector._session = session + session.scalars.return_value = SimpleNamespace(all=lambda: [SimpleNamespace(id="f-1", key="k-1")]) monkeypatch.setattr(vector_factory_module, "UploadFile", SimpleNamespace(id=_Field())) monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt")) - monkeypatch.setattr( - vector_factory_module, - "db", - SimpleNamespace( - session=SimpleNamespace( - scalars=lambda _stmt: SimpleNamespace(all=lambda: [SimpleNamespace(id="f-1", key="k-1")]) - ) - ), - ) monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"abc")) docs = [ @@ -491,12 +487,14 @@ def test_search_by_file_handles_missing_and_existing_upload(vector_factory_modul vector._embeddings = MagicMock() vector._vector_processor = MagicMock() - mock_session = SimpleNamespace(get=lambda _model, _id: None) - monkeypatch.setattr(vector_factory_module, "db", SimpleNamespace(session=mock_session)) + session = MagicMock() + session.get.return_value = None + vector._session = session assert vector.search_by_file("file-1") == [] + session.get.assert_called_once_with(vector_factory_module.UploadFile, "file-1") - mock_session.get = lambda _model, _id: SimpleNamespace(key="blob-key") + session.get.return_value = SimpleNamespace(key="blob-key") monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"file-bytes")) vector._embeddings.embed_multimodal_query.return_value = [0.3, 0.4] vector._vector_processor.search_by_vector.return_value = ["hit"] @@ -504,6 +502,7 @@ def test_search_by_file_handles_missing_and_existing_upload(vector_factory_modul result = vector.search_by_file("file-2", top_k=2) assert result == ["hit"] + session.get.assert_called_with(vector_factory_module.UploadFile, "file-2") payload = vector._embeddings.embed_multimodal_query.call_args.args[0] assert payload["content_type"] == vector_factory_module.DocType.IMAGE assert payload["file_id"] == "file-2" diff --git a/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py b/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py index 007a76aa66c..740e4391cac 100644 --- a/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py +++ b/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py @@ -100,20 +100,18 @@ class TestDatasetDocumentStoreDocs: mock_segment.dataset_id = "test-dataset-id" mock_segment.content = "Test content" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalars.return_value.all.return_value = [mock_segment] + mock_session = MagicMock() + mock_session.scalars.return_value.all.return_value = [mock_segment] - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.docs + result = store.get_docs(mock_session) - assert "node-1" in result - assert isinstance(result["node-1"], Document) + assert "node-1" in result + assert isinstance(result["node-1"], Document) def test_docs_empty_dataset(self): """Test docs property with no segments.""" @@ -121,19 +119,17 @@ class TestDatasetDocumentStoreDocs: mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalars.return_value.all.return_value = [] + mock_session = MagicMock() + mock_session.scalars.return_value.all.return_value = [] - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.docs + result = store.get_docs(mock_session) - assert result == {} + assert result == {} class TestDatasetDocumentStoreAddDocuments: @@ -162,12 +158,10 @@ class TestDatasetDocumentStoreAddDocuments: mock_model_instance.get_text_embedding_num_tokens.return_value = [10] with ( - patch("core.rag.docstore.dataset_docstore.db") as mock_db, patch("core.rag.docstore.dataset_docstore.ModelManager.for_tenant") as mock_manager_class, ): mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = None + mock_session.scalar.return_value = None mock_manager = MagicMock() mock_manager.get_model_instance.return_value = mock_model_instance @@ -181,10 +175,10 @@ class TestDatasetDocumentStoreAddDocuments: document_id="test-doc-id", ) - store.add_documents([mock_doc]) + store.add_documents([mock_doc], session=mock_session) - mock_db.session.add.assert_called() - mock_db.session.commit.assert_called() + mock_session.add.assert_called() + mock_session.flush.assert_called() def test_add_documents_update_existing_document(self): """Test updating existing document with allow_update=True.""" @@ -208,22 +202,20 @@ class TestDatasetDocumentStoreAddDocuments: mock_existing_segment = MagicMock() mock_existing_segment.id = "seg-1" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = 5 + mock_session = MagicMock() + mock_session.scalar.return_value = 5 - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): - with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): + with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_documents([mock_doc]) + store.add_documents([mock_doc], session=mock_session) - mock_db.session.commit.assert_called() + mock_session.flush.assert_called() def test_add_documents_raises_when_not_allowed(self): """Test that adding existing doc without allow_update raises ValueError.""" @@ -243,17 +235,17 @@ class TestDatasetDocumentStoreAddDocuments: mock_doc.children = None mock_existing_segment = MagicMock() + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - with pytest.raises(ValueError, match="already exists"): - store.add_documents([mock_doc], allow_update=False) + with pytest.raises(ValueError, match="already exists"): + store.add_documents([mock_doc], session=mock_session, allow_update=False) def test_add_documents_with_answer_metadata(self): """Test adding document with answer in metadata.""" @@ -273,38 +265,36 @@ class TestDatasetDocumentStoreAddDocuments: mock_doc.attachments = None mock_doc.children = None - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_documents([mock_doc]) + store.add_documents([mock_doc], session=mock_session) - mock_db.session.add.assert_called() + mock_session.add.assert_called() def test_add_documents_with_invalid_document_type(self): """Test that non-Document raises ValueError.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - with pytest.raises(ValueError, match="must be a Document"): - store.add_documents(["not a document"]) + with pytest.raises(ValueError, match="must be a Document"): + store.add_documents(["not a document"], session=mock_session) def test_add_documents_with_none_metadata(self): """Test that document with None metadata raises ValueError.""" @@ -315,16 +305,16 @@ class TestDatasetDocumentStoreAddDocuments: mock_doc = MagicMock(spec=Document) mock_doc.page_content = "Test content" mock_doc.metadata = None + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - with pytest.raises(ValueError, match="metadata must be a dict"): - store.add_documents([mock_doc]) + with pytest.raises(ValueError, match="metadata must be a dict"): + store.add_documents([mock_doc], session=mock_session) def test_add_documents_with_save_child(self): """Test adding documents with save_child=True.""" @@ -350,22 +340,20 @@ class TestDatasetDocumentStoreAddDocuments: mock_doc.attachments = None mock_doc.children = [mock_child] - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_documents([mock_doc], save_child=True) + store.add_documents([mock_doc], session=mock_session, save_child=True) - mock_db.session.add.assert_called() + mock_session.add.assert_called() class TestDatasetDocumentStoreExists: @@ -378,34 +366,34 @@ class TestDatasetDocumentStoreExists: mock_dataset.id = "test-dataset-id" mock_segment = MagicMock() + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.document_exists("doc-1") + result = store.document_exists("doc-1", session=mock_session) - assert result is True + assert result is True def test_document_exists_returns_false(self): """Test document_exists returns False when segment doesn't exist.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.document_exists("doc-1") + result = store.document_exists("doc-1", session=mock_session) - assert result is False + assert result is False class TestDatasetDocumentStoreGetDocument: @@ -423,51 +411,51 @@ class TestDatasetDocumentStoreGetDocument: mock_segment.document_id = "doc-1" mock_segment.dataset_id = "test-dataset-id" mock_segment.content = "Test content" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document("node-1", raise_error=False) + result = store.get_document("node-1", session=mock_session, raise_error=False) - assert isinstance(result, Document) - assert result.page_content == "Test content" + assert isinstance(result, Document) + assert result.page_content == "Test content" def test_get_document_returns_none_when_not_found(self): """Test get_document returns None when not found and raise_error=False.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document("nonexistent", raise_error=False) + result = store.get_document("nonexistent", session=mock_session, raise_error=False) - assert result is None + assert result is None def test_get_document_raises_when_not_found(self): """Test get_document raises ValueError when not found and raise_error=True.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - with pytest.raises(ValueError, match="not found"): - store.get_document("nonexistent", raise_error=True) + with pytest.raises(ValueError, match="not found"): + store.get_document("nonexistent", session=mock_session, raise_error=True) class TestDatasetDocumentStoreDeleteDocument: @@ -480,51 +468,51 @@ class TestDatasetDocumentStoreDeleteDocument: mock_dataset.id = "test-dataset-id" mock_segment = MagicMock() + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - store.delete_document("doc-1") + store.delete_document("doc-1", session=mock_session) - mock_db.session.delete.assert_called_with(mock_segment) - mock_db.session.commit.assert_called() + mock_session.delete.assert_called_with(mock_segment) + mock_session.flush.assert_called() def test_delete_document_returns_none_when_not_found(self): """Test delete_document returns None when not found and raise_error=False.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.delete_document("nonexistent", raise_error=False) + result = store.delete_document("nonexistent", session=mock_session, raise_error=False) - assert result is None + assert result is None def test_delete_document_raises_when_not_found(self): """Test delete_document raises ValueError when not found and raise_error=True.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - with pytest.raises(ValueError, match="not found"): - store.delete_document("nonexistent", raise_error=True) + with pytest.raises(ValueError, match="not found"): + store.delete_document("nonexistent", session=mock_session, raise_error=True) class TestDatasetDocumentStoreHashOperations: @@ -538,35 +526,35 @@ class TestDatasetDocumentStoreHashOperations: mock_segment = MagicMock() mock_segment.index_node_hash = "old-hash" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - store.set_document_hash("doc-1", "new-hash") + store.set_document_hash("doc-1", "new-hash", session=mock_session) - assert mock_segment.index_node_hash == "new-hash" - mock_db.session.commit.assert_called() + assert mock_segment.index_node_hash == "new-hash" + mock_session.flush.assert_called() def test_set_document_hash_returns_none_when_not_found(self): """Test set_document_hash returns None when segment not found.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.set_document_hash("nonexistent", "new-hash") + result = store.set_document_hash("nonexistent", "new-hash", session=mock_session) - assert result is None + assert result is None def test_get_document_hash_success(self): """Test getting document hash successfully.""" @@ -576,34 +564,34 @@ class TestDatasetDocumentStoreHashOperations: mock_segment = MagicMock() mock_segment.index_node_hash = "test-hash" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document_hash("doc-1") + result = store.get_document_hash("doc-1", session=mock_session) - assert result == "test-hash" + assert result == "test-hash" def test_get_document_hash_returns_none_when_not_found(self): """Test get_document_hash returns None when segment not found.""" mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" + mock_session = MagicMock() - with patch("core.rag.docstore.dataset_docstore.db"): - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document_hash("nonexistent") + result = store.get_document_hash("nonexistent", session=mock_session) - assert result is None + assert result is None class TestDatasetDocumentStoreSegment: @@ -617,19 +605,17 @@ class TestDatasetDocumentStoreSegment: mock_segment = MagicMock(spec=DocumentSegment) - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = mock_segment + mock_session = MagicMock() + mock_session.scalar.return_value = mock_segment - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document_segment("doc-1") + result = store.get_document_segment("doc-1", session=mock_session) - assert result == mock_segment + assert result == mock_segment def test_get_document_segment_returns_none(self): """Test getting a non-existent document segment.""" @@ -637,19 +623,17 @@ class TestDatasetDocumentStoreSegment: mock_dataset = MagicMock(spec=Dataset) mock_dataset.id = "test-dataset-id" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + ) - result = store.get_document_segment("nonexistent") + result = store.get_document_segment("nonexistent", session=mock_session) - assert result is None + assert result is None class TestDatasetDocumentStoreMultimodelBinding: @@ -665,19 +649,17 @@ class TestDatasetDocumentStoreMultimodelBinding: mock_attachment = MagicMock(spec=AttachmentDocument) mock_attachment.metadata = {"doc_id": "attachment-1"} - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session + mock_session = MagicMock() - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_multimodel_documents_binding("seg-1", [mock_attachment]) + store.add_multimodel_documents_binding("seg-1", [mock_attachment], session=mock_session) - mock_db.session.add.assert_called() + mock_session.add.assert_called() def test_add_multimodel_documents_binding_without_attachments(self): """Test adding bindings with None attachments.""" @@ -686,19 +668,17 @@ class TestDatasetDocumentStoreMultimodelBinding: mock_dataset.id = "test-dataset-id" mock_dataset.tenant_id = "tenant-1" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session + mock_session = MagicMock() - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_multimodel_documents_binding("seg-1", None) + store.add_multimodel_documents_binding("seg-1", None, session=mock_session) - mock_db.session.add.assert_not_called() + mock_session.add.assert_not_called() def test_add_multimodel_documents_binding_with_empty_list(self): """Test adding bindings with empty list.""" @@ -707,19 +687,17 @@ class TestDatasetDocumentStoreMultimodelBinding: mock_dataset.id = "test-dataset-id" mock_dataset.tenant_id = "tenant-1" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session + mock_session = MagicMock() - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_multimodel_documents_binding("seg-1", []) + store.add_multimodel_documents_binding("seg-1", [], session=mock_session) - mock_db.session.add.assert_not_called() + mock_session.add.assert_not_called() def test_add_multimodel_documents_binding_with_none_document_id(self): """Test that no bindings are added when document_id is None.""" @@ -731,19 +709,17 @@ class TestDatasetDocumentStoreMultimodelBinding: mock_attachment = MagicMock(spec=AttachmentDocument) mock_attachment.metadata = {"doc_id": "attachment-1"} - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session + mock_session = MagicMock() - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id=None, - ) + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id=None, + ) - store.add_multimodel_documents_binding("seg-1", [mock_attachment]) + store.add_multimodel_documents_binding("seg-1", [mock_attachment], session=mock_session) - mock_db.session.add.assert_not_called() + mock_session.add.assert_not_called() class TestDatasetDocumentStoreAddDocumentsUpdateChild: @@ -776,23 +752,21 @@ class TestDatasetDocumentStoreAddDocumentsUpdateChild: mock_existing_segment = MagicMock() mock_existing_segment.id = "seg-1" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = 5 + mock_session = MagicMock() + mock_session.scalar.return_value = 5 - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): - with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): + with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_documents([mock_doc], save_child=True) + store.add_documents([mock_doc], session=mock_session, save_child=True) - mock_db.session.execute.assert_called() - mock_db.session.commit.assert_called() + mock_session.execute.assert_called() + mock_session.flush.assert_called() class TestDatasetDocumentStoreAddDocumentsUpdateAnswer: @@ -819,19 +793,17 @@ class TestDatasetDocumentStoreAddDocumentsUpdateAnswer: mock_existing_segment = MagicMock() mock_existing_segment.id = "seg-1" - with patch("core.rag.docstore.dataset_docstore.db") as mock_db: - mock_session = MagicMock() - mock_db.session = mock_session - mock_db.session.scalar.return_value = 5 + mock_session = MagicMock() + mock_session.scalar.return_value = 5 - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): - with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", - ) + with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_existing_segment): + with patch.object(DatasetDocumentStore, "add_multimodel_documents_binding"): + store = DatasetDocumentStore( + dataset=mock_dataset, + user_id="test-user-id", + document_id="test-doc-id", + ) - store.add_documents([mock_doc]) + store.add_documents([mock_doc], session=mock_session) - mock_db.session.commit.assert_called() + mock_session.flush.assert_called() diff --git a/api/tests/unit_tests/core/rag/extractor/test_extract_processor.py b/api/tests/unit_tests/core/rag/extractor/test_extract_processor.py index 5cf3d20f22f..19d16e1d552 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_extract_processor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_extract_processor.py @@ -126,7 +126,12 @@ class TestExtractProcessorFileRouting: monkeypatch.setattr(processor_module.dify_config, "UNSTRUCTURED_API_KEY", "key") def _run_extract_for_extension( - self, monkeypatch: pytest.MonkeyPatch, extension: str, etl_type: str, is_automatic: bool = False + self, + monkeypatch: pytest.MonkeyPatch, + extension: str, + etl_type: str, + is_automatic: bool = False, + session: object | None = None, ): factory = _patch_all_extractors(monkeypatch) monkeypatch.setattr(processor_module.dify_config, "ETL_TYPE", etl_type) @@ -147,7 +152,7 @@ class TestExtractProcessorFileRouting: ), ) - docs = ExtractProcessor.extract(setting, is_automatic=is_automatic) + docs = ExtractProcessor.extract(setting, is_automatic=is_automatic, session=session) assert len(docs) == 1 assert docs[0].page_content.startswith("extracted-by-") @@ -212,6 +217,16 @@ class TestExtractProcessorFileRouting: assert args[1:] == ("tenant-1", "user-1", "upload-file-1") assert kwargs == {} + @pytest.mark.parametrize("extension", [".pdf", ".docx"]) + def test_extract_passes_session_to_database_backed_file_extractors( + self, monkeypatch: pytest.MonkeyPatch, extension: str + ): + session = object() + + _, _, kwargs = self._run_extract_for_extension(monkeypatch, extension, etl_type="SelfHosted", session=session) + + assert kwargs["session"] is session + def test_extract_requires_upload_file_when_file_path_not_provided(self): setting = SimpleNamespace(datasource_type=DatasourceType.FILE, upload_file=None) diff --git a/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py index f2caf02d5e2..c41a0752b48 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py @@ -61,8 +61,15 @@ def mock_dependencies(monkeypatch: pytest.MonkeyPatch): (b"\x89PNG\r\n\x1a\n some png", "image/png", "png", "test_file_id_png"), ], ) +@pytest.mark.parametrize("inject_session", [False, True]) def test_extract_images_formats( - mock_dependencies, monkeypatch: pytest.MonkeyPatch, image_bytes, expected_mime, expected_ext, file_id + mock_dependencies, + monkeypatch: pytest.MonkeyPatch, + image_bytes, + expected_mime, + expected_ext, + file_id, + inject_session: bool, ): saves = mock_dependencies.saves db_stub = mock_dependencies.db @@ -82,7 +89,12 @@ def test_extract_images_formats( mock_page.get_objects.return_value = [mock_image_obj] - extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id="t1", user_id="u1") + extractor = pe.PdfExtractor( + file_path="test.pdf", + tenant_id="t1", + user_id="u1", + session=db_stub.session if inject_session else None, + ) # We need to handle the import inside _extract_images with patch("pypdfium2.raw", autospec=True) as mock_raw: @@ -97,7 +109,7 @@ def test_extract_images_formats( assert db_stub.session.added[0].size == len(image_bytes) assert db_stub.session.added[0].mime_type == expected_mime assert db_stub.session.added[0].extension == expected_ext - assert db_stub.session.committed is True + assert db_stub.session.committed is not inject_session @pytest.mark.parametrize( diff --git a/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py index 51763ec5f36..00c27504887 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py @@ -111,7 +111,8 @@ def test_init_downloads_via_remote_fetcher(monkeypatch: pytest.MonkeyPatch): extractor.temp_file.close() -def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("inject_session", [False, True]) +def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_session: bool): external_bytes = b"ext-bytes" internal_bytes = b"int-bytes" @@ -129,8 +130,8 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch): self.added = [] self.committed = False - def add(self, obj): - self.added.append(obj) + def add_all(self, objects): + self.added.extend(objects) def commit(self): self.committed = True @@ -177,6 +178,7 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch): extractor = object.__new__(WordExtractor) extractor.tenant_id = "t1" extractor.user_id = "u1" + extractor._session = db_stub.session if inject_session else None image_map = extractor._extract_images_from_docx(doc) @@ -191,7 +193,51 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch): # DB interactions should be recorded assert len(db_stub.session.added) == 2 - assert db_stub.session.committed is True + assert db_stub.session.committed is not inject_session + + +def test_extract_images_does_not_stage_partial_files_on_storage_failure(monkeypatch: pytest.MonkeyPatch): + class HashablePart: + def __init__(self, blob: bytes): + self.blob = blob + + def __hash__(self) -> int: + return id(self) + + first_part = HashablePart(b"first") + second_part = HashablePart(b"second") + doc = SimpleNamespace( + part=SimpleNamespace( + rels={ + "rId1": SimpleNamespace( + is_external=False, + target_ref="word/media/image1.png", + target_part=first_part, + ), + "rId2": SimpleNamespace( + is_external=False, + target_ref="word/media/image2.png", + target_part=second_part, + ), + } + ) + ) + session = MagicMock() + save = MagicMock(side_effect=[None, RuntimeError("storage failure")]) + monkeypatch.setattr(we, "storage", SimpleNamespace(save=save)) + monkeypatch.setattr(we.dify_config, "FILES_URL", "http://files.local", raising=False) + monkeypatch.setattr(we.dify_config, "STORAGE_TYPE", "local", raising=False) + + extractor = object.__new__(WordExtractor) + extractor.tenant_id = "tenant" + extractor.user_id = "user" + extractor._session = session + + with pytest.raises(RuntimeError, match="storage failure"): + extractor._extract_images_from_docx(doc) + + session.add_all.assert_not_called() + session.commit.assert_not_called() def test_extract_images_from_docx_uses_internal_files_url(): @@ -453,6 +499,7 @@ def test_extract_images_handles_invalid_external_cases(monkeypatch: pytest.Monke extractor = object.__new__(WordExtractor) extractor.tenant_id = "tenant" extractor.user_id = "user" + extractor._session = None result = extractor._extract_images_from_docx(doc) 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 f5761b5ba3d..8ca844d0da0 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 @@ -1,4 +1,5 @@ import logging +from contextlib import nullcontext from types import SimpleNamespace from typing import Any from unittest.mock import Mock, patch @@ -55,28 +56,33 @@ class TestParagraphIndexProcessor: def test_extract_forwards_automatic_flag(self, processor: ParagraphIndexProcessor) -> None: extract_setting = Mock() + session = Mock() expected_docs = [Document(page_content="chunk", metadata={})] with patch( "core.rag.index_processor.processor.paragraph_index_processor.ExtractProcessor.extract" ) as mock_extract: mock_extract.return_value = expected_docs - docs = processor.extract(extract_setting, process_rule_mode="hierarchical") + docs = processor.extract(extract_setting, process_rule_mode="hierarchical", session=session) assert docs == expected_docs - mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True) + mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) def test_transform_validates_process_rule(self, processor: ParagraphIndexProcessor) -> None: + session = Mock() with pytest.raises(ValueError, match="No process rule found"): - processor.transform([Document(page_content="text", metadata={})], process_rule=None) + processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) with pytest.raises(ValueError, match="No rules found in process rule"): - processor.transform([Document(page_content="text", metadata={})], process_rule={"mode": "custom"}) + processor.transform( + [Document(page_content="text", metadata={})], process_rule={"mode": "custom"}, session=session + ) def test_transform_validates_segmentation( self, processor: ParagraphIndexProcessor, process_rule: dict[str, Any] ) -> None: rules_without_segmentation = SimpleNamespace(segmentation=None) + session = Mock() with patch( "core.rag.index_processor.processor.paragraph_index_processor.Rule.model_validate", @@ -86,12 +92,14 @@ class TestParagraphIndexProcessor: processor.transform( [Document(page_content="text", metadata={})], process_rule={"mode": "custom", "rules": {"enabled": True}}, + session=session, ) def test_transform_builds_split_documents( self, processor: ParagraphIndexProcessor, process_rule: dict[str, Any] ) -> None: source_document = Document(page_content="source", metadata={"dataset_id": "dataset-1", "document_id": "doc-1"}) + session = Mock() splitter = Mock() splitter.split_documents.return_value = [ Document(page_content=".first", metadata={}), @@ -120,7 +128,7 @@ class TestParagraphIndexProcessor: processor, "_get_content_files", return_value=[AttachmentDocument(page_content="image", metadata={})] ), ): - documents = processor.transform([source_document], process_rule=process_rule) + documents = processor.transform([source_document], process_rule=process_rule, session=session) assert len(documents) == 1 assert documents[0].page_content == "first" @@ -130,6 +138,7 @@ class TestParagraphIndexProcessor: def test_transform_automatic_mode_uses_default_rules(self, processor: ParagraphIndexProcessor) -> None: splitter = Mock() splitter.split_documents.return_value = [Document(page_content="text", metadata={})] + session = Mock() with ( patch( @@ -151,7 +160,11 @@ class TestParagraphIndexProcessor: ), patch.object(processor, "_get_content_files", return_value=[]), ): - processor.transform([Document(page_content="text", metadata={})], process_rule={"mode": "automatic"}) + processor.transform( + [Document(page_content="text", metadata={})], + process_rule={"mode": "automatic"}, + session=session, + ) assert mock_validate.call_count == 1 @@ -160,12 +173,14 @@ class TestParagraphIndexProcessor: ) -> None: docs = [Document(page_content="chunk", metadata={})] multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] + session = Mock() with ( patch("core.rag.index_processor.processor.paragraph_index_processor.Vector") as mock_vector_cls, patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls, ): - processor.load(dataset, docs, multimodal_documents=multimodal_docs) + processor.load(dataset, docs, multimodal_documents=multimodal_docs, session=session) + mock_vector_cls.assert_called_once_with(dataset, session=session) vector = mock_vector_cls.return_value vector.create.assert_called_once_with(docs) vector.create_multimodal.assert_called_once_with(multimodal_docs) @@ -176,22 +191,25 @@ class TestParagraphIndexProcessor: ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="chunk", metadata={})] + session = Mock() + keywords_list = [["k1"], ["k2"]] with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: - processor.load(dataset, docs, keywords_list=["k1", "k2"]) + processor.load(dataset, docs, keywords_list=keywords_list, session=session) - mock_keyword_cls.return_value.add_texts.assert_called_once_with(docs, keywords_list=["k1", "k2"]) + mock_keyword_cls.return_value.add_texts.assert_called_once_with(docs, session, keywords_list=keywords_list) def test_load_uses_keyword_add_texts_without_keywords_when_economy( self, processor: ParagraphIndexProcessor, dataset: Mock ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="chunk", metadata={})] + session = Mock() with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: - processor.load(dataset, docs) + processor.load(dataset, docs, session=session) - mock_keyword_cls.return_value.add_texts.assert_called_once_with(docs) + mock_keyword_cls.return_value.add_texts.assert_called_once_with(docs, session) def test_clean_deletes_summaries_and_vector(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: scalars_result = Mock() @@ -200,22 +218,22 @@ class TestParagraphIndexProcessor: session.scalars.return_value = scalars_result with ( - patch("core.rag.index_processor.processor.paragraph_index_processor.db.session", session), patch( "core.rag.index_processor.processor.paragraph_index_processor.SummaryIndexService.delete_summaries_for_segments" ) as mock_summary, patch("core.rag.index_processor.processor.paragraph_index_processor.Vector") as mock_vector_cls, ): vector = mock_vector_cls.return_value - processor.clean(dataset, ["node-1"], delete_summaries=True) + processor.clean(dataset, ["node-1"], delete_summaries=True, session=session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) + mock_summary.assert_called_once_with(dataset, ["seg-1"], session=session) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_economy_deletes_summaries_and_keywords( self, processor: ParagraphIndexProcessor, dataset: Mock ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY + session = Mock() with ( patch( @@ -223,21 +241,25 @@ class TestParagraphIndexProcessor: ) as mock_summary, patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls, ): - processor.clean(dataset, None, delete_summaries=True) + processor.clean(dataset, None, delete_summaries=True, session=session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) + mock_summary.assert_called_once_with(dataset, None, session=session) mock_keyword_cls.return_value.delete.assert_called_once() def test_clean_deletes_keywords_by_ids(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY + session = Mock() with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: - processor.clean(dataset, ["node-2"], with_keywords=True) + processor.clean(dataset, ["node-2"], with_keywords=True, session=session) - mock_keyword_cls.return_value.delete_by_ids.assert_called_once_with(["node-2"]) + mock_keyword_cls.return_value.delete_by_ids.assert_called_once_with(["node-2"], session) def test_index_list_chunks_high_quality( self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: + session = Mock() + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") with ( patch( "core.rag.index_processor.processor.paragraph_index_processor.helper.generate_text_hash", @@ -251,9 +273,13 @@ class TestParagraphIndexProcessor: ) as mock_store_cls, patch("core.rag.index_processor.processor.paragraph_index_processor.Vector") as mock_vector_cls, ): - processor.index(dataset, dataset_document, ["chunk-1", "chunk-2"]) + mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store") + mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector") + processor.index(dataset, dataset_document, ["chunk-1", "chunk-2"], session) + assert phase_events == ["store", "commit", "vector"] mock_store_cls.return_value.add_documents.assert_called_once() + mock_vector_cls.assert_called_once_with(dataset, session=session) mock_vector_cls.return_value.create.assert_called_once() mock_vector_cls.return_value.create_multimodal.assert_called_once() @@ -261,17 +287,25 @@ class TestParagraphIndexProcessor: self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY + session = Mock() + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") with ( patch( "core.rag.index_processor.processor.paragraph_index_processor.helper.generate_text_hash", return_value="hash", ), patch.object(processor, "_get_content_files", return_value=[]), - patch("core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore"), + patch( + "core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore" + ) as mock_store_cls, patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls, ): - processor.index(dataset, dataset_document, ["chunk-3"]) + mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store") + mock_keyword_cls.return_value.add_texts.side_effect = lambda *_args: phase_events.append("keyword") + processor.index(dataset, dataset_document, ["chunk-3"], session) + assert phase_events == ["store", "commit", "keyword"] mock_keyword_cls.return_value.add_texts.assert_called_once() def test_index_multimodal_structure_handles_files_and_account_lookup( @@ -283,6 +317,8 @@ class TestParagraphIndexProcessor: ) chunk_without_files = SimpleNamespace(content="content-2", files=None) structure = SimpleNamespace(general_chunks=[chunk_with_files, chunk_without_files]) + session = Mock() + account_session = Mock() with ( patch( @@ -296,6 +332,10 @@ class TestParagraphIndexProcessor: patch( "core.rag.index_processor.processor.paragraph_index_processor.AccountService.load_user", return_value=SimpleNamespace(id="user-1"), + ) as load_user, + patch( + "core.rag.index_processor.processor.paragraph_index_processor.session_factory.create_session", + return_value=nullcontext(account_session), ), patch.object( processor, "_get_content_files", return_value=[AttachmentDocument(page_content="img", metadata={})] @@ -303,14 +343,17 @@ class TestParagraphIndexProcessor: patch("core.rag.index_processor.processor.paragraph_index_processor.DatasetDocumentStore"), patch("core.rag.index_processor.processor.paragraph_index_processor.Vector"), ): - processor.index(dataset, dataset_document, {"general_chunks": []}) + processor.index(dataset, dataset_document, {"general_chunks": []}, session) assert mock_files.call_count == 1 + load_user.assert_called_once_with(dataset_document.created_by, account_session) + assert account_session is not session def test_index_multimodal_structure_requires_valid_account( self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: structure = SimpleNamespace(general_chunks=[SimpleNamespace(content="content", files=None)]) + session = Mock() with ( patch( @@ -327,7 +370,7 @@ class TestParagraphIndexProcessor: ), ): with pytest.raises(ValueError, match="Invalid account"): - processor.index(dataset, dataset_document, {"general_chunks": []}) + processor.index(dataset, dataset_document, {"general_chunks": []}, session) def test_format_preview_validates_chunk_shape(self, processor: ParagraphIndexProcessor) -> None: preview = processor.format_preview(["chunk-1", "chunk-2"]) @@ -339,16 +382,34 @@ class TestParagraphIndexProcessor: def test_generate_summary_preview_success_and_failure(self, processor: ParagraphIndexProcessor) -> None: preview_items = [PreviewDetail(content="chunk-1"), PreviewDetail(content="chunk-2")] + session = Mock() + worker_sessions = [Mock(), Mock()] - with patch.object(processor, "generate_summary", return_value=("summary", LLMUsage.empty_usage())): + with ( + patch( + "core.rag.index_processor.processor.paragraph_index_processor.session_factory.create_session", + side_effect=[nullcontext(worker_session) for worker_session in worker_sessions], + ) as create_session, + patch.object( + processor, "generate_summary", return_value=("summary", LLMUsage.empty_usage()) + ) as mock_generate_summary, + ): result = processor.generate_summary_preview( - "tenant-1", preview_items, {"enable": True}, doc_language="English" + "tenant-1", preview_items, {"enable": True}, doc_language="English", session=session ) assert all(item.summary == "summary" for item in result) + call_sessions = [call.kwargs["session"] for call in mock_generate_summary.call_args_list] + assert create_session.call_count == len(preview_items) + assert all(call_session is not session for call_session in call_sessions) + assert {id(call_session) for call_session in call_sessions} == { + id(worker_session) for worker_session in worker_sessions + } with patch.object(processor, "generate_summary", side_effect=RuntimeError("summary failed")): with pytest.raises(ValueError, match="Failed to generate summaries"): - processor.generate_summary_preview("tenant-1", [PreviewDetail(content="chunk-1")], {"enable": True}) + processor.generate_summary_preview( + "tenant-1", [PreviewDetail(content="chunk-1")], {"enable": True}, session=session + ) def test_generate_summary_preview_fallback_without_flask_context(self, processor: ParagraphIndexProcessor) -> None: preview_items = [PreviewDetail(content="chunk-1")] @@ -358,7 +419,7 @@ class TestParagraphIndexProcessor: patch("flask.current_app", fake_current_app), patch.object(processor, "generate_summary", return_value=("summary", LLMUsage.empty_usage())), ): - result = processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}) + result = processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}, session=Mock()) assert result[0].summary == "summary" @@ -374,16 +435,16 @@ class TestParagraphIndexProcessor: patch("concurrent.futures.wait", side_effect=[(set(), {future}), (set(), set())]), ): with pytest.raises(ValueError, match="timeout"): - processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}) + processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}, session=Mock()) future.cancel.assert_called_once() def test_generate_summary_validates_input(self) -> None: with pytest.raises(ValueError, match="must be enabled"): - ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": False}) + ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": False}, session=Mock()) with pytest.raises(ValueError, match="model_name and model_provider_name"): - ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": True}) + ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": True}, session=Mock()) def test_generate_summary_text_only_flow(self, caplog: pytest.LogCaptureFixture) -> None: model_instance = Mock() @@ -413,6 +474,7 @@ class TestParagraphIndexProcessor: "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, document_language="English", + session=Mock(), ) assert summary == "text summary" @@ -454,6 +516,7 @@ class TestParagraphIndexProcessor: "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, segment_id="seg-1", + session=Mock(), ) assert summary == "vision summary" @@ -496,6 +559,7 @@ class TestParagraphIndexProcessor: "tenant-1", "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, + session=Mock(), ) assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) == 1 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 672764e5336..d2c9670b795 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 @@ -1,3 +1,4 @@ +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch @@ -49,23 +50,27 @@ class TestParentChildIndexProcessor: def test_extract_forwards_automatic_flag(self, processor: ParentChildIndexProcessor) -> None: extract_setting = Mock() + session = Mock() expected = [Document(page_content="chunk", metadata={})] with patch( "core.rag.index_processor.processor.parent_child_index_processor.ExtractProcessor.extract" ) as mock_extract: mock_extract.return_value = expected - documents = processor.extract(extract_setting, process_rule_mode="hierarchical") + documents = processor.extract(extract_setting, process_rule_mode="hierarchical", session=session) assert documents == expected - mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True) + mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) def test_transform_validates_process_rule(self, processor: ParentChildIndexProcessor) -> None: + session = MagicMock() with pytest.raises(ValueError, match="No process rule found"): - processor.transform([Document(page_content="text", metadata={})], process_rule=None) + processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) with pytest.raises(ValueError, match="No rules found in process rule"): - processor.transform([Document(page_content="text", metadata={})], process_rule={"mode": "custom"}) + processor.transform( + [Document(page_content="text", metadata={})], process_rule={"mode": "custom"}, session=session + ) def test_transform_paragraph_requires_segmentation(self, processor: ParentChildIndexProcessor) -> None: rules = SimpleNamespace(parent_mode=ParentMode.PARAGRAPH, segmentation=None) @@ -77,6 +82,7 @@ class TestParentChildIndexProcessor: processor.transform( [Document(page_content="text", metadata={})], process_rule={"mode": "custom", "rules": {"enabled": True}}, + session=MagicMock(), ) def test_transform_paragraph_builds_parent_and_child_docs(self, processor: ParentChildIndexProcessor) -> None: @@ -111,6 +117,7 @@ class TestParentChildIndexProcessor: [parent_document], process_rule={"mode": "custom", "rules": {"enabled": True}}, preview=False, + session=MagicMock(), ) assert len(result) == 1 @@ -147,6 +154,7 @@ class TestParentChildIndexProcessor: documents, process_rule={"mode": "custom", "rules": {"enabled": True}}, preview=True, + session=MagicMock(), ) assert len(result) == 10 @@ -180,6 +188,7 @@ class TestParentChildIndexProcessor: docs, process_rule={"mode": "hierarchical", "rules": {"enabled": True}}, preview=True, + session=MagicMock(), ) assert len(result) == 1 @@ -196,11 +205,13 @@ class TestParentChildIndexProcessor: ], ) multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] + session = MagicMock() with patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls: vector = mock_vector_cls.return_value - processor.load(dataset, [parent_doc], multimodal_documents=multimodal_docs) + processor.load(dataset, [parent_doc], multimodal_documents=multimodal_docs, session=session) + mock_vector_cls.assert_called_once_with(dataset, session=session) assert vector.create.call_count == 1 formatted_docs = vector.create.call_args[0][0] assert len(formatted_docs) == 2 @@ -208,11 +219,10 @@ class TestParentChildIndexProcessor: vector.create_multimodal.assert_called_once_with(multimodal_docs) def test_clean_with_precomputed_child_ids(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: - session = Mock() + session = MagicMock() with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, - patch("core.rag.index_processor.processor.parent_child_index_processor.db.session", session), ): vector = mock_vector_cls.return_value processor.clean( @@ -220,65 +230,60 @@ class TestParentChildIndexProcessor: ["node-1"], delete_child_chunks=True, precomputed_child_node_ids=["child-1", "child-2"], + session=session, ) vector.delete_by_ids.assert_called_once_with(["child-1", "child-2"]) session.execute.assert_called() - session.commit.assert_called_once() + session.flush.assert_called_once() def test_clean_queries_child_ids_when_not_precomputed( self, processor: ParentChildIndexProcessor, dataset: Mock ) -> None: execute_result = Mock() execute_result.all.return_value = [("child-1",), (None,), ("child-2",)] - session = Mock() + session = MagicMock() session.execute.return_value = execute_result with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, - patch("core.rag.index_processor.processor.parent_child_index_processor.db.session", session), ): vector = mock_vector_cls.return_value - processor.clean(dataset, ["node-1"], delete_child_chunks=False) + processor.clean(dataset, ["node-1"], delete_child_chunks=False, session=session) vector.delete_by_ids.assert_called_once_with(["child-1", "child-2"]) def test_clean_dataset_wide_cleanup(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: - session = Mock() + session = MagicMock() with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, - patch("core.rag.index_processor.processor.parent_child_index_processor.db.session", session), ): vector = mock_vector_cls.return_value - processor.clean(dataset, None, delete_child_chunks=True) + processor.clean(dataset, None, delete_child_chunks=True, session=session) vector.delete.assert_called_once() session.execute.assert_called() - session.commit.assert_called_once() + session.flush.assert_called_once() def test_clean_deletes_summaries_when_requested(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: scalars_result = Mock() scalars_result.all.return_value = [SimpleNamespace(id="seg-1")] - session = Mock() + session = MagicMock() session.scalars.return_value = scalars_result session_ctx = MagicMock() session_ctx.__enter__.return_value = session session_ctx.__exit__.return_value = False with ( - patch( - "core.rag.index_processor.processor.parent_child_index_processor.session_factory.create_session", - return_value=session_ctx, - ), patch( "core.rag.index_processor.processor.parent_child_index_processor.SummaryIndexService.delete_summaries_for_segments" ) as mock_summary, patch("core.rag.index_processor.processor.parent_child_index_processor.Vector"), ): - processor.clean(dataset, ["node-1"], delete_summaries=True, precomputed_child_node_ids=[]) + processor.clean(dataset, ["node-1"], delete_summaries=True, precomputed_child_node_ids=[], session=session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) + mock_summary.assert_called_once_with(dataset, ["seg-1"], session=session) def test_clean_deletes_all_summaries_when_node_ids_missing( self, processor: ParentChildIndexProcessor, dataset: Mock @@ -289,9 +294,10 @@ class TestParentChildIndexProcessor: ) as mock_summary, patch("core.rag.index_processor.processor.parent_child_index_processor.Vector"), ): - processor.clean(dataset, None, delete_summaries=True) + session = MagicMock() + processor.clean(dataset, None, delete_summaries=True, session=session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) + mock_summary.assert_called_once_with(dataset, None, session=session) def test_split_child_nodes_requires_subchunk_segmentation(self, processor: ParentChildIndexProcessor) -> None: rules = Rule(subchunk_segmentation=None) @@ -336,7 +342,9 @@ class TestParentChildIndexProcessor: ], ) dataset_rule = SimpleNamespace(id="rule-1") - session = Mock() + session = MagicMock() + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") with ( patch( @@ -355,15 +363,17 @@ class TestParentChildIndexProcessor: "core.rag.index_processor.processor.parent_child_index_processor.DatasetDocumentStore" ) as mock_store_cls, patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, - patch("core.rag.index_processor.processor.parent_child_index_processor.db.session", session), ): - processor.index(dataset, dataset_document, {"parent_child_chunks": []}) + mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store") + mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector") + processor.index(dataset, dataset_document, {"parent_child_chunks": []}, session) + assert phase_events == ["store", "commit", "vector"] assert dataset_document.dataset_process_rule_id == "rule-1" session.add.assert_called_once_with(dataset_rule) session.flush.assert_called_once() - session.commit.assert_called_once() mock_store_cls.return_value.add_documents.assert_called_once() + mock_vector_cls.assert_called_once_with(dataset, session=session) assert mock_vector_cls.return_value.create.call_count == 1 mock_vector_cls.return_value.create_multimodal.assert_called_once() @@ -375,7 +385,8 @@ class TestParentChildIndexProcessor: parent_child_chunks=[SimpleNamespace(parent_content="parent", child_contents=["child"], files=None)], ) dataset_rule = SimpleNamespace(id="rule-1") - session = Mock() + session = MagicMock() + account_session = MagicMock() with ( patch( @@ -393,17 +404,22 @@ class TestParentChildIndexProcessor: patch( "core.rag.index_processor.processor.parent_child_index_processor.AccountService.load_user", return_value=SimpleNamespace(id="user-1"), + ) as load_user, + patch( + "core.rag.index_processor.processor.parent_child_index_processor.session_factory.create_session", + return_value=nullcontext(account_session), ), patch.object( processor, "_get_content_files", return_value=[AttachmentDocument(page_content="image", metadata={})] ) as mock_files, patch("core.rag.index_processor.processor.parent_child_index_processor.DatasetDocumentStore"), patch("core.rag.index_processor.processor.parent_child_index_processor.Vector"), - patch("core.rag.index_processor.processor.parent_child_index_processor.db.session", session), ): - processor.index(dataset, dataset_document, {"parent_child_chunks": []}) + processor.index(dataset, dataset_document, {"parent_child_chunks": []}, session) mock_files.assert_called_once() + load_user.assert_called_once_with(dataset_document.created_by, account_session) + assert account_session is not session def test_index_raises_when_account_missing( self, processor: ParentChildIndexProcessor, dataset: Mock, dataset_document: Mock @@ -428,7 +444,7 @@ class TestParentChildIndexProcessor: ), ): with pytest.raises(ValueError, match="Invalid account"): - processor.index(dataset, dataset_document, {"parent_child_chunks": []}) + processor.index(dataset, dataset_document, {"parent_child_chunks": []}, MagicMock()) def test_format_preview_returns_parent_child_structure(self, processor: ParentChildIndexProcessor) -> None: parent_childs = SimpleNamespace( @@ -448,16 +464,30 @@ class TestParentChildIndexProcessor: def test_generate_summary_preview_sets_summaries(self, processor: ParentChildIndexProcessor) -> None: preview_texts = [PreviewDetail(content="chunk-1"), PreviewDetail(content="chunk-2")] + session = MagicMock() + worker_sessions = [MagicMock(), MagicMock()] - with patch( - "core.rag.index_processor.processor.paragraph_index_processor.ParagraphIndexProcessor.generate_summary", - return_value=("summary", None), + with ( + patch( + "core.rag.index_processor.processor.parent_child_index_processor.session_factory.create_session", + side_effect=[nullcontext(worker_session) for worker_session in worker_sessions], + ) as create_session, + patch( + "core.rag.index_processor.processor.paragraph_index_processor.ParagraphIndexProcessor.generate_summary", + return_value=("summary", None), + ) as mock_generate_summary, ): result = processor.generate_summary_preview( - "tenant-1", preview_texts, {"enable": True}, doc_language="English" + "tenant-1", preview_texts, {"enable": True}, doc_language="English", session=session ) assert all(item.summary == "summary" for item in result) + call_sessions = [call.kwargs["session"] for call in mock_generate_summary.call_args_list] + assert create_session.call_count == len(preview_texts) + assert all(call_session is not session for call_session in call_sessions) + assert {id(call_session) for call_session in call_sessions} == { + id(worker_session) for worker_session in worker_sessions + } def test_generate_summary_preview_raises_when_worker_fails(self, processor: ParentChildIndexProcessor) -> None: preview_texts = [PreviewDetail(content="chunk-1")] @@ -467,7 +497,7 @@ class TestParentChildIndexProcessor: side_effect=RuntimeError("summary failed"), ): with pytest.raises(ValueError, match="Failed to generate summaries"): - processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}) + processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=MagicMock()) def test_generate_summary_preview_falls_back_without_flask_context( self, processor: ParentChildIndexProcessor @@ -482,7 +512,9 @@ class TestParentChildIndexProcessor: return_value=("summary", None), ), ): - result = processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}) + result = processor.generate_summary_preview( + "tenant-1", preview_texts, {"enable": True}, session=MagicMock() + ) assert result[0].summary == "summary" @@ -498,6 +530,6 @@ class TestParentChildIndexProcessor: patch("concurrent.futures.wait", side_effect=[(set(), {future}), (set(), set())]), ): with pytest.raises(ValueError, match="timeout"): - processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}) + processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=MagicMock()) future.cancel.assert_called_once() 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 5dde1623d2d..ad7b843f89e 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 @@ -60,27 +60,33 @@ class TestQAIndexProcessor: def test_extract_forwards_automatic_flag(self, processor: QAIndexProcessor) -> None: extract_setting = Mock() + session = Mock() expected_docs = [Document(page_content="chunk", metadata={})] with patch("core.rag.index_processor.processor.qa_index_processor.ExtractProcessor.extract") as mock_extract: mock_extract.return_value = expected_docs - docs = processor.extract(extract_setting, process_rule_mode="automatic") + docs = processor.extract(extract_setting, process_rule_mode="automatic", session=session) assert docs == expected_docs - mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True) + mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) def test_transform_rejects_none_process_rule(self, processor: QAIndexProcessor) -> None: + session = MagicMock() with pytest.raises(ValueError, match="No process rule found"): - processor.transform([Document(page_content="text", metadata={})], process_rule=None) + processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) def test_transform_rejects_missing_rules_key(self, processor: QAIndexProcessor) -> None: + session = MagicMock() with pytest.raises(ValueError, match="No rules found in process rule"): - processor.transform([Document(page_content="text", metadata={})], process_rule={"mode": "custom"}) + processor.transform( + [Document(page_content="text", metadata={})], process_rule={"mode": "custom"}, session=session + ) def test_transform_preview_calls_formatter_once( self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app ) -> None: + session = MagicMock() document = Document(page_content="raw text", metadata={"dataset_id": "dataset-1", "document_id": "doc-1"}) split_node = Document(page_content=".question", metadata={}) splitter = Mock() @@ -114,6 +120,7 @@ class TestQAIndexProcessor: preview=True, tenant_id="tenant-1", doc_language="English", + session=session, ) assert len(result) == 1 @@ -123,6 +130,7 @@ class TestQAIndexProcessor: def test_transform_non_preview_uses_thread_batches( self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app ) -> None: + session = MagicMock() documents = [ Document(page_content="doc-1", metadata={"document_id": "doc-1", "dataset_id": "dataset-1"}), Document(page_content="doc-2", metadata={"document_id": "doc-2", "dataset_id": "dataset-1"}), @@ -157,7 +165,13 @@ class TestQAIndexProcessor: ), ): mock_current_app._get_current_object = Mock(return_value=fake_flask_app) - result = processor.transform(documents, process_rule=process_rule, preview=False, tenant_id="tenant-1") + result = processor.transform( + documents, + process_rule=process_rule, + preview=False, + tenant_id="tenant-1", + session=session, + ) assert len(result) == 2 assert mock_format.call_count == 2 @@ -199,22 +213,25 @@ class TestQAIndexProcessor: processor.format_by_template(csv_file) def test_load_creates_vectors_for_high_quality_dataset(self, processor: QAIndexProcessor, dataset: Mock) -> None: + session = MagicMock() docs = [Document(page_content="Q1", metadata={"answer": "A1"})] multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] with patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls: vector = mock_vector_cls.return_value - processor.load(dataset, docs, multimodal_documents=multimodal_docs) + processor.load(dataset, docs, multimodal_documents=multimodal_docs, session=session) + mock_vector_cls.assert_called_once_with(dataset, session=session) vector.create.assert_called_once_with(docs) vector.create_multimodal.assert_called_once_with(multimodal_docs) def test_load_skips_vector_for_non_high_quality(self, processor: QAIndexProcessor, dataset: Mock) -> None: + session = MagicMock() dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="Q1", metadata={"answer": "A1"})] with patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls: - processor.load(dataset, docs) + processor.load(dataset, docs, session=session) mock_vector_cls.assert_not_called() @@ -224,29 +241,23 @@ class TestQAIndexProcessor: mock_segment = SimpleNamespace(id="seg-1") scalars_result = Mock() scalars_result.all.return_value = [mock_segment] - mock_session = Mock() + mock_session = MagicMock() mock_session.scalars.return_value = scalars_result - session_context = MagicMock() - session_context.__enter__.return_value = mock_session - session_context.__exit__.return_value = False with ( - patch( - "core.rag.index_processor.processor.qa_index_processor.session_factory.create_session", - return_value=session_context, - ), patch( "core.rag.index_processor.processor.qa_index_processor.SummaryIndexService.delete_summaries_for_segments" ) as mock_summary, patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls, ): vector = mock_vector_cls.return_value - processor.clean(dataset, ["node-1"], delete_summaries=True) + processor.clean(dataset, ["node-1"], delete_summaries=True, session=mock_session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=["seg-1"]) + mock_summary.assert_called_once_with(dataset, ["seg-1"], session=mock_session) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_handles_dataset_wide_cleanup(self, processor: QAIndexProcessor, dataset: Mock) -> None: + session = MagicMock() with ( patch( "core.rag.index_processor.processor.qa_index_processor.SummaryIndexService.delete_summaries_for_segments" @@ -254,14 +265,17 @@ class TestQAIndexProcessor: patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls, ): vector = mock_vector_cls.return_value - processor.clean(dataset, None, delete_summaries=True) + processor.clean(dataset, None, delete_summaries=True, session=session) - mock_summary.assert_called_once_with(dataset=dataset, segment_ids=None) + mock_summary.assert_called_once_with(dataset, None, session=session) vector.delete.assert_called_once() def test_index_adds_documents_and_vectors_for_high_quality( self, processor: QAIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: + session = MagicMock() + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") qa_chunks = SimpleNamespace( qa_chunks=[ SimpleNamespace(question="Q1", answer="A1"), @@ -280,14 +294,18 @@ class TestQAIndexProcessor: patch("core.rag.index_processor.processor.qa_index_processor.DatasetDocumentStore") as mock_store_cls, patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls, ): - processor.index(dataset, dataset_document, {"qa_chunks": []}) + mock_store_cls.return_value.add_documents.side_effect = lambda **_kwargs: phase_events.append("store") + mock_vector_cls.return_value.create.side_effect = lambda _documents: phase_events.append("vector") + processor.index(dataset, dataset_document, {"qa_chunks": []}, session) + assert phase_events == ["store", "commit", "vector"] mock_store_cls.return_value.add_documents.assert_called_once() mock_vector_cls.return_value.create.assert_called_once() def test_index_requires_high_quality( self, processor: QAIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: + session = MagicMock() dataset.indexing_technique = IndexTechniqueType.ECONOMY qa_chunks = SimpleNamespace(qa_chunks=[SimpleNamespace(question="Q1", answer="A1")]) @@ -302,7 +320,7 @@ class TestQAIndexProcessor: patch("core.rag.index_processor.processor.qa_index_processor.DatasetDocumentStore"), ): with pytest.raises(ValueError, match="must be high quality"): - processor.index(dataset, dataset_document, {"qa_chunks": []}) + processor.index(dataset, dataset_document, {"qa_chunks": []}, session) def test_format_preview_returns_qa_preview(self, processor: QAIndexProcessor) -> None: qa_chunks = SimpleNamespace(qa_chunks=[SimpleNamespace(question="Q1", answer="A1")]) @@ -319,7 +337,10 @@ class TestQAIndexProcessor: def test_generate_summary_preview_returns_input(self, processor: QAIndexProcessor) -> None: preview_items = [PreviewDetail(content="Q1")] - assert processor.generate_summary_preview("tenant-1", preview_items, {"enable": False}) is preview_items + assert ( + processor.generate_summary_preview("tenant-1", preview_items, {"enable": False}, session=MagicMock()) + is preview_items + ) def test_format_qa_document_ignores_blank_text(self, processor: QAIndexProcessor, fake_flask_app) -> None: all_qa_documents: list[Document] = [] diff --git a/api/tests/unit_tests/core/rag/indexing/test_index_processor.py b/api/tests/unit_tests/core/rag/indexing/test_index_processor.py index a3f284955bc..59d0fb126d6 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/test_index_processor.py @@ -1,4 +1,11 @@ +import datetime +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.index_processor.index_processor import IndexProcessor +from core.workflow.nodes.knowledge_index.protocols import Preview, PreviewItem class TestIndexProcessor: @@ -13,3 +20,96 @@ class TestIndexProcessor: assert len(preview.qa_preview) == 1 assert preview.qa_preview[0].question == "Q1" assert preview.qa_preview[0].answer == "A1" + + def test_index_and_clean_ends_transactions_around_index_io(self) -> None: + document = SimpleNamespace( + id="document-1", + name="Document", + created_at=datetime.datetime(2026, 1, 1), + indexing_latency=None, + indexing_status=None, + completed_at=None, + word_count=0, + need_summary=False, + ) + dataset = SimpleNamespace( + id="dataset-1", + name="Dataset", + chunk_structure="text_model", + summary_index_setting=None, + ) + session = MagicMock() + session.scalar.side_effect = [document, dataset, 3] + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + + index_processor = MagicMock() + index_processor.index.side_effect = lambda *args: phase_events.append("index") + + with patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory: + index_processor_factory.return_value.init_index_processor.return_value = index_processor + IndexProcessor().index_and_clean( + dataset_id=dataset.id, + document_id=document.id, + original_document_id="", + chunks={"general_chunks": ["content"]}, + batch="batch-1", + session=session, + ) + + assert phase_events == ["commit", "index", "commit"] + + def test_preview_summary_workers_use_independent_sessions(self) -> None: + caller_session = MagicMock() + phase_events: list[str] = [] + caller_session.commit.side_effect = lambda: phase_events.append("commit") + caller_session.scalar.return_value = SimpleNamespace( + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + summary_index_setting={"enable": True}, + tenant_id="tenant-1", + ) + worker_sessions = [MagicMock(), MagicMock()] + preview = Preview( + chunk_structure="text_model", + total_segments=2, + preview=[PreviewItem(content="chunk-1"), PreviewItem(content="chunk-2")], + ) + flask_app = SimpleNamespace(app_context=lambda: nullcontext()) + processor = IndexProcessor() + worker_contexts = iter(nullcontext(worker_session) for worker_session in worker_sessions) + + def create_worker_session(): + phase_events.append("worker") + return next(worker_contexts) + + with ( + patch.object(processor, "format_preview", return_value=preview), + patch( + "core.rag.index_processor.index_processor.current_app", + SimpleNamespace(_get_current_object=lambda: flask_app), + ), + patch( + "core.rag.index_processor.index_processor.session_factory.create_session", + side_effect=create_worker_session, + ), + patch( + "core.rag.index_processor.index_processor.ParagraphIndexProcessor.generate_summary", + return_value=("summary", None), + ) as generate_summary, + ): + result = processor.get_preview_output( + chunks=[], + dataset_id="dataset-1", + document_id="", + chunk_structure="text_model", + summary_index_setting={"enable": True}, + session=caller_session, + ) + + assert all(item.summary == "summary" for item in result.preview) + assert phase_events == ["commit", "worker", "worker"] + call_sessions = [call.kwargs["session"] for call in generate_summary.call_args_list] + assert all(call_session is not caller_session for call_session in call_sessions) + assert all( + any(call_session is worker_session for worker_session in worker_sessions) for call_session in call_sessions + ) diff --git a/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py b/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py index fe1109db301..93c039e2ff5 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py +++ b/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py @@ -13,39 +13,50 @@ from core.rag.models.document import AttachmentDocument, Document class _ForwardingBaseIndexProcessor(BaseIndexProcessor): @override - def extract(self, extract_setting, **kwargs): - return super().extract(extract_setting, **kwargs) + def extract(self, extract_setting, *, session, **kwargs): + return super().extract(extract_setting, session=session, **kwargs) @override - def transform(self, documents, current_user=None, **kwargs): - return super().transform(documents, current_user=current_user, **kwargs) + def transform(self, documents, current_user=None, *, session, **kwargs): + return super().transform(documents, current_user=current_user, session=session, **kwargs) @override - def generate_summary_preview(self, tenant_id, preview_texts, summary_index_setting, doc_language=None): + def generate_summary_preview(self, tenant_id, preview_texts, summary_index_setting, doc_language=None, *, session): return super().generate_summary_preview( tenant_id=tenant_id, preview_texts=preview_texts, summary_index_setting=summary_index_setting, doc_language=doc_language, + session=session, ) @override - def load(self, dataset, documents, multimodal_documents=None, with_keywords=True, **kwargs): + def load( + self, + dataset, + documents, + multimodal_documents=None, + with_keywords=True, + *, + session, + **kwargs, + ): return super().load( dataset=dataset, documents=documents, multimodal_documents=multimodal_documents, with_keywords=with_keywords, + session=session, **kwargs, ) @override - def clean(self, dataset, node_ids, with_keywords=True, **kwargs): - return super().clean(dataset=dataset, node_ids=node_ids, with_keywords=with_keywords, **kwargs) + def clean(self, dataset, node_ids, with_keywords=True, *, session, **kwargs): + return super().clean(dataset=dataset, node_ids=node_ids, with_keywords=with_keywords, session=session, **kwargs) @override - def index(self, dataset, document, chunks): - return super().index(dataset=dataset, document=document, chunks=chunks) + def index(self, dataset, document, chunks, session): + return super().index(dataset=dataset, document=document, chunks=chunks, session=session) @override def format_preview(self, chunks): @@ -59,17 +70,19 @@ class TestBaseIndexProcessor: def test_abstract_methods_raise_not_implemented(self, processor: _ForwardingBaseIndexProcessor) -> None: with pytest.raises(NotImplementedError): - processor.extract(Mock()) + processor.extract(Mock(), session=Mock()) with pytest.raises(NotImplementedError): - processor.transform([]) + processor.transform([], session=Mock()) with pytest.raises(NotImplementedError): - processor.generate_summary_preview("tenant", [PreviewDetail(content="c")], {"enable": False}) + processor.generate_summary_preview( + "tenant", [PreviewDetail(content="c")], {"enable": False}, session=Mock() + ) with pytest.raises(NotImplementedError): - processor.load(Mock(), []) + processor.load(Mock(), [], session=Mock()) with pytest.raises(NotImplementedError): - processor.clean(Mock(), None) + processor.clean(Mock(), None, session=Mock()) with pytest.raises(NotImplementedError): - processor.index(Mock(), Mock(), {}) + processor.index(Mock(), Mock(), {}, Mock()) with pytest.raises(NotImplementedError): processor.format_preview([]) @@ -112,7 +125,7 @@ class TestBaseIndexProcessor: def test_get_content_files_without_images_returns_empty(self, processor: _ForwardingBaseIndexProcessor) -> None: document = Document(page_content="no image markdown", metadata={"document_id": "doc-1", "dataset_id": "ds-1"}) - assert processor._get_content_files(document) == [] + assert processor._get_content_files(document, session=Mock()) == [] def test_get_content_files_handles_all_sources_and_duplicates( self, processor: _ForwardingBaseIndexProcessor @@ -133,14 +146,14 @@ class TestBaseIndexProcessor: scalars_result.all.return_value = [upload_a, upload_b, upload_tool, upload_remote] db_session = Mock() db_session.scalars.return_value = scalars_result + current_user = Mock() with ( patch.object(processor, "_extract_markdown_images", return_value=images), patch.object(processor, "_download_tool_file", return_value="tool-upload-id") as mock_tool_download, patch.object(processor, "_download_image", return_value="remote-upload-id") as mock_image_download, - patch("core.rag.index_processor.index_processor_base.db.session", db_session), ): - files = processor._get_content_files(document, current_user=Mock()) + files = processor._get_content_files(document, current_user=current_user, session=db_session) assert len(files) == 5 assert all(isinstance(file, AttachmentDocument) for file in files) @@ -149,7 +162,11 @@ class TestBaseIndexProcessor: assert files[0].metadata["dataset_id"] == "ds-1" assert files[0].metadata["doc_id"] == "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" assert files[1].metadata["doc_id"] == "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" - mock_tool_download.assert_called_once() + mock_tool_download.assert_called_once_with( + "cccccccc-cccc-cccc-cccc-cccccccccccc", + current_user, + session=db_session, + ) mock_image_download.assert_called_once() def test_get_content_files_skips_tool_and_remote_download_without_user( @@ -159,7 +176,7 @@ class TestBaseIndexProcessor: images = ["/files/tools/cccccccc-cccc-cccc-cccc-cccccccccccc.png", "https://example.com/remote.png"] with patch.object(processor, "_extract_markdown_images", return_value=images): - files = processor._get_content_files(document, current_user=None) + files = processor._get_content_files(document, current_user=None, session=Mock()) assert files == [] @@ -173,9 +190,8 @@ class TestBaseIndexProcessor: with ( patch.object(processor, "_extract_markdown_images", return_value=images), - patch("core.rag.index_processor.index_processor_base.db.session", db_session), ): - files = processor._get_content_files(document) + files = processor._get_content_files(document, session=db_session) assert files == [] @@ -258,15 +274,13 @@ class TestBaseIndexProcessor: db_session = Mock() db_session.get.return_value = None - with patch("core.rag.index_processor.index_processor_base.db.session", db_session): - assert processor._download_tool_file("tool-id", current_user=Mock()) is None + assert processor._download_tool_file("tool-id", current_user=Mock(), session=db_session) is None def test_download_tool_file_uploads_file_when_found(self, processor: _ForwardingBaseIndexProcessor) -> None: tool_file = SimpleNamespace(file_key="k1", name="tool.png", mimetype="image/png") db_session = Mock() db_session.get.return_value = tool_file mock_db = Mock() - mock_db.session = db_session mock_db.engine = Mock() upload_result = SimpleNamespace(id="upload-id") @@ -276,7 +290,7 @@ class TestBaseIndexProcessor: patch("services.file_service.FileService") as mock_file_service, ): mock_file_service.return_value.upload_file.return_value = upload_result - result = processor._download_tool_file("tool-id", current_user=Mock()) + result = processor._download_tool_file("tool-id", current_user=Mock(), session=db_session) assert result == "upload-id" mock_load.assert_called_once_with("k1") diff --git a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py index 3f67b9c47ec..bceadf3ee84 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py +++ b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py @@ -56,6 +56,7 @@ from unittest.mock import MagicMock, Mock, patch import pytest from sqlalchemy.orm.exc import ObjectDeletedError +from core.entities.knowledge_entities import PreviewDetail from core.errors.error import ProviderTokenNotInitError from core.indexing_runner import ( DocumentIsDeletedPausedError, @@ -66,8 +67,9 @@ from core.rag.index_processor.constant.index_type import IndexStructureType, Ind from core.rag.models.document import ChildDocument, Document from graphon.model_runtime.entities.model_entities import ModelType from libs.datetime_utils import naive_utc_now -from models.dataset import Dataset, DatasetProcessRule +from models.dataset import Dataset, DatasetProcessRule, DocumentSegment from models.dataset import Document as DatasetDocument +from models.model import Account # ============================================================================ # Helper Functions @@ -260,12 +262,11 @@ class TestIndexingRunnerExtract: def mock_dependencies(self): """Mock all external dependencies for extract tests.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.IndexProcessorFactory") as mock_factory, patch("core.indexing_runner.storage") as mock_storage, ): yield { - "db": mock_db, + "session": MagicMock(), "factory": mock_factory, "storage": mock_storage, } @@ -331,7 +332,9 @@ class TestIndexingRunnerExtract: with patch("core.indexing_runner.select"): with patch("core.indexing_runner.ExtractSetting"): # Act: Call the extract method - result = runner._extract(mock_processor, sample_dataset_document, sample_process_rule) + result = runner._extract( + mock_processor, sample_dataset_document, sample_process_rule, mock_dependencies["session"] + ) # Assert: Verify the extraction results assert len(result) == 2, "Should extract 2 documents from the PDF" @@ -342,6 +345,7 @@ class TestIndexingRunnerExtract: assert result[1].page_content == "Test content 2", "Second document content should match" # Verify the processor was called exactly once (not multiple times) mock_processor.extract.assert_called_once() + assert mock_processor.extract.call_args.kwargs["session"] is mock_dependencies["session"] def test_extract_notion_import_success(self, mock_dependencies, sample_dataset_document, sample_process_rule): """Test successful extraction from Notion import.""" @@ -364,7 +368,9 @@ class TestIndexingRunnerExtract: # Mock update_document_index_status to avoid database calls with patch.object(runner, "_update_document_index_status"): # Act - result = runner._extract(mock_processor, sample_dataset_document, sample_process_rule) + result = runner._extract( + mock_processor, sample_dataset_document, sample_process_rule, mock_dependencies["session"] + ) # Assert assert len(result) == 1 @@ -395,7 +401,9 @@ class TestIndexingRunnerExtract: # Mock update_document_index_status to avoid database calls with patch.object(runner, "_update_document_index_status"): # Act - result = runner._extract(mock_processor, sample_dataset_document, sample_process_rule) + result = runner._extract( + mock_processor, sample_dataset_document, sample_process_rule, mock_dependencies["session"] + ) # Assert assert len(result) == 1 @@ -413,7 +421,7 @@ class TestIndexingRunnerExtract: # Act & Assert with pytest.raises(ValueError, match="no upload file found"): - runner._extract(mock_processor, sample_dataset_document, sample_process_rule) + runner._extract(mock_processor, sample_dataset_document, sample_process_rule, mock_dependencies["session"]) def test_extract_unsupported_data_source(self, mock_dependencies, sample_dataset_document, sample_process_rule): """Test extraction returns empty list for unsupported data sources.""" @@ -424,7 +432,9 @@ class TestIndexingRunnerExtract: mock_processor = MagicMock() # Act - result = runner._extract(mock_processor, sample_dataset_document, sample_process_rule) + result = runner._extract( + mock_processor, sample_dataset_document, sample_process_rule, mock_dependencies["session"] + ) # Assert assert result == [] @@ -445,11 +455,10 @@ class TestIndexingRunnerTransform: def mock_dependencies(self): """Mock all external dependencies for transform tests.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.ModelManager.for_tenant") as mock_model_manager, ): yield { - "db": mock_db, + "session": MagicMock(), "model_manager": mock_model_manager, } @@ -505,7 +514,14 @@ class TestIndexingRunnerTransform: } # Act - result = runner._transform(mock_processor, sample_dataset, sample_text_docs, "English", process_rule) + result = runner._transform( + mock_processor, + sample_dataset, + sample_text_docs, + "English", + process_rule, + session=mock_dependencies["session"], + ) # Assert assert len(result) == 2 @@ -538,7 +554,14 @@ class TestIndexingRunnerTransform: process_rule = {"mode": "automatic", "rules": {}} # Act - result = runner._transform(mock_processor, sample_dataset, sample_text_docs, "English", process_rule) + result = runner._transform( + mock_processor, + sample_dataset, + sample_text_docs, + "English", + process_rule, + session=mock_dependencies["session"], + ) # Assert assert len(result) == 1 @@ -562,7 +585,14 @@ class TestIndexingRunnerTransform: } # Act - result = runner._transform(mock_processor, sample_dataset, sample_text_docs, "Chinese", process_rule) + result = runner._transform( + mock_processor, + sample_dataset, + sample_text_docs, + "Chinese", + process_rule, + session=mock_dependencies["session"], + ) # Assert assert len(result) == 1 @@ -589,7 +619,6 @@ class TestIndexingRunnerLoad: def mock_dependencies(self): """Mock all external dependencies for load tests.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.ModelManager.for_tenant") as mock_model_manager, patch("core.indexing_runner.current_app") as mock_app, patch("core.indexing_runner.threading.Thread") as mock_thread, @@ -597,7 +626,7 @@ class TestIndexingRunnerLoad: ): mock_app._get_current_object = Mock(return_value=Mock()) yield { - "db": mock_db, + "session": MagicMock(), "model_manager": mock_model_manager, "app": mock_app, "thread": mock_thread, @@ -667,7 +696,13 @@ class TestIndexingRunnerLoad: # Mock update_document_index_status to avoid database calls with patch.object(runner, "_update_document_index_status"): # Act - runner._load(mock_processor, sample_dataset, sample_dataset_document, sample_documents) + runner._load( + mock_processor, + sample_dataset, + sample_dataset_document, + sample_documents, + mock_dependencies["session"], + ) # Assert model_manager.get_model_instance.assert_called_once() @@ -692,7 +727,13 @@ class TestIndexingRunnerLoad: # Mock update_document_index_status to avoid database calls with patch.object(runner, "_update_document_index_status"): # Act - runner._load(mock_processor, sample_dataset, sample_dataset_document, sample_documents) + runner._load( + mock_processor, + sample_dataset, + sample_dataset_document, + sample_documents, + mock_dependencies["session"], + ) # Assert # Verify keyword thread was created and joined @@ -737,7 +778,13 @@ class TestIndexingRunnerLoad: # Mock update_document_index_status to avoid database calls with patch.object(runner, "_update_document_index_status"): # Act - runner._load(mock_processor, sample_dataset, sample_dataset_document, sample_documents) + runner._load( + mock_processor, + sample_dataset, + sample_dataset_document, + sample_documents, + mock_dependencies["session"], + ) # Assert # Verify no keyword thread for parent-child index @@ -759,14 +806,13 @@ class TestIndexingRunnerRun: def mock_dependencies(self): """Mock all external dependencies for run tests.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.IndexProcessorFactory") as mock_factory, patch("core.indexing_runner.ModelManager.for_tenant") as mock_model_manager, patch("core.indexing_runner.storage") as mock_storage, patch("core.indexing_runner.threading.Thread") as mock_thread, ): yield { - "db": mock_db, + "session": MagicMock(), "factory": mock_factory, "model_manager": mock_model_manager, "storage": mock_storage, @@ -790,6 +836,33 @@ class TestIndexingRunnerRun: docs.append(doc) return docs + def test_run_in_indexing_status_loads_child_chunks_with_caller_session( + self, mock_dependencies, sample_dataset_documents + ): + runner = IndexingRunner() + dataset_document = sample_dataset_documents[0] + dataset_document.doc_form = IndexStructureType.PARENT_CHILD_INDEX + dataset = Mock(spec=Dataset) + segment = Mock(spec=DocumentSegment) + segment.status = "waiting" + segment.content = "parent" + segment.index_node_id = "parent-node" + segment.index_node_hash = "parent-hash" + segment.document_id = dataset_document.id + segment.dataset_id = dataset_document.dataset_id + segment.get_child_chunks.return_value = [ + SimpleNamespace(content="child", index_node_id="child-node", index_node_hash="child-hash") + ] + session = mock_dependencies["session"] + session.get.side_effect = lambda model, _: dataset_document if model is DatasetDocument else dataset + session.scalars.return_value.all.return_value = [segment] + + with patch.object(runner, "_load") as load: + runner.run_in_indexing_status(dataset_document, session) + + segment.get_child_chunks.assert_called_once_with(session=session) + assert load.call_args.kwargs["documents"][0].children[0].page_content == "child" + def test_run_success_single_document(self, mock_dependencies, sample_dataset_documents): """Test successful run with single document.""" # Arrange @@ -802,15 +875,14 @@ class TestIndexingRunnerRun: mock_dataset.tenant_id = doc.tenant_id mock_dataset.indexing_technique = IndexTechniqueType.ECONOMY - mock_current_user = MagicMock() - mock_current_user.set_tenant_id = MagicMock() + mock_current_user = Mock(spec=Account) get_dispatch = {"Document": doc, "Dataset": mock_dataset, "Account": mock_current_user} - mock_dependencies["db"].session.get.side_effect = lambda model, id: get_dispatch.get(model.__name__) + mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) mock_process_rule = Mock(spec=DatasetProcessRule) mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} - mock_dependencies["db"].session.scalar.return_value = mock_process_rule + mock_dependencies["session"].scalar.return_value = mock_process_rule # Mock processor mock_processor = MagicMock() @@ -841,11 +913,17 @@ class TestIndexingRunnerRun: patch.object(runner, "_load"), ): # Act - runner.run([doc]) + mock_dependencies["session"].commit.reset_mock() + runner.run([doc], mock_dependencies["session"]) # Assert - verify the methods were called # Since we're mocking the internal methods, we just verify no exceptions were raised + mock_current_user.set_tenant_id_with_session.assert_called_once_with( + mock_dataset.tenant_id, + session=mock_dependencies["session"], + ) + mock_dependencies["session"].commit.reset_mock() with ( patch.object(runner, "_extract", return_value=[Document(page_content="Test", metadata={})]) as mock_extract, patch.object( @@ -857,13 +935,14 @@ class TestIndexingRunnerRun: patch.object(runner, "_load") as mock_load, ): # Act - runner.run([doc]) + runner.run([doc], mock_dependencies["session"]) # Assert - verify the methods were called mock_extract.assert_called_once() mock_transform.assert_called_once() mock_load_segments.assert_called_once() mock_load.assert_called_once() + assert mock_dependencies["session"].commit.call_count == 2 mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -872,7 +951,7 @@ class TestIndexingRunnerRun: with patch.object(runner, "_extract", side_effect=DocumentIsPausedError("Document paused")): # Act & Assert with pytest.raises(DocumentIsPausedError): - runner.run([doc]) + runner.run([doc], mock_dependencies["session"]) def test_run_handles_provider_token_error(self, mock_dependencies, sample_dataset_documents): """Test run handles ProviderTokenNotInitError and updates document status.""" @@ -885,22 +964,23 @@ class TestIndexingRunnerRun: mock_dataset.tenant_id = doc.tenant_id get_dispatch = {"Document": doc, "Dataset": mock_dataset} - mock_dependencies["db"].session.get.side_effect = lambda model, id: get_dispatch.get(model.__name__) + mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) mock_process_rule = Mock(spec=DatasetProcessRule) mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} - mock_dependencies["db"].session.scalar.return_value = mock_process_rule + mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor mock_processor.extract.side_effect = ProviderTokenNotInitError("Token not initialized") # Act - runner.run([doc]) + with patch.object(runner, "_extract", side_effect=ProviderTokenNotInitError("Token not initialized")): + runner.run([doc], mock_dependencies["session"]) # Assert # Verify document status was updated to error - assert mock_dependencies["db"].session.commit.called + assert mock_dependencies["session"].flush.called def test_run_handles_object_deleted_error(self, mock_dependencies, sample_dataset_documents): """Test run handles ObjectDeletedError gracefully.""" @@ -913,11 +993,11 @@ class TestIndexingRunnerRun: mock_dataset.tenant_id = doc.tenant_id get_dispatch = {"Document": doc, "Dataset": mock_dataset} - mock_dependencies["db"].session.get.side_effect = lambda model, id: get_dispatch.get(model.__name__) + mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) mock_process_rule = Mock(spec=DatasetProcessRule) mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} - mock_dependencies["db"].session.scalar.return_value = mock_process_rule + mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -925,7 +1005,7 @@ class TestIndexingRunnerRun: # Mock _extract to raise ObjectDeletedError with patch.object(runner, "_extract", side_effect=ObjectDeletedError(state=None, msg="Object deleted")): # Act - runner.run([doc]) + runner.run([doc], mock_dependencies["session"]) # Assert - should not raise, just log warning # No exception should be raised @@ -939,8 +1019,7 @@ class TestIndexingRunnerRun: # Mock database mock_dataset = Mock(spec=Dataset) mock_dataset.indexing_technique = IndexTechniqueType.ECONOMY - mock_current_user = MagicMock() - mock_current_user.set_tenant_id = MagicMock() + mock_current_user = Mock(spec=Account) doc_map = {doc.id: doc for doc in docs} model_dispatch = {"Dataset": mock_dataset, "Account": mock_current_user} @@ -951,11 +1030,11 @@ class TestIndexingRunnerRun: return doc_map.get(id) return model_dispatch.get(name) - mock_dependencies["db"].session.get.side_effect = get_side_effect + mock_dependencies["session"].get.side_effect = get_side_effect mock_process_rule = Mock(spec=DatasetProcessRule) mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} - mock_dependencies["db"].session.scalar.return_value = mock_process_rule + mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -976,11 +1055,12 @@ class TestIndexingRunnerRun: patch.object(runner, "_load"), ): # Act - runner.run(docs) + runner.run(docs, mock_dependencies["session"]) # Assert # Verify extract was called for each document assert mock_extract.call_count == len(docs) + assert mock_current_user.set_tenant_id_with_session.call_count == len(docs) class TestIndexingRunnerRetryLogic: @@ -997,11 +1077,10 @@ class TestIndexingRunnerRetryLogic: def mock_dependencies(self): """Mock all external dependencies.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.redis_client") as mock_redis, ): yield { - "db": mock_db, + "session": MagicMock(), "redis": mock_redis, } @@ -1031,39 +1110,40 @@ class TestIndexingRunnerRetryLogic: mock_document = Mock(spec=DatasetDocument) mock_document.id = document_id - mock_dependencies["db"].session.scalar.return_value = 0 - mock_dependencies["db"].session.get.return_value = mock_document + mock_dependencies["session"].scalar.return_value = 0 + mock_dependencies["session"].get.return_value = mock_document # Act IndexingRunner._update_document_index_status( document_id, "completed", {"tokens": 100, "completed_at": naive_utc_now()}, + session=mock_dependencies["session"], ) # Assert - mock_dependencies["db"].session.commit.assert_called() + mock_dependencies["session"].flush.assert_called() def test_update_document_index_status_paused(self, mock_dependencies): """Test document status update when document is paused.""" # Arrange document_id = str(uuid.uuid4()) - mock_dependencies["db"].session.scalar.return_value = 1 + mock_dependencies["session"].scalar.return_value = 1 # Act & Assert with pytest.raises(DocumentIsPausedError): - IndexingRunner._update_document_index_status(document_id, "completed") + IndexingRunner._update_document_index_status(document_id, "completed", session=mock_dependencies["session"]) def test_update_document_index_status_deleted(self, mock_dependencies): """Test document status update when document is deleted.""" # Arrange document_id = str(uuid.uuid4()) - mock_dependencies["db"].session.scalar.return_value = 0 - mock_dependencies["db"].session.get.return_value = None + mock_dependencies["session"].scalar.return_value = 0 + mock_dependencies["session"].get.return_value = None # Act & Assert with pytest.raises(DocumentIsDeletedPausedError): - IndexingRunner._update_document_index_status(document_id, "completed") + IndexingRunner._update_document_index_status(document_id, "completed", session=mock_dependencies["session"]) class TestIndexingRunnerDocumentCleaning: @@ -1260,11 +1340,10 @@ class TestIndexingRunnerLoadSegments: def mock_dependencies(self): """Mock all external dependencies.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.DatasetDocumentStore") as mock_docstore, ): yield { - "db": mock_db, + "session": MagicMock(), "docstore": mock_docstore, } @@ -1315,7 +1394,9 @@ class TestIndexingRunnerLoadSegments: patch.object(runner, "_update_segments_by_document"), ): # Act - runner._load_segments(sample_dataset, sample_dataset_document, sample_documents) + runner._load_segments( + sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"] + ) # Assert mock_dependencies["docstore"].assert_called_once_with( @@ -1323,7 +1404,9 @@ class TestIndexingRunnerLoadSegments: user_id=sample_dataset_document.created_by, document_id=sample_dataset_document.id, ) - mock_docstore_instance.add_documents.assert_called_once_with(docs=sample_documents, save_child=False) + mock_docstore_instance.add_documents.assert_called_once_with( + docs=sample_documents, save_child=False, session=mock_dependencies["session"] + ) def test_load_segments_parent_child_index( self, mock_dependencies, sample_dataset, sample_dataset_document, sample_documents @@ -1351,10 +1434,14 @@ class TestIndexingRunnerLoadSegments: patch.object(runner, "_update_segments_by_document"), ): # Act - runner._load_segments(sample_dataset, sample_dataset_document, sample_documents) + runner._load_segments( + sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"] + ) # Assert - mock_docstore_instance.add_documents.assert_called_once_with(docs=sample_documents, save_child=True) + mock_docstore_instance.add_documents.assert_called_once_with( + docs=sample_documents, save_child=True, session=mock_dependencies["session"] + ) def test_load_segments_updates_word_count( self, mock_dependencies, sample_dataset, sample_dataset_document, sample_documents @@ -1374,7 +1461,9 @@ class TestIndexingRunnerLoadSegments: patch.object(runner, "_update_segments_by_document"), ): # Act - runner._load_segments(sample_dataset, sample_dataset_document, sample_documents) + runner._load_segments( + sample_dataset, sample_dataset_document, sample_documents, mock_dependencies["session"] + ) # Assert # Verify word count was calculated correctly and passed to status update @@ -1396,11 +1485,10 @@ class TestIndexingRunnerEstimate: def mock_dependencies(self): """Mock all external dependencies.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.IndexProcessorFactory") as mock_factory, ): yield { - "db": mock_db, + "session": MagicMock(), "factory": mock_factory, } @@ -1423,14 +1511,17 @@ class TestIndexingRunnerEstimate: extract_settings=extract_settings, tmp_processing_rule={"mode": "automatic", "rules": {}}, doc_form=IndexStructureType.PARAGRAPH_INDEX, + session=mock_dependencies["session"], ) - def test_indexing_estimate_commits_preview_image_cleanup(self, mock_dependencies): - """Test indexing estimate persists cleanup for preview-only extracted images.""" + def test_indexing_estimate_commits_preview_cleanup_before_summary_workers(self, mock_dependencies): + """Test preview cleanup is visible before summary workers use independent sessions.""" runner = IndexingRunner() tenant_id = str(uuid.uuid4()) mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor + phase_events: list[str] = [] + mock_dependencies["session"].commit.side_effect = lambda: phase_events.append("commit") preview_doc = Document( page_content="![image](http://files.local/files/image-1/file-preview)", @@ -1438,9 +1529,12 @@ class TestIndexingRunnerEstimate: ) mock_processor.extract.return_value = [preview_doc] mock_processor.transform.return_value = [preview_doc] + mock_processor.generate_summary_preview.side_effect = lambda *_args, **_kwargs: ( + phase_events.append("summary") or [PreviewDetail(content=preview_doc.page_content)] + ) image_file = SimpleNamespace(key="image_files/tenant-1/source-file-1/image.png") - mock_dependencies["db"].session.scalar.return_value = image_file + mock_dependencies["session"].scalar.return_value = image_file with ( patch("core.indexing_runner.get_image_upload_file_ids", return_value=["image-1"]), @@ -1452,14 +1546,19 @@ class TestIndexingRunnerEstimate: result = runner.indexing_estimate( tenant_id=tenant_id, extract_settings=[MagicMock()], - tmp_processing_rule={"mode": "automatic", "rules": {}}, + tmp_processing_rule={ + "mode": "automatic", + "rules": {}, + "summary_index_setting": {"enable": True}, + }, doc_form=IndexStructureType.PARAGRAPH_INDEX, + session=mock_dependencies["session"], ) assert result.total_segments == 1 mock_storage.delete.assert_called_once_with(image_file.key) - mock_dependencies["db"].session.delete.assert_called_once_with(image_file) - mock_dependencies["db"].session.commit.assert_called_once() + mock_dependencies["session"].delete.assert_called_once_with(image_file) + assert phase_events == ["commit", "summary"] class TestIndexingRunnerProcessChunk: @@ -1476,11 +1575,10 @@ class TestIndexingRunnerProcessChunk: def mock_dependencies(self): """Mock all external dependencies.""" with ( - patch("core.indexing_runner.db") as mock_db, patch("core.indexing_runner.redis_client") as mock_redis, ): yield { - "db": mock_db, + "session": MagicMock(), "redis": mock_redis, } @@ -1517,7 +1615,11 @@ class TestIndexingRunnerProcessChunk: mock_dependencies["redis"].get.return_value = None # Mock database update for segment status - mock_dependencies["db"].session.execute.return_value = None + mock_dependencies["session"].execute.return_value = None + mock_dependencies["session"].get.side_effect = lambda model, _id: { + Dataset: mock_dataset, + DatasetDocument: mock_dataset_document, + }.get(model) # Create a proper context manager mock mock_context = MagicMock() @@ -1525,15 +1627,25 @@ class TestIndexingRunnerProcessChunk: mock_context.__exit__ = MagicMock(return_value=None) mock_flask_app.app_context.return_value = mock_context - # Act - the method creates its own app_context - tokens = runner._process_chunk( - mock_flask_app, - mock_processor, - chunk_documents, - mock_dataset, - mock_dataset_document, - mock_embedding_instance, - ) + session_context = MagicMock() + session_context.__enter__.return_value = mock_dependencies["session"] + session_context.__exit__.return_value = None + + with ( + patch("core.indexing_runner.session_factory.create_session", return_value=session_context), + patch("core.indexing_runner.IndexProcessorFactory") as mock_factory, + ): + mock_factory.return_value.init_index_processor.return_value = mock_processor + + # Act - the method creates its own app_context and session + tokens = runner._process_chunk( + mock_flask_app, + IndexStructureType.PARAGRAPH_INDEX, + chunk_documents, + mock_dataset.id, + mock_dataset_document.id, + mock_embedding_instance, + ) # Assert assert tokens == 150 @@ -1555,6 +1667,10 @@ class TestIndexingRunnerProcessChunk: # Mock Redis to return paused status mock_dependencies["redis"].get.return_value = "1" + mock_dependencies["session"].get.side_effect = lambda model, _id: { + Dataset: mock_dataset, + DatasetDocument: mock_dataset_document, + }.get(model) # Create a proper context manager mock mock_context = MagicMock() @@ -1562,13 +1678,18 @@ class TestIndexingRunnerProcessChunk: mock_context.__exit__ = MagicMock(return_value=None) mock_flask_app.app_context.return_value = mock_context - # Act & Assert - the method creates its own app_context - with pytest.raises(DocumentIsPausedError): - runner._process_chunk( - mock_flask_app, - mock_processor, - chunk_documents, - mock_dataset, - mock_dataset_document, - mock_embedding_instance, - ) + session_context = MagicMock() + session_context.__enter__.return_value = mock_dependencies["session"] + session_context.__exit__.return_value = None + + with patch("core.indexing_runner.session_factory.create_session", return_value=session_context): + # Act & Assert - the method creates its own app_context and session + with pytest.raises(DocumentIsPausedError): + runner._process_chunk( + mock_flask_app, + IndexStructureType.PARAGRAPH_INDEX, + chunk_documents, + mock_dataset.id, + mock_dataset_document.id, + mock_embedding_instance, + ) diff --git a/api/tests/unit_tests/core/rag/rerank/test_reranker.py b/api/tests/unit_tests/core/rag/rerank/test_reranker.py index 565bb85b634..3ed676bd442 100644 --- a/api/tests/unit_tests/core/rag/rerank/test_reranker.py +++ b/api/tests/unit_tests/core/rag/rerank/test_reranker.py @@ -69,7 +69,7 @@ class TestRerankModelRunner: @pytest.fixture def rerank_runner(self, mock_model_instance): """Create a RerankModelRunner with mocked model instance.""" - return RerankModelRunner(rerank_model_instance=mock_model_instance) + return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) @pytest.fixture def sample_documents(self): @@ -427,7 +427,7 @@ class TestRerankModelRunnerMultimodal: @pytest.fixture def rerank_runner(self, mock_model_instance): - return RerankModelRunner(rerank_model_instance=mock_model_instance) + return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) def test_run_returns_original_documents_for_non_text_query_without_vision_support( self, rerank_runner, mock_model_instance @@ -486,7 +486,7 @@ class TestRerankModelRunnerMultimodal: rerank_result = RerankResult(model="rerank-model", docs=[]) with ( - patch("core.rag.rerank.rerank_model.db.session.get", return_value=SimpleNamespace(key="image-key")), + patch.object(rerank_runner._session, "get", return_value=SimpleNamespace(key="image-key")), patch("core.rag.rerank.rerank_model.storage.load_once", return_value=b"image-bytes") as mock_load_once, patch.object( rerank_runner, @@ -515,7 +515,7 @@ class TestRerankModelRunnerMultimodal: rerank_result = RerankResult(model="rerank-model", docs=[]) with ( - patch("core.rag.rerank.rerank_model.db.session.get", return_value=None), + patch.object(rerank_runner._session, "get", return_value=None), patch.object( rerank_runner, "fetch_text_rerank", @@ -550,7 +550,7 @@ class TestRerankModelRunnerMultimodal: session = MagicMock() session.get.return_value = SimpleNamespace(key="query-image-key") with ( - patch("core.rag.rerank.rerank_model.db.session", session), + patch.object(rerank_runner, "_session", session), patch("core.rag.rerank.rerank_model.storage.load_once", return_value=b"query-image-bytes"), ): result, unique_documents = rerank_runner.fetch_multimodal_rerank( @@ -569,7 +569,7 @@ class TestRerankModelRunnerMultimodal: assert "user" not in invoke_kwargs def test_fetch_multimodal_rerank_raises_when_query_image_not_found(self, rerank_runner): - with patch("core.rag.rerank.rerank_model.db.session.get", return_value=None): + with patch.object(rerank_runner._session, "get", return_value=None): with pytest.raises(ValueError, match="Upload file not found for query"): rerank_runner.fetch_multimodal_rerank( query="missing-upload-id", @@ -1071,6 +1071,7 @@ class TestRerankRunnerFactory: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=mock_model_instance, + session=MagicMock(), ) # Assert: Correct runner type is created @@ -1133,6 +1134,7 @@ class TestRerankRunnerFactory: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL.value, rerank_model_instance=mock_model_instance, + session=MagicMock(), ) # Assert: Runner is created successfully @@ -1197,6 +1199,7 @@ class TestRerankIntegration: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=mock_model_instance, + session=MagicMock(), ) result = runner.run( query="best programming language", @@ -1237,7 +1240,7 @@ class TestRerankIntegration: Document(page_content="Low relevance", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1299,7 +1302,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1339,7 +1342,7 @@ class TestRerankEdgeCases: Document(page_content="Negative score", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking with zero threshold result = runner.run(query="test", documents=documents, score_threshold=0.0) @@ -1375,7 +1378,7 @@ class TestRerankEdgeCases: Document(page_content="Perfect 3", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1416,7 +1419,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test 测试", documents=documents) @@ -1454,7 +1457,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1490,7 +1493,7 @@ class TestRerankEdgeCases: for i in range(num_docs) ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking with top_n result = runner.run(query="test", documents=documents, top_n=10) @@ -1580,7 +1583,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking with empty query result = runner.run(query="", documents=documents) @@ -1633,7 +1636,7 @@ class TestRerankPerformance: for i in range(5) ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1745,7 +1748,7 @@ class TestRerankErrorHandling: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act & Assert: Exception is raised with pytest.raises(RuntimeError, match="Model invocation failed"): @@ -1778,7 +1781,7 @@ class TestRerankErrorHandling: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) # Act & Assert: Should raise IndexError or handle gracefully with pytest.raises(IndexError): diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py index b121fe261c2..bebd9face61 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py @@ -326,6 +326,18 @@ class TestRetrievalService: app.app_context.return_value.__exit__ = Mock() return app + @pytest.fixture + def retrieval_session(self): + session = MagicMock() + session_context = MagicMock() + session_context.__enter__.return_value = session + session_context.__exit__.return_value = None + with ( + patch("core.rag.datasource.retrieval_service.db", SimpleNamespace(engine=Mock())), + patch("core.rag.datasource.retrieval_service.Session", return_value=session_context), + ): + yield session + @pytest.fixture(autouse=True) def mock_thread_pool(self): """ @@ -709,6 +721,7 @@ class TestRetrievalService: mock_data_processor_class, mock_dataset, sample_documents, + retrieval_session, ): """ Test basic hybrid search combining vector and full-text search. @@ -774,13 +787,20 @@ class TestRetrievalService: mock_embedding_search.assert_called_once() mock_fulltext_search.assert_called_once() mock_processor_instance.invoke.assert_called_once() + assert mock_data_processor_class.call_args.kwargs["session"] is retrieval_session @patch("core.rag.datasource.retrieval_service.DataPostProcessor") @patch("core.rag.datasource.retrieval_service.RetrievalService.full_text_index_search") @patch("core.rag.datasource.retrieval_service.RetrievalService.embedding_search") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") def test_hybrid_search_deduplication( - self, mock_get_dataset, mock_embedding_search, mock_fulltext_search, mock_data_processor_class, mock_dataset + self, + mock_get_dataset, + mock_embedding_search, + mock_fulltext_search, + mock_data_processor_class, + mock_dataset, + retrieval_session, ): """ Test that hybrid search properly deduplicates documents. @@ -905,6 +925,7 @@ class TestRetrievalService: doc_ids = [doc.metadata["doc_id"] for doc in results] assert "duplicate_doc" in doc_ids, "Duplicate doc should be present (higher score version)" assert "unique_doc" in doc_ids, "Unique doc should be present" + assert mock_data_processor_class.call_args.kwargs["session"] is retrieval_session # Implicitly verifies that doc1_low (score 0.6) was discarded # in favor of doc1_high (score 0.9) @@ -921,6 +942,7 @@ class TestRetrievalService: mock_data_processor_class, mock_dataset, sample_documents, + retrieval_session, ): """ Test hybrid search with custom weights for score merging. @@ -991,13 +1013,9 @@ class TestRetrievalService: assert len(results) == 3 # Verify DataPostProcessor was created with weights mock_data_processor_class.assert_called_once() - # Check that weights were passed (may be in args or kwargs) call_args = mock_data_processor_class.call_args - if call_args.kwargs: - assert call_args.kwargs.get("weights") == weights - else: - # Weights might be in positional args (position 3) - assert len(call_args.args) >= 4 + assert call_args.args[3] == weights + assert call_args.kwargs["session"] is retrieval_session @pytest.mark.parametrize("empty_query", ["", None]) @patch("core.rag.datasource.retrieval_service.DataPostProcessor") @@ -1011,6 +1029,7 @@ class TestRetrievalService: mock_dataset, sample_documents, empty_query, + retrieval_session, ): """ Regression test for GH #37116: attachment-only hybrid retrieval must use IMAGE_QUERY. @@ -1064,6 +1083,7 @@ class TestRetrievalService: assert invoke_kwargs["query"] == attachment_id, ( "The rerank query must be the attachment_id, not the empty text query" ) + assert mock_data_processor_class.call_args.kwargs["session"] is retrieval_session # ==================== Full-Text Search Tests ==================== @@ -1820,7 +1840,7 @@ class TestRetrievalService: mock_dataset2.provider = "dify" # Act - Call with dataset_count = 2 - with _patched_retriever_session(): + with _patched_retriever_session() as rerank_session: dataset_retrieval._multiple_retrieve_thread( flask_app=mock_flask_app, available_datasets=[mock_dataset, mock_dataset2], @@ -1847,6 +1867,7 @@ class TestRetrievalService: {"reranking_provider_name": "cohere", "reranking_model_name": "rerank-v2"}, None, False, + session=rerank_session, ) # Verify invoke was called with correct parameters @@ -5305,11 +5326,15 @@ class TestInternalHooksCoverage: assert len(all_documents) >= 3 def test_to_dataset_retriever_tool_paths(self, retrieval: DatasetRetrieval) -> None: - dataset_skip_zero = SimpleNamespace(id="d1", provider="dify", available_document_count=0) + dataset_skip_zero = SimpleNamespace( + id="d1", + provider="dify", + get_total_available_documents=Mock(return_value=0), + ) dataset_ok_single = SimpleNamespace( id="d2", provider="dify", - available_document_count=2, + get_total_available_documents=Mock(return_value=2), retrieval_model={"top_k": 2, "score_threshold_enabled": True, "score_threshold": 0.1}, ) single_config = DatasetRetrieveConfigEntity( diff --git a/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py b/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py index 5829098f6b4..007ab09aabc 100644 --- a/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py +++ b/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py @@ -3,7 +3,7 @@ from __future__ import annotations import uuid from contextlib import nullcontext from types import SimpleNamespace -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest from yaml import YAMLError @@ -135,8 +135,8 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources(): {"dataset-1": ["doc-a"]}, {"logical_operator": "and"}, ) - db_session = Mock() - db_session.scalar.return_value = dataset + session = Mock() + session.scalar.return_value = dataset external_documents = [ {"content": "first", "metadata": {"document_id": "doc-a"}, "score": 0.9, "title": "Doc A"}, {"content": "second", "metadata": {"document_id": "doc-b"}, "score": 0.8, "title": "Doc B"}, @@ -152,14 +152,13 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources(): inputs={"x": 1}, ) - with patch.object(single_retriever_module, "db", SimpleNamespace(session=db_session)): - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): - with patch.object( - single_retriever_module.ExternalDatasetService, - "fetch_external_knowledge_retrieval", - return_value=external_documents, - ) as fetch_mock: - result = tool.run(session=MagicMock(), query="hello") + with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): + with patch.object( + single_retriever_module.ExternalDatasetService, + "fetch_external_knowledge_retrieval", + return_value=external_documents, + ) as fetch_mock: + result = tool.run(session=session, query="hello") assert result == "first\nsecond" assert callback.queries == [("hello", "dataset-1")] @@ -168,6 +167,7 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources(): assert [item.position for item in resource_info] == [1, 2] assert resource_info[0].dataset_id == "dataset-1" fetch_mock.assert_called_once() + assert fetch_mock.call_args.kwargs["session"] is session def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_documents(): @@ -181,8 +181,8 @@ def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_do ) dataset_retrieval = Mock() dataset_retrieval.get_metadata_filter_condition.return_value = ({"dataset-1": []}, {"logical_operator": "and"}) - db_session = Mock() - db_session.scalar.return_value = dataset + session = Mock() + session.scalar.return_value = dataset tool = SingleDatasetRetrieverTool( tenant_id="tenant-1", @@ -194,10 +194,9 @@ def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_do inputs={}, ) - with patch.object(single_retriever_module, "db", SimpleNamespace(session=db_session)): - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): - with patch.object(single_retriever_module.RetrievalService, "retrieve") as retrieve_mock: - result = tool.run(session=MagicMock(), query="hello") + with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): + with patch.object(single_retriever_module.RetrievalService, "retrieve") as retrieve_mock: + result = tool.run(session=session, query="hello") assert result == "" retrieve_mock.assert_not_called() @@ -261,9 +260,9 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): lookup_doc_high = SimpleNamespace( id="doc-high", name="Document High", data_source_type="notion", doc_metadata={"lang": "fr"} ) - db_session = Mock() - db_session.scalar.side_effect = [dataset, lookup_doc_low, lookup_doc_high] - db_session.get.return_value = dataset + session = Mock() + session.scalar.side_effect = [dataset, lookup_doc_low, lookup_doc_high] + session.get.return_value = dataset tool = SingleDatasetRetrieverTool( tenant_id="tenant-1", @@ -276,15 +275,14 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): top_k=2, ) - with patch.object(single_retriever_module, "db", SimpleNamespace(session=db_session)): - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): - with patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents): - with patch.object( - single_retriever_module.RetrievalService, - "format_retrieval_documents", - return_value=records, - ): - result = tool.run(session=MagicMock(), query="hello") + with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): + with patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents): + with patch.object( + single_retriever_module.RetrievalService, + "format_retrieval_documents", + return_value=records, + ): + result = tool.run(session=session, query="hello") assert result == "signed high\nsummary low\nquestion:signed low answer:low answer" assert callback.documents == documents @@ -461,12 +459,14 @@ def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(): with patch.object(tool, "_retriever", side_effect=fake_retriever) as retriever_mock: with patch.object(multi_retriever_module, "current_app", fake_current_app): with patch.object(multi_retriever_module.threading, "Thread", _ImmediateThread): - with patch.object(multi_retriever_module, "ModelManager", return_value=model_manager): - with patch.object(multi_retriever_module, "RerankModelRunner", return_value=rerank_runner): - with patch.object(multi_retriever_module, "db", SimpleNamespace(session=db_session)): - result = tool.run(session=MagicMock(), query="hello") + with patch.object(multi_retriever_module.ModelManager, "for_tenant", return_value=model_manager): + with patch.object( + multi_retriever_module, "RerankModelRunner", return_value=rerank_runner + ) as rerank_runner_class: + result = tool.run(session=db_session, query="hello") assert result == "signed one\nquestion:signed two answer:answer two" + rerank_runner_class.assert_called_once_with(model_manager.get_model_instance.return_value, session=db_session) assert retriever_mock.call_count == 2 assert callback.documents == [second_doc, first_doc] assert callback.resources is not None diff --git a/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py b/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py index 5cb6acf50a7..8df8e28bda4 100644 --- a/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py +++ b/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py @@ -559,6 +559,8 @@ def test_resolve_user_from_database_returns_account(monkeypatch: pytest.MonkeyPa """Resolve Account and set tenant in worker context.""" tenant = SimpleNamespace(id="tenant_id") account = SimpleNamespace(id="account_id", current_tenant=None) + set_current_tenant = Mock(side_effect=lambda tenant, *, session: setattr(account, "current_tenant", tenant)) + account.set_current_tenant_with_session = set_current_tenant session = StubSession(scalar_results=[tenant, account]) monkeypatch.setattr("core.tools.workflow_as_tool.tool.session_factory.create_session", lambda: session) @@ -568,6 +570,7 @@ def test_resolve_user_from_database_returns_account(monkeypatch: pytest.MonkeyPa resolved = tool._resolve_user_from_database(user_id="account_id") assert resolved is account assert account.current_tenant is tenant + set_current_tenant.assert_called_once_with(tenant, session=session) assert session.expunge_calls == [account] diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_validators.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_validators.py index 2254cd16d49..0c8169943f9 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_validators.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_validators.py @@ -553,7 +553,7 @@ def test_publish_validation_rejects_missing_or_out_of_scope_knowledge_datasets( captured = {} - def fake_get_datasets_by_ids(ids, tenant_id): + def fake_get_datasets_by_ids(ids, tenant_id, *, session): captured["ids"] = ids captured["tenant_id"] = tenant_id return [], 0 diff --git a/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py b/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py index 0d760a2db75..cefcf04ae02 100644 --- a/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py @@ -1,6 +1,6 @@ import time import uuid -from unittest.mock import Mock +from unittest.mock import MagicMock, Mock import pytest from pytest_mock import MockerFixture @@ -248,6 +248,7 @@ class TestKnowledgeIndexNode: def test_run_preview_mode_success( self, + mocker: MockerFixture, mock_graph_init_params, mock_graph_runtime_state, mock_index_processor, @@ -282,6 +283,13 @@ class TestKnowledgeIndexNode: total_segments=2, ) mock_index_processor.get_preview_output.return_value = mock_preview + session = MagicMock() + session_context = MagicMock() + session_context.__enter__.return_value = session + mocker.patch( + "core.workflow.nodes.knowledge_index.knowledge_index_node.session_factory.create_session", + return_value=session_context, + ) node_id = str(uuid.uuid4()) config = { @@ -302,7 +310,7 @@ class TestKnowledgeIndexNode: # Assert assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED assert result.outputs is not None - assert mock_index_processor.get_preview_output.called + assert mock_index_processor.get_preview_output.call_args.kwargs["session"] is session def test_run_production_mode_success( self, @@ -564,7 +572,9 @@ class TestKnowledgeIndexNode: ) # Act + session = MagicMock() result = node._invoke_knowledge_index( + session=session, dataset_id=dataset_id, document_id=document_id, original_document_id=original_document_id, @@ -577,6 +587,7 @@ class TestKnowledgeIndexNode: # Assert assert mock_summary_index_service.generate_and_vectorize_summary.called assert mock_index_processor.index_and_clean.called + session.commit.assert_called_once() assert result == {"status": "indexed"} def test_version_method(self): @@ -651,7 +662,9 @@ class TestInvokeKnowledgeIndex: ) # Act + session = MagicMock() result = node._invoke_knowledge_index( + session=session, dataset_id=dataset_id, document_id=document_id, original_document_id=original_document_id, @@ -666,6 +679,7 @@ class TestInvokeKnowledgeIndex: dataset_id, document_id, False, summary_setting ) mock_index_processor.index_and_clean.assert_called_once_with( - dataset_id, document_id, original_document_id, chunks, batch, summary_setting + dataset_id, document_id, original_document_id, chunks, batch, summary_setting, session=session ) + session.commit.assert_called_once() assert result == {"status": "indexed"} diff --git a/api/tests/unit_tests/events/event_handlers/test_create_document_index.py b/api/tests/unit_tests/events/event_handlers/test_create_document_index.py index 005053115bc..4b0ca5291aa 100644 --- a/api/tests/unit_tests/events/event_handlers/test_create_document_index.py +++ b/api/tests/unit_tests/events/event_handlers/test_create_document_index.py @@ -90,6 +90,11 @@ def test_handle_runs_indexing_on_success( mock_indexing_runner: MagicMock, caplog: pytest.LogCaptureFixture, ) -> None: + def assert_status_committed(_documents: list[Document], session: Session) -> None: + assert not session.in_transaction() + + mock_indexing_runner.run.side_effect = assert_status_committed + with patch.object(handler_module, "IndexingRunner", return_value=mock_indexing_runner): with caplog.at_level(logging.INFO, logger=handler_module.logger.name): handler_module.handle("dataset-1", document_ids=["doc-1"]) 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 51c1d05edd0..78e441f7eee 100644 --- a/api/tests/unit_tests/events/test_app_event_signals.py +++ b/api/tests/unit_tests/events/test_app_event_signals.py @@ -1,5 +1,7 @@ +import json from collections.abc import Iterator -from unittest.mock import patch +from types import SimpleNamespace +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest @@ -8,7 +10,8 @@ from sqlalchemy.orm import Session from events.app_event import app_was_deleted, app_was_updated from models.account import Account -from models.model import App, AppMode, IconType +from models.dataset import AppDatasetJoin +from models.model import App, AppMode, AppModelConfig, IconType, InstalledApp from services.app_service import AppService @@ -216,3 +219,60 @@ class TestAppWasUpdatedSignal: assert received == [] assert sqlite_session.get(App, app_model.id).enable_api is True # type: ignore[union-attr] + + +class TestAppModelConfigWasUpdatedSignal: + def test_requires_caller_session(self) -> None: + from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle + + with pytest.raises(TypeError, match="session"): + handle(SimpleNamespace(id="app-1"), app_model_config=None) + + def test_reuses_provided_session_without_committing(self) -> None: + from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle + + session = MagicMock() + session.scalars.return_value.all.return_value = [] + app_model_config = AppModelConfig(app_id="app-1", created_by="user-1", updated_by="user-1") + app_model_config.dataset_configs = json.dumps( + { + "retrieval_model": "multiple", + "datasets": {"datasets": [{"dataset": {"id": "dataset-1"}}]}, + } + ) + + handle(SimpleNamespace(id="app-1"), app_model_config=app_model_config, session=session) + + added_join = session.add.call_args.args[0] + assert isinstance(added_join, AppDatasetJoin) + assert added_join.app_id == "app-1" + assert added_join.dataset_id == "dataset-1" + session.commit.assert_not_called() + + +class TestCreateInstalledAppWhenAppCreated: + def test_skips_existing_installation(self) -> None: + from events.event_handlers.create_installed_app_when_app_created import handle + + session = MagicMock() + session.scalar.return_value = "installed-app-1" + + handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session) + + session.add.assert_not_called() + session.flush.assert_not_called() + + def test_adds_missing_installation_without_committing(self) -> None: + from events.event_handlers.create_installed_app_when_app_created import handle + + session = MagicMock() + session.scalar.return_value = None + + handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session) + + installed_app = session.add.call_args.args[0] + assert isinstance(installed_app, InstalledApp) + assert installed_app.app_id == "app-1" + assert installed_app.tenant_id == "tenant-1" + session.flush.assert_called_once_with() + session.commit.assert_not_called() diff --git a/api/tests/unit_tests/fields/test_dataset_fields.py b/api/tests/unit_tests/fields/test_dataset_fields.py index 921e3882a96..caffaae09bf 100644 --- a/api/tests/unit_tests/fields/test_dataset_fields.py +++ b/api/tests/unit_tests/fields/test_dataset_fields.py @@ -1,4 +1,7 @@ -from fields.dataset_fields import DatasetDetailResponse +from types import SimpleNamespace +from unittest.mock import Mock + +from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source def _dataset_detail_payload(**overrides): @@ -179,3 +182,47 @@ def test_dataset_detail_expands_missing_weighted_score_nested_fields(): "embedding_provider_name": None, }, } + + +def test_dataset_detail_response_source_uses_caller_session_for_database_fields(): + session = Mock() + getter_mocks = { + "get_app_count": Mock(return_value=3), + "get_document_count": Mock(return_value=4), + "get_word_count": Mock(return_value=500), + "get_author_name": Mock(return_value="Ada"), + "get_tags": Mock(return_value=[{"id": "tag-1", "name": "Tag", "type": "knowledge"}]), + "get_doc_form": Mock(return_value="paragraph"), + "get_external_knowledge_info": Mock( + return_value={ + "external_knowledge_id": "knowledge-id", + "external_knowledge_api_id": "api-id", + "external_knowledge_api_name": "api", + "external_knowledge_api_endpoint": "https://example.com", + } + ), + "get_doc_metadata": Mock(return_value=[{"id": "metadata-1", "name": "Metadata", "type": "string"}]), + "get_is_published": Mock(return_value=True), + "get_total_documents": Mock(return_value=4), + "get_total_available_documents": Mock(return_value=2), + } + dataset = SimpleNamespace(**_dataset_detail_payload(), **getter_mocks) + + response = DatasetDetailResponse.model_validate( + dataset_detail_response_source(dataset, session=session), + from_attributes=True, + ) + + assert response.app_count == 3 + assert response.document_count == 4 + assert response.word_count == 500 + assert response.author_name == "Ada" + assert response.tags[0].id == "tag-1" + assert response.doc_form == "paragraph" + assert response.external_knowledge_info.external_knowledge_api_id == "api-id" + assert response.doc_metadata[0].id == "metadata-1" + assert response.is_published is True + assert response.total_documents == 4 + assert response.total_available_documents == 2 + for getter in getter_mocks.values(): + getter.assert_called_once_with(session=session) diff --git a/api/tests/unit_tests/fields/test_document_fields.py b/api/tests/unit_tests/fields/test_document_fields.py new file mode 100644 index 00000000000..d3483e2ad75 --- /dev/null +++ b/api/tests/unit_tests/fields/test_document_fields.py @@ -0,0 +1,19 @@ +from unittest.mock import MagicMock + +from fields.document_fields import DocumentWithSession + + +def test_document_with_session_uses_explicit_getters() -> None: + session = MagicMock() + document = MagicMock() + document.get_data_source_detail_dict.return_value = {"source": "detail"} + document.get_hit_count.return_value = 3 + document.get_doc_metadata_details.return_value = [{"name": "author"}] + source = DocumentWithSession(document=document, session=session) + + assert source.data_source_detail_dict == {"source": "detail"} + assert source.hit_count == 3 + assert source.doc_metadata_details == [{"name": "author"}] + document.get_data_source_detail_dict.assert_called_once_with(session=session) + document.get_hit_count.assert_called_once_with(session=session) + document.get_doc_metadata_details.assert_called_once_with(session=session) diff --git a/api/tests/unit_tests/models/test_account_models.py b/api/tests/unit_tests/models/test_account_models.py index 512c043b0c8..d2dc7e9ae5f 100644 --- a/api/tests/unit_tests/models/test_account_models.py +++ b/api/tests/unit_tests/models/test_account_models.py @@ -12,10 +12,11 @@ This test suite covers: import base64 import secrets from datetime import UTC, datetime -from unittest.mock import patch +from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from sqlalchemy.orm import Session from libs.password import compare_password, hash_password, valid_password from models.account import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole @@ -334,6 +335,39 @@ class TestTenantRelationshipIntegrity: # Assert assert tenant_id_none is None + def test_set_current_tenant_with_session_uses_caller_session(self): + account = Account(name="Test User", email="test@example.com") + account.id = str(uuid4()) + tenant = Tenant(name="Test Tenant") + tenant.id = str(uuid4()) + join = MagicMock(role=TenantAccountRole.OWNER) + session = MagicMock(spec=Session) + session.scalar.return_value = join + session.scalars.return_value.one.return_value = tenant + + with patch("models.account.Session") as session_class: + account.set_current_tenant_with_session(tenant, session=session) + + session_class.assert_not_called() + assert account.current_tenant is tenant + assert account.role == TenantAccountRole.OWNER + + def test_set_tenant_id_with_session_uses_caller_session(self): + account = Account(name="Test User", email="test@example.com") + account.id = str(uuid4()) + tenant = Tenant(name="Test Tenant") + tenant.id = str(uuid4()) + join = MagicMock(role=TenantAccountRole.ADMIN) + session = MagicMock(spec=Session) + session.execute.return_value.first.return_value = (tenant, join) + + with patch("models.account.Session") as session_class: + account.set_tenant_id_with_session(tenant.id, session=session) + + session_class.assert_not_called() + assert account.current_tenant is tenant + assert account.role == TenantAccountRole.ADMIN + class TestAccountRolePermissions: """Test suite for account role permissions.""" diff --git a/api/tests/unit_tests/models/test_app_models.py b/api/tests/unit_tests/models/test_app_models.py index 684d7f9fa8e..5ef0f52d3d8 100644 --- a/api/tests/unit_tests/models/test_app_models.py +++ b/api/tests/unit_tests/models/test_app_models.py @@ -17,6 +17,7 @@ from uuid import uuid4 import pytest from sqlalchemy.dialects import postgresql +from sqlalchemy.orm import Session from models.enums import ConversationFromSource from models.model import ( @@ -153,13 +154,14 @@ class TestAppModelValidation: description="", ) - # Mock app_model_config property - with patch.object(App, "app_model_config", new_callable=lambda: property(lambda self: None)): + session = MagicMock() + with patch.object(App, "app_model_config_with_session", return_value=None) as get_model_config: # Act - result = app.desc_or_prompt + result = app.desc_or_prompt_with_session(session=session) # Assert assert result == "" + get_model_config.assert_called_once_with(session=session) def test_app_is_agent_property_false(self): """Test is_agent property returns False when not configured as agent.""" @@ -173,13 +175,64 @@ class TestAppModelValidation: created_by=str(uuid4()), ) - # Mock app_model_config to return None - with patch.object(App, "app_model_config", new_callable=lambda: property(lambda self: None)): + with patch("models.model.db") as mock_db: # Act result = app.is_agent # Assert assert result is False + mock_db.session.assert_called_once_with() + + def test_app_is_agent_with_session_updates_legacy_agent_mode(self): + app = App( + tenant_id=str(uuid4()), + name="Test App", + mode=AppMode.CHAT, + enable_site=True, + enable_api=False, + created_by=str(uuid4()), + ) + app.app_model_config_id = "config-1" + app_model_config = MagicMock(spec=AppModelConfig) + app_model_config.agent_mode = "agent" + app_model_config.agent_mode_dict = {"enabled": True, "strategy": "react"} + session = MagicMock() + session.get.return_value = app_model_config + + result = app.is_agent_with_session(session=session) + + assert result is True + assert app.mode == AppMode.AGENT_CHAT + session.get.assert_called_once_with(AppModelConfig, "config-1") + session.execute.assert_called_once() + session.commit.assert_called_once_with() + + @pytest.mark.parametrize("sqlite_session", [(App, AppModelConfig)], indirect=True) + def test_app_is_agent_with_session_persists_mode_across_sessions(self, sqlite_session: Session): + app = App( + tenant_id=str(uuid4()), + name="Test App", + mode=AppMode.CHAT, + enable_site=True, + enable_api=False, + created_by=str(uuid4()), + ) + sqlite_session.add(app) + sqlite_session.flush() + model_config = AppModelConfig( + app_id=app.id, + agent_mode=json.dumps({"enabled": True, "strategy": "react"}), + ) + sqlite_session.add(model_config) + sqlite_session.flush() + app.app_model_config_id = model_config.id + sqlite_session.commit() + + with Session(sqlite_session.get_bind(), expire_on_commit=False) as migration_session: + assert app.is_agent_with_session(session=migration_session) is True + + sqlite_session.expire_all() + assert sqlite_session.get(App, app.id).mode == AppMode.AGENT_CHAT def test_app_mode_compatible_with_agent(self): """Test mode_compatible_with_agent property.""" @@ -193,13 +246,14 @@ class TestAppModelValidation: created_by=str(uuid4()), ) - # Mock is_agent to return False - with patch.object(App, "is_agent", new_callable=lambda: property(lambda self: False)): + session = MagicMock() + with patch.object(App, "is_agent_with_session", return_value=False) as is_agent: # Act - result = app.mode_compatible_with_agent + result = app.mode_compatible_with_agent_with_session(session=session) # Assert assert result == AppMode.CHAT + is_agent.assert_called_once_with(session=session) def test_deleted_tools_checks_plugin_builtin_providers_through_core_plugin_service(self): """Plugin-backed built-in tools are checked through core PluginService.""" @@ -230,22 +284,19 @@ class TestAppModelValidation: } ), ) - session_context = MagicMock() - session_context.__enter__.return_value = MagicMock() - session_factory = SimpleNamespace(begin=MagicMock(return_value=session_context)) + session = MagicMock() # Act with ( - patch.object(App, "app_model_config", new_callable=lambda: property(lambda self: app_model_config)), - patch("models.model.db", SimpleNamespace(engine=object())), - patch("models.model.sessionmaker", return_value=session_factory), + patch.object(App, "app_model_config_with_session", return_value=app_model_config) as get_model_config, patch("core.tools.tool_manager.ToolManager.get_hardcoded_provider", side_effect=Exception), patch("core.plugin.plugin_service.PluginService.check_tools_existence", return_value=[False]) as exists, ): - result = app.deleted_tools + result = app.deleted_tools_with_session(session=session) # Assert assert result == [{"type": "builtin", "tool_name": "chat", "provider_id": "langgenius/openai/openai"}] + get_model_config.assert_called_once_with(session=session) exists.assert_called_once() assert exists.call_args.args[0] == "tenant-1" assert [str(provider_id) for provider_id in exists.call_args.args[1]] == ["langgenius/openai/openai"] @@ -510,12 +561,67 @@ class TestConversationModel: # Mock first_message to return a message with query mock_message = MagicMock() mock_message.query = "First message query" - with patch.object(Conversation, "first_message", new_callable=lambda: property(lambda self: mock_message)): + session = MagicMock() + with patch.object(Conversation, "first_message_with_session", return_value=mock_message) as get_first_message: # Act - result = conversation.summary_or_query + result = conversation.summary_or_query_with_session(session=session) # Assert assert result == "First message query" + get_first_message.assert_called_once_with(session=session) + + def test_model_config_uses_caller_session_for_annotation_reply(self): + conversation = Conversation( + app_id="app-1", + app_model_config_id="config-1", + mode=AppMode.CHAT, + name="Test Conversation", + status="normal", + from_source=ConversationFromSource.API, + from_end_user_id=str(uuid4()), + model_id="model-1", + model_provider="provider-1", + ) + app_model_config = MagicMock(spec=AppModelConfig) + app_model_config.app_id = "app-1" + app_model_config.to_dict.return_value = {} + session = MagicMock() + session.scalar.return_value = app_model_config + annotation_reply = {"enabled": False} + + with patch("models.model.load_annotation_reply_config", return_value=annotation_reply) as load_config: + result = conversation.model_config_with_session(session=session) + + load_config.assert_called_once_with(session, "app-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert result["model_id"] == "model-1" + assert result["provider"] == "provider-1" + + def test_override_model_config_uses_caller_session_for_annotation_reply(self): + conversation = Conversation( + app_id="app-1", + mode=AppMode.CHAT, + name="Test Conversation", + status="normal", + from_source=ConversationFromSource.API, + from_end_user_id=str(uuid4()), + override_model_configs=json.dumps({"model": {}}), + ) + app_model_config = MagicMock(spec=AppModelConfig) + app_model_config.app_id = "app-1" + app_model_config.to_dict.return_value = {} + session = MagicMock() + annotation_reply = {"enabled": False} + + with ( + patch.object(AppModelConfig, "from_model_config_dict", return_value=app_model_config), + patch("models.model.load_annotation_reply_config", return_value=annotation_reply) as load_config, + ): + conversation.model_config_with_session(session=session) + + load_config.assert_called_once_with(session, "app-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + session.scalar.assert_not_called() def test_conversation_in_debug_mode(self): """Test in_debug_mode property.""" diff --git a/api/tests/unit_tests/models/test_dataset_models.py b/api/tests/unit_tests/models/test_dataset_models.py index 618495e6f2a..0a7a0722933 100644 --- a/api/tests/unit_tests/models/test_dataset_models.py +++ b/api/tests/unit_tests/models/test_dataset_models.py @@ -12,23 +12,28 @@ This test suite covers: import json import pickle from datetime import UTC, datetime -from unittest.mock import Mock, patch +from types import SimpleNamespace +from unittest.mock import Mock, call, patch from urllib.parse import parse_qs, urlparse from uuid import uuid4 import pytest -from core.rag.index_processor.constant.index_type import IndexTechniqueType +from core.rag.entities import ParentMode +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from extensions.storage.storage_type import StorageType +from models.account import Account from models.dataset import ( AppDatasetJoin, ChildChunk, Dataset, DatasetKeywordTable, DatasetProcessRule, + DatasetQuery, Document, DocumentSegment, Embedding, + ExternalKnowledgeApis, ExternalKnowledgeBindings, ) from models.enums import ( @@ -86,6 +91,79 @@ class TestDatasetModelValidation: assert dataset.embedding_model == "text-embedding-ada-002" assert dataset.embedding_model_provider == "openai" + def test_session_aware_dataset_getters_use_caller_session(self): + dataset = Dataset( + tenant_id=str(uuid4()), + name="Test Dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=str(uuid4()), + ) + dataset.id = str(uuid4()) + account = Mock() + process_rule = Mock() + session = Mock() + session.get.return_value = account + session.scalar.side_effect = [process_rule, IndexStructureType.PARAGRAPH_INDEX] + + assert dataset.get_created_by_account(session=session) is account + assert dataset.get_latest_process_rule(session=session) is process_rule + assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX + + session.get.assert_called_once_with(Account, dataset.created_by) + assert session.scalar.call_count == 2 + + def test_get_dataset_keyword_table_uses_caller_session(self): + dataset = Dataset( + tenant_id=str(uuid4()), + name="Test Dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=str(uuid4()), + ) + dataset.id = str(uuid4()) + keyword_table = Mock() + session = Mock() + session.scalar.return_value = keyword_table + + result = dataset.get_dataset_keyword_table(session=session) + + assert result is keyword_table + session.scalar.assert_called_once() + + def test_dataset_detail_getters_use_caller_session(self): + dataset = Dataset( + tenant_id=str(uuid4()), + name="Test Dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=str(uuid4()), + provider="vendor", + built_in_field_enabled=False, + ) + dataset.id = str(uuid4()) + account = Mock(name="account", name_value="Ada") + account.name = "Ada" + session = Mock() + session.get.return_value = account + session.scalar.side_effect = [2, 1, 3, 2, 500, IndexStructureType.PARAGRAPH_INDEX] + session.scalars.return_value.all.return_value = [] + + with patch("models.dataset.db") as mock_db: + assert dataset.get_total_documents(session=session) == 2 + assert dataset.get_total_available_documents(session=session) == 1 + assert dataset.get_app_count(session=session) == 3 + assert dataset.get_document_count(session=session) == 2 + assert dataset.get_word_count(session=session) == 500 + assert dataset.get_author_name(session=session) == "Ada" + assert dataset.get_tags(session=session) == [] + assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX + assert dataset.get_external_knowledge_info(session=session) is None + assert dataset.get_doc_metadata(session=session) == [] + assert dataset.get_is_published(session=session) is False + + assert session.scalar.call_count == 6 + assert session.scalars.call_count == 2 + mock_db.session.scalar.assert_not_called() + mock_db.session.scalars.assert_not_called() + def test_dataset_indexing_technique_validation(self): """Test dataset indexing technique values.""" # Arrange & Act @@ -200,10 +278,79 @@ class TestDatasetModelValidation: binding.external_knowledge_id = "knowledge-1" binding.external_knowledge_api_id = str(uuid4()) + session = Mock() + session.scalar.side_effect = [binding, None] with patch("models.dataset.db") as mock_db: - mock_db.session.scalar.side_effect = [binding, None] + assert dataset.get_external_knowledge_info(session=session) is None - assert dataset.external_knowledge_info is None + assert session.scalar.call_count == 2 + mock_db.session.scalar.assert_not_called() + + def test_external_knowledge_api_dataset_bindings_use_caller_session(self): + external_api = ExternalKnowledgeApis( + tenant_id=str(uuid4()), + created_by=str(uuid4()), + updated_by=None, + name="External API", + description="", + settings=None, + ) + binding = Mock(dataset_id="dataset-1") + dataset = SimpleNamespace(id="dataset-1", name="Dataset") + session = Mock() + session.scalars.side_effect = [ + Mock(all=Mock(return_value=[binding])), + Mock(all=Mock(return_value=[dataset])), + ] + + with patch("models.dataset.db") as mock_db: + result = external_api.get_dataset_bindings(session=session) + + assert result == [{"id": "dataset-1", "name": "Dataset"}] + assert session.scalars.call_count == 2 + mock_db.session.scalars.assert_not_called() + + def test_dataset_query_get_queries_uses_caller_session(self): + dataset_query = DatasetQuery( + dataset_id=str(uuid4()), + content=json.dumps([{"content_type": "image_query", "content": "file-1"}]), + source="hit_testing", + source_app_id=None, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + ) + upload_file = SimpleNamespace( + id="file-1", + name="image.png", + size=10, + extension="png", + mime_type="image/png", + ) + session = Mock() + session.scalar.return_value = upload_file + + with ( + patch("models.dataset.db") as mock_db, + patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"), + ): + queries = dataset_query.get_queries(session=session) + + assert queries == [ + { + "content_type": "image_query", + "content": "file-1", + "file_info": { + "id": "file-1", + "name": "image.png", + "size": 10, + "extension": "png", + "mime_type": "image/png", + "source_url": "signed-url", + }, + } + ] + session.scalar.assert_called_once() + mock_db.session.scalar.assert_not_called() def test_dataset_retrieval_model_dict_property(self): """Test retrieval_model_dict property with default values.""" @@ -301,6 +448,30 @@ class TestDocumentModelRelationships: assert "notion_import" in Document.DATA_SOURCES assert "website_crawl" in Document.DATA_SOURCES + def test_session_aware_document_getters_use_caller_session(self): + document = Document( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch_001", + name="test.pdf", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + dataset_process_rule_id=str(uuid4()), + ) + process_rule = Mock() + session = Mock() + session.get.return_value = process_rule + session.scalar.side_effect = [3, 7] + + assert document.get_dataset_process_rule(session=session) is process_rule + assert document.get_segment_count(session=session) == 3 + assert document.get_hit_count(session=session) == 7 + + session.get.assert_called_once_with(DatasetProcessRule, document.dataset_process_rule_id) + assert session.scalar.call_count == 2 + def test_document_display_status_queuing(self): """Test document display_status property for queuing state.""" # Arrange @@ -504,6 +675,24 @@ class TestDocumentModelRelationships: # Assert assert result == {} + def test_document_get_dataset_uses_caller_session(self): + document = Document( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch_001", + name="test.pdf", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + ) + dataset = Mock(spec=Dataset) + session = Mock() + session.get.return_value = dataset + + assert document.get_dataset(session=session) is dataset + session.get.assert_called_once_with(Dataset, document.dataset_id) + def test_document_average_segment_length(self): """Test average_segment_length property calculation.""" # Arrange @@ -552,6 +741,85 @@ class TestDocumentModelRelationships: class TestDocumentSegmentIndexing: """Test suite for DocumentSegment model indexing and operations.""" + def test_get_child_chunks_uses_caller_session(self): + segment = DocumentSegment( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + document_id=str(uuid4()), + position=1, + content="Test content", + word_count=2, + tokens=5, + created_by=str(uuid4()), + ) + document = Mock(spec=Document) + process_rule = Mock(mode="hierarchical", rules_dict={"parent_mode": "paragraph"}) + document.get_dataset_process_rule.return_value = process_rule + child_chunk = Mock(spec=ChildChunk) + session = Mock() + session.get.return_value = document + session.scalars.return_value.all.return_value = [child_chunk] + + with patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode="paragraph")): + result = segment.get_child_chunks(session=session) + + assert result == [child_chunk] + session.get.assert_called_once_with(Document, segment.document_id) + document.get_dataset_process_rule.assert_called_once_with(session=session) + session.scalars.assert_called_once() + + def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self): + segment = DocumentSegment( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + document_id=str(uuid4()), + position=1, + content="Test content", + word_count=2, + tokens=5, + created_by=str(uuid4()), + ) + document = Mock(spec=Document) + document.get_dataset_process_rule.return_value = Mock( + mode="hierarchical", + rules_dict={"parent_mode": ParentMode.FULL_DOC}, + ) + session = Mock() + session.get.return_value = document + child_chunk = Mock(spec=ChildChunk) + session.scalars.return_value.all.return_value = [child_chunk] + + with patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode=ParentMode.FULL_DOC)): + result = segment.get_child_chunks(session=session) + response_result = segment.get_child_chunks(session=session, include_full_doc=False) + + assert result == [child_chunk] + assert response_result == [] + session.scalars.assert_called_once() + + def test_relationship_getters_use_caller_session(self): + segment = DocumentSegment( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + document_id=str(uuid4()), + position=1, + content="Test content", + word_count=2, + tokens=5, + created_by=str(uuid4()), + ) + dataset = Mock(spec=Dataset) + document = Mock(spec=Document) + session = Mock() + session.get.side_effect = [dataset, document] + + assert segment.get_dataset(session=session) is dataset + assert segment.get_document(session=session) is document + assert session.get.call_args_list == [ + call(Dataset, segment.dataset_id), + call(Document, segment.document_id), + ] + def test_document_segment_creation_with_required_fields(self): """Test creating a document segment with all required fields.""" # Arrange @@ -742,11 +1010,11 @@ class TestDocumentSegmentIndexing: monkeypatch.setattr("models.dataset.dify_config.FILES_URL", "https://files.example.com") monkeypatch.setattr("models.dataset.dify_config.CONSOLE_API_URL", "https://console.example.com") - with patch("models.dataset.db") as mock_db: - mock_db.session.execute.return_value.all.return_value = [(Mock(), attachment)] + session = Mock() + session.execute.return_value.all.return_value = [(Mock(), attachment)] - # Act - attachments = segment.attachments + # Act + attachments = segment.get_attachments(session=session) # Assert assert len(attachments) == 1 @@ -758,6 +1026,7 @@ class TestDocumentSegmentIndexing: assert query["timestamp"] == ["1700000000"] assert query["nonce"] == ["01010101010101010101010101010101"] assert query["sign"][0] + session.execute.assert_called_once() def test_document_segment_error_tracking(self): """Test document segment error tracking.""" @@ -985,6 +1254,21 @@ class TestDatasetKeywordTable: # Assert assert keyword_table.data_source_type == "file" + def test_get_keyword_table_dict_from_database_uses_caller_session(self): + dataset = Mock(tenant_id="tenant-1") + session = Mock() + session.scalar.return_value = dataset + keyword_table = DatasetKeywordTable( + dataset_id="dataset-1", + keyword_table=json.dumps({"__data__": {"table": {"keyword": ["node-1"]}}}), + data_source_type="database", + ) + + result = keyword_table.get_keyword_table_dict(session=session) + + assert result == {"__data__": {"table": {"keyword": {"node-1"}}}} + session.scalar.assert_called_once() + class TestAppDatasetJoin: """Test suite for AppDatasetJoin model.""" diff --git a/api/tests/unit_tests/models/test_model.py b/api/tests/unit_tests/models/test_model.py index a87dd7f15ad..bb3206713d7 100644 --- a/api/tests/unit_tests/models/test_model.py +++ b/api/tests/unit_tests/models/test_model.py @@ -1,5 +1,6 @@ import importlib import types +from unittest.mock import MagicMock, patch import pytest @@ -121,3 +122,19 @@ def test_inputs_restore_external_remote_url_file_mappings(owner_cls: type[Conver assert restored_file.transfer_method == FileTransferMethod.REMOTE_URL assert restored_file.remote_url == "https://example.com/report.pdf" + + +def test_message_inputs_resolve_file_tenant_with_caller_session() -> None: + message = Message(app_id="app-1") + message.inputs = {"file": _build_local_file_mapping("upload-1")} + session = MagicMock() + session.scalar.return_value = "tenant-1" + + with patch( + "models.model.build_file_from_input_mapping", + side_effect=lambda **kwargs: kwargs["tenant_resolver"](), + ): + inputs = message.inputs_with_session(session=session) + + assert inputs["file"] == "tenant-1" + session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/models/test_workflow.py b/api/tests/unit_tests/models/test_workflow.py index d80c0f45ff8..12f45e21beb 100644 --- a/api/tests/unit_tests/models/test_workflow.py +++ b/api/tests/unit_tests/models/test_workflow.py @@ -10,6 +10,7 @@ from factories.variable_factory import build_segment from graphon.file import File, FileTransferMethod, FileType from graphon.variables import FloatVariable, IntegerVariable, SecretVariable, StringVariable from graphon.variables.segments import IntegerSegment, Segment +from models.account import Account from models.workflow import ( Workflow, WorkflowDraftVariable, @@ -134,6 +135,57 @@ def test_to_dict(): assert workflow_dict["environment_variables"][1]["value"] == "text" +def test_workflow_account_getters_use_caller_session(): + workflow = Workflow( + tenant_id="tenant_id", + app_id="app_id", + type="workflow", + version="draft", + graph="{}", + features="{}", + created_by="created-account-id", + environment_variables=[], + conversation_variables=[], + ) + workflow.updated_by = "updated-account-id" + created_account = mock.Mock(spec=Account) + updated_account = mock.Mock(spec=Account) + session = mock.Mock() + session.get.side_effect = [created_account, updated_account] + + with mock.patch("models.workflow.db") as mock_db: + assert workflow.get_created_by_account(session=session) is created_account + assert workflow.get_updated_by_account(session=session) is updated_account + + assert session.get.call_args_list == [ + mock.call(Account, "created-account-id"), + mock.call(Account, "updated-account-id"), + ] + mock_db.session.get.assert_not_called() + + +def test_workflow_tool_published_getter_uses_caller_session(): + workflow = Workflow( + tenant_id="tenant_id", + app_id="app_id", + type="workflow", + version="draft", + graph="{}", + features="{}", + created_by="account_id", + environment_variables=[], + conversation_variables=[], + ) + session = mock.Mock() + session.execute.return_value.scalar_one.return_value = True + + with mock.patch("models.workflow.db") as mock_db: + assert workflow.get_tool_published(session=session) is True + + session.execute.assert_called_once() + mock_db.session.execute.assert_not_called() + + def test_normalize_environment_variable_mappings_converts_full_mask_to_hidden_value(): normalized = Workflow.normalize_environment_variable_mappings( [ diff --git a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py index b74f5a53ec4..926a04f53f9 100644 --- a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py +++ b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py @@ -608,17 +608,21 @@ def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(mon }, } ) - monkeypatch.setattr( - "services.agent.dsl_service.get_tenant_knowledge_dataset_rows", - Mock(return_value={"existing": SimpleNamespace(id="existing")}), - ) + session = Mock() + get_dataset_rows = Mock(return_value={"existing": SimpleNamespace(id="existing")}) + monkeypatch.setattr("services.agent.dsl_service.get_tenant_knowledge_dataset_rows", get_dataset_rows) - resolved, warnings = AgentDslService(Mock())._resolve_package_soul( + resolved, warnings = AgentDslService(session)._resolve_package_soul( tenant_id="tenant-1", package=make_portable_agent_package(_agent(), soul), package_path="agent_packages.agent_1", ) + get_dataset_rows.assert_called_once_with( + session=session, + tenant_id="tenant-1", + dataset_ids=["existing", "missing"], + ) datasets = resolved.knowledge.sets[0].datasets assert datasets[0].id == "existing" assert datasets[1].id is not None 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 2480a6ddf40..e19fa36dca9 100644 --- a/api/tests/unit_tests/services/agent/test_agent_services.py +++ b/api/tests/unit_tests/services/agent/test_agent_services.py @@ -136,11 +136,12 @@ def test_get_published_agent_soul_for_app_returns_none_without_backing_agent(): def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() 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", session=composer_service.db.session + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1" ) assert result["binding"] is None @@ -156,6 +157,7 @@ def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPa def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() binding = SimpleNamespace( agent_id="agent-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, @@ -180,13 +182,14 @@ def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest. ) result = AgentComposerService.load_workflow_composer( - tenant_id="tenant-1", app_id="app-1", node_id="node-1", session=composer_service.db.session + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1" ) assert result == {"agent": "agent-1", "version": "version-1"} def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() binding = SimpleNamespace( agent_id="agent-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, @@ -212,17 +215,18 @@ def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest ) result = AgentComposerService.load_workflow_composer( + session=session, tenant_id="tenant-1", 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"} def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() binding = SimpleNamespace( agent_id="inline-agent-1", binding_type=WorkflowAgentBindingType.INLINE_AGENT, @@ -255,17 +259,18 @@ def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest ) result = AgentComposerService.load_workflow_composer( + session=session, tenant_id="tenant-1", 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"} def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() captured: dict[str, object] = {} class FakeRosterService: @@ -282,20 +287,23 @@ def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch agent = SimpleNamespace(id="inline-agent-1", scope=AgentScope.WORKFLOW_ONLY) debug_conversation_id = AgentComposerService._workflow_inline_debug_conversation_id( + session=session, tenant_id="tenant-1", binding=binding, agent=agent, account_id="account-1", - session="session-1", ) assert debug_conversation_id == "debug-conversation-1" assert captured["tenant_id"] == "tenant-1" assert captured["agent_id"] == "inline-agent-1" assert captured["account_id"] == "account-1" + assert captured["commit"] is False def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() + class UnexpectedRosterService: def __init__(self, session): raise AssertionError("roster service should not be used") @@ -304,37 +312,38 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p assert ( AgentComposerService._workflow_inline_debug_conversation_id( + session=session, tenant_id="tenant-1", binding=SimpleNamespace(binding_type=WorkflowAgentBindingType.ROSTER_AGENT), agent=SimpleNamespace(id="agent-1", scope=AgentScope.ROSTER), account_id="account-1", - session="session-1", ) is None ) assert ( AgentComposerService._workflow_inline_debug_conversation_id( + session=session, tenant_id="tenant-1", 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 ) def test_load_workflow_composer_rejects_preview_without_binding(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) with pytest.raises(AgentVersionNotFoundError): AgentComposerService.load_workflow_composer( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", snapshot_id="preview-version", - session=composer_service.db.session, ) @@ -358,7 +367,7 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, calls = [] serialize_calls = [] - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) @@ -393,22 +402,23 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, ) result = AgentComposerService.save_workflow_composer( + session=session, 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": []} assert result == {"state": "ok"} assert calls assert serialize_calls[0]["account_id"] == "account-1" - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_save_workflow_composer_rejects_agent_app_variant(): + session = FakeSession() payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -419,12 +429,12 @@ def test_save_workflow_composer_rejects_agent_app_variant(): with pytest.raises(ValueError): AgentComposerService.save_workflow_composer( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", payload=payload, - session=composer_service.db.session, ) @@ -472,7 +482,7 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. fake_session = FakeSession(scalar=[None]) saved_draft = SimpleNamespace(id="draft-1", config_snapshot_dict={"prompt": {"system_prompt": "x"}}) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **kwargs: saved_draft) monkeypatch.setattr(AgentComposerService, "load_agent_composer", lambda **kwargs: {"loaded": True}) @@ -486,11 +496,11 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. ) result = AgentComposerService.save_agent_app_composer( + session=session, 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": []} @@ -498,10 +508,11 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. assert fake_session.added[0].name == "Analyst" assert fake_session.added[0].active_config_snapshot_id is None assert fake_session.added[0].active_config_is_published is False - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="version-1", @@ -521,14 +532,13 @@ 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", session=composer_service.db.session - ) + result = AgentComposerService.load_agent_app_composer(session=session, tenant_id="tenant-1", app_id="app-1") assert result["save_options"] == [ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION.value] def test_save_agent_app_composer_rejects_version_save_strategy(): + session = FakeSession() payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -539,11 +549,11 @@ def test_save_agent_app_composer_rejects_version_save_strategy(): with pytest.raises(InvalidComposerConfigError, match="Use the publish endpoint"): AgentComposerService.save_agent_app_composer( + session=session, tenant_id="tenant-1", app_id="app-1", account_id="account-1", payload=payload, - session=composer_service.db.session, ) @@ -559,7 +569,7 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey fake_session = FakeSession(scalar=[agent]) saved = {} - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr( AgentComposerService, @@ -577,11 +587,11 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey ) result = AgentComposerService.save_agent_app_composer( + session=session, 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": []} @@ -590,7 +600,7 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey assert saved["agent_soul"].model_dump(mode="json") == _agent_soul_with_model().model_dump(mode="json") assert agent.active_config_is_published is False assert fake_session._scalar == [] - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snapshot(monkeypatch: pytest.MonkeyPatch): @@ -605,7 +615,7 @@ def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snaps active_version = SimpleNamespace(config_snapshot_dict=agent_soul.model_dump(mode="json")) fake_session = FakeSession(scalar=[agent], scalars=[[AgentConfigRevisionOperation.PUBLISH_DRAFT]]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr( AgentComposerService, @@ -623,11 +633,15 @@ 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, session=fake_session + session=session, + tenant_id="tenant-1", + app_id="app-1", + account_id="account-1", + payload=payload, ) assert agent.active_config_is_published is True - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.MonkeyPatch): @@ -659,7 +673,6 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke def fail_validate_knowledge_datasets(**_kwargs): raise AssertionError("knowledge datasets must not be validated when Agent Soul has no model") - monkeypatch.setattr(composer_service.db, "session", fake_session) monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", fail_validate_knowledge_datasets) monkeypatch.setattr(AgentComposerService, "_create_config_version", fail_create_config_version) @@ -677,6 +690,7 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke assert agent.active_config_snapshot_id == "version-1" assert agent.active_config_is_published is False assert draft.base_snapshot_id == "version-1" + assert fake_session.flushes == 0 assert fake_session.commits == 0 @@ -704,7 +718,7 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. fake_session = FakeSession(scalar=[agent, draft]) created: dict[str, object] = {} - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **kwargs: None) monkeypatch.setattr( @@ -715,11 +729,11 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. monkeypatch.setattr(AgentComposerService, "_serialize_version", lambda _version: {"id": _version.id}) result = AgentComposerService.publish_agent_app_draft( + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", version_note="ship it", - session=composer_service.db.session, ) assert result["result"] == "success" @@ -730,7 +744,7 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. assert agent.active_config_snapshot_id == "version-2" assert agent.active_config_has_model is True assert agent.active_config_is_published is True - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkeypatch: pytest.MonkeyPatch): @@ -757,13 +771,13 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey ) fake_session = FakeSession(scalar=[agent, normal_draft, None]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session checked_out = AgentComposerService.checkout_agent_app_build_draft( + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", - session=composer_service.db.session, ) build_draft = fake_session.added[0] @@ -772,20 +786,20 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey assert checked_out["draft"]["account_id"] == "account-1" assert checked_out["draft"]["base_snapshot_id"] == "version-1" assert checked_out["agent_soul"] == normal_draft.config_snapshot_dict - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 active_version = SimpleNamespace(config_snapshot_dict=build_draft.config_snapshot_dict) fake_session = FakeSession( scalar=[agent, build_draft, normal_draft, active_version], scalars=[[AgentConfigRevisionOperation.PUBLISH_DRAFT]], ) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session applied = AgentComposerService.apply_agent_app_build_draft( + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", - session=composer_service.db.session, ) assert applied["result"] == "success" @@ -793,7 +807,7 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey assert normal_draft.config_snapshot_dict == build_draft.config_snapshot_dict assert agent.active_config_is_published is True assert fake_session.deleted == [build_draft] - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 @pytest.mark.parametrize( @@ -906,19 +920,19 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( ) active_version = SimpleNamespace(config_snapshot_dict=active_agent_soul.model_dump(mode="json")) fake_session = FakeSession(scalar=[agent, build_draft, normal_draft, active_version]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session AgentComposerService.apply_agent_app_build_draft( + session=session, 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 assert agent.active_config_is_published is False assert fake_session.deleted == [build_draft] - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatch): @@ -926,7 +940,7 @@ def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatc SimpleNamespace(app_id="app-1", workflow_id="workflow-1", node_id="node-1"), SimpleNamespace(app_id="app-1", workflow_id="workflow-1", node_id="node-2"), ] - monkeypatch.setattr(composer_service.db, "session", FakeSession(scalars=[bindings])) + session = FakeSession(scalars=[bindings]) # Candidates assembly is covered in test_composer_candidates.py; here we stub # the IO loaders and assert the response envelope per variant (ENG-615). @@ -938,20 +952,17 @@ 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( + session=session, 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", - session=composer_service.db.session, + session=session, tenant_id="tenant-1", agent_id="agent-1", user_id="account-1" ) impact = AgentComposerService.calculate_impact( - tenant_id="tenant-1", current_snapshot_id="version-1", session=composer_service.db.session + session=session, tenant_id="tenant-1", current_snapshot_id="version-1" ) assert workflow_candidates["variant"] == "workflow" @@ -964,6 +975,7 @@ def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatc def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -990,7 +1002,7 @@ def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pyt monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) state = AgentComposerService._serialize_workflow_state( - binding=binding, agent=agent, version=version, session=composer_service.db.session + session=session, binding=binding, agent=agent, version=version ) assert state["soul_lock"]["locked"] is True @@ -1007,6 +1019,7 @@ def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pyt def test_serialize_workflow_state_passes_user_declared_outputs_through_effective(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -1031,7 +1044,7 @@ def test_serialize_workflow_state_passes_user_declared_outputs_through_effective monkeypatch.setattr(AgentComposerService, "calculate_impact", lambda **kwargs: {"workflow_node_count": 1}) state = AgentComposerService._serialize_workflow_state( - binding=binding, agent=agent, version=version, session=composer_service.db.session + session=session, binding=binding, agent=agent, version=version ) # When the user has declared outputs, effective_declared_outputs is the same @@ -1045,6 +1058,7 @@ def test_serialize_workflow_state_passes_user_declared_outputs_through_effective def test_serialize_workflow_state_includes_inline_debug_conversation_message_state( monkeypatch: pytest.MonkeyPatch, ): + session = FakeSession() binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -1078,11 +1092,11 @@ def test_serialize_workflow_state_includes_inline_debug_conversation_message_sta ) state = AgentComposerService._serialize_workflow_state( + session=session, binding=binding, agent=agent, version=version, account_id="account-1", - session=composer_service.db.session, ) assert state["debug_conversation_id"] == "debug-conversation-1" @@ -1092,7 +1106,7 @@ def test_serialize_workflow_state_includes_inline_debug_conversation_message_sta def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session workflow_agent = SimpleNamespace(id="inline-agent-1", active_config_snapshot_id="inline-version-1") roster_agent = SimpleNamespace( id="roster-agent-1", @@ -1157,6 +1171,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk existing_binding = WorkflowAgentNodeBinding(agent_id="inline-agent-1", current_snapshot_id="inline-version-1") updated_binding = AgentComposerService._save_node_job_only( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1164,9 +1179,9 @@ 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( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1174,9 +1189,9 @@ 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( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1184,9 +1199,9 @@ 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( + session=session, tenant_id="tenant-1", account_id="account-1", binding=WorkflowAgentNodeBinding( @@ -1198,14 +1213,13 @@ 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( + session=session, 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" @@ -1222,6 +1236,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk assert create_roster_calls[1]["role"] == "Copied role" assert create_roster_calls[1]["icon"] == "copied" assert create_roster_calls[1]["icon_background"] == "#E0F2FE" + copy_drive_calls[0].pop("session", None) assert copy_drive_calls == [ { "tenant_id": "tenant-1", @@ -1230,14 +1245,13 @@ 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, } ] def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session inline_agent = SimpleNamespace( id="inline-agent-1", scope=AgentScope.WORKFLOW_ONLY, @@ -1290,6 +1304,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch ) updated_binding = AgentComposerService._save_node_job_only( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1297,7 +1312,6 @@ 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" @@ -1310,7 +1324,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session created_agent = SimpleNamespace(id="inline-agent-1", active_config_snapshot_id="inline-version-1") captured: dict[str, object] = {} @@ -1343,6 +1357,7 @@ def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pyte ) updated_binding = AgentComposerService._save_node_job_only( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1350,7 +1365,6 @@ 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 @@ -1369,6 +1383,7 @@ def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pyte def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_id(): + session = FakeSession() binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1393,6 +1408,7 @@ def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_i with pytest.raises(ValueError, match="Start from Scratch"): AgentComposerService._save_node_job_only( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1400,13 +1416,12 @@ 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, ) def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session current_snapshot = AgentConfigSnapshot( id="inline-version-1", tenant_id="tenant-1", @@ -1441,6 +1456,7 @@ def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypat with pytest.raises(ValueError, match="workflow-only agent"): AgentComposerService._save_node_job_only( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -1448,7 +1464,6 @@ 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, ) @@ -1456,7 +1471,7 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n monkeypatch: pytest.MonkeyPatch, ): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session workflow = SimpleNamespace(id="workflow-1") node_job = WorkflowNodeJobConfig(workflow_prompt="keep this node task") binding = WorkflowAgentNodeBinding( @@ -1529,13 +1544,13 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n ) state = AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", 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 @@ -1549,10 +1564,11 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n drive_kwargs = captured["drive"] assert drive_kwargs["source_agent_id"] == "roster-agent-1" assert drive_kwargs["target_agent_id"] == "inline-agent-1" - assert fake_session.commits == 1 + assert fake_session.flushes >= 1 def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr( AgentComposerService, @@ -1590,13 +1606,13 @@ def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkey with pytest.raises(AgentVersionConflictError): AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", source_agent_id="roster-agent-1", source_snapshot_id="roster-version-1", - session=composer_service.db.session, ) @@ -1628,7 +1644,7 @@ def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(mo config_snapshot='{"prompt":{"system_prompt":"inline"}}', ) serialize_calls = [] - monkeypatch.setattr(composer_service.db, "session", FakeSession()) + session = FakeSession() monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: inline_binding) monkeypatch.setattr(AgentComposerService, "_get_agent_if_present", lambda **kwargs: inline_agent) @@ -1641,13 +1657,13 @@ def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(mo monkeypatch.setattr(AgentComposerService, "_serialize_workflow_state", serialize_workflow_state) state = AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", 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} @@ -1695,6 +1711,7 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( source_status: AgentStatus, expected_message: str, ): + session = FakeSession() binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1721,12 +1738,12 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( with pytest.raises(InvalidComposerConfigError, match=expected_message): AgentComposerService.copy_workflow_composer_from_roster( + session=session, tenant_id="tenant-1", app_id="app-1", node_id="node-1", account_id="account-1", source_agent_id="roster-agent-1", - session=composer_service.db.session, ) @@ -1764,7 +1781,7 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest mime_type="application/pdf", ) fake_session = FakeSession(scalars=[[skill_row, script_row, file_row], []]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session agent_soul = AgentSoulConfig.model_validate( { "prompt": { @@ -1777,13 +1794,13 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest ) AgentComposerService._copy_agent_drive_rows( + session=session, tenant_id="tenant-1", source_agent_id="roster-agent-1", target_agent_id="inline-agent-1", 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)] @@ -1800,16 +1817,16 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "No drive mentions."}}) AgentComposerService._copy_agent_drive_rows( + session=session, tenant_id="tenant-1", source_agent_id="roster-agent-1", target_agent_id="inline-agent-1", account_id="account-1", agent_soul=agent_soul, - session=composer_service.db.session, ) assert fake_session.added == [] @@ -1827,16 +1844,16 @@ def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.Mo mime_type="application/pdf", ) fake_session = FakeSession(scalars=[[source_row], ["files/qna.pdf"]]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "[§file:files/qna.pdf:qna.pdf§]"}}) AgentComposerService._copy_agent_drive_rows( + session=session, tenant_id="tenant-1", source_agent_id="roster-agent-1", 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)] == [] @@ -1886,7 +1903,7 @@ def test_drive_copy_scopes_include_declared_output_benchmark_files(): def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session created_apps = [] hidden_backing_apps = [] backing_agent = Agent( @@ -1900,7 +1917,7 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes ) class FakeAppService: - def create_app(self, tenant_id, params, account, session): + def create_app(self, tenant_id, params, account, *, session): created_apps.append((tenant_id, params, account)) return SimpleNamespace(id="app-agent-1") @@ -1932,22 +1949,22 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes ) workflow_agent = AgentComposerService._create_workflow_only_agent( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", 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( + session=session, tenant_id="tenant-1", account_id="account-1", name="Ready Agent", 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" @@ -1967,24 +1984,24 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes 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)) + session = SimpleNamespace(get=lambda model, account_id: account) - assert AgentComposerService._require_account(account_id="account-1", session=composer_service.db.session) is account + assert AgentComposerService._require_account(session=session, account_id="account-1") 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)) + session = SimpleNamespace(get=lambda model, account_id: None) with pytest.raises(ValueError, match="Account not found"): - AgentComposerService._require_account(account_id="missing-account", session=composer_service.db.session) + AgentComposerService._require_account(session=session, account_id="missing-account") def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session class FakeAppService: - def create_app(self, tenant_id, params, account, session): + def create_app(self, tenant_id, params, account, *, session): raise IntegrityError("insert apps", params, Exception("duplicate")) monkeypatch.setattr(composer_service, "AppService", FakeAppService) @@ -1992,13 +2009,13 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte with pytest.raises(AgentNameConflictError): AgentComposerService._create_roster_agent_for_composer( + session=session, tenant_id="tenant-1", account_id="account-1", name="Duplicate Agent", agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, - session=composer_service.db.session, ) assert fake_session.rollbacks == 1 @@ -2006,10 +2023,10 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession() - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session class FakeAppService: - def create_app(self, tenant_id, params, account, session): + def create_app(self, tenant_id, params, account, *, session): return SimpleNamespace(id="app-agent-1") class FakeAgentRosterService: @@ -2025,13 +2042,13 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa with pytest.raises(AgentNotFoundError): AgentComposerService._create_roster_agent_for_composer( + session=session, tenant_id="tenant-1", account_id="account-1", name="Missing Backing Agent", agent_soul=_agent_soul_with_model(), operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, - session=composer_service.db.session, ) @@ -2045,15 +2062,15 @@ def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeyp ) snapshot = SimpleNamespace(config_snapshot_dict=agent_soul) fake_session = FakeSession(scalars=[[AgentConfigRevisionOperation.CREATE_VERSION]]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **kwargs: snapshot) assert ( AgentComposerService._agent_soul_matches_active_config( + session=session, tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, - session=fake_session, ) is False ) @@ -2069,15 +2086,15 @@ def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monke ) snapshot = SimpleNamespace(config_snapshot_dict=agent_soul) fake_session = FakeSession(scalars=[[AgentConfigRevisionOperation.PUBLISH_DRAFT]]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **kwargs: snapshot) assert ( AgentComposerService._agent_soul_matches_active_config( + session=session, tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, - session=fake_session, ) is True ) @@ -2098,19 +2115,20 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa None, ] ) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "new"}}) version = AgentComposerService._create_config_version( + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note="note", - session=composer_service.db.session, ) updated_snapshot = AgentComposerService._update_current_version( + session=session, current_snapshot=AgentConfigSnapshot( id="version-1", tenant_id="tenant-1", @@ -2122,39 +2140,32 @@ 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(session=session, tenant_id="tenant-1", app_id="app-1") with pytest.raises(ValueError): - AgentComposerService._get_draft_workflow( - tenant_id="tenant-1", app_id="missing", session=composer_service.db.session - ) + AgentComposerService._get_draft_workflow(session=session, tenant_id="tenant-1", app_id="missing") assert ( - AgentComposerService._require_agent( - tenant_id="tenant-1", agent_id="agent-1", session=composer_service.db.session - ).id - == "agent-1" + AgentComposerService._require_agent(session=session, 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, 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 - ) + AgentComposerService._require_agent(session=session, tenant_id="tenant-1", agent_id=None) + assert AgentComposerService._get_agent_if_present(session=session, tenant_id="tenant-1", agent_id="agent-1") is None assert ( AgentComposerService._require_version( - tenant_id="tenant-1", agent_id="agent-1", version_id="version-1", session=composer_service.db.session + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + version_id="version-1", ).id == "version-1" ) with pytest.raises(composer_service.AgentVersionNotFoundError): AgentComposerService._require_version( - tenant_id="tenant-1", agent_id="agent-1", version_id="missing", session=composer_service.db.session + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + version_id="missing", ) assert version.version == 2 @@ -2164,7 +2175,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatch): fake_session = FakeSession(scalar=[2]) - monkeypatch.setattr(composer_service.db, "session", fake_session) + session = fake_session payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.WORKFLOW.value, @@ -2189,11 +2200,11 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc ) result = AgentComposerService._save_to_current_version( + session=session, tenant_id="tenant-1", account_id="account-1", binding=binding, payload=payload, - session=composer_service.db.session, ) assert result.updated_by == "account-1" @@ -2202,6 +2213,7 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc AgentComposerService._require_binding(None) with pytest.raises(ValueError): AgentComposerService._save_as_new_agent( + session=session, tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", @@ -2214,7 +2226,6 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc "save_strategy": ComposerSaveStrategy.SAVE_AS_NEW_AGENT.value, } ), - session=composer_service.db.session, ) @@ -3452,10 +3463,12 @@ class TestAgentAppBackingAgent: use_icon_as_answer_icon=True, tracing="{}", app_model_config=source_config, + app_model_config_with_session=lambda *, session: source_config, ) target_app = SimpleNamespace( id="target-app", app_model_config=target_config, + app_model_config_with_session=lambda *, session: target_config, enable_site=True, enable_api=True, use_icon_as_answer_icon=False, @@ -3513,7 +3526,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object, session: object) -> object: + def create_app(self, tenant_id: str, params, account: object, *, session) -> object: captured["tenant_id"] = tenant_id captured["params"] = params captured["account"] = account @@ -3582,7 +3595,7 @@ class TestAgentAppBackingAgent: captured: dict[str, object] = {} class FakeAppService: - def create_app(self, tenant_id: str, params, account: object, session: object) -> object: + def create_app(self, tenant_id: str, params, account: object, *, session) -> object: captured["params"] = params return target_app @@ -3644,7 +3657,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, session: object) -> object: + def create_app(self, tenant_id: str, params, account: object, *, session) -> object: return target_app access_mode_updates = [] @@ -4633,7 +4646,7 @@ def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): (placeholder semantics), never reach the UUID-typed dataset query (E2E 500).""" captured = {} - def fake_get_datasets_by_ids(ids, tenant_id): + def fake_get_datasets_by_ids(ids, tenant_id, *, session): captured["ids"] = ids return [], 0 @@ -4643,46 +4656,22 @@ def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(dataset_service_module.DatasetService, "get_datasets_by_ids", fake_get_datasets_by_ids) valid = "550e8400-e29b-41d4-a716-446655440000" - rows = get_tenant_knowledge_dataset_rows(tenant_id="tenant-1", dataset_ids=["9999dead-beef", valid]) + rows = get_tenant_knowledge_dataset_rows( + session=FakeSession(), tenant_id="tenant-1", dataset_ids=["9999dead-beef", valid] + ) assert rows == {} assert captured["ids"] == [valid] # all-malformed input never touches the DB captured.clear() - assert get_tenant_knowledge_dataset_rows(tenant_id="tenant-1", dataset_ids=["nope"]) == {} + assert get_tenant_knowledge_dataset_rows(session=FakeSession(), tenant_id="tenant-1", dataset_ids=["nope"]) == {} assert captured == {} -@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): +def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pytest.MonkeyPatch): captured = {"calls": 0} - def fake_get_datasets_by_ids(ids, tenant_id): + def fake_get_datasets_by_ids(ids, tenant_id, *, session): captured["calls"] += 1 captured["ids"] = ids captured["tenant_id"] = tenant_id @@ -4704,49 +4693,23 @@ def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pyte "retrieval": {"mode": "multiple", "top_k": 4}, } ] - }, + } } ) with pytest.raises(InvalidComposerConfigError, match="not-a-uuid"): - AgentComposerService.validate_knowledge_datasets(tenant_id="tenant-1", agent_soul=agent_soul) + AgentComposerService.validate_knowledge_datasets( + session=FakeSession(), tenant_id="tenant-1", agent_soul=agent_soul + ) assert captured == {"calls": 0} -@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 -): +def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets(monkeypatch: pytest.MonkeyPatch): captured = {} missing_dataset_id = "550e8400-e29b-41d4-a716-446655440000" - def fake_get_datasets_by_ids(ids, tenant_id): + def fake_get_datasets_by_ids(ids, tenant_id, *, session): captured["ids"] = ids captured["tenant_id"] = tenant_id return [], 0 @@ -4767,12 +4730,14 @@ def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets( "retrieval": {"mode": "multiple", "top_k": 4}, } ] - }, + } } ) with pytest.raises(InvalidComposerConfigError, match=missing_dataset_id): - AgentComposerService.validate_knowledge_datasets(tenant_id="tenant-1", agent_soul=agent_soul) + AgentComposerService.validate_knowledge_datasets( + session=FakeSession(), tenant_id="tenant-1", agent_soul=agent_soul + ) assert captured == {"ids": [missing_dataset_id], "tenant_id": "tenant-1"} @@ -4791,7 +4756,6 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte import services.dataset_service as dataset_service_module - monkeypatch.setattr(composer_service.db, "session", fake_session) monkeypatch.setattr( dataset_service_module.DatasetService, "get_datasets_by_ids", @@ -4827,11 +4791,11 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte ) result = AgentComposerService.save_agent_composer( + session=fake_session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", payload=payload, - session=fake_session, ) assert result["loaded"] is True @@ -4840,7 +4804,7 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte assert saved["agent_soul"].knowledge.sets[0].retrieval.model is None assert saved["agent_soul"].knowledge.sets[0].metadata_filtering.mode == "automatic" assert saved["agent_soul"].knowledge.sets[0].metadata_filtering.metadata_model_config is None - assert fake_session.commits == 1 + assert fake_session.flushes == 1 def test_workspace_dify_tools_returns_provider_and_tool_granularities(monkeypatch: pytest.MonkeyPatch): @@ -4901,7 +4865,6 @@ def _drive_soul(**overrides): def _patch_drive_keys(monkeypatch, existing_keys): - import services.agent.composer_service as composer_service_module captured: dict[str, object] = {} @@ -4909,18 +4872,18 @@ def _patch_drive_keys(monkeypatch, existing_keys): captured["stmt"] = stmt return list(existing_keys) - monkeypatch.setattr(composer_service_module.db, "session", type("S", (), {"scalars": staticmethod(fake_scalars)})()) + captured["session"] = type("S", (), {"scalars": staticmethod(fake_scalars)})() return captured def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyPatch): - _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md"]) + session = _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md"])["session"] findings = AgentComposerService._drive_mention_findings( + session=session, 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")] @@ -4929,27 +4892,28 @@ def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyP def test_drive_mention_findings_clean_when_all_keys_exist(monkeypatch: pytest.MonkeyPatch): - _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md", "files/sample.pdf"]) + session = _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md", "files/sample.pdf"])["session"] assert ( AgentComposerService._drive_mention_findings( + session=session, tenant_id="tenant-1", agent_id="agent-1", prompt=_drive_soul().prompt.system_prompt, - session=composer_service.db.session, ) == [] ) def test_drive_mention_findings_skips_prompt_without_drive_mentions(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() # No drive-backed mention at all -> no DB roundtrip, no findings. soul = _drive_soul(prompt={"system_prompt": "Use [§knowledge:kb-1:Docs§]."}) findings = AgentComposerService._drive_mention_findings( + session=session, tenant_id="tenant-1", agent_id="agent-1", prompt=soul.prompt.system_prompt, - session=composer_service.db.session, ) assert findings == [] @@ -4959,7 +4923,7 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c ): from services.entities.agent_entities import ComposerSavePayload - _patch_drive_keys(monkeypatch, existing_keys=[]) + session = _patch_drive_keys(monkeypatch, existing_keys=[])["session"] payload = ComposerSavePayload.model_validate( { "variant": "agent_app", @@ -4969,10 +4933,7 @@ 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", - session=composer_service.db.session, + session=session, tenant_id="tenant-1", payload=payload, agent_id="agent-1" ) codes = {w["code"] for w in findings["warnings"]} @@ -4983,7 +4944,7 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c } # without agent context the drive check is skipped entirely findings_no_agent = AgentComposerService.collect_validation_findings( - tenant_id="tenant-1", payload=payload, session=composer_service.db.session + session=session, tenant_id="tenant-1", payload=payload ) assert all(w["code"] != "mention_target_missing" for w in findings_no_agent["warnings"]) @@ -4994,18 +4955,12 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c def test_resolve_bound_agent_id_queries_active_roster_agent(monkeypatch: pytest.MonkeyPatch): from types import SimpleNamespace - 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", session=composer_service.db.session - ) - == "agent-9" - ) + session = SimpleNamespace(scalar=lambda stmt: "agent-9") + assert AgentComposerService.resolve_bound_agent_id(session=session, tenant_id="t-1", app_id="app-1") == "agent-9" def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(monkeypatch: pytest.MonkeyPatch): + session = FakeSession() from types import SimpleNamespace def boom(cls, **kwargs): @@ -5013,9 +4968,7 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", classmethod(boom)) assert ( - AgentComposerService.resolve_workflow_node_agent_id( - tenant_id="t", app_id="a", node_id="n", session=composer_service.db.session - ) + AgentComposerService.resolve_workflow_node_agent_id(session=session, tenant_id="t", app_id="a", node_id="n") is None ) @@ -5024,9 +4977,7 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon ) 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", session=composer_service.db.session - ) + AgentComposerService.resolve_workflow_node_agent_id(session=session, tenant_id="t", app_id="a", node_id="n") is None ) @@ -5036,9 +4987,7 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon 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", session=composer_service.db.session - ) + AgentComposerService.resolve_workflow_node_agent_id(session=session, tenant_id="t", app_id="a", node_id="n") == "agent-7" ) @@ -5062,7 +5011,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( agent_id="agent-1", current_snapshot_id="version-1", ) - monkeypatch.setattr(composer_service.db, "session", FakeSession()) + session = FakeSession() monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) @@ -5083,7 +5032,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, session=None): + def fake_collect(cls, *, session, tenant_id, payload, agent_id=None): guarded["tenant_id"] = tenant_id guarded["agent_id"] = agent_id return {"warnings": [{"code": "mention_target_missing", "id": "files/sample.pdf"}]} @@ -5091,12 +5040,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( + session=session, 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 == { @@ -5125,7 +5074,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only( agent_id="agent-1", current_snapshot_id="version-1", ) - monkeypatch.setattr(composer_service.db, "session", FakeSession()) + session = FakeSession() monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) @@ -5146,19 +5095,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, session=None): + def fake_collect(cls, *, session, tenant_id, payload, agent_id=None): captured["agent_id"] = agent_id return {"warnings": []} monkeypatch.setattr(AgentComposerService, "collect_validation_findings", classmethod(fake_collect)) result = AgentComposerService.save_workflow_composer( + session=session, 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/dataset_service_test_helpers.py b/api/tests/unit_tests/services/dataset_service_test_helpers.py index 806f1e8d91b..d3dfe2fce79 100644 --- a/api/tests/unit_tests/services/dataset_service_test_helpers.py +++ b/api/tests/unit_tests/services/dataset_service_test_helpers.py @@ -180,6 +180,7 @@ class DatasetServiceUnitDataFactory: dataset.embedding_model = embedding_model dataset.built_in_field_enabled = built_in_field_enabled dataset.doc_form = doc_form + dataset.get_doc_form.return_value = doc_form dataset.enable_api = enable_api dataset.updated_by = None dataset.updated_at = None @@ -288,6 +289,7 @@ def _make_dataset( dataset.data_source_type = data_source_type dataset.indexing_technique = indexing_technique dataset.latest_process_rule = latest_process_rule + dataset.get_latest_process_rule.return_value = latest_process_rule dataset.embedding_model_provider = "provider" dataset.embedding_model = "embedding-model" dataset.summary_index_setting = None diff --git a/api/tests/unit_tests/services/document_service_validation.py b/api/tests/unit_tests/services/document_service_validation.py index 71df8c4e20d..8fad09e4d4c 100644 --- a/api/tests/unit_tests/services/document_service_validation.py +++ b/api/tests/unit_tests/services/document_service_validation.py @@ -176,6 +176,7 @@ class DocumentValidationTestDataFactory: dataset.id = dataset_id dataset.tenant_id = tenant_id dataset.doc_form = doc_form + dataset.get_doc_form.return_value = doc_form dataset.indexing_technique = indexing_technique dataset.embedding_model_provider = embedding_model_provider dataset.embedding_model = embedding_model @@ -327,9 +328,10 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) doc_form = IndexStructureType.PARAGRAPH_INDEX + session = Mock() # Act (should not raise) - DatasetService.check_doc_form(dataset, doc_form) + DatasetService.check_doc_form(dataset, doc_form, session=session) # Assert # No exception should be raised @@ -349,9 +351,10 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=None) doc_form = IndexStructureType.PARAGRAPH_INDEX + session = Mock() # Act (should not raise) - DatasetService.check_doc_form(dataset, doc_form) + DatasetService.check_doc_form(dataset, doc_form, session=session) # Assert # No exception should be raised @@ -371,10 +374,11 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) doc_form = IndexStructureType.PARENT_CHILD_INDEX # Different form + session = Mock() # Act & Assert with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"): - DatasetService.check_doc_form(dataset, doc_form) + DatasetService.check_doc_form(dataset, doc_form, session=session) def test_check_doc_form_different_form_types_error(self): """ @@ -390,10 +394,11 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form="knowledge_card") doc_form = IndexStructureType.PARAGRAPH_INDEX # Different form + session = Mock() # Act & Assert with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"): - DatasetService.check_doc_form(dataset, doc_form) + DatasetService.check_doc_form(dataset, doc_form, session=session) # ============================================================================ diff --git a/api/tests/unit_tests/services/hit_service.py b/api/tests/unit_tests/services/hit_service.py index ffeb158e37a..0257fd43676 100644 --- a/api/tests/unit_tests/services/hit_service.py +++ b/api/tests/unit_tests/services/hit_service.py @@ -587,6 +587,7 @@ class TestHitTestingServiceCompactRetrieveResponse: HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 1", score=0.95), HitTestingTestDataFactory.create_retrieval_record_mock(content="Doc 2", score=0.85), ] + session = MagicMock() with patch( "services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True @@ -594,14 +595,16 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = mock_records # Act - result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) + result = HitTestingService.compact_retrieve_response(query, documents, session=session) # Assert assert result["query"]["content"] == query assert len(result["records"]) == 2 assert result["records"][0]["content"] == "Doc 1" assert result["records"][0]["score"] == 0.95 - mock_format.assert_called_once_with(documents) + mock_format.assert_called_once() + assert mock_format.call_args.args[0] is not session + assert mock_format.call_args.args[1] == documents def test_compact_retrieve_response_empty_documents(self): """ @@ -613,6 +616,7 @@ class TestHitTestingServiceCompactRetrieveResponse: # Arrange query = "test query" documents = [] + session = MagicMock() with patch( "services.hit_testing_service.RetrievalService.format_retrieval_documents", autospec=True @@ -620,11 +624,14 @@ class TestHitTestingServiceCompactRetrieveResponse: mock_format.return_value = [] # Act - result = HitTestingService.compact_retrieve_response(query, documents, session=MagicMock()) + result = HitTestingService.compact_retrieve_response(query, documents, session=session) # Assert assert result["query"]["content"] == query assert result["records"] == [] + mock_format.assert_called_once() + assert mock_format.call_args.args[0] is not session + assert mock_format.call_args.args[1] == documents class TestHitTestingServiceCompactExternalRetrieveResponse: diff --git a/api/tests/unit_tests/services/plugin/test_plugin_migration.py b/api/tests/unit_tests/services/plugin/test_plugin_migration.py index d94ab540dfd..27b9749bf11 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_migration.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_migration.py @@ -33,6 +33,26 @@ def test_fetch_latest_package_identifier_calls_marketplace_when_enabled(mocker: assert result == "langgenius/openai:1.0.0@abc" +def test_extract_app_tables_checks_agent_mode_with_its_session(mocker: MockerFixture) -> None: + app = mocker.MagicMock(app_model_config_id=None, mode="chat") + app.is_agent_with_session.return_value = False + apps_result = mocker.MagicMock() + apps_result.all.return_value = [app] + configs_result = mocker.MagicMock() + configs_result.all.return_value = [] + session = mocker.MagicMock() + session.scalars.side_effect = [apps_result, configs_result] + session_context = mocker.MagicMock() + session_context.__enter__.return_value = session + mocker.patch(f"{MIGRATION_MODULE}.Session", return_value=session_context) + mocker.patch(f"{MIGRATION_MODULE}.db") + + result = PluginMigration.extract_app_tables("tenant-1") + + assert result == [] + app.is_agent_with_session.assert_called_once_with(session=session) + + class TestHandlePluginInstanceInstall: def test_raises_when_disabled_and_map_nonempty(self) -> None: with patch(f"{MIGRATION_MODULE}.dify_config") as mock_cfg: 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 c8992585653..d35ffcdb525 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 @@ -82,6 +82,7 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke assert result == "stream-events" update_status_mock.assert_called_once_with("doc-1", session=session_mock) + assert generator_instance.generate.call_args.kwargs["session"] is session_mock def test_update_document_status_updates_existing_document(mocker: MockerFixture) -> None: @@ -126,6 +127,7 @@ def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: assert result == "stream-iter" generator_instance.single_iteration_generate.assert_called_once() + assert generator_instance.single_iteration_generate.call_args.kwargs["session"] is session # --- generate_single_loop --- @@ -147,3 +149,4 @@ def test_generate_single_loop_delegates(mocker: MockerFixture) -> None: assert result == "stream-loop" generator_instance.single_loop_generate.assert_called_once() + assert generator_instance.single_loop_generate.call_args.kwargs["session"] is session 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 ba4c8d6f540..47f16c76b87 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 @@ -2424,3 +2424,16 @@ def test_get_pipeline_returns_pipeline_when_found( result = rag_pipeline_service.service.get_pipeline("t1", "d1") assert result is pipeline + + +def test_get_pipeline_by_id_uses_provided_session() -> None: + pipeline = _make_pipeline() + session = Mock() + session.scalar.return_value = pipeline + + result = RagPipelineService.get_pipeline_by_id("p1", "t1", session=session) + + assert result is pipeline + statement = session.scalar.call_args.args[0] + where_clauses = statement.whereclause.clauses + assert [clause.right.value for clause in where_clauses] == ["p1", "t1"] diff --git a/api/tests/unit_tests/services/test_account_service.py b/api/tests/unit_tests/services/test_account_service.py index dcb11a4f83e..763f8c0bd23 100644 --- a/api/tests/unit_tests/services/test_account_service.py +++ b/api/tests/unit_tests/services/test_account_service.py @@ -468,13 +468,13 @@ class TestAccountService: sqlite_session.commit() with ( - patch.object(Account, "set_tenant_id") as mock_set_tenant_id, + patch.object(Account, "set_tenant_id_with_session") as mock_set_tenant_id, patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active, ): result = AccountService.load_user(account.id, sqlite_session) assert result is account - mock_set_tenant_id.assert_called_once_with(tenant.id) + mock_set_tenant_id.assert_called_once_with(tenant.id, session=sqlite_session) mock_refresh_last_active.assert_called_once_with(account, sqlite_session) def test_load_user_not_found(self, sqlite_session: Session): @@ -512,7 +512,7 @@ class TestAccountService: sqlite_session.commit() with ( - patch.object(Account, "set_tenant_id") as mock_set_tenant_id, + patch.object(Account, "set_tenant_id_with_session") as mock_set_tenant_id, patch("services.account_service.naive_utc_now") as mock_naive_utc_now, patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active, ): @@ -524,10 +524,34 @@ class TestAccountService: assert result is account assert available_tenant_join.current is True assert available_tenant_join.last_opened_at == mock_now - mock_set_tenant_id.assert_called_once_with(tenant.id) + mock_set_tenant_id.assert_called_once_with(tenant.id, session=sqlite_session) mock_refresh_last_active.assert_called_once_with(account, sqlite_session) + def test_load_user_keeps_tenant_accessible_with_expiring_session(self, sqlite_session: Session): + account = Account(name="Test User", email="test@example.com") + tenant = Tenant(name="Test Workspace") + sqlite_session.add_all([account, tenant]) + sqlite_session.flush() + sqlite_session.add( + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + current=False, + ) + ) + sqlite_session.commit() + account_id = account.id + tenant_id = tenant.id + + with Session(sqlite_session.get_bind()) as expiring_session: + with patch.object(AccountService, "_refresh_account_last_active"): + result = AccountService.load_user(account_id, expiring_session) + + assert result is not None + assert result.current_tenant_id == tenant_id + def test_load_user_no_tenants(self, sqlite_session: Session): """Test user loading when user has no tenants at all.""" account = Account(name="Test User", email="test@example.com") @@ -794,7 +818,7 @@ class TestTenantService: ) assert tenant_account_join is not None assert tenant_account_join.role == TenantAccountRole.OWNER - assert mock_account.current_tenant == tenant + mock_account.set_current_tenant_with_session.assert_called_once_with(tenant, session=sqlite_session) mock_tenant_was_created.assert_called_once_with(tenant) mock_rsa_dependencies.assert_called_once_with(tenant.id) @@ -940,7 +964,19 @@ class TestTenantService: assert tenant_join.current is True assert tenant_join.last_opened_at == mock_now assert other_tenant_join.current is False - mock_account.set_tenant_id.assert_called_once_with(tenant.id) + mock_account.set_tenant_id_with_session.assert_called_once_with(tenant.id, session=sqlite_session) + + def test_switch_tenant_commits_changes(self): + account = TestAccountAssociatedDataFactory.create_account_mock() + tenant_join = TestAccountAssociatedDataFactory.create_tenant_join_mock( + tenant_id="tenant-456", account_id="user-123", current=False + ) + session = MagicMock() + session.scalar.return_value = tenant_join + + TenantService.switch_tenant(account, "tenant-456", session=session) + + session.commit.assert_called_once_with() @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) def test_switch_tenant_no_tenant_id(self, sqlite_session: Session): diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index 79bbb5873ac..87863697cfb 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -13,6 +13,7 @@ import pytest from werkzeug.datastructures import FileStorage from werkzeug.exceptions import NotFound +from models.dataset import DatasetCollectionBinding from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation from services.annotation_service import AppAnnotationService from services.app_ref_service import AnnotationRef, AppRef @@ -97,13 +98,13 @@ class TestAppAnnotationServiceUpInsert: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1", session=mock_db.session) + AppAnnotationService.up_insert_app_annotation_from_message(args, "app-1", 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.""" @@ -115,13 +116,13 @@ class TestAppAnnotationServiceUpInsert: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = app # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, 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.""" @@ -133,13 +134,13 @@ class TestAppAnnotationServiceUpInsert: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.side_effect = [app, None] # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, 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.""" @@ -150,18 +151,17 @@ class TestAppAnnotationServiceUpInsert: app = _make_app() annotation = _make_annotation("ann-1") message = _make_message(message_id="msg-1", app_id=app.id) - message.annotation = annotation setting = _make_setting() with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.add_annotation_to_index_task") as mock_task, ): - mock_db.session.scalar.side_effect = [app, message, setting] + mock_db.session.scalar.side_effect = [app, message, annotation, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, mock_db.session) # Assert assert result == annotation @@ -188,30 +188,25 @@ class TestAppAnnotationServiceUpInsert: app = _make_app() message = _make_message(message_id="msg-1", app_id=app.id) message.annotation = None - annotation_instance = _make_annotation("ann-1") with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, - patch("services.annotation_service.MessageAnnotation", return_value=annotation_instance) as mock_cls, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.add_annotation_to_index_task") as mock_task, ): - mock_db.session.scalar.side_effect = [app, message, None] + mock_db.session.scalar.side_effect = [app, message, None, None] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, mock_db.session) # Assert - assert result == annotation_instance - mock_cls.assert_called_once_with( - app_id=app.id, - conversation_id=message.conversation_id, - message_id=message.id, - content="hello", - question="q1", - account_id=current_user.id, - ) - mock_db.session.add.assert_called_once_with(annotation_instance) + assert result.app_id == app.id + assert result.conversation_id == message.conversation_id + assert result.message_id == message.id + assert result.content == "hello" + assert result.question == "q1" + assert result.account_id == current_user.id + mock_db.session.add.assert_called_once_with(result) mock_db.session.commit.assert_called_once() mock_task.delay.assert_not_called() @@ -225,13 +220,13 @@ class TestAppAnnotationServiceUpInsert: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = app # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, 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.""" @@ -245,14 +240,14 @@ class TestAppAnnotationServiceUpInsert: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.MessageAnnotation", return_value=annotation_instance) as mock_cls, patch("services.annotation_service.add_annotation_to_index_task") as mock_task, ): mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, session=mock_db.session) + result = AppAnnotationService.up_insert_app_annotation_from_message(args, app.id, mock_db.session) # Assert assert result == annotation_instance @@ -377,13 +372,13 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "", session=mock_db.session) + AppAnnotationService.get_annotation_list_by_app_id("app-1", 1, 10, "", 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.""" @@ -394,7 +389,7 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.paginate_query") as mock_paginate, patch("libs.helper.escape_like_pattern", return_value="safe"), ): @@ -402,9 +397,7 @@ class TestAppAnnotationServiceListAndExport: mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, 1, 10, "keyword", session=mock_db.session - ) + items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "keyword", mock_db.session) # Assert assert items == ["a1"] @@ -419,16 +412,14 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.paginate_query") as mock_paginate, ): mock_db.session.scalar.return_value = app mock_paginate.return_value = pagination # Act - items, total = AppAnnotationService.get_annotation_list_by_app_id( - app.id, 1, 10, "", session=mock_db.session - ) + items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "", mock_db.session) # Assert assert items == ["a1", "a2"] @@ -448,14 +439,14 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.CSVSanitizer.sanitize_value", side_effect=lambda v: f"safe:{v}"), ): mock_db.session.scalar.return_value = app mock_db.session.scalars.return_value.all.return_value = [annotation1, annotation2] # Act - result = AppAnnotationService.export_annotation_list_by_app_id(app.id, session=mock_db.session) + result = AppAnnotationService.export_annotation_list_by_app_id(app.id, mock_db.session) # Assert assert result == [annotation1, annotation2] @@ -471,13 +462,13 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.export_annotation_list_by_app_id("app-1", session=mock_db.session) + AppAnnotationService.export_annotation_list_by_app_id("app-1", mock_db.session) class TestAppAnnotationServiceDirectManipulation: @@ -491,13 +482,13 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.insert_app_annotation_directly(args, "app-1", session=mock_db.session) + AppAnnotationService.insert_app_annotation_directly(args, "app-1", mock_db.session) def test_insert_app_annotation_directly_should_raise_value_error_when_question_missing(self) -> None: """Test missing question raises ValueError.""" @@ -508,13 +499,13 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = app # Act & Assert with pytest.raises(ValueError): - AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) + AppAnnotationService.insert_app_annotation_directly(args, app.id, mock_db.session) def test_insert_app_annotation_directly_should_create_annotation_and_index(self) -> None: """Test insert creates annotation and triggers index task.""" @@ -528,14 +519,14 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.MessageAnnotation", return_value=annotation_instance) as mock_cls, patch("services.annotation_service.add_annotation_to_index_task") as mock_task, ): mock_db.session.scalar.side_effect = [app, setting] # Act - result = AppAnnotationService.insert_app_annotation_directly(args, app.id, session=mock_db.session) + result = AppAnnotationService.insert_app_annotation_directly(args, app.id, mock_db.session) # Assert assert result == annotation_instance @@ -564,7 +555,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None @@ -586,7 +577,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = annotation @@ -608,7 +599,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.update_annotation_to_index_task") as mock_task, ): mock_db.session.scalar.side_effect = [annotation, setting] @@ -645,7 +636,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.delete_annotation_index_task") as mock_task, ): mock_db.session.scalar.side_effect = [annotation, setting] @@ -679,7 +670,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None @@ -695,7 +686,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.execute.return_value.all.return_value = [] @@ -718,7 +709,7 @@ class TestAppAnnotationServiceDirectManipulation: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.delete_annotation_index_task") as mock_task, ): # First execute().all() for multi-column query, subsequent execute() calls for deletes @@ -757,13 +748,13 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.batch_import_app_annotations("app-1", file, session=mock_db.session) + AppAnnotationService.batch_import_app_annotations("app-1", file, mock_db.session) def test_batch_import_app_annotations_should_return_error_when_columns_invalid(self) -> None: """Test invalid column count returns error message.""" @@ -775,7 +766,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -785,7 +776,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -800,7 +791,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch( "configs.dify_config", new=SimpleNamespace(ANNOTATION_IMPORT_MAX_RECORDS=5, ANNOTATION_IMPORT_MIN_RECORDS=1), @@ -809,7 +800,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -826,7 +817,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch("services.annotation_service.FeatureService.get_features", return_value=features), patch( @@ -837,7 +828,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -853,7 +844,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -863,7 +854,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -883,7 +874,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -893,7 +884,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -909,7 +900,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -919,7 +910,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -935,7 +926,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -945,7 +936,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -961,7 +952,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch( "configs.dify_config", @@ -971,7 +962,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -991,7 +982,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch("services.annotation_service.FeatureService.get_features", return_value=features), patch( @@ -1002,7 +993,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert error_msg = cast(str, result["error_msg"]) @@ -1020,7 +1011,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch("services.annotation_service.FeatureService.get_features", return_value=features), patch("services.annotation_service.batch_import_annotations_task") as mock_task, @@ -1035,7 +1026,7 @@ class TestAppAnnotationServiceBatchImport: mock_db.session.scalar.return_value = app # Act - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert assert result == {"job_id": "uuid-3", "job_status": "waiting", "record_count": 1} @@ -1058,7 +1049,7 @@ class TestAppAnnotationServiceBatchImport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.pd.read_csv", return_value=df), patch("services.annotation_service.FeatureService.get_features", return_value=features), patch("services.annotation_service.redis_client") as mock_redis, @@ -1075,7 +1066,7 @@ class TestAppAnnotationServiceBatchImport: # Act with caplog.at_level(logging.DEBUG): - result = AppAnnotationService.batch_import_app_annotations(app.id, file, session=mock_db.session) + result = AppAnnotationService.batch_import_app_annotations(app.id, file, mock_db.session) # Assert assert result["error_msg"] == "An error occurred while processing the file: boom" @@ -1093,7 +1084,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Arrange app = _make_app() - with patch("services.annotation_service.db") as mock_db: + with patch("services.annotation_service.db", create=True) as mock_db: mock_db.session.scalar.return_value = None # Act & Assert @@ -1112,7 +1103,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.paginate_query") as mock_paginate, ): mock_db.session.scalar.return_value = app @@ -1124,7 +1115,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: _make_annotation_ref(app, annotation.id), 1, 10, - session=mock_db.session, + mock_db.session, ) # Assert @@ -1136,11 +1127,11 @@ class TestAppAnnotationServiceHitHistoryAndSettings: def test_get_annotation_by_id_should_return_none_when_missing(self) -> None: """Test get_annotation_by_id returns None when not found.""" # Arrange - with patch("services.annotation_service.db") as mock_db: + with patch("services.annotation_service.db", create=True) as mock_db: mock_db.session.get.return_value = None # Act - result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) + result = AppAnnotationService.get_annotation_by_id("ann-1", mock_db.session) # Assert assert result is None @@ -1149,11 +1140,11 @@ class TestAppAnnotationServiceHitHistoryAndSettings: """Test get_annotation_by_id returns annotation when found.""" # Arrange annotation = _make_annotation("ann-1") - with patch("services.annotation_service.db") as mock_db: + with patch("services.annotation_service.db", create=True) as mock_db: mock_db.session.get.return_value = annotation # Act - result = AppAnnotationService.get_annotation_by_id("ann-1", session=mock_db.session) + result = AppAnnotationService.get_annotation_by_id("ann-1", mock_db.session) # Assert assert result == annotation @@ -1162,7 +1153,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: """Test add_annotation_history updates hit count and creates history.""" # Arrange with ( - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.AppAnnotationHitHistory") as mock_history_cls, ): # Act @@ -1183,7 +1174,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: mock_db.session.execute.assert_called_once() mock_history_cls.assert_called_once() mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + mock_db.session.flush.assert_called_once() def test_get_app_annotation_setting_by_app_id_should_return_embedding_model_when_detail_exists(self) -> None: """Test setting detail returns embedding model info.""" @@ -1191,21 +1182,25 @@ class TestAppAnnotationServiceHitHistoryAndSettings: tenant_id = "tenant-1" app = _make_app() setting = _make_setting(with_detail=True) + detail = setting.collection_binding_detail + setting.collection_binding_detail = None with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.side_effect = [app, setting] + mock_db.session.get.return_value = detail # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, mock_db.session) # Assert assert result["enabled"] is True embedding_model = cast(dict[str, Any], result["embedding_model"]) assert embedding_model["embedding_provider_name"] == "provider-a" assert embedding_model["embedding_model_name"] == "model-a" + mock_db.session.get.assert_called_once_with(DatasetCollectionBinding, setting.collection_binding_id) def test_get_app_annotation_setting_by_app_id_should_raise_not_found_when_app_missing(self) -> None: """Test missing app raises NotFound.""" @@ -1214,13 +1209,13 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.get_app_annotation_setting_by_app_id("app-1", session=mock_db.session) + AppAnnotationService.get_app_annotation_setting_by_app_id("app-1", 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.""" @@ -1228,19 +1223,22 @@ class TestAppAnnotationServiceHitHistoryAndSettings: tenant_id = "tenant-1" app = _make_app() setting = _make_setting(with_detail=False) + setting.collection_binding_detail = SimpleNamespace(provider_name="wrong", model_name="wrong") with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.side_effect = [app, setting] + mock_db.session.get.return_value = None # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, mock_db.session) # Assert assert result["enabled"] is True assert result["embedding_model"] == {} + mock_db.session.get.assert_called_once_with(DatasetCollectionBinding, setting.collection_binding_id) def test_get_app_annotation_setting_by_app_id_should_return_disabled_when_setting_missing(self) -> None: """Test missing setting returns disabled payload.""" @@ -1250,12 +1248,12 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.side_effect = [app, None] # Act - result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, session=mock_db.session) + result = AppAnnotationService.get_app_annotation_setting_by_app_id(app.id, mock_db.session) # Assert assert result == {"enabled": False} @@ -1267,27 +1265,29 @@ class TestAppAnnotationServiceHitHistoryAndSettings: current_user = _make_user("user-1") app = _make_app() setting = _make_setting(with_detail=True) + detail = setting.collection_binding_detail + setting.collection_binding_detail = None args = {"score_threshold": 0.8} with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.naive_utc_now", return_value="now"), ): mock_db.session.scalar.side_effect = [app, setting] + mock_db.session.get.return_value = detail # Act - result = AppAnnotationService.update_app_annotation_setting( - app.id, setting.id, args, session=mock_db.session - ) + result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args, mock_db.session) # Assert assert result["enabled"] is True assert result["score_threshold"] == 0.8 embedding_model = cast(dict[str, Any], result["embedding_model"]) assert embedding_model["embedding_provider_name"] == "provider-a" + mock_db.session.get.assert_called_once_with(DatasetCollectionBinding, setting.collection_binding_id) mock_db.session.add.assert_called_once_with(setting) - mock_db.session.commit.assert_called_once() + mock_db.session.flush.assert_called_once() def test_update_app_annotation_setting_should_return_empty_embedding_model_when_detail_missing(self) -> None: """Test update returns empty embedding_model when collection detail is absent.""" @@ -1296,24 +1296,25 @@ class TestAppAnnotationServiceHitHistoryAndSettings: current_user = _make_user("user-1") app = _make_app() setting = _make_setting(with_detail=False) + setting.collection_binding_detail = SimpleNamespace(provider_name="wrong", model_name="wrong") args = {"score_threshold": 0.7} with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.naive_utc_now", return_value="now"), ): mock_db.session.scalar.side_effect = [app, setting] + mock_db.session.get.return_value = None # Act - result = AppAnnotationService.update_app_annotation_setting( - app.id, setting.id, args, session=mock_db.session - ) + result = AppAnnotationService.update_app_annotation_setting(app.id, setting.id, args, mock_db.session) # Assert assert result["enabled"] is True assert result["score_threshold"] == 0.7 assert result["embedding_model"] == {} + mock_db.session.get.assert_called_once_with(DatasetCollectionBinding, setting.collection_binding_id) def test_update_app_annotation_setting_should_raise_not_found_when_app_missing(self) -> None: """Test update raises NotFound when app is missing.""" @@ -1322,7 +1323,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None @@ -1340,7 +1341,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.side_effect = [app, None] @@ -1366,7 +1367,7 @@ class TestAppAnnotationServiceClearAll: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, patch("services.annotation_service.delete_annotation_index_task") as mock_task, ): # scalar calls: app lookup, annotation_setting lookup @@ -1381,7 +1382,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, session=mock_db.session) + result = AppAnnotationService.clear_all_annotations(app.id, mock_db.session) # Assert assert result == {"result": "success"} @@ -1399,10 +1400,10 @@ class TestAppAnnotationServiceClearAll: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), - patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.db", create=True) as mock_db, ): mock_db.session.scalar.return_value = None # Act & Assert with pytest.raises(NotFound): - AppAnnotationService.clear_all_annotations("app-1", session=mock_db.session) + AppAnnotationService.clear_all_annotations("app-1", mock_db.session) diff --git a/api/tests/unit_tests/services/test_app_dsl_service.py b/api/tests/unit_tests/services/test_app_dsl_service.py index 051c0a0d885..d80cb12d83e 100644 --- a/api/tests/unit_tests/services/test_app_dsl_service.py +++ b/api/tests/unit_tests/services/test_app_dsl_service.py @@ -1,8 +1,12 @@ +from types import SimpleNamespace +from typing import cast from unittest.mock import Mock import pytest from sqlalchemy.orm import Session +from models import App, AppMode +from models.model import AppModelConfig, IconType from services.app_dsl_service import AppDslService from services.entities.dsl_entities import ImportStatus @@ -63,3 +67,102 @@ def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes( assert result.status == ImportStatus.FAILED assert "utf-8" in result.error assert not sqlite_session.in_transaction() + + +def test_create_or_update_app_loads_existing_model_config_with_service_session() -> None: + session = Mock() + session.get.return_value = Mock() + service = AppDslService(session=session) + app = cast( + App, + SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + app_model_config_id="config-1", + name="Existing app", + description="", + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + ), + ) + + result = service._create_or_update_app( + app=app, + data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, + account=Mock(id="account-1"), + ) + + assert result is app + session.get.assert_called_once_with(AppModelConfig, "config-1") + + +def test_create_or_update_app_flushes_new_model_config_before_signal(monkeypatch: pytest.MonkeyPatch) -> None: + events: list[str] = [] + session = Mock() + session.add.side_effect = lambda _config: events.append("add") + session.flush.side_effect = lambda: events.append("flush") + signal = Mock() + signal.send.side_effect = lambda *_args, **_kwargs: events.append("signal") + monkeypatch.setattr("services.app_dsl_service.app_model_config_was_updated", signal) + app = cast( + App, + SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + app_model_config_id=None, + name="Existing app", + description="", + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + ), + ) + + AppDslService(session=session)._create_or_update_app( + app=app, + data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, + account=Mock(id="account-1"), + ) + + assert events == ["add", "flush", "signal"] + assert signal.send.call_args.kwargs["session"] is session + session.commit.assert_not_called() + + +def test_export_dsl_loads_model_config_and_annotation_reply_with_request_session( + monkeypatch: pytest.MonkeyPatch, +) -> None: + model_config = {"model": {}, "agent_mode": {"tools": []}} + app_model_config = Mock(app_id="app-1") + app_model_config.to_dict.return_value = model_config + session = Mock() + session.get.return_value = app_model_config + annotation_reply = {"enabled": False} + load_annotation_reply_config = Mock(return_value=annotation_reply) + monkeypatch.setattr("services.app_dsl_service.load_annotation_reply_config", load_annotation_reply_config) + monkeypatch.setattr( + "services.app_dsl_service.DependenciesAnalysisService.generate_dependencies", + Mock(return_value=[]), + ) + app = cast( + App, + SimpleNamespace( + id="app-1", + tenant_id="tenant-1", + app_model_config_id="config-1", + mode=AppMode.CHAT, + name="Chat app", + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + description="", + use_icon_as_answer_icon=False, + ), + ) + + AppDslService.export_dsl(app, session=session) + + session.get.assert_called_once_with(AppModelConfig, "config-1") + load_annotation_reply_config.assert_called_once_with(session, "app-1") + app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) 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 552403e0d0b..507977f287f 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -69,6 +69,7 @@ def _make_app(mode: AppMode | str, *, max_active_requests: int = 0, is_agent: bo app.tenant_id = "tenant-id" app.max_active_requests = max_active_requests app.is_agent = is_agent + app.is_agent_with_session.return_value = is_agent return app @@ -277,16 +278,42 @@ class TestGenerate: side_effect=lambda x: x, ) app = _make_app(AppMode.CHAT, is_agent=True) + session = MagicMock() result = AppGenerateService.generate( app_model=app, user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=session, ) assert result == {"result": "agent-via-flag"} gen_spy.assert_called_once() + app.is_agent_with_session.assert_called_once_with(session=session) + + # -- AGENT -------------------------------------------------------------- + def test_agent_mode_passes_session(self, mocker: MockerFixture): + gen_spy = mocker.patch( + "services.app_generate_service.AgentAppGenerator.generate", + return_value={"result": "agent"}, + ) + mocker.patch( + "services.app_generate_service.AgentAppGenerator.convert_to_event_stream", + side_effect=lambda x: x, + ) + session = MagicMock() + + result = AppGenerateService.generate( + app_model=_make_app(AppMode.AGENT), + user=_make_user(), + args={"inputs": {}}, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, + session=session, + ) + + assert result == {"result": "agent"} + assert gen_spy.call_args.kwargs["session"] is session # -- CHAT --------------------------------------------------------------- def test_chat_mode(self, mocker: MockerFixture): @@ -325,17 +352,19 @@ class TestGenerate: side_effect=lambda x: x, ) + session = MagicMock() result = AppGenerateService.generate( 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(), + session=session, ) assert result == {"result": "advanced-blocking"} call_kwargs = gen_spy.call_args.kwargs assert call_kwargs.get("streaming") is False + assert call_kwargs["session"] is session retrieve_spy.assert_not_called() # -- ADVANCED_CHAT streaming -------------------------------------------- @@ -384,13 +413,14 @@ class TestGenerate: side_effect=lambda x: x, ) + session = MagicMock() result = AppGenerateService.generate( app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, invoke_from=InvokeFrom.SERVICE_API, streaming=False, - session=MagicMock(), + session=session, ) assert result == {"result": "workflow-blocking"} call_kwargs = gen_spy.call_args.kwargs @@ -731,14 +761,16 @@ class TestGenerateSingleIteration: return_value={"event": "iteration"}, ) app = _make_app(AppMode.ADVANCED_CHAT) + session = MagicMock() result = AppGenerateService.generate_single_iteration( app_model=app, user=_make_user(), node_id="n1", args={"k": "v"}, - session=MagicMock(), + session=session, ) iter_spy.assert_called_once() + assert iter_spy.call_args.kwargs["session"] is session assert result == {"event": "iteration"} def test_workflow_mode(self, mocker: MockerFixture): @@ -753,14 +785,16 @@ class TestGenerateSingleIteration: return_value={"event": "wf-iteration"}, ) app = _make_app(AppMode.WORKFLOW) + session = MagicMock() result = AppGenerateService.generate_single_iteration( app_model=app, user=_make_user(), node_id="n1", args={"k": "v"}, - session=MagicMock(), + session=session, ) iter_spy.assert_called_once() + assert iter_spy.call_args.kwargs["session"] is session assert result == {"event": "wf-iteration"} def test_invalid_mode_raises(self, mocker: MockerFixture): @@ -787,14 +821,16 @@ class TestGenerateSingleLoop: return_value={"event": "loop"}, ) app = _make_app(AppMode.ADVANCED_CHAT) + session = MagicMock() result = AppGenerateService.generate_single_loop( app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), - session=MagicMock(), + session=session, ) loop_spy.assert_called_once() + assert loop_spy.call_args.kwargs["session"] is session assert result == {"event": "loop"} def test_workflow_mode(self, mocker: MockerFixture): @@ -809,14 +845,16 @@ class TestGenerateSingleLoop: return_value={"event": "wf-loop"}, ) app = _make_app(AppMode.WORKFLOW) + session = MagicMock() result = AppGenerateService.generate_single_loop( app_model=app, user=_make_user(), node_id="n1", args=MagicMock(), - session=MagicMock(), + session=session, ) loop_spy.assert_called_once() + assert loop_spy.call_args.kwargs["session"] is session assert result == {"event": "wf-loop"} def test_invalid_mode_raises(self, mocker: MockerFixture): diff --git a/api/tests/unit_tests/services/test_app_model_config_service.py b/api/tests/unit_tests/services/test_app_model_config_service.py index d4b4bf14a36..06119d33ba3 100644 --- a/api/tests/unit_tests/services/test_app_model_config_service.py +++ b/api/tests/unit_tests/services/test_app_model_config_service.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest @@ -50,20 +50,24 @@ class TestAppModelConfigService: mock_agent_validate = mock_config_managers["agent"] mock_completion_validate = mock_config_managers["completion"] - result = AppModelConfigService.validate_configuration(tenant_id=tenant_id, config=config, app_mode=app_mode) + session = MagicMock() + + result = AppModelConfigService.validate_configuration( + tenant_id=tenant_id, config=config, app_mode=app_mode, session=session + ) assert result == {"manager": selected_manager} if selected_manager == "chat": - mock_chat_validate.assert_called_once_with(tenant_id, config) + mock_chat_validate.assert_called_once_with(tenant_id, config, session) mock_agent_validate.assert_not_called() mock_completion_validate.assert_not_called() elif selected_manager == "agent": - mock_agent_validate.assert_called_once_with(tenant_id, config) + mock_agent_validate.assert_called_once_with(tenant_id, config, session) mock_chat_validate.assert_not_called() mock_completion_validate.assert_not_called() else: - mock_completion_validate.assert_called_once_with(tenant_id, config) + mock_completion_validate.assert_called_once_with(tenant_id, config, session) mock_chat_validate.assert_not_called() mock_agent_validate.assert_not_called() @@ -81,6 +85,7 @@ class TestAppModelConfigService: tenant_id=tenant_id, config=config, app_mode=AppMode.WORKFLOW, + session=MagicMock(), ) mock_chat_validate.assert_not_called() diff --git a/api/tests/unit_tests/services/test_app_service.py b/api/tests/unit_tests/services/test_app_service.py index 36914679d3e..0fba9e82b32 100644 --- a/api/tests/unit_tests/services/test_app_service.py +++ b/api/tests/unit_tests/services/test_app_service.py @@ -1,14 +1,69 @@ from __future__ import annotations +from collections.abc import Callable from types import SimpleNamespace +from typing import cast from unittest.mock import MagicMock, patch import pytest from sqlalchemy.exc import IntegrityError -from models.model import App +from models import Account +from models.model import App, AppMode, AppModelConfig +from models.workflow import Workflow from services.agent.errors import AgentNameConflictError -from services.app_service import AppService +from services.app_service import AppService, CreateAppParams + + +class TestCreateAppTransactionBoundary: + def test_commits_database_state_before_external_side_effects(self) -> None: + session = MagicMock() + account = MagicMock(spec=Account, id="account-1", current_tenant_id="tenant-1") + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + + with ( + patch( + "services.app_service.app_was_created.send", + side_effect=lambda *_args, **_kwargs: phase_events.append("signal"), + ), + patch( + "services.app_service.enterprise_rbac_service.try_sync_creator_access_policy_member_bindings", + side_effect=lambda *_args: phase_events.append("external"), + ), + patch( + "services.app_service.FeatureService.get_system_features", + return_value=SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), + ), + patch("services.app_service.dify_config.BILLING_ENABLED", False), + ): + AppService().create_app( + "tenant-1", + CreateAppParams(name="Workflow", mode=AppMode.WORKFLOW.value), + account, + session=session, + ) + + assert phase_events == ["commit", "signal", "commit", "external"] + + +@pytest.mark.parametrize( + "update_status", + [AppService.update_app_site_status, AppService.update_app_api_status], +) +def test_app_status_updates_commit_before_signal(update_status: Callable[..., App]) -> None: + app = cast(App, SimpleNamespace(enable_site=False, enable_api=False)) + session = MagicMock() + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + + with ( + patch("services.app_service.current_user", SimpleNamespace(id="account-1")), + patch("services.app_service.app_was_updated.send", side_effect=lambda *_args: phase_events.append("signal")), + ): + update_status(AppService(), app, True, session=session) + + assert phase_events == ["commit", "signal"] class TestOpenapiVisibilityHelpers: @@ -29,14 +84,14 @@ class TestOpenapiVisibilityHelpers: sentinel_app.status = "archived" # explicitly NOT "normal" mock_session.get.return_value = sentinel_app - assert AppService.get_app_by_id("app-uuid", session=mock_session) is sentinel_app + assert AppService.get_app_by_id("app-uuid", 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("missing", session=mock_session) is None + assert AppService.get_app_by_id("missing", mock_session) is None def test_get_visible_app_by_id_returns_app_when_visible(self): mock_session = MagicMock() @@ -45,7 +100,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("app-uuid", session=mock_session) is app + assert AppService.get_visible_app_by_id("app-uuid", mock_session) is app mock_session.get.assert_called_once_with(App, "app-uuid") @@ -53,7 +108,7 @@ class TestOpenapiVisibilityHelpers: mock_session = MagicMock() mock_session.get.return_value = None - assert AppService.get_visible_app_by_id("missing", session=mock_session) is None + assert AppService.get_visible_app_by_id("missing", 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 +120,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("app-uuid", session=mock_session) is None + assert AppService.get_visible_app_by_id("app-uuid", 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 +133,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("app-uuid", session=mock_session) is None + assert AppService.get_visible_app_by_id("app-uuid", 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 @@ -113,7 +168,7 @@ class TestOpenapiVisibilityHelpers: """ mock_session = MagicMock() - assert AppService.find_visible_apps_by_ids([], session=mock_session) == [] + assert AppService.find_visible_apps_by_ids([], mock_session) == [] mock_session.execute.assert_not_called() def test_find_visible_apps_by_ids_passes_through_visibility_gate(self): @@ -127,13 +182,63 @@ 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(["a", "b"], session=mock_session) + out = AppService.find_visible_apps_by_ids(["a", "b"], mock_session) assert out == rows gate.assert_called_once() mock_session.execute.assert_called_once() +class TestAppMeta: + def test_loads_workflow_with_caller_session(self): + session = MagicMock() + session.get.return_value = SimpleNamespace(graph_dict={"nodes": []}) + app = cast(App, SimpleNamespace(mode=AppMode.WORKFLOW, workflow_id="workflow-1")) + + assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}} + + session.get.assert_called_once_with(Workflow, "workflow-1") + + def test_loads_app_model_config_with_caller_session(self): + session = MagicMock() + session.get.return_value = SimpleNamespace(agent_mode_dict={"tools": []}) + app = cast(App, SimpleNamespace(mode=AppMode.CHAT, app_model_config_id="config-1")) + + assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}} + + session.get.assert_called_once_with(AppModelConfig, "config-1") + + +class TestGetApp: + def test_legacy_agent_detection_uses_caller_session(self): + session = MagicMock() + app = MagicMock(spec=App) + app.mode = AppMode.CHAT + app.is_agent_with_session.return_value = False + account = MagicMock(spec=Account) + account.current_tenant_id = "tenant-1" + + with patch("services.app_service.current_user", account): + assert AppService().get_app(app, session=session) is app + + app.is_agent_with_session.assert_called_once_with(session=session) + app.app_model_config_with_session.assert_not_called() + + def test_agent_model_config_uses_caller_session(self): + session = MagicMock() + app = MagicMock(spec=App) + app.mode = AppMode.AGENT_CHAT + app.app_model_config_with_session.return_value = None + account = MagicMock(spec=Account) + account.current_tenant_id = "tenant-1" + + with patch("services.app_service.current_user", account): + assert AppService().get_app(app, session=session) is app + + app.is_agent_with_session.assert_not_called() + app.app_model_config_with_session.assert_called_once_with(session=session) + + class TestAgentAppType: """S1: new ``AppMode.AGENT`` app type wiring.""" diff --git a/api/tests/unit_tests/services/test_audio_service.py b/api/tests/unit_tests/services/test_audio_service.py index e5b0fdb077f..b9f44bc212b 100644 --- a/api/tests/unit_tests/services/test_audio_service.py +++ b/api/tests/unit_tests/services/test_audio_service.py @@ -138,6 +138,8 @@ class AudioServiceTestDataFactory: app.tenant_id = tenant_id app.workflow = kwargs.get("workflow") app.app_model_config = kwargs.get("app_model_config") + app.workflow_with_session.return_value = app.workflow + app.app_model_config_with_session.return_value = app.app_model_config for key, value in kwargs.items(): setattr(app, key, value) return app @@ -240,7 +242,7 @@ class TestAudioServiceASR: mock_model_manager.get_default_model_instance.return_value = mock_model_instance # Act - result = AudioService.transcript_asr(app_model=app, file=file, end_user="user-123") + result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="user-123") # Assert assert result == {"text": "Transcribed text"} @@ -267,7 +269,7 @@ class TestAudioServiceASR: mock_model_manager.get_default_model_instance.return_value = mock_model_instance # Act - result = AudioService.transcript_asr(app_model=app, file=file) + result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) # Assert assert result == {"text": "Workflow transcribed text"} @@ -288,7 +290,7 @@ class TestAudioServiceASR: mock_model_instance.invoke_speech2text.return_value = "Published Agent transcript" mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance - result = AudioService.transcript_asr(app_model=app, file=file, end_user="end-user-1") + result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="end-user-1") assert result == {"text": "Published Agent transcript"} mock_roster_service_class.return_value.get_published_agent_soul_for_app.assert_called_once_with( @@ -312,7 +314,7 @@ class TestAudioServiceASR: mock_model_instance.invoke_speech2text.return_value = "Legacy Agent transcript" mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance - result = AudioService.transcript_asr(app_model=app, file=file) + result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) assert result == {"text": "Legacy Agent transcript"} @@ -331,6 +333,7 @@ class TestAudioServiceASR: app_model=app, agent_soul=agent_soul, file=file, + session=MagicMock(), end_user="account-1", ) @@ -351,7 +354,7 @@ class TestAudioServiceASR: file = factory.create_file_storage_mock() with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file) + AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock()) @patch("services.audio_service.ModelManager.for_tenant", autospec=True) def test_transcript_agent_asr_preserves_legacy_feature_fallback( @@ -369,6 +372,7 @@ class TestAudioServiceASR: app_model=app, agent_soul=AgentSoulConfig(), file=file, + session=MagicMock(), ) assert result == {"text": "Legacy feature transcript"} @@ -381,7 +385,7 @@ class TestAudioServiceASR: agent_soul = AgentSoulConfig.model_validate({"app_features": {"speech_to_text": {"enabled": False}}}) with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file) + AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock()) def test_transcript_asr_raises_error_when_feature_disabled_chat_mode(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when speech-to-text is disabled in CHAT mode.""" @@ -395,7 +399,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) def test_transcript_asr_raises_error_when_feature_disabled_workflow_mode( self, factory: AudioServiceTestDataFactory @@ -411,7 +415,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) def test_transcript_asr_raises_error_when_workflow_missing(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when workflow is missing in WORKFLOW mode.""" @@ -424,7 +428,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) def test_transcript_asr_raises_error_when_no_file_uploaded(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when no file is uploaded.""" @@ -437,7 +441,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(NoAudioUploadedServiceError): - AudioService.transcript_asr(app_model=app, file=None) + AudioService.transcript_asr(app_model=app, file=None, session=MagicMock()) def test_transcript_asr_raises_error_for_unsupported_audio_type(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error for unsupported audio file types.""" @@ -451,7 +455,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(UnsupportedAudioTypeServiceError): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) def test_transcript_asr_raises_error_for_large_file(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when file exceeds size limit (30MB).""" @@ -467,7 +471,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(AudioTooLargeServiceError, match="Audio size larger than 30 mb"): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) @patch("services.audio_service.ModelManager.for_tenant", autospec=True) def test_transcript_asr_raises_error_when_no_model_instance( @@ -488,7 +492,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(ProviderNotSupportSpeechToTextServiceError): - AudioService.transcript_asr(app_model=app, file=file) + AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) @pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) 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 dcb6250a8f0..4ce8d46d702 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -33,14 +33,20 @@ class TestDatasetServiceValidation: ) def test_check_doc_form_allows_matching_or_missing_dataset_doc_form(self, dataset_doc_form, incoming_doc_form): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(doc_form=dataset_doc_form) + session = MagicMock() - DatasetService.check_doc_form(dataset, incoming_doc_form) + DatasetService.check_doc_form(dataset, incoming_doc_form, session=session) + + dataset.get_doc_form.assert_called_once_with(session=session) def test_check_doc_form_rejects_mismatched_doc_form(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(doc_form="qa_model") + session = MagicMock() with pytest.raises(ValueError, match="doc_form is different"): - DatasetService.check_doc_form(dataset, "text_model") + DatasetService.check_doc_form(dataset, "text_model", session=session) + + dataset.get_doc_form.assert_called_once_with(session=session) def test_check_dataset_model_setting_skips_non_high_quality_datasets(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(indexing_technique="economy") @@ -167,17 +173,34 @@ class TestDatasetServiceValidation: DatasetService.check_is_multimodal_model("tenant-1", "provider", "embedding-model") +class TestDatasetServiceRetrieval: + def test_get_dataset_for_tenant_scopes_query_by_dataset_and_tenant(self): + session = MagicMock() + expected_dataset = DatasetServiceUnitDataFactory.create_dataset_mock( + dataset_id="dataset-1", tenant_id="tenant-1" + ) + session.scalar.return_value = expected_dataset + + result = DatasetService.get_dataset_for_tenant("dataset-1", "tenant-1", session=session) + + assert result is expected_dataset + statement = session.scalar.call_args.args[0].compile() + assert "datasets.id" in str(statement) + assert "datasets.tenant_id" in str(statement) + assert "dataset-1" in statement.params.values() + assert "tenant-1" in statement.params.values() + + class TestDatasetServiceRetrievalPermissions: """Unit tests for dataset list permission branching.""" def test_get_datasets_filters_by_maintainer_and_rbac_overrides(self): - mock_db = MagicMock() + session = MagicMock() explicit_session = MagicMock() explicit_session.scalars.return_value.all.return_value = [] user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( @@ -197,19 +220,18 @@ class TestDatasetServiceRetrievalPermissions: ) explicit_session.scalars.assert_called_once() - mock_db.session.scalars.assert_not_called() + session.scalars.assert_not_called() select_stmt = mock_paginate.call_args.args[0] visibility_clause = str(select_stmt._where_criteria[1]) assert "maintainer" in visibility_clause assert "IN" in visibility_clause def test_get_datasets_filters_only_by_rbac_overrides_without_manage_own_permission(self): - mock_db = MagicMock() - mock_db.session.scalars.return_value.all.return_value = [] + session = MagicMock() + session.scalars.return_value.all.return_value = [] user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( @@ -221,7 +243,7 @@ class TestDatasetServiceRetrievalPermissions: DatasetService.get_datasets( page=1, per_page=20, - session=mock_db.session, + session=session, tenant_id="tenant-1", user=user, accessible_dataset_ids=["dataset-shared"], @@ -233,11 +255,10 @@ class TestDatasetServiceRetrievalPermissions: assert "IN" in visibility_clause def test_get_datasets_by_ids_applies_rbac_visibility(self): - mock_db = MagicMock() + session = MagicMock() user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), ): @@ -248,6 +269,7 @@ class TestDatasetServiceRetrievalPermissions: user=user, accessible_dataset_ids=["dataset-shared", "dataset-not-requested"], include_own_datasets=True, + session=session, ) select_stmt = mock_paginate.call_args.args[0] @@ -260,14 +282,13 @@ class TestDatasetServiceRetrievalPermissions: assert all("dataset-not-requested" not in value for value in list_params) def test_get_datasets_rbac_include_all_uses_workspace_permission(self): - mock_db = MagicMock() - mock_db.session.scalars.return_value.all.return_value = [] + session = MagicMock() + session.scalars.return_value.all.return_value = [] user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) mock_permissions = SimpleNamespace(workspace=SimpleNamespace(permission_keys=["dataset.create_and_management"])) with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( @@ -279,42 +300,40 @@ class TestDatasetServiceRetrievalPermissions: DatasetService.get_datasets( page=1, per_page=20, - session=mock_db.session, + session=session, tenant_id="tenant-1", user=user, include_all=True, ) - mock_db.session.scalars.assert_called_once() + session.scalars.assert_called_once() mock_paginate.assert_called_once() select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 1 def test_get_datasets_rbac_without_user_returns_empty_result(self): - mock_db = MagicMock() - mock_db.session.scalars.return_value.all.return_value = [] + session = MagicMock() + session.scalars.return_value.all.return_value = [] with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), ): mock_paginate.return_value = SimpleNamespace(items=[], total=0) - DatasetService.get_datasets(page=1, per_page=20, session=mock_db.session, tenant_id="tenant-1", user=None) + DatasetService.get_datasets(page=1, per_page=20, session=session, tenant_id="tenant-1", user=None) - mock_db.session.scalars.assert_not_called() + session.scalars.assert_not_called() mock_paginate.assert_called_once() select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 2 def test_get_datasets_legacy_owner_include_all_keeps_full_access(self): - mock_db = MagicMock() - mock_db.session.scalars.return_value.all.return_value = [] + session = MagicMock() + session.scalars.return_value.all.return_value = [] user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.OWNER) with ( - patch("services.dataset_service.db", mock_db), patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", False), ): @@ -322,13 +341,13 @@ class TestDatasetServiceRetrievalPermissions: DatasetService.get_datasets( page=1, per_page=20, - session=mock_db.session, + session=session, tenant_id="tenant-1", user=user, include_all=True, ) - mock_db.session.scalars.assert_called_once() + session.scalars.assert_called_once() mock_paginate.assert_called_once() select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 1 @@ -338,22 +357,20 @@ class TestDatasetServiceCreationAndUpdate: """Unit tests for dataset creation and update helpers.""" def test_create_empty_dataset_raises_when_name_already_exists(self): + session = MagicMock() account = SimpleNamespace(id="user-1") - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = object() + session.scalar.return_value = object() - with pytest.raises(DatasetNameDuplicateError, match="Dataset with name Dataset already exists"): - DatasetService.create_empty_dataset( - "tenant-1", "Dataset", None, "economy", account, session=mock_db.session - ) + with pytest.raises(DatasetNameDuplicateError, match="Dataset with name Dataset already exists"): + DatasetService.create_empty_dataset("tenant-1", "Dataset", None, "economy", account, session=session) def test_create_empty_dataset_uses_default_embedding_model_for_high_quality_dataset(self): + session = MagicMock() account = SimpleNamespace(id="user-1") default_embedding_model = SimpleNamespace(provider="provider", model_name="default-embedding") with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.select"), patch( "services.dataset_service.Dataset", @@ -362,7 +379,7 @@ class TestDatasetServiceCreationAndUpdate: patch("services.dataset_service.ModelManager") as model_manager_cls, patch.object(DatasetService, "check_embedding_model_setting") as check_embedding, ): - mock_db.session.scalar.return_value = None + session.scalar.return_value = None model_manager_cls.for_tenant.return_value.get_default_model_instance.return_value = default_embedding_model dataset = DatasetService.create_empty_dataset( @@ -371,7 +388,7 @@ class TestDatasetServiceCreationAndUpdate: description="Description", indexing_technique="high_quality", account=account, - session=mock_db.session, + session=session, ) assert dataset.embedding_model_provider == "provider" @@ -383,15 +400,15 @@ class TestDatasetServiceCreationAndUpdate: model_type=ModelType.TEXT_EMBEDDING, ) check_embedding.assert_not_called() - mock_db.session.commit.assert_called_once() + session.commit.assert_called_once() def test_create_empty_dataset_creates_external_binding_for_high_quality_dataset(self): + session = MagicMock() account = SimpleNamespace(id="user-1") retrieval_model = _make_retrieval_model() embedding_model = SimpleNamespace(provider="provider", model_name="embedding-model") with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.select"), patch( "services.dataset_service.Dataset", @@ -406,7 +423,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "check_embedding_model_setting") as check_embedding, patch.object(DatasetService, "check_reranking_model_setting") as check_reranking, ): - mock_db.session.scalar.return_value = None + session.scalar.return_value = None model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model dataset = DatasetService.create_empty_dataset( @@ -423,7 +440,7 @@ class TestDatasetServiceCreationAndUpdate: embedding_model_name="embedding-model", retrieval_model=retrieval_model, summary_index_setting={"enable": True}, - session=mock_db.session, + session=session, ) assert dataset.embedding_model_provider == "provider" @@ -439,10 +456,11 @@ class TestDatasetServiceCreationAndUpdate: external_knowledge_id="knowledge-1", created_by="user-1", ) - assert mock_db.session.add.call_count == 2 - mock_db.session.commit.assert_called_once() + assert session.add.call_count == 2 + session.commit.assert_called_once() def test_create_empty_rag_pipeline_dataset_raises_for_duplicate_name(self): + session = MagicMock() entity = RagPipelineDatasetCreateEntity( name="Existing Dataset", description="Description", @@ -450,13 +468,13 @@ class TestDatasetServiceCreationAndUpdate: permission=DatasetPermissionEnum.ALL_TEAM, ) - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = object() + session.scalar.return_value = object() - with pytest.raises(DatasetNameDuplicateError, match="Existing Dataset already exists"): - DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) + with pytest.raises(DatasetNameDuplicateError, match="Existing Dataset already exists"): + DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, session) def test_create_empty_rag_pipeline_dataset_generates_name_and_creates_dataset(self): + session = MagicMock() entity = RagPipelineDatasetCreateEntity( name="", description="Description", @@ -473,27 +491,27 @@ class TestDatasetServiceCreationAndUpdate: return SimpleNamespace(id="dataset-1", **kwargs) with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.select"), patch("services.dataset_service.current_user", SimpleNamespace(id="user-1")), patch("services.dataset_service.generate_incremental_name", return_value="Untitled 2") as generate_name, patch("services.dataset_service.Pipeline", side_effect=pipeline_factory), patch("services.dataset_service.Dataset", side_effect=dataset_factory), ): - mock_db.session.scalars.return_value.all.return_value = [ + session.scalars.return_value.all.return_value = [ SimpleNamespace(name="Untitled"), SimpleNamespace(name="Untitled 1"), ] - dataset = DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) + dataset = DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, session) assert entity.name == "Untitled 2" assert dataset.pipeline_id == "pipeline-1" assert dataset.runtime_mode == "rag_pipeline" generate_name.assert_called_once_with(["Untitled", "Untitled 1"], "Untitled") - mock_db.session.commit.assert_called_once() + session.commit.assert_called_once() def test_create_empty_rag_pipeline_dataset_requires_current_user_id(self): + session = MagicMock() entity = RagPipelineDatasetCreateEntity( name="Dataset", description="Description", @@ -502,13 +520,12 @@ class TestDatasetServiceCreationAndUpdate: ) with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.current_user", SimpleNamespace(id=None)), ): - mock_db.session.scalar.return_value = None + session.scalar.return_value = None with pytest.raises(ValueError, match="Current user or current user id not found"): - DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, mock_db.session) + DatasetService.create_empty_rag_pipeline_dataset("tenant-1", entity, session) def test_update_dataset_raises_when_dataset_is_missing(self): session = MagicMock() @@ -568,14 +585,15 @@ class TestDatasetServiceCreationAndUpdate: update_internal.assert_called_once_with(dataset, {"name": dataset.name}, user, session) def test_has_dataset_same_name_returns_true_when_query_matches(self): - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = object() + session = MagicMock() + session.scalar.return_value = object() - result = DatasetService._has_dataset_same_name("tenant-1", "dataset-1", "Dataset", mock_db.session) + result = DatasetService._has_dataset_same_name("tenant-1", "dataset-1", "Dataset", session) assert result is True def test_update_external_dataset_updates_dataset_and_binding(self): + session = MagicMock() dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") user = SimpleNamespace(id="user-1") now = object() @@ -586,7 +604,6 @@ class TestDatasetServiceCreationAndUpdate: "services.dataset_service.ExternalDatasetService.get_external_knowledge_api", return_value=object() ) as get_external_knowledge_api, patch("services.dataset_service.naive_utc_now", return_value=now), - patch("services.dataset_service.db") as mock_db, ): result = DatasetService._update_external_dataset( dataset, @@ -600,7 +617,7 @@ class TestDatasetServiceCreationAndUpdate: "external_knowledge_api_id": "api-1", }, user, - mock_db.session, + session, ) assert result is dataset @@ -611,10 +628,10 @@ 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("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() + get_external_knowledge_api.assert_called_once_with("api-1", dataset.tenant_id, session=session) + update_binding.assert_called_once_with("dataset-1", "knowledge-1", "api-1", session) + session.add.assert_called_once_with(dataset) + session.commit.assert_called_once() @pytest.mark.parametrize( ("payload", "message"), @@ -630,6 +647,7 @@ class TestDatasetServiceCreationAndUpdate: DatasetService._update_external_dataset(dataset, payload, SimpleNamespace(id="user-1"), MagicMock()) def test_update_external_dataset_rejects_cross_tenant_external_api_id(self): + session = MagicMock() dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") with ( @@ -638,7 +656,6 @@ class TestDatasetServiceCreationAndUpdate: side_effect=ValueError("api template not found"), ) as get_external_knowledge_api, patch.object(DatasetService, "_update_external_knowledge_binding") as update_binding, - patch("services.dataset_service.db") as mock_db, ): with pytest.raises(ValueError, match="api template not found"): DatasetService._update_external_dataset( @@ -648,12 +665,12 @@ class TestDatasetServiceCreationAndUpdate: "external_knowledge_api_id": "foreign-api", }, SimpleNamespace(id="user-1"), - mock_db.session, + session, ) - get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id, session=mock_db.session) + get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id, session=session) update_binding.assert_not_called() - mock_db.session.commit.assert_not_called() + session.commit.assert_not_called() def test_update_external_knowledge_binding_updates_changed_binding_values(self): binding = SimpleNamespace(external_knowledge_id="old-knowledge", external_knowledge_api_id="old-api") @@ -673,6 +690,7 @@ class TestDatasetServiceCreationAndUpdate: DatasetService._update_external_knowledge_binding("dataset-1", "knowledge-1", "api-1", session) def test_update_internal_dataset_updates_fields_and_dispatches_regeneration_tasks(self): + session = MagicMock() dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") user = SimpleNamespace(id="user-1") now = object() @@ -692,14 +710,13 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_handle_indexing_technique_change", return_value="update"), patch.object(DatasetService, "_update_pipeline_knowledge_base_node_data") as update_pipeline, patch("services.dataset_service.naive_utc_now", return_value=now), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.deal_dataset_vector_index_task") as vector_task, patch("services.dataset_service.regenerate_summary_index_task") as regenerate_task, ): - result = DatasetService._update_internal_dataset(dataset, update_payload.copy(), user, mock_db.session) + result = DatasetService._update_internal_dataset(dataset, update_payload.copy(), user, session) assert result is dataset - updated_values = mock_db.session.execute.call_args.args[0].compile().params + updated_values = session.execute.call_args.args[0].compile().params assert updated_values["name"] == "Updated Dataset" assert updated_values["description"] is None assert updated_values["retrieval_model"] == {"top_k": 4} @@ -711,9 +728,9 @@ class TestDatasetServiceCreationAndUpdate: assert "external_knowledge_api_id" not in updated_values assert "external_knowledge_id" not in updated_values assert "external_retrieval_model" not in updated_values - mock_db.session.commit.assert_called_once() - mock_db.session.refresh.assert_called_once_with(dataset) - update_pipeline.assert_called_once_with(dataset, "user-1", mock_db.session) + session.commit.assert_called_once() + session.refresh.assert_called_once_with(dataset) + update_pipeline.assert_called_once_with(dataset, "user-1", session) vector_task.delay.assert_called_once_with("dataset-1", "update") regenerate_task.delay.assert_called_once_with( "dataset-1", @@ -722,24 +739,25 @@ class TestDatasetServiceCreationAndUpdate: ) def test_update_pipeline_knowledge_base_node_data_returns_early_for_non_pipeline_dataset(self): + session = MagicMock() dataset = SimpleNamespace(runtime_mode="workflow", pipeline_id="pipeline-1") - with patch("services.dataset_service.db") as mock_db: - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", session) - mock_db.session.get.assert_not_called() + session.get.assert_not_called() def test_update_pipeline_knowledge_base_node_data_returns_when_pipeline_is_missing(self): + session = MagicMock() dataset = SimpleNamespace(runtime_mode="rag_pipeline", pipeline_id="pipeline-1") - with patch("services.dataset_service.db") as mock_db: - mock_db.session.get.return_value = None + session.get.return_value = None - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", session) - mock_db.session.commit.assert_not_called() + session.commit.assert_not_called() def test_update_pipeline_knowledge_base_node_data_updates_published_and_draft_workflows(self): + session = MagicMock() dataset = SimpleNamespace( id="dataset-1", runtime_mode="rag_pipeline", @@ -768,38 +786,37 @@ class TestDatasetServiceCreationAndUpdate: rag_pipeline_service.get_draft_workflow.return_value = draft_workflow with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.RagPipelineService", return_value=rag_pipeline_service), patch("services.dataset_service.Workflow.new", return_value=new_workflow) as workflow_new, ): - mock_db.session.get.return_value = pipeline + session.get.return_value = pipeline - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", session) published_graph = json.loads(workflow_new.call_args.kwargs["graph"]) assert published_graph["nodes"][0]["data"]["embedding_model"] == "embedding-model" assert published_graph["nodes"][0]["data"]["summary_index_setting"] == {"enable": True} assert json.loads(draft_workflow.graph)["nodes"][0]["data"]["embedding_model_provider"] == "provider" - mock_db.session.add.assert_any_call(new_workflow) - mock_db.session.add.assert_any_call(draft_workflow) - mock_db.session.commit.assert_called_once() + session.add.assert_any_call(new_workflow) + session.add.assert_any_call(draft_workflow) + session.commit.assert_called_once() def test_update_pipeline_knowledge_base_node_data_rolls_back_when_update_fails(self): + session = MagicMock() dataset = SimpleNamespace(runtime_mode="rag_pipeline", pipeline_id="pipeline-1") pipeline = SimpleNamespace(id="pipeline-1", tenant_id="tenant-1") rag_pipeline_service = MagicMock() rag_pipeline_service.get_published_workflow.side_effect = RuntimeError("boom") with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.RagPipelineService", return_value=rag_pipeline_service), ): - mock_db.session.get.return_value = pipeline + session.get.return_value = pipeline with pytest.raises(RuntimeError, match="boom"): - DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", mock_db.session) + DatasetService._update_pipeline_knowledge_base_node_data(dataset, "user-1", session) - mock_db.session.rollback.assert_called_once() + session.rollback.assert_called_once() def test_handle_indexing_technique_change_returns_none_without_indexing_technique(self): filtered_data: dict[str, object] = {} @@ -1418,7 +1435,7 @@ class TestDatasetCollectionBindingService: class TestDatasetPermissionService: """Unit tests for dataset partial-member management helpers.""" - def test_update_partial_member_list_rolls_back_on_exception(self): + def test_update_partial_member_list_does_not_rollback_caller_session_on_exception(self): session = MagicMock() session.add_all.side_effect = RuntimeError("boom") @@ -1430,7 +1447,7 @@ class TestDatasetPermissionService: session, ) - session.rollback.assert_called_once() + session.rollback.assert_not_called() def test_check_permission_requires_dataset_editor(self): user = SimpleNamespace(is_dataset_editor=False, is_dataset_operator=False) @@ -1481,11 +1498,11 @@ class TestDatasetPermissionService: user, dataset, "partial_members", [{"user_id": "user-1"}], session=session ) - def test_clear_partial_member_list_rolls_back_on_exception(self): + def test_clear_partial_member_list_does_not_rollback_caller_session_on_exception(self): session = MagicMock() session.execute.side_effect = RuntimeError("boom") with pytest.raises(RuntimeError, match="boom"): DatasetPermissionService.clear_partial_member_list("dataset-1", session) - session.rollback.assert_called_once() + session.rollback.assert_not_called() 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 44619a29e83..0c9969815e5 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -10,6 +10,7 @@ from .dataset_service_test_helpers import ( DatasetService, DatasetServiceUnitDataFactory, DataSource, + Document, DocumentIndexingError, DocumentService, FileInfo, @@ -81,6 +82,18 @@ class TestDocumentServiceDisplayStatus: query.where.assert_called_once() +class TestDocumentServiceRetrieval: + def test_get_document_by_id_uses_provided_session(self): + session = MagicMock() + expected_document = DatasetServiceUnitDataFactory.create_document_mock(document_id="document-1") + session.get.return_value = expected_document + + result = DocumentService.get_document_by_id("document-1", session=session) + + assert result is expected_document + session.get.assert_called_once_with(Document, "document-1") + + class TestDocumentServiceMutations: """Unit tests for DocumentService mutation and orchestration helpers.""" @@ -106,26 +119,26 @@ class TestDocumentServiceMutations: assert DocumentService.check_archived(document) is expected def test_delete_documents_limits_query_and_cleanup_to_dataset_ref(self): + session = MagicMock() dataset = _make_dataset(dataset_id="dataset-1", tenant_id="tenant-1") dataset.doc_form = "paragraph_index" document = _make_document(document_id="doc-1", dataset_id=dataset.id, tenant_id=dataset.tenant_id) document.data_source_info_dict = {} with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.batch_clean_document_task") as clean_task, ): - mock_db.session.scalars.return_value.all.return_value = [document] + session.scalars.return_value.all.return_value = [document] dataset_ref = DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id) DocumentService.delete_documents( dataset_ref, ["doc-1", "other-doc"], dataset.doc_form, - mock_db.session, + session, ) - stmt = mock_db.session.scalars.call_args.args[0] + stmt = session.scalars.call_args.args[0] compiled = stmt.compile() statement = str(compiled) assert "documents.id IN" in statement @@ -134,8 +147,8 @@ class TestDocumentServiceMutations: assert ["doc-1", "other-doc"] in compiled.params.values() assert dataset.tenant_id in compiled.params.values() assert dataset.id in compiled.params.values() - mock_db.session.delete.assert_called_once_with(document) - mock_db.session.commit.assert_called_once() + session.delete.assert_called_once_with(document) + session.commit.assert_called_once() clean_task.delay.assert_called_once_with(["doc-1"], dataset.id, dataset.doc_form, []) def test_rename_document_raises_when_dataset_is_missing(self, rename_account_context): @@ -169,6 +182,7 @@ class TestDocumentServiceMutations: DocumentService.rename_document(dataset.id, document.id, "New Name", session) def test_rename_document_updates_document_metadata_and_upload_file_name(self, rename_account_context): + session = MagicMock() dataset = DatasetServiceUnitDataFactory.create_dataset_mock( built_in_field_enabled=True, tenant_id="tenant-1", @@ -183,16 +197,15 @@ class TestDocumentServiceMutations: with ( patch.object(DatasetService, "get_dataset", return_value=dataset), patch.object(DocumentService, "get_document", return_value=document), - patch("services.dataset_service.db") as mock_db, ): - result = DocumentService.rename_document(dataset.id, document.id, "New Name", mock_db.session) + result = DocumentService.rename_document(dataset.id, document.id, "New Name", session) assert result is document assert document.name == "New Name" assert document.doc_metadata[BuiltInField.document_name] == "New Name" - mock_db.session.add.assert_called_once_with(document) - mock_db.session.execute.assert_called() - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(document) + session.execute.assert_called() + session.flush.assert_called_once() def test_recover_document_raises_when_document_is_not_paused(self): document = DatasetServiceUnitDataFactory.create_document_mock(is_paused=False) @@ -222,6 +235,7 @@ class TestDocumentServiceMutations: DocumentService.sync_website_document("dataset-1", document, session) def test_sync_website_document_updates_status_sets_cache_and_dispatches_task(self): + session = MagicMock() document = DatasetServiceUnitDataFactory.create_document_mock( document_id="doc-1", data_source_info_dict={"mode": "crawl"}, @@ -230,17 +244,16 @@ class TestDocumentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.sync_website_document_indexing_task") as sync_task, ): mock_redis.get.return_value = None - DocumentService.sync_website_document("dataset-1", document, mock_db.session) + DocumentService.sync_website_document("dataset-1", document, session) assert document.indexing_status == "waiting" assert '"mode": "scrape"' in document.data_source_info - mock_db.session.add.assert_called_once_with(document) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(document) + session.commit.assert_called_once() mock_redis.setex.assert_called_once_with("document_doc-1_is_sync", 600, 1) sync_task.delay.assert_called_once_with("dataset-1", "doc-1") @@ -260,6 +273,7 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: def test_save_document_without_dataset_id_creates_high_quality_dataset_with_default_retrieval_model( self, account_context ): + session = MagicMock() knowledge_config = KnowledgeConfig( indexing_technique="high_quality", data_source=DataSource( @@ -294,13 +308,12 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: patch.object( DocumentService, "save_document_with_dataset_id", return_value=([first_document], "batch-1") ) as save_document, - patch("services.dataset_service.db") as mock_db, ): dataset, documents, batch = DocumentService.save_document_without_dataset_id( tenant_id="tenant-1", knowledge_config=knowledge_config, account=account_context, - session=mock_db.session, + session=session, ) assert dataset is created_dataset @@ -321,11 +334,12 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: created_dataset, knowledge_config, account_context, - session=mock_db.session, + session=session, ) - assert mock_db.session.commit.call_count == 1 + assert session.flush.call_count == 2 def test_save_document_without_dataset_id_uses_provided_retrieval_model(self, account_context): + session = MagicMock() retrieval_model = RetrievalModel( search_method=RetrievalMethod.SEMANTIC_SEARCH, reranking_enable=True, @@ -360,13 +374,12 @@ class TestDocumentServiceSaveDocumentWithoutDatasetId: "save_document_with_dataset_id", return_value=([SimpleNamespace(name="Doc")], "batch-1"), ), - patch("services.dataset_service.db") as mock_db, ): DocumentService.save_document_without_dataset_id( "tenant-1", knowledge_config, account_context, - mock_db.session, + session, ) assert created_dataset.retrieval_model == retrieval_model.model_dump() @@ -465,6 +478,7 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: ) def test_update_document_with_dataset_id_upload_file_process_rule_and_name_override(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document = _make_document() document.dataset_process_rule_id = "old-rule" @@ -493,17 +507,16 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: patch.object(DocumentService, "get_document", return_value=document), patch.object(DatasetService, "check_dataset_model_setting"), patch("services.dataset_service.DatasetProcessRule", return_value=created_process_rule), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.document_indexing_update_task") as update_task, ): - mock_db.session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") + session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") result = DocumentService.update_document_with_dataset_id( dataset, document_data, account_context, - session=mock_db.session, + session=session, ) assert result is document @@ -520,11 +533,12 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: assert document.updated_at == "now" assert document.created_from == "web" assert document.doc_form == IndexStructureType.QA_INDEX - assert mock_db.session.commit.call_count == 3 - mock_db.session.execute.assert_called() + assert session.commit.call_count == 3 + session.execute.assert_called() update_task.delay.assert_called_once_with(document.dataset_id, document.id) def test_update_document_with_dataset_id_notion_import_requires_binding(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document = SimpleNamespace(display_status="available", id="doc-1", dataset_id="dataset-1") document_data = KnowledgeConfig( @@ -547,19 +561,19 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: with ( patch.object(DocumentService, "get_document", return_value=document), patch.object(DatasetService, "check_dataset_model_setting"), - patch("services.dataset_service.db") as mock_db, ): - mock_db.session.scalar.return_value = None + session.scalar.return_value = None with pytest.raises(ValueError, match="Data source binding not found"): DocumentService.update_document_with_dataset_id( dataset, document_data, account_context, - session=mock_db.session, + session=session, ) def test_update_document_with_dataset_id_website_crawl_updates_segments_and_dispatches_task(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document = _make_document() document_data = KnowledgeConfig( @@ -582,7 +596,6 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: with ( patch.object(DocumentService, "get_document", return_value=document), patch.object(DatasetService, "check_dataset_model_setting"), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.document_indexing_update_task") as update_task, ): @@ -590,7 +603,7 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: dataset, document_data, account_context, - session=mock_db.session, + session=session, ) assert result is document @@ -601,7 +614,7 @@ class TestDocumentServiceUpdateDocumentWithDatasetId: ) assert document.name == "" assert document.doc_form == IndexStructureType.PARENT_CHILD_INDEX - mock_db.session.execute.assert_called() + session.execute.assert_called() update_task.delay.assert_called_once_with("dataset-1", "doc-1") @@ -868,6 +881,8 @@ class TestDocumentServiceSaveDocumentWithDatasetId: session=session, ) + dataset.get_latest_process_rule.assert_called_once_with(session=session) + def test_save_document_with_dataset_id_rejects_invalid_indexing_technique(self, account_context): dataset = _make_dataset(indexing_technique=None) knowledge_config = SimpleNamespace( @@ -904,6 +919,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: assert batch == "" def test_save_document_with_dataset_id_upload_file_creates_and_reindexes_documents(self, account_context): + session = MagicMock() dataset = _make_dataset() dataset_process_rule = SimpleNamespace(id="rule-1") knowledge_config = _make_upload_knowledge_config(file_ids=["file-1", "file-2"]) @@ -915,7 +931,6 @@ class TestDocumentServiceSaveDocumentWithDatasetId: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch.object(DocumentService, "get_documents_position", return_value=4), patch.object(DocumentService, "build_document", return_value=created_document) as build_document, patch("services.dataset_service.DocumentIndexingTaskProxy") as document_proxy_cls, @@ -925,7 +940,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: patch("services.dataset_service.secrets.randbelow", return_value=23), ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalars.return_value.all.side_effect = [ + session.scalars.return_value.all.side_effect = [ [upload_file_a, upload_file_b], [duplicate_document], ] @@ -935,7 +950,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config, account_context, dataset_process_rule=dataset_process_rule, - session=mock_db.session, + session=session, ) assert documents == [duplicate_document, created_document] @@ -965,6 +980,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: def test_save_document_with_dataset_id_notion_import_truncates_names_and_cleans_removed_pages( self, account_context ): + session = MagicMock() dataset = _make_dataset() dataset_process_rule = SimpleNamespace(id="rule-1") notion_page_name = "a" * 300 @@ -1002,21 +1018,20 @@ class TestDocumentServiceSaveDocumentWithDatasetId: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch.object(DocumentService, "get_documents_position", return_value=1), patch.object(DocumentService, "build_document", return_value=created_document) as build_document, patch("services.dataset_service.clean_notion_document_task") as clean_task, patch("services.dataset_service.DocumentIndexingTaskProxy") as document_proxy_cls, ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalars.return_value.all.return_value = [existing_keep, existing_remove] + session.scalars.return_value.all.return_value = [existing_keep, existing_remove] documents, _ = DocumentService.save_document_with_dataset_id( dataset, knowledge_config, account_context, dataset_process_rule=dataset_process_rule, - session=mock_db.session, + session=session, ) assert created_document in documents @@ -1026,6 +1041,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: document_proxy_cls.return_value.delay.assert_called_once() def test_save_document_with_dataset_id_website_crawl_truncates_long_urls(self, account_context): + session = MagicMock() dataset = _make_dataset() dataset_process_rule = SimpleNamespace(id="rule-1") long_url = "https://example.com/" + ("a" * 260) @@ -1052,7 +1068,6 @@ class TestDocumentServiceSaveDocumentWithDatasetId: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch.object(DocumentService, "get_documents_position", return_value=2), patch.object( DocumentService, @@ -1068,7 +1083,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: knowledge_config, account_context, dataset_process_rule=dataset_process_rule, - session=mock_db.session, + session=session, ) assert documents == [first_document, second_document] @@ -1109,50 +1124,50 @@ class TestDocumentServiceBatchUpdateStatus: assert result["async_task"]["args"] == [document.id] def test_batch_update_document_status_rejects_indexing_documents(self): + session = MagicMock() dataset = _make_dataset() document = _make_document(name="Busy document") with ( patch.object(DocumentService, "get_document", return_value=document), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, ): mock_redis.get.return_value = "1" with pytest.raises(DocumentIndexingError, match="Busy document is being indexed"): DocumentService.batch_update_document_status( - dataset, [document.id], "archive", SimpleNamespace(id="user-1"), mock_db.session + dataset, [document.id], "archive", SimpleNamespace(id="user-1"), session ) - mock_db.session.commit.assert_not_called() + session.flush.assert_not_called() def test_batch_update_document_status_rolls_back_when_commit_fails(self): + session = MagicMock() dataset = _make_dataset() document = _make_document(enabled=False) with ( patch.object(DocumentService, "get_document", return_value=document), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, ): mock_redis.get.return_value = None - mock_db.session.commit.side_effect = RuntimeError("commit failed") + session.commit.side_effect = RuntimeError("commit failed") with pytest.raises(RuntimeError, match="commit failed"): DocumentService.batch_update_document_status( - dataset, [document.id], "enable", SimpleNamespace(id="user-1"), mock_db.session + dataset, [document.id], "enable", SimpleNamespace(id="user-1"), session ) - mock_db.session.rollback.assert_called_once() + session.rollback.assert_called_once() def test_batch_update_document_status_raises_async_task_error_after_commit(self): + session = MagicMock() dataset = _make_dataset() document = _make_document(enabled=False) with ( patch.object(DocumentService, "get_document", return_value=document), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.add_document_to_index_task") as add_task, ): mock_redis.get.return_value = None @@ -1160,10 +1175,10 @@ class TestDocumentServiceBatchUpdateStatus: with pytest.raises(RuntimeError, match="task failed"): DocumentService.batch_update_document_status( - dataset, [document.id], "enable", SimpleNamespace(id="user-1"), mock_db.session + dataset, [document.id], "enable", SimpleNamespace(id="user-1"), session ) - mock_db.session.commit.assert_called_once() + session.commit.assert_called_once() mock_redis.setex.assert_called_once_with(f"document_{document.id}_indexing", 600, 1) @@ -1180,14 +1195,15 @@ class TestDocumentServiceTenantAndUpdateEdges: yield account def test_get_tenant_documents_count_returns_query_count(self, account_context): - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = 12 + session = MagicMock() + session.scalar.return_value = 12 - result = DocumentService.get_tenant_documents_count(session=mock_db.session) + result = DocumentService.get_tenant_documents_count(session) assert result == 12 def test_update_document_with_dataset_id_uses_automatic_process_rule_payload(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document = _make_document() document_data = KnowledgeConfig( @@ -1214,19 +1230,18 @@ class TestDocumentServiceTenantAndUpdateEdges: patch.object(DocumentService, "get_document", return_value=document), patch("services.dataset_service.DatasetProcessRule") as process_rule_cls, patch.object(DatasetService, "check_dataset_model_setting"), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.document_indexing_update_task") as update_task, ): process_rule_cls.AUTOMATIC_RULES = DatasetProcessRule.AUTOMATIC_RULES process_rule_cls.return_value = created_process_rule - mock_db.session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") + session.scalar.return_value = SimpleNamespace(id="file-1", name="upload.txt") result = DocumentService.update_document_with_dataset_id( dataset, document_data, account_context, - session=mock_db.session, + session=session, ) assert result is document @@ -1238,7 +1253,7 @@ class TestDocumentServiceTenantAndUpdateEdges: "rules": json.dumps(DatasetProcessRule.AUTOMATIC_RULES), "created_by": "user-1", } - assert mock_db.session.commit.call_count == 3 + assert session.commit.call_count == 3 update_task.delay.assert_called_once_with("dataset-1", "doc-1") def test_update_document_with_dataset_id_requires_upload_file_info(self, account_context): @@ -1263,6 +1278,7 @@ class TestDocumentServiceTenantAndUpdateEdges: ) def test_update_document_with_dataset_id_raises_when_upload_file_is_missing(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document_data = KnowledgeConfig( original_document_id="doc-1", @@ -1278,16 +1294,15 @@ class TestDocumentServiceTenantAndUpdateEdges: with ( patch.object(DocumentService, "get_document", return_value=_make_document()), patch.object(DatasetService, "check_dataset_model_setting"), - patch("services.dataset_service.db") as mock_db, ): - mock_db.session.scalar.return_value = None + session.scalar.return_value = None with pytest.raises(FileNotExistsError): DocumentService.update_document_with_dataset_id( dataset, document_data, account_context, - session=mock_db.session, + session=session, ) def test_update_document_with_dataset_id_requires_notion_info_list(self, account_context): @@ -1312,6 +1327,7 @@ class TestDocumentServiceTenantAndUpdateEdges: ) def test_update_document_with_dataset_id_notion_import_updates_page_info(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") document = _make_document() document_data = KnowledgeConfig( @@ -1338,17 +1354,16 @@ class TestDocumentServiceTenantAndUpdateEdges: with ( patch.object(DocumentService, "get_document", return_value=document), patch.object(DatasetService, "check_dataset_model_setting"), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.document_indexing_update_task") as update_task, ): - mock_db.session.scalar.return_value = SimpleNamespace(id="binding-1") + session.scalar.return_value = SimpleNamespace(id="binding-1") result = DocumentService.update_document_with_dataset_id( dataset, document_data, account_context, - session=mock_db.session, + session=session, ) assert result is document @@ -1379,6 +1394,7 @@ class TestDocumentServiceSaveWithoutDatasetBilling: yield account def test_save_document_without_dataset_id_counts_notion_pages_for_quota(self, account_context): + session = MagicMock() knowledge_config = KnowledgeConfig( indexing_technique="economy", data_source=DataSource( @@ -1418,13 +1434,12 @@ class TestDocumentServiceSaveWithoutDatasetBilling: "save_document_with_dataset_id", return_value=([SimpleNamespace(name="Doc")], "batch-1"), ), - patch("services.dataset_service.db") as mock_db, ): DocumentService.save_document_without_dataset_id( "tenant-1", knowledge_config, account_context, - mock_db.session, + session, ) check_quota.assert_called_once_with(3, features) @@ -1683,6 +1698,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: assert dataset.retrieval_model == knowledge_config.retrieval_model.model_dump() def test_save_document_with_dataset_id_creates_custom_process_rule_for_new_upload_document(self, account_context): + session = MagicMock() dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config( file_ids=["file-1"], @@ -1700,7 +1716,6 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.DatasetProcessRule") as process_rule_cls, patch.object(DocumentService, "get_documents_position", return_value=3), patch.object(DocumentService, "build_document", return_value=created_document), @@ -1710,13 +1725,13 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: ): mock_redis.lock.return_value = _make_lock_context() process_rule_cls.return_value = created_process_rule - mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] + session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] documents, batch = DocumentService.save_document_with_dataset_id( dataset, knowledge_config, account_context, - session=mock_db.session, + session=session, ) assert documents == [created_document] @@ -1733,6 +1748,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: def test_save_document_with_dataset_id_creates_automatic_process_rule_for_new_upload_document( self, account_context ): + session = MagicMock() dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config( file_ids=["file-1"], @@ -1744,7 +1760,6 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.DatasetProcessRule") as process_rule_cls, patch.object(DocumentService, "get_documents_position", return_value=1), patch.object(DocumentService, "build_document", return_value=created_document), @@ -1755,13 +1770,13 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: mock_redis.lock.return_value = _make_lock_context() process_rule_cls.AUTOMATIC_RULES = DatasetProcessRule.AUTOMATIC_RULES process_rule_cls.return_value = created_process_rule - mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] + session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] DocumentService.save_document_with_dataset_id( dataset, knowledge_config, account_context, - session=mock_db.session, + session=session, ) assert process_rule_cls.call_args.kwargs == { @@ -1770,11 +1785,12 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: "rules": json.dumps(DatasetProcessRule.AUTOMATIC_RULES), "created_by": "user-1", } - assert mock_db.session.flush.call_count >= 2 + assert session.flush.call_count >= 2 def test_save_document_with_dataset_id_creates_fallback_automatic_process_rule_when_latest_is_missing( self, account_context ): + session = MagicMock() dataset = _make_dataset(latest_process_rule=None) knowledge_config = _make_upload_knowledge_config(file_ids=["file-1"], process_rule=None) created_process_rule = SimpleNamespace(id="rule-fallback") @@ -1783,7 +1799,6 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.DatasetProcessRule") as process_rule_cls, patch.object(DocumentService, "get_documents_position", return_value=1), patch.object(DocumentService, "build_document", return_value=created_document), @@ -1794,15 +1809,16 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: mock_redis.lock.return_value = _make_lock_context() process_rule_cls.AUTOMATIC_RULES = DatasetProcessRule.AUTOMATIC_RULES process_rule_cls.return_value = created_process_rule - mock_db.session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] + session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] DocumentService.save_document_with_dataset_id( dataset, knowledge_config, account_context, - session=mock_db.session, + session=session, ) + dataset.get_latest_process_rule.assert_called_once_with(session=session) assert process_rule_cls.call_args.kwargs == { "dataset_id": "dataset-1", "mode": "automatic", @@ -1811,26 +1827,26 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: } def test_save_document_with_dataset_id_raises_when_upload_file_lookup_is_incomplete(self, account_context): + session = MagicMock() dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config(file_ids=["file-1", "file-2"]) with ( patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch.object(DocumentService, "get_documents_position", return_value=1), patch("services.dataset_service.time.strftime", return_value="20260101010101"), patch("services.dataset_service.secrets.randbelow", return_value=23), ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalars.return_value.all.return_value = [SimpleNamespace(id="file-1", name="file.txt")] + session.scalars.return_value.all.return_value = [SimpleNamespace(id="file-1", name="file.txt")] with pytest.raises(FileNotExistsError, match="One or more files not found"): DocumentService.save_document_with_dataset_id( dataset, knowledge_config, account_context, - session=mock_db.session, + session=session, ) def test_save_document_with_dataset_id_requires_notion_info_list_for_notion_import(self, account_context): 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 c94093d59b7..b53cd152f07 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -74,26 +74,26 @@ class TestSegmentServiceChildChunks: yield account def test_create_child_chunk_assigns_next_position_and_commits(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") document = _make_document() segment = _make_segment() with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.uuid.uuid4", return_value="node-1"), patch("services.dataset_service.helper.generate_text_hash", return_value="hash-1"), patch("services.dataset_service.VectorService") as vector_service, ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalar.return_value = 2 + session.scalar.return_value = 2 child_chunk = SegmentService.create_child_chunk( "child content", segment, document, dataset, - mock_db.session, + session, ) assert isinstance(child_chunk, ChildChunk) @@ -101,33 +101,34 @@ class TestSegmentServiceChildChunks: assert child_chunk.index_node_id == "node-1" assert child_chunk.index_node_hash == "hash-1" assert child_chunk.word_count == len("child content") - mock_db.session.add.assert_called_once_with(child_chunk) - vector_service.create_child_chunk_vector.assert_called_once_with(child_chunk, dataset) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(child_chunk) + vector_service.create_child_chunk_vector.assert_called_once_with(child_chunk, dataset, session=session) + session.commit.assert_called() def test_create_child_chunk_rolls_back_and_raises_on_vector_failure(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") document = _make_document() segment = _make_segment() with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.uuid.uuid4", return_value="node-1"), patch("services.dataset_service.helper.generate_text_hash", return_value="hash-1"), patch("services.dataset_service.VectorService") as vector_service, ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalar.return_value = None + session.scalar.return_value = None vector_service.create_child_chunk_vector.side_effect = RuntimeError("vector failed") with pytest.raises(ChildChunkIndexingError, match="vector failed"): - SegmentService.create_child_chunk("child content", segment, document, dataset, mock_db.session) + SegmentService.create_child_chunk("child content", segment, document, dataset, session) - mock_db.session.rollback.assert_called_once() - mock_db.session.commit.assert_not_called() + session.rollback.assert_called_once() + session.commit.assert_not_called() def test_update_child_chunks_updates_deletes_and_creates_records(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") document = _make_document() segment = _make_segment() @@ -154,13 +155,12 @@ class TestSegmentServiceChildChunks: existing_a.id = "child-a" existing_b.id = "child-b" with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.uuid.uuid4", return_value="node-new"), patch("services.dataset_service.helper.generate_text_hash", return_value="hash-new"), patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.VectorService") as vector_service, ): - mock_db.session.scalars.return_value.all.return_value = [existing_a, existing_b] + session.scalars.return_value.all.return_value = [existing_a, existing_b] result = SegmentService.update_child_chunks( [ @@ -170,36 +170,36 @@ class TestSegmentServiceChildChunks: segment, document, dataset, - mock_db.session, + session, ) assert [chunk.position for chunk in result] == [1, 3] assert existing_a.content == "updated content" assert existing_a.updated_by == account_context.id assert existing_a.updated_at == "now" - mock_db.session.bulk_save_objects.assert_called_once_with([existing_a]) - mock_db.session.delete.assert_called_once_with(existing_b) + session.bulk_save_objects.assert_called_once_with([existing_a]) + session.delete.assert_called_once_with(existing_b) new_chunk = result[1] assert isinstance(new_chunk, ChildChunk) assert new_chunk.position == 3 assert new_chunk.index_node_id == "node-new" vector_service.update_child_chunk_vector.assert_called_once_with( - [new_chunk], [existing_a], [existing_b], dataset + [new_chunk], [existing_a], [existing_b], dataset, session=session ) - mock_db.session.commit.assert_called_once() + session.commit.assert_called() def test_update_child_chunks_rolls_back_on_vector_failure(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") document = _make_document() segment = _make_segment() existing_chunk = _make_child_chunk() with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.VectorService") as vector_service, ): - mock_db.session.scalars.return_value.all.return_value = [existing_chunk] + session.scalars.return_value.all.return_value = [existing_chunk] vector_service.update_child_chunk_vector.side_effect = RuntimeError("vector failed") with pytest.raises(ChildChunkIndexingError, match="vector failed"): @@ -208,23 +208,23 @@ class TestSegmentServiceChildChunks: segment, document, dataset, - mock_db.session, + session, ) - mock_db.session.rollback.assert_called_once() + session.rollback.assert_called_once() def test_update_child_chunk_updates_vector_and_commits(self, account_context): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") child_chunk = _make_child_chunk() with ( patch("services.dataset_service.current_user", SimpleNamespace(id="user-1")), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.VectorService") as vector_service, ): result = SegmentService.update_child_chunk( - "new content", child_chunk, _make_segment(), _make_document(), dataset, mock_db.session + "new content", child_chunk, _make_segment(), _make_document(), dataset, session ) assert result is child_chunk @@ -232,25 +232,27 @@ class TestSegmentServiceChildChunks: assert child_chunk.word_count == len("new content") assert child_chunk.updated_by == "user-1" assert child_chunk.updated_at == "now" - mock_db.session.add.assert_called_once_with(child_chunk) - vector_service.update_child_chunk_vector.assert_called_once_with([], [child_chunk], [], dataset) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(child_chunk) + vector_service.update_child_chunk_vector.assert_called_once_with( + [], [child_chunk], [], dataset, session=session + ) + session.commit.assert_called() def test_delete_child_chunk_raises_delete_index_error_on_vector_failure(self): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") child_chunk = _make_child_chunk() with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, ): vector_service.delete_child_chunk_vector.side_effect = RuntimeError("delete failed") with pytest.raises(ChildChunkDeleteIndexError, match="delete failed"): - SegmentService.delete_child_chunk(child_chunk, dataset, mock_db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, session) - mock_db.session.delete.assert_called_once_with(child_chunk) - mock_db.session.rollback.assert_called_once() + session.delete.assert_called_once_with(child_chunk) + session.rollback.assert_called_once() class TestSegmentServiceQueries: @@ -266,10 +268,10 @@ class TestSegmentServiceQueries: yield account def test_get_child_chunks_applies_keyword_filter_and_paginate(self, account_context): + session = MagicMock() paginated = SimpleNamespace(items=["chunk"], total=1) with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like, ): @@ -282,6 +284,7 @@ class TestSegmentServiceQueries: page=2, limit=10, keyword="needle", + session=session, ) assert result is paginated @@ -289,27 +292,28 @@ class TestSegmentServiceQueries: mock_paginate.assert_called_once() def test_get_child_chunk_by_id_returns_only_child_chunk_instances(self): + session = MagicMock() child_chunk = _make_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_id("child-a", "tenant-1", mock_db.session) + session.scalar.return_value = child_chunk + result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1", session) assert result is child_chunk - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = SimpleNamespace() - result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1", mock_db.session) + session.scalar.return_value = SimpleNamespace() + result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1", session) assert result is None def test_get_child_chunk_by_segment_ref_uses_full_ownership_chain(self): + session = MagicMock() child_chunk = _make_child_chunk() segment_ref = _make_segment_ref() session = MagicMock() session.scalar.return_value = child_chunk - result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref, session) + session.scalar.return_value = child_chunk + result = SegmentService.get_child_chunk_by_segment_ref("child-a", segment_ref, session=session) assert result is child_chunk stmt = session.scalar.call_args.args[0] @@ -321,10 +325,10 @@ class TestSegmentServiceQueries: assert "child_chunks.segment_id = 'segment-1'" in sql def test_get_segments_uses_status_and_keyword_filters(self): + session = MagicMock() paginated = SimpleNamespace(items=["segment"], total=1) with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like, ): @@ -337,6 +341,7 @@ class TestSegmentServiceQueries: keyword="needle", page=1, limit=20, + session=session, ) assert items == ["segment"] @@ -345,6 +350,7 @@ class TestSegmentServiceQueries: mock_paginate.assert_called_once() def test_get_segment_by_id_returns_only_document_segment_instances(self): + session = MagicMock() segment = DocumentSegment( tenant_id="tenant-1", dataset_id="dataset-1", @@ -356,19 +362,18 @@ class TestSegmentServiceQueries: created_by="user-1", ) segment.id = "segment-1" - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = segment - result = SegmentService.get_segment_by_id("segment-1", "tenant-1", mock_db.session) + session.scalar.return_value = segment + result = SegmentService.get_segment_by_id("segment-1", "tenant-1", session) assert result is segment - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalar.return_value = SimpleNamespace() - result = SegmentService.get_segment_by_id("segment-1", "tenant-1", mock_db.session) + session.scalar.return_value = SimpleNamespace() + result = SegmentService.get_segment_by_id("segment-1", "tenant-1", session) assert result is None def test_get_segment_by_ref_uses_full_ownership_chain(self): + session = MagicMock() segment = DocumentSegment( tenant_id="tenant-1", dataset_id="dataset-1", @@ -384,7 +389,8 @@ class TestSegmentServiceQueries: session = MagicMock() session.scalar.return_value = segment - result = SegmentService.get_segment_by_ref(segment_ref, session) + session.scalar.return_value = segment + result = SegmentService.get_segment_by_ref(segment_ref, session=session) assert result is segment stmt = session.scalar.call_args.args[0] @@ -395,6 +401,7 @@ class TestSegmentServiceQueries: assert "document_segments.document_id = 'doc-1'" in sql def test_get_segments_by_document_and_dataset_returns_scalars_result(self): + session = MagicMock() segment = DocumentSegment( tenant_id="tenant-1", dataset_id="dataset-1", @@ -407,19 +414,18 @@ class TestSegmentServiceQueries: ) segment.id = "segment-1" - with patch("services.dataset_service.db") as mock_db: - mock_db.session.scalars.return_value.all.return_value = [segment] + session.scalars.return_value.all.return_value = [segment] - result = SegmentService.get_segments_by_document_and_dataset( - document_id="doc-1", - dataset_id="dataset-1", - session=mock_db.session, - status="completed", - enabled=True, - ) + result = SegmentService.get_segments_by_document_and_dataset( + document_id="doc-1", + dataset_id="dataset-1", + session=session, + status="completed", + enabled=True, + ) assert result == [segment] - mock_db.session.scalars.assert_called_once() + session.scalars.assert_called_once() class TestSegmentServiceValidation: @@ -465,6 +471,7 @@ class TestSegmentServiceMutations: yield account def test_create_segment_creates_bindings_and_marks_segment_error_on_vector_failure(self, account_context): + session = MagicMock() dataset = _make_dataset(indexing_technique="economy") document = _make_document( dataset_id=dataset.id, @@ -482,7 +489,6 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", return_value="hash-1"), patch("services.dataset_service.uuid.uuid4", return_value="node-1"), @@ -490,27 +496,27 @@ class TestSegmentServiceMutations: ): mock_redis.lock.return_value = _make_lock_context() - mock_db.session.scalar.return_value = 2 - mock_db.session.get.return_value = refreshed_segment + session.scalar.return_value = 2 + session.get.return_value = refreshed_segment def add_side_effect(obj): if obj.__class__.__name__ == "DocumentSegment" and getattr(obj, "id", None) is None: obj.id = "segment-1" - mock_db.session.add.side_effect = add_side_effect + session.add.side_effect = add_side_effect vector_service.create_segments_vector.side_effect = RuntimeError("vector failed") result = SegmentService.create_segment( args=args, document=document, dataset=dataset, - session=mock_db.session, + session=session, ) created_segment = vector_service.create_segments_vector.call_args.args[1][0] attachment_bindings = [ call.args[0] - for call in mock_db.session.add.call_args_list + for call in session.add.call_args_list if call.args and call.args[0].__class__.__name__ == "SegmentAttachmentBinding" ] @@ -524,9 +530,10 @@ class TestSegmentServiceMutations: assert document.word_count == len("question") + len("answer") assert len(attachment_bindings) == 2 assert {binding.attachment_id for binding in attachment_bindings} == {"att-1", "att-2"} - assert mock_db.session.commit.call_count == 3 + assert session.commit.call_count == 3 def test_multi_create_segment_high_quality_marks_segments_error_when_vector_creation_fails(self, account_context): + session = MagicMock() dataset = _make_dataset(indexing_technique="high_quality") document = _make_document( dataset_id=dataset.id, @@ -543,7 +550,6 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", side_effect=["hash-1", "hash-2"]), @@ -552,10 +558,10 @@ class TestSegmentServiceMutations: ): mock_redis.lock.return_value = _make_lock_context() model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - mock_db.session.scalar.return_value = 1 + session.scalar.return_value = 1 vector_service.create_segments_vector.side_effect = RuntimeError("vector failed") - result = SegmentService.multi_create_segment(segments, document, dataset, mock_db.session) + result = SegmentService.multi_create_segment(segments, document, dataset, session) assert result assert len(result) == 2 @@ -566,11 +572,12 @@ 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, mock_db.session + [["k1"], None], result, dataset, document.doc_form, session=session ) - mock_db.session.commit.assert_called_once() + session.commit.assert_called() def test_update_segment_disables_enabled_segment_and_dispatches_index_cleanup(self, account_context): + session = MagicMock() segment = _make_segment(enabled=True) document = _make_document() dataset = _make_dataset() @@ -578,20 +585,19 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.disable_segment_from_index_task") as disable_task, ): mock_redis.get.return_value = None - result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) + result = SegmentService.update_segment(args, segment, document, dataset, session) assert result is segment assert segment.enabled is False assert segment.disabled_at == "now" assert segment.disabled_by == account_context.id - mock_db.session.add.assert_called_once_with(segment) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(segment) + session.commit.assert_called() mock_redis.setex.assert_called_once_with(f"segment_{segment.id}_indexing", 600, 1) disable_task.delay.assert_called_once_with(segment.id) @@ -622,6 +628,7 @@ class TestSegmentServiceMutations: ) def test_update_segment_updates_keywords_for_same_content_segment(self, account_context): + session = MagicMock() segment = _make_segment(content="same content", keywords=["old"]) document = _make_document(doc_form=IndexStructureType.PARAGRAPH_INDEX, word_count=20) dataset = _make_dataset() @@ -630,20 +637,20 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, ): mock_redis.get.return_value = None - mock_db.session.get.return_value = refreshed_segment + session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) + result = SegmentService.update_segment(args, segment, document, dataset, session) 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, mock_db.session) + vector_service.update_segment_vector.assert_called_once_with(["new"], segment, dataset, session=session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_regenerates_child_chunks_and_updates_manual_summary(self, account_context): + session = MagicMock() segment = _make_segment(content="same content", word_count=len("same content")) document = _make_document( doc_form=IndexStructureType.PARENT_CHILD_INDEX, @@ -662,7 +669,6 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.summary_index_service.SummaryIndexService.update_summary_for_segment") as update_summary, @@ -671,11 +677,11 @@ class TestSegmentServiceMutations: model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model_instance # get calls: processing_rule, then refreshed_segment - mock_db.session.get.side_effect = [processing_rule, refreshed_segment] + session.get.side_effect = [processing_rule, refreshed_segment] # scalar call: existing_summary - mock_db.session.scalar.return_value = existing_summary + session.scalar.return_value = existing_summary - result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) + result = SegmentService.update_segment(args, segment, document, dataset, session) assert result is refreshed_segment vector_service.generate_child_chunks.assert_called_once_with( @@ -684,13 +690,14 @@ class TestSegmentServiceMutations: dataset, embedding_model_instance, processing_rule, - mock_db.session, True, + session=session, ) - 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) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_auto_regenerates_summary_after_content_change(self, account_context): + session = MagicMock() segment = _make_segment(content="old", word_count=3) document = _make_document(doc_form=IndexStructureType.PARAGRAPH_INDEX, word_count=10) dataset = _make_dataset(indexing_technique="high_quality") @@ -703,7 +710,6 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", return_value="hash-1"), @@ -715,21 +721,22 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = None model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - mock_db.session.scalar.return_value = existing_summary - mock_db.session.get.return_value = refreshed_segment + session.scalar.return_value = existing_summary + session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) + result = SegmentService.update_segment(args, segment, document, dataset, session) assert result is refreshed_segment assert segment.content == "new content" assert segment.index_node_hash == "hash-1" 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}, session=mock_db.session) - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) + vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset, session=session) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_regenerates_summary_when_manual_summary_is_unchanged(self, account_context): + session = MagicMock() segment = _make_segment(content="old", word_count=3) document = _make_document(doc_form=IndexStructureType.PARAGRAPH_INDEX, word_count=10) dataset = _make_dataset(indexing_technique="high_quality") @@ -742,7 +749,6 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", return_value="hash-2"), @@ -755,30 +761,30 @@ class TestSegmentServiceMutations: mock_redis.get.return_value = None model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - mock_db.session.scalar.return_value = existing_summary - mock_db.session.get.return_value = refreshed_segment + session.scalar.return_value = existing_summary + session.get.return_value = refreshed_segment - result = SegmentService.update_segment(args, segment, document, dataset, mock_db.session) + result = SegmentService.update_segment(args, segment, document, dataset, session) assert result is refreshed_segment - generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=mock_db.session) + generate_summary.assert_called_once_with(segment, dataset, {"enable": True}, session=session) update_summary.assert_not_called() - vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_delete_segment_removes_index_and_updates_document_word_count(self): + session = MagicMock() segment = _make_segment(word_count=4, index_node_id="parent-node") document = _make_document(word_count=10) dataset = _make_dataset() with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.delete_segment_from_index_task") as delete_task, ): mock_redis.get.return_value = None - mock_db.session.scalars.return_value.all.return_value = ["child-1", "child-2"] + session.scalars.return_value.all.return_value = ["child-1", "child-2"] - SegmentService.delete_segment(segment, document, dataset, mock_db.session) + SegmentService.delete_segment(segment, document, dataset, session) assert document.word_count == 6 mock_redis.setex.assert_called_once_with(f"segment_{segment.id}_delete_indexing", 600, 1) @@ -789,9 +795,9 @@ class TestSegmentServiceMutations: [segment.id], ["child-1", "child-2"], ) - mock_db.session.delete.assert_called_once_with(segment) - mock_db.session.add.assert_called_once_with(document) - mock_db.session.commit.assert_called_once() + session.delete.assert_called_once_with(segment) + session.add.assert_called_once_with(document) + session.commit.assert_called() def test_delete_segment_rejects_when_delete_is_already_in_progress(self): segment = _make_segment() @@ -805,13 +811,13 @@ class TestSegmentServiceMutations: SegmentService.delete_segment(segment, document, dataset, MagicMock()) def test_delete_segments_removes_records_and_clamps_document_word_count(self): + session = MagicMock() dataset = _make_dataset() document = _make_document(word_count=3) current_user = SimpleNamespace(current_tenant_id="tenant-1") with ( patch("services.dataset_service.current_user", current_user), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.delete_segment_from_index_task") as delete_task, ): # execute().all() for segments_info (multi-column) @@ -820,14 +826,14 @@ class TestSegmentServiceMutations: ("node-1", "segment-1", 2), ("node-2", "segment-2", 5), ] - mock_db.session.execute.return_value = execute_result + session.execute.return_value = execute_result # scalars() for child_node_ids - mock_db.session.scalars.return_value.all.return_value = ["child-1"] + session.scalars.return_value.all.return_value = ["child-1"] - SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset, mock_db.session) + SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset, session) assert document.word_count == 0 - mock_db.session.add.assert_called_once_with(document) + session.add.assert_called_once_with(document) delete_task.delay.assert_called_once_with( ["node-1", "node-2"], dataset.id, @@ -835,9 +841,10 @@ class TestSegmentServiceMutations: ["segment-1", "segment-2"], ["child-1"], ) - mock_db.session.commit.assert_called_once() + session.commit.assert_called() def test_update_segments_status_enables_only_segments_without_indexing_cache(self): + session = MagicMock() dataset = _make_dataset() document = _make_document() segment_a = _make_segment(segment_id="segment-a", enabled=False) @@ -847,26 +854,24 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.current_user", current_user), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.enable_segments_to_index_task") as enable_task, ): - mock_db.session.scalars.return_value.all.return_value = [segment_a, segment_b] + session.scalars.return_value.all.return_value = [segment_a, segment_b] mock_redis.get.side_effect = [None, "1"] - SegmentService.update_segments_status( - ["segment-a", "segment-b"], "enable", dataset, document, mock_db.session - ) + SegmentService.update_segments_status(["segment-a", "segment-b"], "enable", dataset, document, session) assert segment_a.enabled is True assert segment_a.disabled_at is None assert segment_a.disabled_by is None assert segment_b.enabled is False - mock_db.session.add.assert_called_once_with(segment_a) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(segment_a) + session.commit.assert_called() enable_task.delay.assert_called_once_with(["segment-a"], dataset.id, document.id) def test_update_segments_status_disables_only_segments_without_indexing_cache(self): + session = MagicMock() dataset = _make_dataset() document = _make_document() segment_a = _make_segment(segment_id="segment-a", enabled=True) @@ -876,23 +881,20 @@ class TestSegmentServiceMutations: with ( patch("services.dataset_service.current_user", current_user), patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.disable_segments_from_index_task") as disable_task, ): - mock_db.session.scalars.return_value.all.return_value = [segment_a, segment_b] + session.scalars.return_value.all.return_value = [segment_a, segment_b] mock_redis.get.side_effect = [None, "1"] - SegmentService.update_segments_status( - ["segment-a", "segment-b"], "disable", dataset, document, mock_db.session - ) + SegmentService.update_segments_status(["segment-a", "segment-b"], "disable", dataset, document, session) assert segment_a.enabled is False assert segment_a.disabled_at == "now" assert segment_a.disabled_by == current_user.id assert segment_b.enabled is True - mock_db.session.add.assert_called_once_with(segment_a) - mock_db.session.commit.assert_called_once() + session.add.assert_called_once_with(segment_a) + session.commit.assert_called() disable_task.delay.assert_called_once_with(["segment-a"], dataset.id, document.id) @@ -900,12 +902,12 @@ class TestSegmentServiceChildChunkTailHelpers: """Unit tests for the remaining child-chunk helper branches.""" def test_update_child_chunk_rolls_back_on_vector_failure(self): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") child_chunk = _make_child_chunk() with ( patch("services.dataset_service.current_user", SimpleNamespace(id="user-1")), - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.VectorService") as vector_service, ): @@ -913,25 +915,25 @@ class TestSegmentServiceChildChunkTailHelpers: with pytest.raises(ChildChunkIndexingError, match="vector failed"): SegmentService.update_child_chunk( - "new content", child_chunk, SimpleNamespace(), SimpleNamespace(), dataset, mock_db.session + "new content", child_chunk, SimpleNamespace(), SimpleNamespace(), dataset, session ) - mock_db.session.rollback.assert_called_once() - mock_db.session.commit.assert_not_called() + session.rollback.assert_called_once() + session.commit.assert_not_called() def test_delete_child_chunk_commits_after_successful_vector_delete(self): + session = MagicMock() dataset = SimpleNamespace(id="dataset-1") child_chunk = _make_child_chunk() with ( - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, ): - SegmentService.delete_child_chunk(child_chunk, dataset, mock_db.session) + SegmentService.delete_child_chunk(child_chunk, dataset, session) - mock_db.session.delete.assert_called_once_with(child_chunk) - vector_service.delete_child_chunk_vector.assert_called_once_with(child_chunk, dataset) - mock_db.session.commit.assert_called_once() + session.delete.assert_called_once_with(child_chunk) + vector_service.delete_child_chunk_vector.assert_called_once_with(child_chunk, dataset, session=session) + session.commit.assert_called() class TestSegmentServiceAdditionalRegenerationBranches: @@ -947,6 +949,7 @@ class TestSegmentServiceAdditionalRegenerationBranches: yield account def test_update_segment_same_content_updates_answer_and_document_word_count_for_qa_segments(self, account_context): + session = MagicMock() segment = _make_segment(content="question", word_count=8) document = _make_document(doc_form=IndexStructureType.QA_INDEX, word_count=20) dataset = _make_dataset() @@ -954,18 +957,17 @@ class TestSegmentServiceAdditionalRegenerationBranches: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.VectorService") as vector_service, ): mock_redis.get.return_value = None - mock_db.session.get.return_value = refreshed_segment + session.get.return_value = refreshed_segment result = SegmentService.update_segment( SegmentUpdateArgs(content="question", answer="new answer"), segment, document, dataset, - mock_db.session, + session, ) assert result is refreshed_segment @@ -973,9 +975,10 @@ 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, mock_db.session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_content_change_uses_answer_when_counting_tokens_for_qa_segments(self, account_context): + session = MagicMock() segment = _make_segment(content="old", word_count=3) document = _make_document(doc_form=IndexStructureType.QA_INDEX, word_count=10) dataset = _make_dataset(indexing_technique="high_quality") @@ -985,7 +988,6 @@ class TestSegmentServiceAdditionalRegenerationBranches: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", return_value="hash-qa"), @@ -993,15 +995,15 @@ class TestSegmentServiceAdditionalRegenerationBranches: ): mock_redis.get.return_value = None model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model - mock_db.session.scalar.return_value = None - mock_db.session.get.return_value = refreshed_segment + session.scalar.return_value = None + session.get.return_value = refreshed_segment result = SegmentService.update_segment( SegmentUpdateArgs(content="new question", answer="new answer", keywords=["kw-1"]), segment, document, dataset, - mock_db.session, + session, ) assert result is refreshed_segment @@ -1009,12 +1011,13 @@ class TestSegmentServiceAdditionalRegenerationBranches: assert segment.answer == "new answer" 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, mock_db.session) + vector_service.update_segment_vector.assert_called_once_with(["kw-1"], segment, dataset, session=session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_content_change_parent_child_uses_default_embedding_and_ignores_summary_failures( self, account_context ): + session = MagicMock() segment = _make_segment(content="old", word_count=3) document = _make_document( doc_form=IndexStructureType.PARENT_CHILD_INDEX, @@ -1029,7 +1032,6 @@ class TestSegmentServiceAdditionalRegenerationBranches: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.ModelManager") as model_manager_cls, patch("services.dataset_service.VectorService") as vector_service, patch("services.dataset_service.helper.generate_text_hash", return_value="hash-parent"), @@ -1041,16 +1043,16 @@ class TestSegmentServiceAdditionalRegenerationBranches: update_summary.side_effect = RuntimeError("summary failed") # get calls: processing_rule, then refreshed_segment - mock_db.session.get.side_effect = [processing_rule, refreshed_segment] + session.get.side_effect = [processing_rule, refreshed_segment] # scalar call: existing_summary - mock_db.session.scalar.return_value = existing_summary + session.scalar.return_value = existing_summary result = SegmentService.update_segment( SegmentUpdateArgs(content="new parent content", regenerate_child_chunks=True, summary="new summary"), segment, document, dataset, - mock_db.session, + session, ) assert result is refreshed_segment @@ -1064,15 +1066,16 @@ class TestSegmentServiceAdditionalRegenerationBranches: dataset, embedding_model_instance, processing_rule, - mock_db.session, True, + session=session, ) - 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) + update_summary.assert_called_once_with(segment, dataset, "new summary", session=session) + vector_service.update_multimodel_vector.assert_called_once_with(segment, [], dataset, session=session) def test_update_segment_same_content_parent_child_marks_segment_error_for_non_high_quality_dataset( self, account_context ): + session = MagicMock() segment = _make_segment(content="same content", word_count=len("same content")) document = _make_document( doc_form=IndexStructureType.PARENT_CHILD_INDEX, @@ -1083,19 +1086,18 @@ class TestSegmentServiceAdditionalRegenerationBranches: with ( patch("services.dataset_service.redis_client") as mock_redis, - patch("services.dataset_service.db") as mock_db, patch("services.dataset_service.naive_utc_now", return_value="now"), patch("services.dataset_service.VectorService") as vector_service, ): mock_redis.get.return_value = None - mock_db.session.get.return_value = refreshed_segment + session.get.return_value = refreshed_segment result = SegmentService.update_segment( SegmentUpdateArgs(content="same content", regenerate_child_chunks=True), segment, document, dataset, - mock_db.session, + session, ) assert result is refreshed_segment 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 9dff74f8dd5..ffc52b5c369 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -162,7 +162,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=page, per_page=per_page, tenant_id=tenant_id + page=page, per_page=per_page, tenant_id=tenant_id, session=MagicMock() ) # Assert @@ -190,7 +190,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id=tenant_id, search=search + page=1, per_page=10, tenant_id=tenant_id, search=search, session=MagicMock() ) # Assert @@ -211,7 +211,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id="tenant-123" + page=1, per_page=10, tenant_id="tenant-123", session=MagicMock() ) # Assert @@ -233,7 +233,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id="tenant-123" + page=1, per_page=10, tenant_id="tenant-123", session=MagicMock() ) # Assert @@ -255,7 +255,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=10, per_page=10, tenant_id="tenant-123" + page=10, per_page=10, tenant_id="tenant-123", session=MagicMock() ) # Assert @@ -280,7 +280,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id="tenant-123", search="PRODUCTION" + page=1, per_page=10, tenant_id="tenant-123", search="PRODUCTION", session=MagicMock() ) # Assert @@ -302,7 +302,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id="tenant-123", search="v2.0" + page=1, per_page=10, tenant_id="tenant-123", search="v2.0", session=MagicMock() ) # Assert @@ -323,7 +323,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=100, tenant_id="tenant-123" + page=1, per_page=100, tenant_id="tenant-123", session=MagicMock() ) # Assert @@ -348,7 +348,7 @@ class TestExternalDatasetServiceGetAPIs: # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( - page=1, per_page=10, tenant_id="tenant-123" + page=1, per_page=10, tenant_id="tenant-123", session=MagicMock() ) # Assert @@ -463,7 +463,8 @@ class TestExternalDatasetServiceCreateAPI: assert result.updated_by == user_id mock_check.assert_called_once_with(args["settings"]) mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + mock_db.session.flush.assert_called_once() + mock_db.session.commit.assert_not_called() @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") @@ -901,7 +902,8 @@ class TestExternalDatasetServiceUpdateAPI: assert result.description == "Updated description" assert result.updated_by == user_id assert result.updated_at == current_time - mock_db.session.commit.assert_called_once() + mock_db.session.flush.assert_called_once() + mock_db.session.commit.assert_not_called() @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_preserve_hidden_api_key( @@ -1005,7 +1007,8 @@ class TestExternalDatasetServiceDeleteAPI: # Assert mock_db.session.delete.assert_called_once_with(existing_api) - mock_db.session.commit.assert_called_once() + mock_db.session.flush.assert_called_once() + mock_db.session.commit.assert_not_called() @patch("services.external_knowledge_service.db") def test_delete_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): @@ -1546,6 +1549,7 @@ class TestExternalDatasetServiceCreateDataset: assert result.provider == "external" assert result.created_by == user_id mock_db.session.add.assert_called() + mock_db.session.flush.assert_called_once() mock_db.session.commit.assert_called_once() @patch("services.external_knowledge_service.db") diff --git a/api/tests/unit_tests/services/test_message_service.py b/api/tests/unit_tests/services/test_message_service.py index 13f340e9f4a..dfdaf7b40e2 100644 --- a/api/tests/unit_tests/services/test_message_service.py +++ b/api/tests/unit_tests/services/test_message_service.py @@ -70,6 +70,8 @@ class TestMessageServiceFactory: message.query = query message.answer = answer message.created_at = created_at or datetime.now() + message.user_feedback_with_session.return_value = None + message.admin_feedback_with_session.return_value = None return message @@ -754,6 +756,7 @@ class TestMessageServiceFeedback: user = factory.create_end_user_mock() message = factory.create_message_mock() message.user_feedback = None + message.user_feedback_with_session.return_value = None mock_get_message.return_value = message # Act @@ -787,6 +790,7 @@ class TestMessageServiceFeedback: message = factory.create_message_mock() feedback = MagicMock(spec=MessageFeedback) message.admin_feedback = feedback + message.admin_feedback_with_session.return_value = feedback mock_get_message.return_value = message # Act @@ -816,6 +820,7 @@ class TestMessageServiceFeedback: message = factory.create_message_mock() feedback = MagicMock() message.user_feedback = feedback + message.user_feedback_with_session.return_value = feedback mock_get_message.return_value = message # Act @@ -1197,6 +1202,7 @@ class TestMessageServiceSuggestedQuestions: "model_id": None, "provider": None, } + conversation.model_config_with_session.return_value = conversation.model_config mock_conversation_service.get_conversation.return_value = conversation mock_memory.return_value.get_history_prompt_text.return_value = "histories" diff --git a/api/tests/unit_tests/services/test_metadata_service_session_boundary.py b/api/tests/unit_tests/services/test_metadata_service_session_boundary.py new file mode 100644 index 00000000000..7832cdf435b --- /dev/null +++ b/api/tests/unit_tests/services/test_metadata_service_session_boundary.py @@ -0,0 +1,107 @@ +from datetime import datetime +from unittest.mock import MagicMock, patch + +from core.rag.index_processor.constant.built_in_field import BuiltInField +from models import Account +from services.dataset_service import DocumentService +from services.entities.knowledge_entities.knowledge_entities import ( + DocumentMetadataOperation, + MetadataArgs, + MetadataOperationData, +) +from services.metadata_service import MetadataService + + +def _account() -> Account: + account = Account(name="User", email="user@example.com") + account.id = "account-1" + return account + + +def test_create_metadata_flushes_without_committing_caller_session() -> None: + session = MagicMock() + session.scalar.return_value = None + + metadata = MetadataService.create_metadata( + "dataset-1", + MetadataArgs(type="string", name="author"), + _account(), + "tenant-1", + session=session, + ) + + assert metadata.name == "author" + session.flush.assert_called_once_with() + session.commit.assert_not_called() + session.rollback.assert_not_called() + + +def _document() -> MagicMock: + document = MagicMock() + document.id = "document-1" + document.name = "Document" + document.doc_metadata = {} + document.data_source_type = "upload_file" + document.upload_date = datetime(2026, 1, 1) + document.last_update_date = datetime(2026, 1, 2) + document.uploader = "global-session uploader" + document.get_uploader.return_value = "caller-session uploader" + return document + + +def test_enable_built_in_field_uses_caller_session_for_uploader() -> None: + session = MagicMock() + dataset = MagicMock(id="dataset-1", built_in_field_enabled=False) + document = _document() + + with ( + patch.object(MetadataService, "knowledge_base_metadata_lock_check"), + patch.object(DocumentService, "get_working_documents_by_dataset_id", return_value=[document]), + patch("services.metadata_service.redis_client.delete"), + ): + MetadataService.enable_built_in_field(dataset, session) + + assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader" + document.get_uploader.assert_called_once_with(session=session) + + +def test_update_documents_metadata_uses_caller_session_for_uploader() -> None: + session = MagicMock() + dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=True) + document = _document() + metadata_args = MetadataOperationData( + operation_data=[ + DocumentMetadataOperation(document_id=document.id, metadata_list=[], partial_update=False), + ] + ) + + with ( + patch.object(MetadataService, "knowledge_base_metadata_lock_check"), + patch.object(DocumentService, "get_document", return_value=document), + patch("services.metadata_service.redis_client.delete"), + ): + MetadataService.update_documents_metadata( + dataset, + metadata_args, + _account(), + "tenant-1", + session=session, + ) + + assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader" + document.get_uploader.assert_called_once_with(session=session) + + +def test_get_dataset_metadatas_uses_caller_session() -> None: + session = MagicMock() + session.scalar.return_value = 2 + dataset = MagicMock(id="dataset-1", built_in_field_enabled=False) + dataset.get_doc_metadata.return_value = [{"id": "metadata-1", "name": "author", "type": "string"}] + + result = MetadataService.get_dataset_metadatas(dataset, session) + + assert result == { + "doc_metadata": [{"id": "metadata-1", "name": "author", "type": "string", "count": 2}], + "built_in_field_enabled": False, + } + dataset.get_doc_metadata.assert_called_once_with(session=session) diff --git a/api/tests/unit_tests/services/test_snippet_generate_service.py b/api/tests/unit_tests/services/test_snippet_generate_service.py index 098787a6d69..63ccbdb351e 100644 --- a/api/tests/unit_tests/services/test_snippet_generate_service.py +++ b/api/tests/unit_tests/services/test_snippet_generate_service.py @@ -1,4 +1,5 @@ import json +from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import Mock @@ -26,6 +27,10 @@ def _workflow(graph: dict) -> Workflow: ) +def _session_maker(session: object | None = None) -> Mock: + return Mock(return_value=nullcontext(session or Mock())) + + def test_filter_virtual_start_events_keeps_blocking_response_unchanged(): response = {"data": {"outputs": {"text": "ok"}}} @@ -287,11 +292,13 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch): ) monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class) + session = Mock() result = SnippetGenerateService.generate_single_iteration( snippet=snippet, user=user, node_id="iteration-1", args={"inputs": {"items": [1]}}, + session_maker=_session_maker(session), ) assert list(result) == ["event"] @@ -302,6 +309,7 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch): assert kwargs["node_id"] == "iteration-1" assert kwargs["user"] is user assert kwargs["streaming"] is True + assert kwargs["session"] is session workflow_generator_class.convert_to_event_stream.assert_called_once_with(response) @@ -317,6 +325,7 @@ def test_generate_single_iteration_raises_when_draft_workflow_missing(monkeypatc user=SimpleNamespace(id="user-1"), node_id="iteration-1", args={"inputs": {}}, + session_maker=_session_maker(), ) @@ -335,11 +344,13 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch): ) monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class) + session = Mock() result = SnippetGenerateService.generate_single_loop( snippet=snippet, user=user, node_id="loop-1", args=SimpleNamespace(inputs={"items": [1]}), + session_maker=_session_maker(session), ) assert list(result) == ["event"] @@ -350,6 +361,7 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch): assert kwargs["node_id"] == "loop-1" assert kwargs["user"] is user assert kwargs["streaming"] is True + assert kwargs["session"] is session workflow_generator_class.convert_to_event_stream.assert_called_once_with(response) @@ -365,6 +377,7 @@ def test_generate_single_loop_raises_when_draft_workflow_missing(monkeypatch): user=SimpleNamespace(id="user-1"), node_id="loop-1", args=SimpleNamespace(inputs={}), + session_maker=_session_maker(), ) 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 d9482fdbe42..4af3b4cdad0 100644 --- a/api/tests/unit_tests/services/test_summary_index_service.py +++ b/api/tests/unit_tests/services/test_summary_index_service.py @@ -11,6 +11,7 @@ from unittest.mock import MagicMock import pytest from sqlalchemy import create_engine, select +from sqlalchemy.exc import SAWarning from sqlalchemy.orm import sessionmaker import services.summary_index_service as summary_module @@ -55,8 +56,10 @@ def _segment(*, has_document: bool = True) -> MagicMock: doc.doc_language = "en" doc.doc_form = IndexStructureType.PARAGRAPH_INDEX segment.document = doc + segment.get_document.return_value = doc else: segment.document = None + segment.get_document.return_value = None return segment @@ -97,14 +100,17 @@ def test_generate_summary_for_segment_passes_document_language(monkeypatch: pyte segment = _segment(has_document=True) dataset = _dataset() + session = MagicMock() - content, got_usage = SummaryIndexService.generate_summary_for_segment(segment, dataset, {"a": 1}) + content, got_usage = SummaryIndexService.generate_summary_for_segment(segment, dataset, {"a": 1}, session=session) assert content == "sum" assert got_usage is usage paragraph_module.ParagraphIndexProcessor.generate_summary.assert_called_once() _, kwargs = paragraph_module.ParagraphIndexProcessor.generate_summary.call_args assert kwargs["document_language"] == "en" + assert kwargs["session"] is session + segment.get_document.assert_called_once_with(session=session) def test_generate_summary_for_segment_raises_when_empty(monkeypatch: pytest.MonkeyPatch) -> None: @@ -118,7 +124,7 @@ def test_generate_summary_for_segment_raises_when_empty(monkeypatch: pytest.Monk ) with pytest.raises(ValueError, match="Generated summary is empty"): - SummaryIndexService.generate_summary_for_segment(_segment(), _dataset(), {"a": 1}) + SummaryIndexService.generate_summary_for_segment(_segment(), _dataset(), {"a": 1}, session=MagicMock()) def test_create_summary_record_updates_existing_and_reenables() -> None: @@ -237,8 +243,9 @@ def test_vectorize_summary_without_session_creates_record_when_missing(monkeypat SummaryIndexService.vectorize_summary(summary, segment, dataset, session=None) - # One context for success path, no error handler session. + # Vector initialization and the record update both obtain local sessions. create_session_mock.assert_called() + assert all(call.kwargs["session"] is session for call in vector_cls.call_args_list) session.add.assert_called() session.commit.assert_called_once() assert summary.status == SummaryStatus.COMPLETED @@ -331,16 +338,20 @@ def test_generate_and_vectorize_summary_success(monkeypatch: pytest.MonkeyPatch) session = MagicMock() session.scalar.return_value = record + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") - monkeypatch.setattr( - SummaryIndexService, "generate_summary_for_segment", MagicMock(return_value=("sum", MagicMock(total_tokens=0))) + generate_summary = MagicMock( + side_effect=lambda *_args, **_kwargs: phase_events.append("generate") or ("sum", MagicMock(total_tokens=0)) ) - monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(return_value=None)) + monkeypatch.setattr(SummaryIndexService, "generate_summary_for_segment", generate_summary) + vectorize_summary = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("vectorize")) + monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_summary) 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() + assert phase_events == ["commit", "generate", "vectorize", "commit"] def test_generate_and_vectorize_summary_vectorize_failure_sets_error(monkeypatch: pytest.MonkeyPatch) -> None: @@ -361,6 +372,49 @@ def test_generate_and_vectorize_summary_vectorize_failure_sets_error(monkeypatch assert record.status == SummaryStatus.ERROR # Outer exception handler overwrites the error with the raw exception message. assert record.error == "boom" + session.rollback.assert_called_once() + + +def test_generate_and_vectorize_summary_rolls_back_failed_transaction_before_recording_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + dataset = _dataset() + segment = _segment() + record = _summary_record(summary_content="") + session = MagicMock() + rolled_back = False + scalar_calls = 0 + + def scalar(*_args, **_kwargs): + nonlocal scalar_calls + scalar_calls += 1 + if scalar_calls > 1 and not rolled_back: + raise PendingRollbackError("rollback required") + return record + + def rollback() -> None: + nonlocal rolled_back + rolled_back = True + + session.scalar.side_effect = scalar + session.rollback.side_effect = rollback + session.flush.side_effect = RuntimeError("flush failed") + monkeypatch.setattr( + SummaryIndexService, + "generate_summary_for_segment", + MagicMock(return_value=("sum", MagicMock(total_tokens=0))), + ) + vectorize_summary = MagicMock() + monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_summary) + + with pytest.raises(RuntimeError, match="flush failed"): + SummaryIndexService.generate_and_vectorize_summary(segment, dataset, {"enable": True}, session=session) + + assert rolled_back is True + assert record.status == SummaryStatus.ERROR + assert record.error == "flush failed" + vectorize_summary.assert_not_called() + assert session.commit.call_count == 2 def test_vectorize_summary_updates_existing_record_found_by_chunk_id(monkeypatch: pytest.MonkeyPatch) -> None: @@ -451,7 +505,10 @@ def test_vectorize_summary_session_enter_returns_none_triggers_runtime_error(mon error_session = MagicMock() error_session.scalar.return_value = summary - create_session_mock = MagicMock(side_effect=[_BadContext(), _SessionContext(error_session)]) + vector_session = MagicMock() + create_session_mock = MagicMock( + side_effect=[_SessionContext(vector_session), _BadContext(), _SessionContext(error_session)] + ) monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) with pytest.raises(RuntimeError, match="Session should not be None"): @@ -480,7 +537,10 @@ def test_vectorize_summary_created_record_becomes_none_triggers_guard(monkeypatc error_session = MagicMock() error_session.scalar.return_value = summary - create_session_mock = MagicMock(side_effect=[_SessionContext(session), _SessionContext(error_session)]) + vector_session = MagicMock() + create_session_mock = MagicMock( + side_effect=[_SessionContext(vector_session), _SessionContext(session), _SessionContext(error_session)] + ) monkeypatch.setattr(summary_module, "session_factory", SimpleNamespace(create_session=create_session_mock)) # Force the created record to be None so the "should not be None" guard triggers. @@ -808,29 +868,19 @@ def test_delete_summaries_for_segments_deletes_vectors_and_records(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))), - ) - SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids=[summary.chunk_id]) + SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids=[summary.chunk_id], session=session) vector_instance.delete_by_ids.assert_called_once_with(["n1"]) session.delete.assert_called_once_with(summary) - session.commit.assert_called_once() + session.flush.assert_called_once() def test_delete_summaries_for_segments_no_summaries_noop(monkeypatch: pytest.MonkeyPatch) -> None: dataset = _dataset() session = MagicMock() session.scalars.return_value.all.return_value = [] - monkeypatch.setattr( - summary_module, - "session_factory", - SimpleNamespace(create_session=MagicMock(return_value=_SessionContext(session))), - ) - SummaryIndexService.delete_summaries_for_segments(dataset) - session.commit.assert_not_called() + SummaryIndexService.delete_summaries_for_segments(dataset, session=session) + session.flush.assert_not_called() def test_update_summary_for_segment_skip_conditions() -> None: @@ -838,7 +888,7 @@ def test_update_summary_for_segment_skip_conditions() -> None: economy_dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY) 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 + seg.get_document.return_value.doc_form = IndexStructureType.QA_INDEX assert SummaryIndexService.update_summary_for_segment(seg, _dataset(), "x", session=session) is None @@ -933,14 +983,62 @@ def test_update_summary_for_segment_existing_vectorize_failure_returns_error_rec segment = _segment() record = _summary_record(summary_content="old", node_id="n1") - session = MagicMock() + session = MagicMock(is_active=True) session.scalar.return_value = record monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) + assert out is record + assert out.summary_content == "new" assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") + session.rollback.assert_not_called() + session.add.assert_called_with(record) + session.commit.assert_called_once() + + +def test_update_summary_for_segment_failed_flush_persists_new_content(monkeypatch: pytest.MonkeyPatch) -> None: + engine = create_engine("sqlite+pysqlite:///:memory:") + DocumentSegmentSummary.__table__.create(engine) + session_maker = sessionmaker(bind=engine, expire_on_commit=False) + + with session_maker() as session: + record = DocumentSegmentSummary( + dataset_id="dataset-1", + document_id="doc-1", + chunk_id="seg-1", + summary_content="old", + status=SummaryStatus.COMPLETED, + ) + record.id = "sum-1" + session.add(record) + session.commit() + + def fail_flush(*_args, **_kwargs) -> None: + duplicate = DocumentSegmentSummary( + dataset_id="dataset-1", + document_id="doc-1", + chunk_id="seg-2", + summary_content="duplicate", + ) + duplicate.id = record.id + session.add(duplicate) + session.flush() + + segment = _segment() + segment.get_document.return_value = SimpleNamespace(doc_form=IndexStructureType.PARAGRAPH_INDEX) + monkeypatch.setattr(SummaryIndexService, "vectorize_summary", fail_flush) + + with pytest.warns(SAWarning, match="conflicts with persistent instance"): + out = SummaryIndexService.update_summary_for_segment(segment, _dataset(), "new", session=session) + session.expire_all() + persisted = session.get(DocumentSegmentSummary, record.id) + + assert out is record + assert persisted is not None + assert persisted.summary_content == "new" + assert persisted.status == SummaryStatus.ERROR def test_update_summary_for_segment_new_record_success(monkeypatch: pytest.MonkeyPatch) -> None: @@ -971,6 +1069,7 @@ def test_update_summary_for_segment_outer_exception_sets_error_and_reraises(monk SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert record.status == SummaryStatus.ERROR assert record.error == "flush boom" + session.rollback.assert_called_once() session.commit.assert_called() @@ -1030,18 +1129,19 @@ def test_update_summary_for_segment_creates_new_and_vectorize_fails_returns_erro dataset = _dataset() segment = _segment() - session = MagicMock() + session = MagicMock(is_active=False) session.scalar.return_value = None created = _summary_record(summary_content="new", node_id=None) monkeypatch.setattr(SummaryIndexService, "create_summary_record", MagicMock(return_value=created)) - - vectorize_mock = MagicMock(side_effect=RuntimeError("boom")) - monkeypatch.setattr(SummaryIndexService, "vectorize_summary", vectorize_mock) + monkeypatch.setattr(SummaryIndexService, "vectorize_summary", MagicMock(side_effect=RuntimeError("boom"))) out = SummaryIndexService.update_summary_for_segment(segment, dataset, "new", session=session) assert out.status == SummaryStatus.ERROR assert "Vectorization failed" in (out.error or "") + session.rollback.assert_called_once() + session.add.assert_called_with(created) + session.commit.assert_called_once() def test_get_segments_summaries_empty_list() -> None: diff --git a/api/tests/unit_tests/services/test_vector_service.py b/api/tests/unit_tests/services/test_vector_service.py index 3659b85228b..eb7bd57e720 100644 --- a/api/tests/unit_tests/services/test_vector_service.py +++ b/api/tests/unit_tests/services/test_vector_service.py @@ -49,6 +49,7 @@ def _make_dataset( dataset.is_multimodal = is_multimodal dataset.embedding_model_provider = embedding_model_provider dataset.embedding_model = embedding_model + dataset.get_doc_form.return_value = doc_form return dataset @@ -72,6 +73,7 @@ def _make_segment( segment.index_node_id = index_node_id segment.index_node_hash = index_node_hash segment.attachments = attachments or [] + segment.get_attachments.return_value = attachments or [] return segment @@ -98,7 +100,9 @@ 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, MagicMock()) + VectorService.create_segments_vector( + [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock() + ) index_processor.load.assert_called_once() args, kwargs = index_processor.load.call_args @@ -123,7 +127,10 @@ 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, MagicMock()) + session = MagicMock() + VectorService.create_segments_vector( + [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=session + ) assert index_processor.load.call_count == 2 first_args, first_kwargs = index_processor.load.call_args_list[0] @@ -136,6 +143,7 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk assert second_args[1] == [] assert len(second_args[2]) == 2 assert second_kwargs["with_keywords"] is False + segment.get_attachments.assert_called_once_with(session=session) def test_create_segments_vector_with_no_segments_does_not_load(monkeypatch: pytest.MonkeyPatch) -> None: @@ -145,7 +153,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, MagicMock()) + VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock()) index_processor.load.assert_not_called() @@ -207,22 +215,12 @@ 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, - db_mock.session, + None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=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, - db_mock.session, - False, + segment, dataset_document, dataset, embedding_model_instance, processing_rule, False, session=db_mock.session ) index_processor.load.assert_not_called() @@ -263,11 +261,7 @@ 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, - db_mock.session, + None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=db_mock.session ) model_manager_instance.get_default_model_instance.assert_called_once() @@ -295,7 +289,7 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - db_mock.session, + session=db_mock.session, ) assert any(r.levelno >= logging.WARNING for r in caplog.records) index_processor.load.assert_not_called() @@ -315,7 +309,7 @@ def test_create_segments_vector_parent_child_missing_processing_rule_raises(monk [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - db_mock.session, + session=db_mock.session, ) @@ -336,7 +330,7 @@ def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - db_mock.session, + session=db_mock.session, ) @@ -345,10 +339,13 @@ def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.Monk segment = _make_segment() vector_instance = MagicMock() - monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) + vector_cls = MagicMock(return_value=vector_instance) + monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.update_segment_vector(["k"], segment, dataset) + session = MagicMock() + VectorService.update_segment_vector(["k"], segment, dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=session) vector_instance.delete_by_ids.assert_called_once_with([segment.index_node_id]) vector_instance.add_texts.assert_called_once() add_args, add_kwargs = vector_instance.add_texts.call_args @@ -363,12 +360,14 @@ def test_update_segment_vector_economy_uses_keyword_with_keywords_list(monkeypat keyword_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance)) - VectorService.update_segment_vector(["a", "b"], segment, dataset) + session = MagicMock() + VectorService.update_segment_vector(["a", "b"], segment, dataset, session=session) - keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id]) + keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id], session) keyword_instance.add_texts.assert_called_once() args, kwargs = keyword_instance.add_texts.call_args assert len(args[0]) == 1 + assert args[1] is session assert kwargs["keywords_list"] == [["a", "b"]] @@ -379,9 +378,12 @@ def test_update_segment_vector_economy_uses_keyword_without_keywords_list(monkey keyword_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance)) - VectorService.update_segment_vector(None, segment, dataset) + session = MagicMock() + VectorService.update_segment_vector(None, segment, dataset, session=session) keyword_instance.add_texts.assert_called_once() - _, kwargs = keyword_instance.add_texts.call_args + args, kwargs = keyword_instance.add_texts.call_args + assert len(args[0]) == 1 + assert args[1] is session assert "keywords_list" not in kwargs @@ -410,7 +412,9 @@ 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) - session = MagicMock() + db_mock = MagicMock() + db_mock.session.add = MagicMock() + db_mock.session.flush = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -418,16 +422,16 @@ 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, + session=db_mock.session, ) index_processor.clean.assert_called_once() _, 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 session.add.call_count == 2 - session.commit.assert_called_once() + assert db_mock.session.add.call_count == 2 + db_mock.session.flush.assert_called_once() def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest.MonkeyPatch) -> None: @@ -446,7 +450,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)) - session = MagicMock() + db_mock = MagicMock() VectorService.generate_child_chunks( segment=segment, @@ -454,13 +458,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, + session=db_mock.session, ) index_processor.load.assert_not_called() - session.add.assert_not_called() - session.commit.assert_called_once() + db_mock.session.add.assert_not_called() + db_mock.session.flush.assert_called_once() def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.MonkeyPatch) -> None: @@ -473,9 +477,12 @@ def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.M child_chunk.dataset_id = "dataset-1" vector_instance = MagicMock() - monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) + vector_cls = MagicMock(return_value=vector_instance) + monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.create_child_chunk_vector(child_chunk, dataset) + session = MagicMock() + VectorService.create_child_chunk_vector(child_chunk, dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=session) vector_instance.add_texts.assert_called_once() @@ -491,7 +498,7 @@ def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) child_chunk.document_id = "doc-1" child_chunk.dataset_id = "dataset-1" - VectorService.create_child_chunk_vector(child_chunk, dataset) + VectorService.create_child_chunk_vector(child_chunk, dataset, session=MagicMock()) vector_cls.assert_not_called() @@ -516,10 +523,13 @@ def test_update_child_chunk_vector_high_quality_updates_vector(monkeypatch: pyte del_chunk.index_node_id = "did" vector_instance = MagicMock() - monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) + vector_cls = MagicMock(return_value=vector_instance) + monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset) + session = MagicMock() + VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=session) vector_instance.delete_by_ids.assert_called_once_with(["uid", "did"]) vector_instance.add_texts.assert_called_once() docs = vector_instance.add_texts.call_args.args[0] @@ -530,7 +540,7 @@ def test_update_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY) vector_cls = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.update_child_chunk_vector([], [], [], dataset) + VectorService.update_child_chunk_vector([], [], [], dataset, session=MagicMock()) vector_cls.assert_not_called() @@ -540,9 +550,12 @@ def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch child_chunk.index_node_id = "cid" vector_instance = MagicMock() - monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) + vector_cls = MagicMock(return_value=vector_instance) + monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.delete_child_chunk_vector(child_chunk, dataset) + session = MagicMock() + VectorService.delete_child_chunk_vector(child_chunk, dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=session) vector_instance.delete_by_ids.assert_called_once_with(["cid"]) @@ -560,7 +573,7 @@ def test_update_multimodel_vector_returns_when_not_high_quality(monkeypatch: pyt monkeypatch.setattr(vector_service_module, "Vector", vector_cls) VectorService.update_multimodel_vector( - segment=segment, attachment_ids=["a"], dataset=dataset, session=db_mock.session + session=db_mock.session, segment=segment, attachment_ids=["a"], dataset=dataset ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -575,7 +588,7 @@ def test_update_multimodel_vector_returns_when_no_actual_change(monkeypatch: pyt monkeypatch.setattr(vector_service_module, "Vector", vector_cls) VectorService.update_multimodel_vector( - segment=segment, attachment_ids=["b", "a"], dataset=dataset, session=db_mock.session + session=db_mock.session, segment=segment, attachment_ids=["b", "a"], dataset=dataset ) vector_cls.assert_not_called() db_mock.session.query.assert_not_called() @@ -595,10 +608,10 @@ def test_update_multimodel_vector_deletes_bindings_and_commits_on_empty_new_ids( VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=db_mock.session) - vector_cls.assert_called_once_with(dataset=dataset) + vector_cls.assert_called_once_with(dataset=dataset, session=db_mock.session) vector_instance.delete_by_ids.assert_called_once_with(["old-1", "old-2"]) db_mock.session.execute.assert_called_once() - db_mock.session.commit.assert_called_once() + db_mock.session.flush.assert_called_once() db_mock.session.add_all.assert_not_called() vector_instance.add_texts.assert_not_called() @@ -612,10 +625,10 @@ def test_update_multimodel_vector_commits_when_no_upload_files_found(monkeypatch db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) VectorService.update_multimodel_vector( - segment=segment, attachment_ids=["new-1"], dataset=dataset, session=db_mock.session + session=db_mock.session, segment=segment, attachment_ids=["new-1"], dataset=dataset ) - db_mock.session.commit.assert_called_once() + db_mock.session.flush.assert_called_once() db_mock.session.add_all.assert_not_called() vector_instance.add_texts.assert_not_called() @@ -638,10 +651,7 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up with caplog.at_level(logging.WARNING, logger="services.vector_service"): VectorService.update_multimodel_vector( - segment=segment, - attachment_ids=["file-1", "missing"], - dataset=dataset, - session=db_mock.session, + session=db_mock.session, segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset ) assert any(r.levelno >= logging.WARNING for r in caplog.records) db_mock.session.add_all.assert_called_once() @@ -654,7 +664,7 @@ def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_up assert len(documents) == 1 assert documents[0].page_content == "img.png" assert documents[0].metadata["doc_id"] == "file-1" - db_mock.session.commit.assert_called_once() + db_mock.session.flush.assert_called_once() def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops( @@ -673,13 +683,13 @@ def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops monkeypatch.setattr(vector_service_module, "select", MagicMock()) VectorService.update_multimodel_vector( - segment=segment, attachment_ids=["file-1"], dataset=dataset, session=db_mock.session + session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset ) vector_instance.delete_by_ids.assert_not_called() vector_instance.add_texts.assert_not_called() db_mock.session.add_all.assert_called_once() - db_mock.session.commit.assert_called_once() + db_mock.session.flush.assert_called_once() def test_update_multimodel_vector_rolls_back_and_reraises_on_error( @@ -692,7 +702,7 @@ def test_update_multimodel_vector_rolls_back_and_reraises_on_error( 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")]) - db_mock.session.commit.side_effect = RuntimeError("boom") + db_mock.session.flush.side_effect = RuntimeError("boom") monkeypatch.setattr( vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) ) @@ -702,7 +712,7 @@ 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, session=db_mock.session + session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset ) assert any(r.levelno >= logging.ERROR for r in caplog.records) diff --git a/api/tests/unit_tests/services/test_workflow_run_service.py b/api/tests/unit_tests/services/test_workflow_run_service.py index c21b764a9cd..fcfa9992cd1 100644 --- a/api/tests/unit_tests/services/test_workflow_run_service.py +++ b/api/tests/unit_tests/services/test_workflow_run_service.py @@ -34,11 +34,13 @@ def _end_user(**kwargs: Any) -> EndUser: return cast(EndUser, SimpleNamespace(**kwargs)) -def _fake_session_returning_messages(messages: list[Any]) -> SimpleNamespace: - """A stand-in db session whose scalars(...).all() returns the given messages.""" - scalars_result = MagicMock() - scalars_result.all.return_value = messages - return SimpleNamespace(scalars=MagicMock(return_value=scalars_result)) +def _fake_session_factory_returning_messages(messages: list[Any]) -> tuple[MagicMock, MagicMock]: + """Build a session factory whose session returns the given messages.""" + session = MagicMock() + session.scalars.return_value.all.return_value = messages + session_factory = MagicMock() + session_factory.return_value.__enter__.return_value = session + return session_factory, session class TestWorkflowRunServiceInitialization: @@ -125,17 +127,15 @@ class TestWorkflowRunServiceQueries: repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, ) -> None: - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + message = SimpleNamespace(id="msg-1", conversation_id="conv-1", workflow_run_id="run-1") + session_factory, session = _fake_session_factory_returning_messages([message]) + service = WorkflowRunService(session_factory=session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") run_with_message = SimpleNamespace(id="run-1", status="running") run_without_message = SimpleNamespace(id="run-2", status="succeeded") pagination = SimpleNamespace(data=[run_with_message, run_without_message]) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) - message = SimpleNamespace(id="msg-1", conversation_id="conv-1", workflow_run_id="run-1") - fake_session = _fake_session_returning_messages([message]) - monkeypatch.setattr(service_module, "db", SimpleNamespace(session=fake_session)) - result = service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={"limit": "2"}) assert result is pagination @@ -146,7 +146,8 @@ class TestWorkflowRunServiceQueries: assert not hasattr(result.data[1], "message_id") assert result.data[1].id == "run-2" # Messages are batch-loaded in a single query, not one per run. - fake_session.scalars.assert_called_once() + session_factory.assert_called_once_with() + session.scalars.assert_called_once() def test_get_paginate_advanced_chat_workflow_runs_batch_loads_messages_without_n_plus_one( self, @@ -158,19 +159,18 @@ class TestWorkflowRunServiceQueries: Previously the deprecated WorkflowRun.message property issued one query per run (N+1); they are now batch-loaded in a single query. """ - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + session_factory, session = _fake_session_factory_returning_messages([]) + service = WorkflowRunService(session_factory=session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") runs = [SimpleNamespace(id=f"run-{i}", status="succeeded") for i in range(5)] pagination = SimpleNamespace(data=runs) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) - fake_session = _fake_session_returning_messages([]) - monkeypatch.setattr(service_module, "db", SimpleNamespace(session=fake_session)) - service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={}) # Exactly one message query for the whole page, independent of run count. - assert fake_session.scalars.call_count == 1 + session_factory.assert_called_once_with() + assert session.scalars.call_count == 1 def test_get_workflow_run_should_delegate_to_repository_by_tenant_and_app( self, 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 f471e4aeb56..60c8b5213f6 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 @@ -3,7 +3,7 @@ from __future__ import annotations import json from types import SimpleNamespace from typing import Any, cast -from unittest.mock import MagicMock +from unittest.mock import MagicMock, call import pytest @@ -356,7 +356,8 @@ def test__convert_to_answer_node() -> None: def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(converter: WorkflowConverter) -> None: - app_model = _app_model(app_model_config=None) + app_model = _app_model(app_model_config_id=None) + session = MagicMock() with pytest.raises(ValueError, match="App model config is required"): converter.convert_to_workflow( @@ -366,9 +367,11 @@ def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(conve icon_type="emoji", icon="robot", icon_background="#fff", - session=MagicMock(), + session=session, ) + session.get.assert_not_called() + @pytest.mark.parametrize( ("source_mode", "expected_mode"), @@ -391,9 +394,16 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( monkeypatch.setattr(converter, "convert_app_model_config_to_workflow", MagicMock(return_value=workflow)) monkeypatch.setattr(converter_module, "App", FakeApp) - db_session = SimpleNamespace(add=MagicMock(), flush=MagicMock(), commit=MagicMock()) + app_model_config = _app_model_config(id="config-1") + phase_events: list[str] = [] + db_session = SimpleNamespace( + add=MagicMock(), + flush=MagicMock(), + commit=MagicMock(side_effect=lambda: phase_events.append("commit")), + get=MagicMock(return_value=app_model_config), + ) - send_mock = MagicMock() + send_mock = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal")) monkeypatch.setattr(converter_module.app_was_created, "send", send_mock) account = _account(id="account-1") @@ -409,7 +419,7 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( api_rpm=10, api_rph=100, is_public=False, - app_model_config=_app_model_config(id="config-1"), + app_model_config_id="config-1", ) new_app = converter.convert_to_workflow( @@ -431,8 +441,10 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( assert workflow.app_id == "new-app-id" db_session.add.assert_called_once() db_session.flush.assert_called_once() - db_session.commit.assert_called_once() - send_mock.assert_called_once_with(new_app, account=account) + assert phase_events == ["commit", "signal", "commit"] + assert db_session.commit.call_count == 2 + db_session.get.assert_called_once_with(AppModelConfig, "config-1") + send_mock.assert_called_once_with(new_app, account=account, session=db_session) def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_and_features( @@ -594,44 +606,68 @@ def test_convert_to_app_config_should_route_to_correct_manager( agent_result = SimpleNamespace(kind="agent") chat_result = SimpleNamespace(kind="chat") completion_result = SimpleNamespace(kind="completion") - monkeypatch.setattr( - converter_module.AgentChatAppConfigManager, "get_app_config", MagicMock(return_value=agent_result) - ) - monkeypatch.setattr(converter_module.ChatAppConfigManager, "get_app_config", MagicMock(return_value=chat_result)) - monkeypatch.setattr( - converter_module.CompletionAppConfigManager, - "get_app_config", - MagicMock(return_value=completion_result), - ) + agent_get_app_config = MagicMock(return_value=agent_result) + chat_get_app_config = MagicMock(return_value=chat_result) + completion_get_app_config = MagicMock(return_value=completion_result) + load_annotation_reply = MagicMock(return_value={"enabled": False}) + monkeypatch.setattr(converter_module.AgentChatAppConfigManager, "get_app_config", agent_get_app_config) + monkeypatch.setattr(converter_module.ChatAppConfigManager, "get_app_config", chat_get_app_config) + monkeypatch.setattr(converter_module.CompletionAppConfigManager, "get_app_config", completion_get_app_config) + monkeypatch.setattr(converter_module, "load_annotation_reply_config", load_annotation_reply) + session = MagicMock() + agent_mode_app = _app_model(mode=AppMode.AGENT_CHAT, is_agent_with_session=MagicMock(return_value=False)) + agent_flag_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=True)) + chat_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=False)) + completion_app = _app_model(mode=AppMode.COMPLETION, is_agent_with_session=MagicMock(return_value=False)) + agent_mode_config = _app_model_config(id="cfg-1", app_id="app-1") + agent_flag_config = _app_model_config(id="cfg-2", app_id="app-2") + chat_config = _app_model_config(id="cfg-3", app_id="app-3") + completion_config = _app_model_config(id="cfg-4", app_id="app-4") from_agent_mode = converter._convert_to_app_config( - app_model=_app_model(mode=AppMode.AGENT_CHAT, is_agent=False), - app_model_config=_app_model_config(id="cfg-1"), + app_model=agent_mode_app, + app_model_config=agent_mode_config, + session=session, ) from_agent_flag = converter._convert_to_app_config( - app_model=_app_model(mode=AppMode.CHAT, is_agent=True), - app_model_config=_app_model_config(id="cfg-2"), + app_model=agent_flag_app, + app_model_config=agent_flag_config, + session=session, ) from_chat_mode = converter._convert_to_app_config( - app_model=_app_model(mode=AppMode.CHAT, is_agent=False), - app_model_config=_app_model_config(id="cfg-3"), + app_model=chat_app, + app_model_config=chat_config, + session=session, ) from_completion_mode = converter._convert_to_app_config( - app_model=_app_model(mode=AppMode.COMPLETION, is_agent=False), - app_model_config=_app_model_config(id="cfg-4"), + app_model=completion_app, + app_model_config=completion_config, + session=session, ) assert from_agent_mode is agent_result assert from_agent_flag is agent_result assert from_chat_mode is chat_result assert from_completion_mode is completion_result + agent_flag_app.is_agent_with_session.assert_called_once_with(session=session) + load_annotation_reply.assert_has_calls( + [call(session, "app-1"), call(session, "app-2"), call(session, "app-3"), call(session, "app-4")] + ) + assert all( + manager_call.kwargs["annotation_reply"] == {"enabled": False} + for manager in (agent_get_app_config, chat_get_app_config, completion_get_app_config) + for manager_call in manager.call_args_list + ) def test_convert_to_app_config_should_raise_for_invalid_app_mode(converter: WorkflowConverter) -> None: - app_model = _app_model(mode=AppMode.WORKFLOW, is_agent=False) + app_model = _app_model(mode=AppMode.WORKFLOW, is_agent_with_session=MagicMock(return_value=False)) + session = MagicMock() with pytest.raises(ValueError, match="Invalid app mode"): - converter._convert_to_app_config(app_model=app_model, app_model_config=_app_model_config(id="cfg")) + converter._convert_to_app_config( + app_model=app_model, app_model_config=_app_model_config(id="cfg"), session=session + ) def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_id( diff --git a/api/tests/unit_tests/tasks/test_clean_dataset_task.py b/api/tests/unit_tests/tasks/test_clean_dataset_task.py index 7ce897eb029..826276086ba 100644 --- a/api/tests/unit_tests/tasks/test_clean_dataset_task.py +++ b/api/tests/unit_tests/tasks/test_clean_dataset_task.py @@ -459,5 +459,6 @@ class TestIndexProcessorParameters: assert call_args[0][1] is None # Verify keyword arguments + assert call_args[1]["session"] is mock_db_session.session assert call_args[1]["with_keywords"] is True assert call_args[1]["delete_child_chunks"] is True diff --git a/api/tests/unit_tests/tasks/test_dataset_indexing_task.py b/api/tests/unit_tests/tasks/test_dataset_indexing_task.py index 69b7a156ca8..39ff472ab7e 100644 --- a/api/tests/unit_tests/tasks/test_dataset_indexing_task.py +++ b/api/tests/unit_tests/tasks/test_dataset_indexing_task.py @@ -105,17 +105,30 @@ def mock_db_session(): try: where = stmt.whereclause if where is None: - return None - # Both single-clause and AND-clause-list cases - clauses = list(getattr(where, "clauses", [where])) - for clause in clauses: - left = getattr(clause, "left", None) - right = getattr(clause, "right", None) - if left is not None and right is not None: - if getattr(left, "key", None) == "id": - return getattr(right, "value", None) + clauses = [] + else: + try: + clauses = list(where.clauses) + except AttributeError: + clauses = [where] except Exception: - pass + return None + + for clause in clauses: + try: + left = clause.left + right = clause.right + except AttributeError: + continue + try: + key = left.key + except AttributeError: + continue + if key == "id": + try: + return right.value + except AttributeError: + return None return None def _scalar_side_effect(stmt): @@ -127,7 +140,6 @@ def mock_db_session(): docs = session._shared_data.get("documents", []) if not docs: return None - # When the WHERE clause filters by id, return the matching document queried_id = _extract_id_from_where(stmt) if queried_id: doc_map = {d.id: d for d in docs} @@ -514,7 +526,7 @@ class TestBatchProcessing: _document_indexing(dataset_id, document_ids) # Assert - IndexingRunner should still be called with empty list - mock_indexing_runner.run.assert_called_once_with([]) + mock_indexing_runner.run.assert_called_once_with([], mock_db_session) # ============================================================================ @@ -1606,16 +1618,24 @@ class TestDocumentIndexingTaskSummaryFlow: session2 = MagicMock() session2.begin.return_value = nullcontext() session3 = MagicMock() + session4 = MagicMock() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=phase1_docs)) session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock( + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=phase1_docs)) + session4.scalar.return_value = dataset + session4.scalars.return_value = MagicMock( all=MagicMock(return_value=[doc_eligible, doc_skip_form, doc_skip_status]) ) create_session_mock = MagicMock( - side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)] + side_effect=[ + _SessionContext(session1), + _SessionContext(session2), + _SessionContext(session3), + _SessionContext(session4), + ] ) monkeypatch.setattr("tasks.document_indexing_task.session_factory.create_session", create_session_mock) @@ -1659,15 +1679,25 @@ class TestDocumentIndexingTaskSummaryFlow: session2 = MagicMock() session2.begin.return_value = nullcontext() session3 = MagicMock() + session4 = MagicMock() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[doc_eligible])) + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session4.scalar.return_value = dataset + session4.scalars.return_value = MagicMock(all=MagicMock(return_value=[doc_eligible])) monkeypatch.setattr( "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), + MagicMock( + side_effect=[ + _SessionContext(session1), + _SessionContext(session2), + _SessionContext(session3), + _SessionContext(session4), + ] + ), ) features = SimpleNamespace( @@ -1698,13 +1728,23 @@ class TestDocumentIndexingTaskSummaryFlow: session2 = MagicMock() session2.begin.return_value = nullcontext() session3 = MagicMock() + session4 = MagicMock() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3.scalar.return_value = None # dataset not found on second query + session3.scalar.return_value = dataset + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session4.scalar.return_value = None # dataset not found after indexing monkeypatch.setattr( "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), + MagicMock( + side_effect=[ + _SessionContext(session1), + _SessionContext(session2), + _SessionContext(session3), + _SessionContext(session4), + ] + ), ) features = SimpleNamespace( @@ -1720,7 +1760,7 @@ class TestDocumentIndexingTaskSummaryFlow: _document_indexing("dataset-1", ["doc-1"]) # Assert - session3.scalar.assert_called() + session4.scalar.assert_called() def test_should_skip_summary_when_not_high_quality(self, monkeypatch: pytest.MonkeyPatch) -> None: """Test summary generation skipped when indexing_technique is not high_quality.""" @@ -1735,14 +1775,24 @@ class TestDocumentIndexingTaskSummaryFlow: session2 = MagicMock() session2.begin.return_value = nullcontext() session3 = MagicMock() + session4 = MagicMock() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) session3.scalar.return_value = dataset + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session4.scalar.return_value = dataset monkeypatch.setattr( "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), + MagicMock( + side_effect=[ + _SessionContext(session1), + _SessionContext(session2), + _SessionContext(session3), + _SessionContext(session4), + ] + ), ) features = SimpleNamespace( @@ -1773,8 +1823,13 @@ class TestDocumentIndexingTaskSummaryFlow: session2.begin.return_value = nullcontext() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session3 = MagicMock() + session3.scalar.return_value = dataset + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - create_session_mock = MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2)]) + create_session_mock = MagicMock( + side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)] + ) monkeypatch.setattr("tasks.document_indexing_task.session_factory.create_session", create_session_mock) features = SimpleNamespace( @@ -1807,10 +1862,13 @@ class TestDocumentIndexingTaskSummaryFlow: session2.begin.return_value = nullcontext() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session3 = MagicMock() + session3.scalar.return_value = dataset + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) monkeypatch.setattr( "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2)]), + MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), ) features = SimpleNamespace( @@ -1855,15 +1913,25 @@ class TestDocumentIndexingTaskSummaryFlow: session2 = MagicMock() session2.begin.return_value = nullcontext() session3 = MagicMock() + session4 = MagicMock() session1.scalar.return_value = dataset session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[_FalseyDocument("missing-doc")])) + session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) + session4.scalar.return_value = dataset + session4.scalars.return_value = MagicMock(all=MagicMock(return_value=[_FalseyDocument("missing-doc")])) monkeypatch.setattr( "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), + MagicMock( + side_effect=[ + _SessionContext(session1), + _SessionContext(session2), + _SessionContext(session3), + _SessionContext(session4), + ] + ), ) features = SimpleNamespace( diff --git a/api/tests/unit_tests/tasks/test_document_indexing_sync_task.py b/api/tests/unit_tests/tasks/test_document_indexing_sync_task.py index e5782899e3b..6787d7cc6e1 100644 --- a/api/tests/unit_tests/tasks/test_document_indexing_sync_task.py +++ b/api/tests/unit_tests/tasks/test_document_indexing_sync_task.py @@ -239,18 +239,14 @@ class TestDataSourceInfoSerialization: mock_runner = MagicMock() mock_runner_class.return_value = mock_runner - # DB session mock — shared across all ``session_factory.create_session()`` calls + # DB session mock — shared across all ``session_factory.create_session()`` calls. session = MagicMock() session.scalars.return_value.all.return_value = [] - # All .first() calls are now session.scalar() — ordered by call sequence: - # session 1: document + dataset, session 2: dataset (clean), session 3: document (update), - # session 4: document (indexing) session.scalar.side_effect = [ mock_document, mock_dataset, + mock_document, mock_dataset, - mock_document, - mock_document, ] begin_cm = MagicMock() diff --git a/api/tests/unit_tests/tasks/test_document_indexing_update_task.py b/api/tests/unit_tests/tasks/test_document_indexing_update_task.py index b73275b97d0..59eaf497960 100644 --- a/api/tests/unit_tests/tasks/test_document_indexing_update_task.py +++ b/api/tests/unit_tests/tasks/test_document_indexing_update_task.py @@ -117,7 +117,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3] + session3.scalar.side_effect = [doc_s3, dataset] runner = MagicMock() processor = MagicMock() @@ -151,7 +151,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.return_value = dataset # dataset.indexing_technique == "economy" + session3.scalar.side_effect = [doc_s1, dataset] # dataset.indexing_technique == "economy" runner = MagicMock() processor = MagicMock() @@ -184,7 +184,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.return_value = dataset + session3.scalar.side_effect = [doc_s1, dataset] runner = MagicMock() processor = MagicMock() @@ -217,7 +217,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.return_value = dataset + session3.scalar.side_effect = [doc_s1, dataset] runner = MagicMock() processor = MagicMock() @@ -251,7 +251,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3] + session3.scalar.side_effect = [doc_s3, dataset] runner = MagicMock() processor = MagicMock() @@ -285,7 +285,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3] + session3.scalar.side_effect = [doc_s3, dataset] runner = MagicMock() processor = MagicMock() @@ -384,7 +384,7 @@ class TestUpdateTaskSummaryGeneration: # Session 3: dataset is None session3 = MagicMock() - session3.scalar.return_value = None + session3.scalar.side_effect = [doc_s1, None] runner = MagicMock() processor = MagicMock() @@ -425,7 +425,7 @@ class TestUpdateTaskSummaryGeneration: need_summary=True, ) session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3_error] + session3.scalar.side_effect = [doc_s3_error, dataset] runner = MagicMock() processor = MagicMock() @@ -458,7 +458,7 @@ class TestUpdateTaskSummaryGeneration: session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3] + session3.scalar.side_effect = [doc_s3, dataset] runner = MagicMock() processor = MagicMock() @@ -493,11 +493,8 @@ class TestUpdateTaskSummaryGeneration: seg = SimpleNamespace(index_node_id="node-1") session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[seg])) - # Session 2: segment deletion - session2 = _session_with_begin() - session3 = MagicMock() - session3.scalar.side_effect = [dataset, doc_s3] + session3.scalar.side_effect = [doc_s3, dataset] runner = MagicMock() processor = MagicMock() @@ -506,7 +503,6 @@ class TestUpdateTaskSummaryGeneration: monkeypatch, sessions=[ _SessionContext(session1), - _SessionContext(session2), _SessionContext(session3), ], runner=runner, diff --git a/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py new file mode 100644 index 00000000000..af7ecb8c570 --- /dev/null +++ b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py @@ -0,0 +1,193 @@ +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from core.rag.index_processor.constant.index_type import IndexStructureType +from models.enums import IndexingStatus, SegmentStatus +from tasks.enable_segment_to_index_task import enable_segment_to_index_task +from tasks.enable_segments_to_index_task import enable_segments_to_index_task + + +def test_enable_segment_commits_index_rows_after_loading() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = SimpleNamespace( + id="segment-1", + status=SegmentStatus.COMPLETED, + content="content", + index_node_id="node-1", + index_node_hash="hash-1", + document_id=document.id, + dataset_id=dataset.id, + get_dataset=MagicMock(return_value=dataset), + get_document=MagicMock(return_value=document), + ) + session = MagicMock() + session.scalar.return_value = segment + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + index_processor = MagicMock() + index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") + enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) + + with ( + patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, + patch( + "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", + enable_summaries, + ), + patch("tasks.enable_segment_to_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = index_processor + enable_segment_to_index_task.run(segment.id) + + assert phase_events == ["load", "commit", "summary"] + + +def test_enable_segment_rolls_back_before_error_compensation() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = SimpleNamespace( + id="segment-1", + status=SegmentStatus.COMPLETED, + content="content", + index_node_id="node-1", + index_node_hash="hash-1", + document_id=document.id, + dataset_id=dataset.id, + enabled=True, + disabled_at=None, + error=None, + get_dataset=MagicMock(return_value=dataset), + get_document=MagicMock(return_value=document), + ) + phase_events: list[str] = [] + session = MagicMock() + session.scalar.return_value = segment + session.rollback.side_effect = lambda: phase_events.append("rollback") + + def commit() -> None: + assert segment.enabled is False + assert segment.status == SegmentStatus.ERROR + assert segment.error == "load failed" + phase_events.append("commit") + + session.commit.side_effect = commit + index_processor = MagicMock() + + def fail_load(*_args, **_kwargs) -> None: + phase_events.append("load") + raise RuntimeError("load failed") + + index_processor.load.side_effect = fail_load + + with ( + patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, + patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, + patch("tasks.enable_segment_to_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = index_processor + enable_segment_to_index_task.run(segment.id) + + assert phase_events == ["load", "rollback", "commit"] + enable_summaries.assert_not_called() + + +def test_enable_segments_commits_index_rows_after_loading() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status="completed", + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = SimpleNamespace( + id="segment-1", + content="content", + index_node_id="node-1", + index_node_hash="hash-1", + document_id=document.id, + dataset_id=dataset.id, + ) + session = MagicMock() + session.scalar.side_effect = [dataset, document] + session.scalars.return_value.all.return_value = [segment] + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + index_processor = MagicMock() + index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") + enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) + + with ( + patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, + patch( + "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", + enable_summaries, + ), + patch("tasks.enable_segments_to_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = index_processor + enable_segments_to_index_task.run([segment.id], dataset.id, document.id) + + assert phase_events == ["load", "commit", "summary"] + + +def test_enable_segments_rolls_back_before_error_compensation() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status="completed", + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = SimpleNamespace( + id="segment-1", + content="content", + index_node_id="node-1", + index_node_hash="hash-1", + document_id=document.id, + dataset_id=dataset.id, + ) + phase_events: list[str] = [] + session = MagicMock() + session.scalar.side_effect = [dataset, document] + session.scalars.return_value.all.return_value = [segment] + session.rollback.side_effect = lambda: phase_events.append("rollback") + session.execute.side_effect = lambda *_args, **_kwargs: phase_events.append("compensate") + session.commit.side_effect = lambda: phase_events.append("commit") + index_processor = MagicMock() + + def fail_load(*_args, **_kwargs) -> None: + phase_events.append("load") + raise RuntimeError("load failed") + + index_processor.load.side_effect = fail_load + + with ( + patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, + patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, + patch("tasks.enable_segments_to_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = index_processor + enable_segments_to_index_task.run([segment.id], dataset.id, document.id) + + assert phase_events == ["load", "rollback", "compensate", "commit"] + enable_summaries.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_resume_agent_app_task.py b/api/tests/unit_tests/tasks/test_resume_agent_app_task.py index 189f6c1512e..3b1c887e8ca 100644 --- a/api/tests/unit_tests/tasks/test_resume_agent_app_task.py +++ b/api/tests/unit_tests/tasks/test_resume_agent_app_task.py @@ -49,18 +49,19 @@ def test_resume_happy_path_account_user_sets_tenant_and_runs(mocker: MockerFixtu conversation = MagicMock(from_account_id="acct-1", from_end_user_id=None, invoke_from=InvokeFrom.WEB_APP) account = MagicMock() app = MagicMock(tenant_id="tenant-1") - _wire_db(mocker, form=_form(), app=app, conversation=conversation, account=account) + db = _wire_db(mocker, form=_form(), app=app, conversation=conversation, account=account) gen = mocker.patch(f"{MODULE}.AgentAppGenerator") mod.resume_agent_app_execution(conversation_id="conv-1", form_id="form-1") - account.set_tenant_id.assert_called_once_with("tenant-1") + account.set_tenant_id_with_session.assert_called_once_with("tenant-1", session=db.session.return_value) gen.return_value.resume_after_form_submission.assert_called_once() kwargs = gen.return_value.resume_after_form_submission.call_args.kwargs assert kwargs["conversation_id"] == "conv-1" assert kwargs["user"] is account assert kwargs["app_model"] is app assert kwargs["invoke_from"] == InvokeFrom.WEB_APP + assert kwargs["session"] is db.session.return_value def test_resume_end_user_path(mocker: MockerFixture): diff --git a/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py new file mode 100644 index 00000000000..3454035fe5f --- /dev/null +++ b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py @@ -0,0 +1,106 @@ +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from models.enums import SegmentStatus +from tasks.delete_segment_from_index_task import delete_segment_from_index_task +from tasks.disable_segment_from_index_task import disable_segment_from_index_task +from tasks.disable_segments_from_index_task import disable_segments_from_index_task + + +def test_disable_segment_commits_index_cleanup() -> None: + dataset = SimpleNamespace(id="dataset-1") + document = SimpleNamespace(enabled=True, archived=False, indexing_status="completed", doc_form="text_model") + segment = SimpleNamespace( + id="segment-1", + status=SegmentStatus.COMPLETED, + index_node_id="node-1", + disabled_by="user-1", + get_dataset=MagicMock(return_value=dataset), + get_document=MagicMock(return_value=document), + ) + session = MagicMock() + session.scalar.return_value = segment + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + processor = MagicMock() + processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") + disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) + + with ( + patch( + "tasks.disable_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session) + ), + patch("tasks.disable_segment_from_index_task.IndexProcessorFactory") as processor_factory, + patch( + "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", + disable_summaries, + ), + patch("tasks.disable_segment_from_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = processor + disable_segment_from_index_task.run(segment.id) + + assert phase_events == ["clean", "commit", "summary"] + + +def test_disable_segments_commits_index_cleanup() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status="completed", + doc_form="text_model", + ) + segment = SimpleNamespace(id="segment-1", index_node_id="node-1", disabled_by="user-1") + session = MagicMock() + session.scalar.side_effect = [dataset, document] + session.scalars.return_value.all.return_value = [segment] + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + processor = MagicMock() + processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") + disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) + + with ( + patch( + "tasks.disable_segments_from_index_task.session_factory.create_session", return_value=nullcontext(session) + ), + patch("tasks.disable_segments_from_index_task.IndexProcessorFactory") as processor_factory, + patch( + "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", + disable_summaries, + ), + patch("tasks.disable_segments_from_index_task.redis_client.delete"), + ): + processor_factory.return_value.init_index_processor.return_value = processor + disable_segments_from_index_task.run([segment.id], dataset.id, document.id) + + assert phase_events == ["clean", "commit", "summary"] + + +def test_delete_segment_commits_index_cleanup_without_attachments() -> None: + dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) + document = SimpleNamespace( + id="document-1", + enabled=True, + archived=False, + indexing_status="completed", + doc_form="text_model", + ) + session = MagicMock() + session.scalar.side_effect = [dataset, document] + phase_events: list[str] = [] + session.commit.side_effect = lambda: phase_events.append("commit") + processor = MagicMock() + processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") + + with ( + patch("tasks.delete_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session)), + patch("tasks.delete_segment_from_index_task.IndexProcessorFactory") as processor_factory, + ): + processor_factory.return_value.init_index_processor.return_value = processor + delete_segment_from_index_task.run(["node-1"], dataset.id, document.id, ["segment-1"]) + + assert phase_events == ["clean", "commit"] diff --git a/api/tests/unit_tests/tasks/test_workflow_execute_task.py b/api/tests/unit_tests/tasks/test_workflow_execute_task.py index a3fd70f205f..d5a32e082ee 100644 --- a/api/tests/unit_tests/tasks/test_workflow_execute_task.py +++ b/api/tests/unit_tests/tasks/test_workflow_execute_task.py @@ -714,6 +714,7 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monk "tasks.app_generate.workflow_execute_task.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", lambda **kwargs: MagicMock(), ) + session = MagicMock() _resume_advanced_chat( app_model=SimpleNamespace(id="app-id"), @@ -728,10 +729,12 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs(monk pause_state_config=MagicMock(), workflow_run_id="workflow-run-id", workflow_run=SimpleNamespace(triggered_from="app_run"), + session=session, ) resumed_entity = generator_instance.resume.call_args.kwargs["application_generate_entity"] assert resumed_entity.stream is True + assert generator_instance.resume.call_args.kwargs["session"] is session publish_streaming_response.assert_called_once_with( response_stream, "workflow-run-id", diff --git a/packages/contracts/generated/api/console/apps/orpc.gen.ts b/packages/contracts/generated/api/console/apps/orpc.gen.ts index 7af9c84a97f..d479394c028 100644 --- a/packages/contracts/generated/api/console/apps/orpc.gen.ts +++ b/packages/contracts/generated/api/console/apps/orpc.gen.ts @@ -2205,7 +2205,7 @@ export const messages = { } /** - * Modify app model config + * Modify the app model config and dataset joins in one request transaction * * Update application model configuration */ @@ -2216,7 +2216,7 @@ export const post28 = oc method: 'POST', operationId: 'postAppsByAppIdModelConfig', path: '/apps/{app_id}/model-config', - summary: 'Modify app model config', + summary: 'Modify the app model config and dataset joins in one request transaction', tags: ['console'], }) .input(