diff --git a/.github/workflows/style.yml b/.github/workflows/style.yml index 0c34de4a7b9..438dd125ce4 100644 --- a/.github/workflows/style.yml +++ b/.github/workflows/style.yml @@ -22,6 +22,7 @@ jobs: uses: actions/checkout@df4cb1c069e1874edd31b4311f1884172cec0e10 # v6.0.3 with: persist-credentials: false + fetch-depth: 0 - name: Check changed files id: changed-files diff --git a/api/controllers/console/app/completion.py b/api/controllers/console/app/completion.py index 0d9c6009329..c408569703a 100644 --- a/api/controllers/console/app/completion.py +++ b/api/controllers/console/app/completion.py @@ -7,6 +7,7 @@ from uuid import UUID from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, InternalServerError, NotFound import services @@ -22,7 +23,7 @@ from controllers.console.app.error import ( ProviderNotInitializeError, ProviderQuotaExceededError, ) -from controllers.console.app.wraps import get_app_model +from controllers.console.app.wraps import get_app_model, with_session from controllers.console.wraps import ( RBACPermission, RBACResourceScope, @@ -155,7 +156,8 @@ class CompletionMessageApi(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) @get_app_model(mode=AppMode.COMPLETION) - def post(self, current_user: Account, app_model: App): + @with_session + 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) @@ -164,7 +166,12 @@ class CompletionMessageApi(Resource): try: response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=streaming + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.DEBUGGER, + streaming=streaming, ) # response-contract:ignore compact_generate_response @@ -231,8 +238,11 @@ class ChatMessageApi(Resource): @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]) - def post(self, current_tenant_id: str, current_user: Account, app_model: App): - return _create_chat_message(current_tenant_id=current_tenant_id, current_user=current_user, app_model=app_model) + @with_session + 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 + ) @console_ns.route("/agent//chat-messages") @@ -251,9 +261,11 @@ class AgentChatMessageApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, agent_id: UUID): + @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) return _create_chat_message( + session=session, current_tenant_id=current_tenant_id, current_user=current_user, app_model=app_model, @@ -276,9 +288,11 @@ class AgentBuildChatFinalizeApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, agent_id: UUID): + @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) return _create_build_chat_finalization_message( + session=session, current_tenant_id=current_tenant_id, current_user=current_user, app_model=app_model, @@ -340,6 +354,7 @@ def _resolve_current_user_agent_debug_conversation_id( def _create_chat_message( *, + session: Session, current_user: Account, app_model: App, current_tenant_id: str | None = None, @@ -374,6 +389,7 @@ def _create_chat_message( args["external_trace_id"] = external_trace_id return _generate_chat_message_response( + session=session, current_user=current_user, app_model=app_model, args=args, @@ -382,7 +398,7 @@ def _create_chat_message( def _create_build_chat_finalization_message( - *, current_user: Account, app_model: App, current_tenant_id: str, agent_id: str + *, 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( current_tenant_id=current_tenant_id, @@ -403,6 +419,7 @@ def _create_build_chat_finalization_message( args["external_trace_id"] = external_trace_id response = _generate_chat_message( + session=session, current_user=current_user, app_model=app_model, args=args, @@ -458,6 +475,7 @@ def _drain_streaming_generate_response(response: RateLimitGenerator | Generator[ def _generate_chat_message( *, + session: Session, current_user: Account, app_model: App, args: dict[str, Any], @@ -465,6 +483,7 @@ def _generate_chat_message( ): try: return AppGenerateService.generate( + session=session, app_model=app_model, user=current_user, args=args, @@ -499,12 +518,14 @@ def _generate_chat_message( def _generate_chat_message_response( *, + session: Session, current_user: Account, app_model: App, args: dict[str, Any], streaming: bool, ): response = _generate_chat_message( + session=session, current_user=current_user, app_model=app_model, args=args, diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 2f2f451ed36..609dbfb82c5 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -27,7 +27,7 @@ from controllers.console.app.error import ( DraftWorkflowNotSync, ) from controllers.console.app.permission_keys import get_app_permission_keys -from controllers.console.app.wraps import get_app_model +from controllers.console.app.wraps import get_app_model, with_session from controllers.console.wraps import ( RBACPermission, RBACResourceScope, @@ -631,7 +631,8 @@ class AdvancedChatDraftWorkflowRunApi(Resource): @get_app_model(mode=[AppMode.ADVANCED_CHAT]) @with_current_user @edit_permission_required - def post(self, current_user: Account, app_model: App): + @with_session + def post(self, session: Session, current_user: Account, app_model: App): """ Run draft workflow """ @@ -644,7 +645,12 @@ class AdvancedChatDraftWorkflowRunApi(Resource): try: response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.DEBUGGER, + streaming=True, ) return helper.compact_generate_response(response) @@ -1045,7 +1051,8 @@ class DraftWorkflowRunApi(Resource): @get_app_model(mode=[AppMode.WORKFLOW]) @with_current_user @edit_permission_required - def post(self, current_user: Account, app_model: App): + @with_session + def post(self, session: Session, current_user: Account, app_model: App): """ Run draft workflow """ @@ -1057,6 +1064,7 @@ class DraftWorkflowRunApi(Resource): try: response = AppGenerateService.generate( + session=session, app_model=app_model, user=current_user, args=args, @@ -1590,7 +1598,8 @@ class DraftWorkflowTriggerRunApi(Resource): @get_app_model(mode=[AppMode.WORKFLOW]) @with_current_user @edit_permission_required - def post(self, current_user: Account, app_model: App): + @with_session + def post(self, session: Session, current_user: Account, app_model: App): """ Poll for trigger events and execute full workflow when event arrives """ @@ -1618,6 +1627,7 @@ class DraftWorkflowTriggerRunApi(Resource): workflow_args[SKIP_PREPARE_USER_INPUTS_KEY] = True return helper.compact_generate_response( AppGenerateService.generate( + session=session, app_model=app_model, user=current_user, args=workflow_args, @@ -1740,7 +1750,8 @@ class DraftWorkflowTriggerRunAllApi(Resource): @with_current_user @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) - def post(self, current_user: Account, app_model: App): + @with_session + def post(self, session: Session, current_user: Account, app_model: App): """ Full workflow debug when the start node is a trigger """ @@ -1772,6 +1783,7 @@ class DraftWorkflowTriggerRunAllApi(Resource): workflow_args[SKIP_PREPARE_USER_INPUTS_KEY] = True response = AppGenerateService.generate( + session=session, app_model=app_model, user=current_user, args=workflow_args, diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index dbf546532ee..0bef535d82b 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -6,6 +6,7 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator, model_validator from sqlalchemy import func, select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -15,6 +16,7 @@ from controllers.common.schema import query_params_from_model, register_response 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, @@ -538,7 +540,8 @@ class DatasetListApi(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): payload = DatasetCreatePayload.model_validate(console_ns.payload or {}) # The role of the current user in the ta table must be admin, owner, or editor, or dataset_operator @@ -552,6 +555,7 @@ class DatasetListApi(Resource): try: dataset = DatasetService.create_empty_dataset( + session=session, tenant_id=current_tenant_id, name=payload.name, description=payload.description, @@ -561,21 +565,20 @@ class DatasetListApi(Resource): provider=payload.provider, external_knowledge_api_id=payload.external_knowledge_api_id, external_knowledge_id=payload.external_knowledge_id, - session=db.session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() permission_keys_map = enterprise_rbac_service.RBACService.DatasetPermissions.batch_get( - str(current_tenant_id), + current_tenant_id, current_user.id, - [str(dataset.id)], + [dataset.id], ) item = DatasetDetailWithPartialMembersResponse.model_validate(dataset, from_attributes=True).model_dump( mode="json" ) - item["permission_keys"] = permission_keys_map.get(str(dataset.id), []) + item["permission_keys"] = permission_keys_map.get(dataset.id, []) return item, 201 @@ -607,7 +610,7 @@ class DatasetApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) permissions = enterprise_rbac_service.RBACService.MyPermissions.get( - str(current_tenant_id), + current_tenant_id, current_user.id, dataset_id=dataset_id_str, ) @@ -660,7 +663,8 @@ class DatasetApi(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): + @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) if dataset is None: @@ -681,16 +685,16 @@ class DatasetApi(Resource): # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not dify_config.RBAC_ENABLED: DatasetPermissionService.check_permission( - current_user, dataset, payload.permission, payload.partial_member_list, db.session + session, current_user, dataset, payload.permission, payload.partial_member_list ) - dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user, db.session) + dataset = DatasetService.update_dataset(session, dataset_id_str, payload_data, current_user) if dataset is None: raise NotFound("Dataset not found.") permission_keys_map = enterprise_rbac_service.RBACService.DatasetPermissions.batch_get( - str(current_tenant_id), + current_tenant_id, current_user.id, [dataset_id_str], ) diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index 7a3c746b80f..5b036641d4d 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -4,6 +4,7 @@ from uuid import UUID from flask import request from flask_restx import Resource, fields, marshal from pydantic import BaseModel, Field, RootModel +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, InternalServerError, NotFound import services @@ -15,6 +16,7 @@ from controllers.common.schema import ( register_schema_models, ) 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, @@ -207,7 +209,8 @@ class ExternalApiTemplateListApi(Resource): ) @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): payload = ExternalKnowledgeApiPayload.model_validate(console_ns.payload or {}) ExternalDatasetService.validate_api_list(payload.settings) @@ -218,7 +221,10 @@ class ExternalApiTemplateListApi(Resource): try: external_knowledge_api = ExternalDatasetService.create_external_knowledge_api( - tenant_id=current_tenant_id, user_id=current_user.id, args=payload.model_dump() + tenant_id=current_tenant_id, + user_id=current_user.id, + args=payload.model_dump(), + session=session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -241,10 +247,11 @@ class ExternalApiTemplateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, current_tenant_id: str, external_knowledge_api_id: UUID): + @with_session + def get(self, session: Session, current_tenant_id: str, external_knowledge_api_id: UUID): external_knowledge_api_id_str = str(external_knowledge_api_id) external_knowledge_api = ExternalDatasetService.get_external_knowledge_api( - external_knowledge_api_id_str, current_tenant_id + external_knowledge_api_id=external_knowledge_api_id_str, tenant_id=current_tenant_id, session=session ) if external_knowledge_api is None: raise NotFound("API template not found.") @@ -262,7 +269,8 @@ class ExternalApiTemplateApi(Resource): @console_ns.expect(console_ns.models[ExternalKnowledgeApiPayload.__name__]) @with_current_user @with_current_tenant_id - def patch(self, current_tenant_id: str, current_user: Account, external_knowledge_api_id: UUID): + @with_session + def patch(self, session: Session, current_tenant_id: str, current_user: Account, external_knowledge_api_id: UUID): external_knowledge_api_id_str = str(external_knowledge_api_id) payload = ExternalKnowledgeApiPayload.model_validate(console_ns.payload or {}) @@ -273,6 +281,7 @@ class ExternalApiTemplateApi(Resource): user_id=current_user.id, external_knowledge_api_id=external_knowledge_api_id_str, args=payload.model_dump(), + session=session, ) return external_knowledge_api.to_dict(), 200 @@ -283,13 +292,14 @@ class ExternalApiTemplateApi(Resource): @console_ns.response(204, "External knowledge API deleted successfully") @with_current_user @with_current_tenant_id - def delete(self, current_tenant_id: str, current_user: Account, external_knowledge_api_id: UUID): + @with_session + def delete(self, session: Session, current_tenant_id: str, current_user: Account, external_knowledge_api_id: UUID): external_knowledge_api_id_str = str(external_knowledge_api_id) if not (current_user.has_edit_permission or current_user.is_dataset_operator): raise Forbidden() - ExternalDatasetService.delete_external_knowledge_api(current_tenant_id, external_knowledge_api_id_str) + ExternalDatasetService.delete_external_knowledge_api(session, current_tenant_id, external_knowledge_api_id_str) return "", 204 @@ -303,11 +313,14 @@ class ExternalApiUseCheckApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, current_tenant_id: str, external_knowledge_api_id: UUID): + @with_session + def get(self, session: Session, current_tenant_id: str, external_knowledge_api_id: UUID): external_knowledge_api_id_str = str(external_knowledge_api_id) external_knowledge_api_is_using, count = ExternalDatasetService.external_knowledge_api_use_check( - external_knowledge_api_id_str, current_tenant_id + session, + external_knowledge_api_id_str, + current_tenant_id, ) return {"is_using": external_knowledge_api_is_using, "count": count}, 200 @@ -327,7 +340,8 @@ class ExternalDatasetCreateApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EXTERNAL_CONNECT) @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, or editor payload = ExternalDatasetCreatePayload.model_validate(console_ns.payload or {}) args = payload.model_dump(exclude_none=True) @@ -341,6 +355,7 @@ class ExternalDatasetCreateApi(Resource): tenant_id=current_tenant_id, user_id=current_user.id, args=args, + session=session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -375,7 +390,8 @@ class ExternalKnowledgeHitTestingApi(Resource): @account_initialization_required @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_PIPELINE_TEST) - 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) if dataset is None: @@ -391,7 +407,7 @@ class ExternalKnowledgeHitTestingApi(Resource): try: response = HitTestingService.external_retrieve( - session=db.session, + session=session, dataset=dataset, query=payload.query, account=current_user, diff --git a/api/controllers/console/datasets/hit_testing.py b/api/controllers/console/datasets/hit_testing.py index 739f0250333..b30aadfbbc7 100644 --- a/api/controllers/console/datasets/hit_testing.py +++ b/api/controllers/console/datasets/hit_testing.py @@ -3,8 +3,10 @@ from __future__ import annotations from uuid import UUID from flask_restx import Resource +from sqlalchemy.orm import Session from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required from fields.hit_testing_fields import HitTestingResponse from libs.helper import dump_response @@ -45,7 +47,10 @@ class HitTestingApi(Resource, DatasetsHitTestingBase): @with_current_tenant_id @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_PIPELINE_TEST) - def post(self, current_user: Account, current_tenant_id: str, dataset_id: UUID) -> dict[str, object]: + @with_session + def post( + self, session: Session, current_user: Account, current_tenant_id: str, dataset_id: UUID + ) -> dict[str, object]: dataset_id_str = str(dataset_id) dataset = self.get_and_validate_dataset(dataset_id_str, current_user, current_tenant_id) @@ -54,5 +59,5 @@ class HitTestingApi(Resource, DatasetsHitTestingBase): return dump_response( HitTestingResponse, - self.perform_hit_testing(dataset, args, current_user, current_tenant_id), + self.perform_hit_testing(session, dataset, args, current_user, current_tenant_id), ) diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index 6464e435fc2..cc02a990168 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -2,6 +2,7 @@ import logging from typing import Any, cast from pydantic import BaseModel, Field +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, InternalServerError, NotFound import services @@ -108,6 +109,7 @@ class DatasetsHitTestingBase: @staticmethod def perform_hit_testing( + session: Session, dataset: Dataset, args: dict[str, Any], current_user: Account | None = None, @@ -116,7 +118,7 @@ class DatasetsHitTestingBase: try: current_user, _ = resolve_account_fallback(current_user, current_tenant_id) response = HitTestingService.retrieve( - session=db.session, + session=session, dataset=dataset, query=cast(str, args.get("query")), account=current_user, diff --git a/api/controllers/console/explore/completion.py b/api/controllers/console/explore/completion.py index 729be6a909b..ef761c8febd 100644 --- a/api/controllers/console/explore/completion.py +++ b/api/controllers/console/explore/completion.py @@ -3,6 +3,7 @@ from typing import Any, Literal from uuid import UUID from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError, NotFound import services @@ -16,6 +17,7 @@ from controllers.console.app.error import ( ProviderNotInitializeError, ProviderQuotaExceededError, ) +from controllers.console.app.wraps import with_session from controllers.console.explore.error import NotChatAppError, NotCompletionAppError from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import with_current_user, with_current_user_id @@ -85,7 +87,8 @@ class CompletionApi(InstalledAppResource): @console_ns.expect(console_ns.models[CompletionMessageExplorePayload.__name__]) @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp): + @with_session + def post(self, session: Session, current_user: Account, installed_app: InstalledApp): app_model = installed_app.app if app_model is None: raise AppUnavailableError() @@ -103,7 +106,12 @@ class CompletionApi(InstalledAppResource): try: response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=streaming + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=streaming, ) return helper.compact_generate_response(response) @@ -161,7 +169,8 @@ class ChatApi(InstalledAppResource): @console_ns.expect(console_ns.models[ChatMessagePayload.__name__]) @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp): + @with_session + def post(self, session: Session, current_user: Account, installed_app: InstalledApp): app_model = installed_app.app if app_model is None: raise AppUnavailableError() @@ -179,7 +188,12 @@ class ChatApi(InstalledAppResource): try: response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=True + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=True, ) return helper.compact_generate_response(response) diff --git a/api/controllers/console/explore/message.py b/api/controllers/console/explore/message.py index 2be550b2f28..891966bee5b 100644 --- a/api/controllers/console/explore/message.py +++ b/api/controllers/console/explore/message.py @@ -4,6 +4,7 @@ from uuid import UUID from flask import request from pydantic import BaseModel, TypeAdapter +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError, NotFound from controllers.common.controller_schemas import MessageFeedbackPayload, MessageListQuery @@ -17,6 +18,7 @@ from controllers.console.app.error import ( ProviderNotInitializeError, ProviderQuotaExceededError, ) +from controllers.console.app.wraps import with_session from controllers.console.explore.error import ( AppSuggestedQuestionsAfterAnswerDisabledError, NotChatAppError, @@ -88,8 +90,8 @@ class MessageListApi(InstalledAppResource): pagination = MessageService.pagination_by_first_id( app_model, current_user, - str(args.conversation_id), - str(args.first_id) if args.first_id else None, + args.conversation_id, + args.first_id or None, args.limit, ) adapter = TypeAdapter(ExploreMessageListItem) @@ -144,7 +146,8 @@ class MessageMoreLikeThisApi(InstalledAppResource): @console_ns.doc(params=query_params_from_model(MoreLikeThisQuery)) @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @with_current_user - def get(self, current_user: Account, installed_app: InstalledApp, message_id: UUID): + @with_session + def get(self, session: Session, current_user: Account, installed_app: InstalledApp, message_id: UUID): app_model = installed_app.app if app_model is None: raise AppUnavailableError() @@ -159,6 +162,7 @@ class MessageMoreLikeThisApi(InstalledAppResource): try: response = AppGenerateService.generate_more_like_this( + session=session, app_model=app_model, user=current_user, message_id=message_id_str, diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index 8b40ce43814..4488f645702 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -6,6 +6,7 @@ from flask import request from flask_restx import Resource, fields, marshal, marshal_with from pydantic import AliasChoices, BaseModel, Field, field_validator from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, InternalServerError, NotFound import services @@ -37,7 +38,7 @@ from controllers.console.app.error import ( ProviderQuotaExceededError, UnsupportedAudioTypeError, ) -from controllers.console.app.wraps import get_app_model_with_trial +from controllers.console.app.wraps import get_app_model_with_trial, with_session from controllers.console.explore.error import ( AppSuggestedQuestionsAfterAnswerDisabledError, NotChatAppError, @@ -458,7 +459,8 @@ class TrialAppWorkflowRunApi(TrialAppResource): @console_ns.expect(console_ns.models[WorkflowRunRequest.__name__]) @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @with_current_user - def post(self, current_user: Account, trial_app): + @with_session + def post(self, session: Session, current_user: Account, trial_app): """ Run workflow """ @@ -475,7 +477,12 @@ class TrialAppWorkflowRunApi(TrialAppResource): app_id = app_model.id user_id = current_user.id response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=True + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=True, ) RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) return helper.compact_generate_response(response) @@ -525,7 +532,8 @@ class TrialChatApi(TrialAppResource): @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @trial_feature_enable @with_current_user - def post(self, current_user: Account, trial_app): + @with_session + def post(self, session: Session, current_user: Account, trial_app): app_model = trial_app app_mode = AppMode.value_of(app_model.mode) if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}: @@ -548,7 +556,12 @@ class TrialChatApi(TrialAppResource): user_id = current_user.id response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=True + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=True, ) RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) return helper.compact_generate_response(response) @@ -721,7 +734,8 @@ class TrialCompletionApi(TrialAppResource): @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @trial_feature_enable @with_current_user - def post(self, current_user: Account, trial_app): + @with_session + def post(self, session: Session, current_user: Account, trial_app): app_model = trial_app if app_model.mode != "completion": raise NotCompletionAppError() @@ -738,7 +752,12 @@ class TrialCompletionApi(TrialAppResource): user_id = current_user.id response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=streaming + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=streaming, ) RecommendedAppService.add_trial_app_record(db.session, app_id, user_id) diff --git a/api/controllers/console/explore/workflow.py b/api/controllers/console/explore/workflow.py index f4e9e3cd7ef..e41e3e18b2c 100644 --- a/api/controllers/console/explore/workflow.py +++ b/api/controllers/console/explore/workflow.py @@ -1,5 +1,6 @@ import logging +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError from controllers.common.controller_schemas import WorkflowRunPayload @@ -11,6 +12,7 @@ from controllers.console.app.error import ( ProviderNotInitializeError, ProviderQuotaExceededError, ) +from controllers.console.app.wraps import with_session from controllers.console.explore.error import NotWorkflowAppError from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import with_current_user @@ -44,7 +46,8 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): @console_ns.expect(console_ns.models[WorkflowRunPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[GeneratedAppResponse.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp): + @with_session + def post(self, session: Session, current_user: Account, installed_app: InstalledApp): """ Run workflow """ @@ -59,7 +62,12 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): args = payload.model_dump(exclude_none=True) try: response = AppGenerateService.generate( - app_model=app_model, user=current_user, args=args, invoke_from=InvokeFrom.EXPLORE, streaming=True + session=session, + app_model=app_model, + user=current_user, + args=args, + invoke_from=InvokeFrom.EXPLORE, + streaming=True, ) return helper.compact_generate_response(response) diff --git a/api/controllers/inner_api/agent/tools.py b/api/controllers/inner_api/agent/tools.py index 8317a64bc28..be11c47486c 100644 --- a/api/controllers/inner_api/agent/tools.py +++ b/api/controllers/inner_api/agent/tools.py @@ -2,11 +2,12 @@ from flask_restx import Resource from pydantic import ValidationError +from sqlalchemy.orm import Session from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.inner_api import inner_api_ns from controllers.inner_api.wraps import agent_inner_api_only -from extensions.ext_database import db from libs.exception import BaseHTTPException from services.agent_tool_inner_service import AgentToolInnerService from services.entities.agent_tool_inner import AgentToolInvokeRequest, AgentToolInvokeResponse @@ -46,7 +47,8 @@ class AgentToolInvokeApi(Resource): @inner_api_ns.doc("inner_agent_tool_invoke") @inner_api_ns.expect(inner_api_ns.models[AgentToolInvokeRequest.__name__]) @inner_api_ns.response(200, "Tool invoked successfully", inner_api_ns.models[AgentToolInvokeResponse.__name__]) - def post(self) -> dict[str, object]: + @with_session + def post(self, session: Session) -> dict[str, object]: try: payload = AgentToolInvokeRequest.model_validate(inner_api_ns.payload or {}) except ValidationError as exc: @@ -57,7 +59,7 @@ class AgentToolInvokeApi(Resource): ) from exc try: - response = AgentToolInnerService().invoke(payload, session=db.session()) + response = AgentToolInnerService().invoke(session=session, request=payload) except AgentToolInnerServiceError as exc: raise AgentToolInvokeHttpError( error_code=exc.error_code, diff --git a/api/controllers/inner_api/knowledge/retrieval.py b/api/controllers/inner_api/knowledge/retrieval.py index e34dedea286..9d83007aa3d 100644 --- a/api/controllers/inner_api/knowledge/retrieval.py +++ b/api/controllers/inner_api/knowledge/retrieval.py @@ -9,12 +9,13 @@ app/dataset validation remains in the service layer. from flask_restx import Resource from pydantic import ValidationError +from sqlalchemy.orm import Session from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.inner_api import inner_api_ns from controllers.inner_api.wraps import plugin_inner_api_only from core.workflow.nodes.knowledge_retrieval import exc as retrieval_exc -from extensions.ext_database import db from libs.exception import BaseHTTPException from services.entities.knowledge_retrieval_inner import InnerKnowledgeRetrieveRequest, InnerKnowledgeRetrieveResponse from services.errors.knowledge_retrieval import ExternalKnowledgeRetrievalError, InnerKnowledgeRetrievalServiceError @@ -70,7 +71,8 @@ class InnerKnowledgeRetrieveApi(Resource): 500: "Unexpected knowledge retrieval failure", } ) - def post(self) -> dict[str, object]: + @with_session + def post(self, session: Session) -> dict[str, object]: """Validate the payload, run retrieval, and return workflow-style sources.""" try: payload = InnerKnowledgeRetrieveRequest.model_validate(inner_api_ns.payload or {}) @@ -82,7 +84,7 @@ class InnerKnowledgeRetrieveApi(Resource): ) from exc try: - response = InnerKnowledgeRetrievalService().retrieve(payload, session=db.session) + response = InnerKnowledgeRetrievalService().retrieve(payload, session=session) except InnerKnowledgeRetrievalServiceError as exc: raise InnerKnowledgeRetrievalHttpError( error_code=exc.error_code, diff --git a/api/controllers/inner_api/plugin/plugin.py b/api/controllers/inner_api/plugin/plugin.py index d385ca57379..e1de82c8b68 100644 --- a/api/controllers/inner_api/plugin/plugin.py +++ b/api/controllers/inner_api/plugin/plugin.py @@ -1,5 +1,7 @@ from flask_restx import Resource +from sqlalchemy.orm import Session +from controllers.console.app.wraps import with_session from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns from controllers.inner_api.plugin.wraps import get_user_tenant, plugin_data @@ -236,10 +238,12 @@ class PluginInvokeToolApi(Resource): 404: "Service not available", } ) - def post(self, user_model: Account | EndUser, tenant_model: Tenant, payload: RequestInvokeTool): + @with_session + def post(self, session: Session, user_model: Account | EndUser, tenant_model: Tenant, payload: RequestInvokeTool): def generator(): return PluginToolBackwardsInvocation.convert_to_event_stream( PluginToolBackwardsInvocation.invoke_tool( + session=session, tenant_id=tenant_model.id, user_id=user_model.id, tool_type=ToolProviderType.value_of(payload.tool_type), @@ -334,8 +338,10 @@ class PluginInvokeAppApi(Resource): 404: "Service not available", } ) - def post(self, user_model: Account | EndUser, tenant_model: Tenant, payload: RequestInvokeApp): + @with_session + def post(self, session: Session, user_model: Account | EndUser, tenant_model: Tenant, payload: RequestInvokeApp): response = PluginAppBackwardsInvocation.invoke_app( + session=session, app_id=payload.app_id, user_id=user_model.id, tenant_id=tenant_model.id, diff --git a/api/controllers/mcp/mcp.py b/api/controllers/mcp/mcp.py index cda6b915018..3830c9585a5 100644 --- a/api/controllers/mcp/mcp.py +++ b/api/controllers/mcp/mcp.py @@ -236,7 +236,6 @@ class MCPAppApi(Resource): if not end_user and isinstance(mcp_request.root, mcp_types.InitializeRequest): client_info = mcp_request.root.params.clientInfo client_name = f"{client_info.name}@{client_info.version}" - with sessionmaker(db.engine, expire_on_commit=False).begin() as create_session: - end_user = self._create_end_user(client_name, app.tenant_id, app.id, mcp_server.id, create_session) + end_user = self._create_end_user(client_name, app.tenant_id, app.id, mcp_server.id, session) - return handle_mcp_request(app, mcp_request, user_input_form, mcp_server, end_user, request_id) + return handle_mcp_request(session, app, mcp_request, user_input_form, mcp_server, end_user, request_id) diff --git a/api/controllers/openapi/app_run.py b/api/controllers/openapi/app_run.py index a22534ae82c..6074c7c0e02 100644 --- a/api/controllers/openapi/app_run.py +++ b/api/controllers/openapi/app_run.py @@ -8,6 +8,7 @@ from contextlib import contextmanager from typing import Any from flask_restx import Resource +from sqlalchemy.orm import Session from werkzeug.exceptions import ( BadRequest, HTTPException, @@ -20,6 +21,7 @@ from werkzeug.exceptions import ( import services from controllers.common.fields import EventStreamResponse from controllers.common.wraps import RBACPermission, RBACResourceScope +from controllers.console.app.wraps import with_session from controllers.openapi import openapi_ns from controllers.openapi._audit import emit_app_run from controllers.openapi._contract import accepts, returns @@ -92,8 +94,9 @@ def _translate_service_errors() -> Generator[None, None, None]: raise CompletionRequestError(e.description) -def _generate(app: App, caller: Any, args: dict[str, Any], streaming: bool): +def _generate(app: App, caller: Any, args: dict[str, Any], streaming: bool, session: Session): return AppGenerateService.generate( + session=session, app_model=app, user=caller, args=args, @@ -102,31 +105,31 @@ def _generate(app: App, caller: Any, args: dict[str, Any], streaming: bool): ) -def _run_chat(app: App, caller: Any, payload: AppRunRequest): +def _run_chat(app: App, caller: Any, payload: AppRunRequest, session: Session): if not payload.query or not payload.query.strip(): raise UnprocessableEntity("query_required_for_chat") args = payload.model_dump(exclude_none=True) with _translate_service_errors(): - return _generate(app, caller, args, streaming=True) + return _generate(app, caller, args, streaming=True, session=session) -def _run_completion(app: App, caller: Any, payload: AppRunRequest): +def _run_completion(app: App, caller: Any, payload: AppRunRequest, session: Session): args = payload.model_dump(exclude_none=True) args["auto_generate_name"] = False args.setdefault("query", "") with _translate_service_errors(): - return _generate(app, caller, args, streaming=True) + return _generate(app, caller, args, streaming=True, session=session) -def _run_workflow(app: App, caller: Any, payload: AppRunRequest): +def _run_workflow(app: App, caller: Any, payload: AppRunRequest, session: Session): if payload.query is not None: raise UnprocessableEntity("query_not_supported_for_workflow") args = payload.model_dump(exclude={"query", "conversation_id", "auto_generate_name"}, exclude_none=True) with _translate_service_errors(): - return _generate(app, caller, args, streaming=True) + return _generate(app, caller, args, streaming=True, session=session) -_DISPATCH: dict[AppMode, Callable[[App, Any, AppRunRequest], Any]] = { +_DISPATCH: dict[AppMode, Callable[[App, Any, AppRunRequest, Session], Any]] = { AppMode.CHAT: _run_chat, AppMode.AGENT_CHAT: _run_chat, AppMode.ADVANCED_CHAT: _run_chat, @@ -143,7 +146,8 @@ class AppRunApi(Resource): ) @openapi_ns.response(200, "Run result (SSE stream)", openapi_ns.models[EventStreamResponse.__name__]) @accepts(body=AppRunRequest) - def post(self, app_id: str, *, auth_data: AuthData, body: AppRunRequest): + @with_session + def post(self, session: Session, app_id: str, *, auth_data: AuthData, body: AppRunRequest): app_model, caller, caller_kind = auth_data.require_app_context() handler = _DISPATCH.get(app_model.mode) @@ -151,7 +155,7 @@ class AppRunApi(Resource): raise UnprocessableEntity("mode_not_runnable") try: - stream_obj = handler(app_model, caller, body) + stream_obj = handler(app_model, caller, body, session) except HTTPException: raise except Exception: diff --git a/api/controllers/service_api/app/completion.py b/api/controllers/service_api/app/completion.py index 1468f3d776f..900d46a0f0f 100644 --- a/api/controllers/service_api/app/completion.py +++ b/api/controllers/service_api/app/completion.py @@ -6,11 +6,13 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from pydantic.json_schema import SkipJsonSchema +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, InternalServerError, NotFound import services from controllers.common.fields import GeneratedAppResponse, SimpleResultResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.service_api import service_api_ns from controllers.service_api.app.error import ( AppUnavailableError, @@ -203,7 +205,8 @@ class CompletionApi(Resource): service_api_ns.models[GeneratedAppResponse.__name__], ) @validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON, required=True)) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): """Create a completion for the given prompt. This endpoint generates a completion based on the provided inputs and query. @@ -229,6 +232,7 @@ class CompletionApi(Resource): try: response = AppGenerateService.generate( + session=session, app_model=app_model, user=end_user, args=args, @@ -352,7 +356,8 @@ class ChatApi(Resource): service_api_ns.models[GeneratedAppResponse.__name__], ) @validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON, required=True)) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): """Send a message in a chat conversation. This endpoint handles chat messages for chat, agent chat, and advanced chat applications. @@ -376,7 +381,12 @@ class ChatApi(Resource): try: response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=streaming + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=streaming, ) return helper.compact_generate_response(response) diff --git a/api/controllers/service_api/app/workflow.py b/api/controllers/service_api/app/workflow.py index 091b79fefbd..1d961051acd 100644 --- a/api/controllers/service_api/app/workflow.py +++ b/api/controllers/service_api/app/workflow.py @@ -8,12 +8,13 @@ from flask import request from flask_restx import Resource, fields from pydantic import BaseModel, Field, field_validator from pydantic.json_schema import SkipJsonSchema -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, InternalServerError, NotFound from controllers.common.controller_schemas import WorkflowRunPayload as WorkflowRunPayloadBase from controllers.common.fields import GeneratedAppResponse, SimpleResultResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.service_api import service_api_ns from controllers.service_api.app.error import ( CompletionRequestError, @@ -341,7 +342,8 @@ class WorkflowRunApi(Resource): service_api_ns.models[GeneratedAppResponse.__name__], ) @validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON, required=True)) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): """Execute a workflow. Runs a workflow with the provided inputs and returns the results. @@ -363,7 +365,12 @@ class WorkflowRunApi(Resource): try: response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=streaming + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=streaming, ) return helper.compact_generate_response(response) @@ -448,7 +455,8 @@ class WorkflowRunByIdApi(Resource): service_api_ns.models[GeneratedAppResponse.__name__], ) @validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON, required=True)) - def post(self, app_model: App, end_user: EndUser, workflow_id: str): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser, workflow_id: str): """Run specific workflow by ID. Executes a specific workflow version identified by its ID. @@ -473,7 +481,12 @@ class WorkflowRunByIdApi(Resource): try: response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=streaming + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=streaming, ) return helper.compact_generate_response(response) diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index 0d52b7a25e6..56836f56895 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -12,6 +12,7 @@ from pydantic import ( field_validator, model_validator, ) +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -23,6 +24,7 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) +from controllers.console.app.wraps 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 @@ -481,7 +483,8 @@ class DatasetListApi(DatasetApiResource): service_api_ns.models[DatasetDetailResponse.__name__], ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id): + @with_session + def post(self, session: Session, tenant_id): """Resource for creating datasets.""" payload = DatasetCreatePayload.model_validate(service_api_ns.payload or {}) @@ -506,6 +509,7 @@ class DatasetListApi(DatasetApiResource): try: assert isinstance(current_user, Account) dataset = DatasetService.create_empty_dataset( + session=session, tenant_id=tenant_id, name=payload.name, description=payload.description, @@ -519,7 +523,6 @@ class DatasetListApi(DatasetApiResource): embedding_model_name=payload.embedding_model, retrieval_model=payload.retrieval_model, summary_index_setting=payload.summary_index_setting, - session=db.session, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -634,7 +637,8 @@ class DatasetApi(DatasetApiResource): service_api_ns.models[DatasetDetailWithPartialMembersResponse.__name__], ) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def patch(self, _, dataset_id: UUID): + @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) if dataset is None: @@ -680,7 +684,7 @@ class DatasetApi(DatasetApiResource): db.session, ) - dataset = DatasetService.update_dataset(dataset_id_str, update_data, current_user, db.session) + dataset = DatasetService.update_dataset(session, dataset_id_str, update_data, current_user) if dataset is None: raise NotFound("Dataset not found.") diff --git a/api/controllers/service_api/dataset/hit_testing.py b/api/controllers/service_api/dataset/hit_testing.py index 31881e86322..86a64829e22 100644 --- a/api/controllers/service_api/dataset/hit_testing.py +++ b/api/controllers/service_api/dataset/hit_testing.py @@ -1,6 +1,9 @@ from uuid import UUID +from sqlalchemy.orm import Session + from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.console.datasets.hit_testing_base import DatasetsHitTestingBase, HitTestingPayload from controllers.service_api import service_api_ns from controllers.service_api.wraps import DatasetApiResource, cloud_edition_billing_rate_limit_check @@ -51,15 +54,15 @@ class HitTestingApi(DatasetApiResource, DatasetsHitTestingBase): @service_api_ns.response(404, "Dataset not found") @service_api_ns.expect(service_api_ns.models[HitTestingPayload.__name__]) @cloud_edition_billing_rate_limit_check("knowledge", "dataset") - def post(self, tenant_id: str, dataset_id: UUID) -> dict[str, object]: + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID) -> dict[str, object]: """Perform hit testing on a dataset. Tests retrieval performance for the specified dataset. """ dataset_id_str = str(dataset_id) - dataset = self.get_and_validate_dataset(dataset_id_str) args = self.parse_args(service_api_ns.payload) self.hit_testing_args_check(args) - return dump_response(HitTestingResponse, self.perform_hit_testing(dataset, args)) + return dump_response(HitTestingResponse, self.perform_hit_testing(session, dataset, args)) diff --git a/api/controllers/web/completion.py b/api/controllers/web/completion.py index 7871b411c4b..2c852e208a5 100644 --- a/api/controllers/web/completion.py +++ b/api/controllers/web/completion.py @@ -2,11 +2,13 @@ import logging from typing import Any, Literal from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, InternalServerError, NotFound import services from controllers.common.fields import GeneratedAppResponse, SimpleResultResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.web import web_ns from controllers.web.error import ( AppUnavailableError, @@ -107,7 +109,8 @@ class CompletionApi(WebApiResource): } ) @web_ns.response(200, "Success", web_ns.models[GeneratedAppResponse.__name__]) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): if app_model.mode != AppMode.COMPLETION: raise NotCompletionAppError() @@ -119,7 +122,12 @@ class CompletionApi(WebApiResource): try: response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.WEB_APP, streaming=streaming + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.WEB_APP, + streaming=streaming, ) return helper.compact_generate_response(response) @@ -191,7 +199,8 @@ class ChatApi(WebApiResource): } ) @web_ns.response(200, "Success", web_ns.models[GeneratedAppResponse.__name__]) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): app_mode = AppMode.value_of(app_model.mode) if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT}: raise NotChatAppError() @@ -210,7 +219,12 @@ class ChatApi(WebApiResource): ) response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.WEB_APP, streaming=streaming + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.WEB_APP, + streaming=streaming, ) return helper.compact_generate_response(response) diff --git a/api/controllers/web/message.py b/api/controllers/web/message.py index 65ef02471a9..691eba05491 100644 --- a/api/controllers/web/message.py +++ b/api/controllers/web/message.py @@ -4,11 +4,13 @@ from uuid import UUID from flask import request from pydantic import BaseModel, Field, TypeAdapter +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError, NotFound from controllers.common.controller_schemas import MessageFeedbackPayload, MessageListQuery from controllers.common.fields import GeneratedAppResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.web import web_ns from controllers.web.error import ( AppMoreLikeThisDisabledError, @@ -162,7 +164,8 @@ class MessageMoreLikeThisApi(WebApiResource): } ) @web_ns.response(200, "Success", web_ns.models[GeneratedAppResponse.__name__]) - def get(self, app_model: App, end_user: EndUser, message_id: UUID): + @with_session + def get(self, session: Session, app_model: App, end_user: EndUser, message_id: UUID): if app_model.mode != "completion": raise NotCompletionAppError() @@ -175,6 +178,7 @@ class MessageMoreLikeThisApi(WebApiResource): try: response = AppGenerateService.generate_more_like_this( + session=session, app_model=app_model, user=end_user, message_id=message_id_str, diff --git a/api/controllers/web/workflow.py b/api/controllers/web/workflow.py index 06d9c02fedc..6aa5b495946 100644 --- a/api/controllers/web/workflow.py +++ b/api/controllers/web/workflow.py @@ -1,10 +1,12 @@ import logging +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError from controllers.common.controller_schemas import WorkflowRunPayload from controllers.common.fields import GeneratedAppResponse, SimpleResultResponse from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.app.wraps import with_session from controllers.web import web_ns from controllers.web.error import ( CompletionRequestError, @@ -52,7 +54,8 @@ class WorkflowRunApi(WebApiResource): } ) @web_ns.response(200, "Success", web_ns.models[GeneratedAppResponse.__name__]) - def post(self, app_model: App, end_user: EndUser): + @with_session + def post(self, session: Session, app_model: App, end_user: EndUser): """ Run workflow """ @@ -65,7 +68,12 @@ class WorkflowRunApi(WebApiResource): try: response = AppGenerateService.generate( - app_model=app_model, user=end_user, args=args, invoke_from=InvokeFrom.WEB_APP, streaming=True + session=session, + app_model=app_model, + user=end_user, + args=args, + invoke_from=InvokeFrom.WEB_APP, + streaming=True, ) return helper.compact_generate_response(response) diff --git a/api/core/agent/base_agent_runner.py b/api/core/agent/base_agent_runner.py index 55a31563d69..850858bfcef 100644 --- a/api/core/agent/base_agent_runner.py +++ b/api/core/agent/base_agent_runner.py @@ -5,6 +5,7 @@ from decimal import Decimal from typing import Union, cast from sqlalchemy import func, select +from sqlalchemy.orm import Session from core.agent.entities import AgentEntity, AgentToolEntity from core.app.app_config.features.file_upload.manager import FileUploadConfigManager @@ -51,6 +52,7 @@ class BaseAgentRunner(AppRunner): def __init__( self, *, + session: Session, tenant_id: str, application_generate_entity: AgentChatAppGenerateEntity, conversation: Conversation, @@ -88,6 +90,7 @@ class BaseAgentRunner(AppRunner): invoke_from=self.application_generate_entity.invoke_from, ) self.dataset_tools = DatasetRetrieverTool.get_dataset_tools( + session=session, tenant_id=tenant_id, dataset_ids=app_config.dataset.dataset_ids if app_config.dataset else [], retrieve_config=app_config.dataset.retrieve_config if app_config.dataset else None, diff --git a/api/core/agent/cot_agent_runner.py b/api/core/agent/cot_agent_runner.py index 9c9fa1092f6..8fcf42ce67d 100644 --- a/api/core/agent/cot_agent_runner.py +++ b/api/core/agent/cot_agent_runner.py @@ -4,6 +4,8 @@ from abc import ABC, abstractmethod from collections.abc import Generator, Mapping, Sequence from typing import Any, TypedDict +from sqlalchemy.orm import Session + from core.agent.base_agent_runner import BaseAgentRunner from core.agent.entities import AgentScratchpadUnit from core.agent.errors import AgentMaxIterationError @@ -46,6 +48,7 @@ class CotAgentRunner(BaseAgentRunner, ABC): def run( self, + session: Session, message: Message, query: str, inputs: Mapping[str, str], @@ -221,6 +224,7 @@ class CotAgentRunner(BaseAgentRunner, ABC): function_call_state = True # action is tool call, invoke tool tool_invoke_response, tool_invoke_meta = self._handle_invoke_action( + session=session, action=scratchpad.action, tool_instances=tool_instances, message_file_ids=message_file_ids, @@ -287,6 +291,7 @@ class CotAgentRunner(BaseAgentRunner, ABC): def _handle_invoke_action( self, + session: Session, action: AgentScratchpadUnit.Action, tool_instances: Mapping[str, Tool], message_file_ids: list[str], @@ -317,6 +322,7 @@ class CotAgentRunner(BaseAgentRunner, ABC): # invoke tool tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke( + session=session, tool=tool_instance, tool_parameters=tool_call_args, user_id=self.user_id, diff --git a/api/core/agent/fc_agent_runner.py b/api/core/agent/fc_agent_runner.py index 29de0b8b1c4..9db5fa08f75 100644 --- a/api/core/agent/fc_agent_runner.py +++ b/api/core/agent/fc_agent_runner.py @@ -4,6 +4,8 @@ from collections.abc import Generator from copy import deepcopy from typing import Any, Union +from sqlalchemy.orm import Session + from core.agent.base_agent_runner import BaseAgentRunner from core.agent.errors import AgentMaxIterationError from core.app.apps.base_app_queue_manager import PublishFrom @@ -32,7 +34,9 @@ logger = logging.getLogger(__name__) class FunctionCallAgentRunner(BaseAgentRunner): - def run(self, message: Message, query: str, **kwargs: Any) -> Generator[LLMResultChunk, None, None]: + def run( + self, session: Session, message: Message, query: str, **kwargs: Any + ) -> Generator[LLMResultChunk, None, None]: """ Run FunctionCall agent application """ @@ -167,7 +171,7 @@ class FunctionCallAgentRunner(BaseAgentRunner): for content in result.message.content: response += content.data else: - response += str(result.message.content) + response += result.message.content if not result.message.content: result.message.content = "" @@ -238,6 +242,7 @@ class FunctionCallAgentRunner(BaseAgentRunner): else: # invoke tool tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke( + session=session, tool=tool_instance, tool_parameters=tool_call_args, user_id=self.user_id, diff --git a/api/core/app/apps/agent_chat/app_generator.py b/api/core/app/apps/agent_chat/app_generator.py index 9ad724cc89c..d640bcdc863 100644 --- a/api/core/app/apps/agent_chat/app_generator.py +++ b/api/core/app/apps/agent_chat/app_generator.py @@ -7,6 +7,7 @@ from typing import Any, Literal, overload from flask import Flask, current_app from pydantic import ValidationError +from sqlalchemy.orm import Session from configs import dify_config from constants import UUID_NIL @@ -228,6 +229,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, + session: Session, context: contextvars.Context, application_generate_entity: AgentChatAppGenerateEntity, queue_manager: AppQueueManager, @@ -253,6 +255,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator): # chatbot app runner = AgentChatAppRunner() runner.run( + session=session, application_generate_entity=application_generate_entity, queue_manager=queue_manager, conversation=conversation, diff --git a/api/core/app/apps/agent_chat/app_runner.py b/api/core/app/apps/agent_chat/app_runner.py index 5f9c75129b5..6bbc20388dd 100644 --- a/api/core/app/apps/agent_chat/app_runner.py +++ b/api/core/app/apps/agent_chat/app_runner.py @@ -2,6 +2,7 @@ import logging from typing import cast from sqlalchemy import select +from sqlalchemy.orm import Session from core.agent.cot_chat_agent_runner import CotChatAgentRunner from core.agent.cot_completion_agent_runner import CotCompletionAgentRunner @@ -31,6 +32,7 @@ class AgentChatAppRunner(AppRunner): def run( self, + session: Session, application_generate_entity: AgentChatAppGenerateEntity, queue_manager: AppQueueManager, conversation: Conversation, @@ -217,6 +219,7 @@ class AgentChatAppRunner(AppRunner): raise ValueError(f"Invalid agent strategy: {agent_entity.strategy}") runner = runner_cls( + session=session, tenant_id=app_config.tenant_id, application_generate_entity=application_generate_entity, conversation=conversation_result, @@ -232,6 +235,7 @@ class AgentChatAppRunner(AppRunner): ) invoke_result = runner.run( + session=session, message=message, query=query, inputs=inputs, diff --git a/api/core/app/apps/chat/app_generator.py b/api/core/app/apps/chat/app_generator.py index 4b8701da8e6..4873168b885 100644 --- a/api/core/app/apps/chat/app_generator.py +++ b/api/core/app/apps/chat/app_generator.py @@ -7,6 +7,7 @@ from typing import Any, Literal, overload from flask import Flask, copy_current_request_context, current_app from pydantic import ValidationError +from sqlalchemy.orm import Session from configs import dify_config from constants import UUID_NIL @@ -36,6 +37,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -46,6 +48,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -56,6 +59,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -65,6 +69,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -197,6 +202,7 @@ 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, @@ -223,6 +229,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, + session: Session, application_generate_entity: ChatAppGenerateEntity, queue_manager: AppQueueManager, conversation_id: str, @@ -246,6 +253,7 @@ class ChatAppGenerator(MessageBasedAppGenerator): # chatbot app runner = ChatAppRunner() runner.run( + session=session, application_generate_entity=application_generate_entity, queue_manager=queue_manager, conversation=conversation, diff --git a/api/core/app/apps/chat/app_runner.py b/api/core/app/apps/chat/app_runner.py index 9c2eaf60dc7..1e3be128296 100644 --- a/api/core/app/apps/chat/app_runner.py +++ b/api/core/app/apps/chat/app_runner.py @@ -2,6 +2,7 @@ import logging from typing import cast from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.base_app_runner import AppRunner @@ -31,6 +32,7 @@ class ChatAppRunner(AppRunner): def run( self, + session: Session, application_generate_entity: ChatAppGenerateEntity, queue_manager: AppQueueManager, conversation: Conversation, @@ -163,6 +165,7 @@ class ChatAppRunner(AppRunner): dataset_retrieval = DatasetRetrieval(application_generate_entity) context, retrieved_files = dataset_retrieval.retrieve( + session=session, app_id=app_record.id, user_id=application_generate_entity.user_id, tenant_id=app_record.tenant_id, diff --git a/api/core/app/apps/completion/app_generator.py b/api/core/app/apps/completion/app_generator.py index 9f29b8df29d..5096f323354 100644 --- a/api/core/app/apps/completion/app_generator.py +++ b/api/core/app/apps/completion/app_generator.py @@ -8,6 +8,7 @@ from typing import Any, Literal, overload from flask import Flask, copy_current_request_context, current_app from pydantic import ValidationError from sqlalchemy import select +from sqlalchemy.orm import Session from configs import dify_config from core.app.app_config.easy_ui_based_app.model_config.converter import ModelConfigConverter @@ -36,6 +37,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -46,6 +48,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -56,6 +59,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): @overload def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -65,6 +69,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def generate( self, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -175,6 +180,7 @@ 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, @@ -200,6 +206,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def _generate_worker( self, flask_app: Flask, + session: Session, application_generate_entity: CompletionAppGenerateEntity, queue_manager: AppQueueManager, message_id: str, @@ -220,6 +227,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): # chatbot app runner = CompletionAppRunner() runner.run( + session=session, application_generate_entity=application_generate_entity, queue_manager=queue_manager, message=message, @@ -245,6 +253,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): def generate_more_like_this( self, + session: Session, app_model: App, message_id: str, user: Account | EndUser, @@ -343,6 +352,7 @@ 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 38ef672ae22..3be70c860f9 100644 --- a/api/core/app/apps/completion/app_runner.py +++ b/api/core/app/apps/completion/app_runner.py @@ -2,6 +2,7 @@ import logging from typing import cast from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.apps.base_app_queue_manager import AppQueueManager from core.app.apps.base_app_runner import AppRunner @@ -28,7 +29,11 @@ class CompletionAppRunner(AppRunner): """ def run( - self, application_generate_entity: CompletionAppGenerateEntity, queue_manager: AppQueueManager, message: Message + self, + session: Session, + application_generate_entity: CompletionAppGenerateEntity, + queue_manager: AppQueueManager, + message: Message, ): """ Run application @@ -123,6 +128,7 @@ class CompletionAppRunner(AppRunner): dataset_retrieval = DatasetRetrieval(application_generate_entity) context, retrieved_files = dataset_retrieval.retrieve( + session=session, app_id=app_record.id, user_id=application_generate_entity.user_id, tenant_id=app_record.tenant_id, diff --git a/api/core/mcp/server/streamable_http.py b/api/core/mcp/server/streamable_http.py index 08b4ed0e19c..3bb75e485a8 100644 --- a/api/core/mcp/server/streamable_http.py +++ b/api/core/mcp/server/streamable_http.py @@ -3,6 +3,8 @@ import logging from collections.abc import Mapping from typing import Any, NotRequired, TypedDict, cast +from sqlalchemy.orm import Session + from configs import dify_config from core.app.entities.app_invoke_entities import InvokeFrom from core.app.features.rate_limiting.rate_limit import RateLimitGenerator @@ -26,6 +28,7 @@ class ToolArgumentsDict(TypedDict): def handle_mcp_request( + session: Session, app: App, request: mcp_types.ClientRequest, user_input_form: list[VariableEntity], @@ -82,7 +85,7 @@ def handle_mcp_request( ) ) case mcp_types.CallToolRequest(): - return create_success_response(handle_call_tool(app, request, user_input_form, end_user)) + return create_success_response(handle_call_tool(session, app, request, user_input_form, end_user)) case mcp_types.PingRequest(): return create_success_response(handle_ping()) case _: @@ -137,6 +140,7 @@ def handle_list_tools( def handle_call_tool( + session: Session, app: App, request: mcp_types.ClientRequest, user_input_form: list[VariableEntity], @@ -150,6 +154,7 @@ def handle_call_tool( raise ValueError("End user not found") response = AppGenerateService.generate( + session, app, end_user, args, diff --git a/api/core/plugin/backwards_invocation/app.py b/api/core/plugin/backwards_invocation/app.py index d022b002f72..046be355daf 100644 --- a/api/core/plugin/backwards_invocation/app.py +++ b/api/core/plugin/backwards_invocation/app.py @@ -3,6 +3,7 @@ from collections.abc import Generator, Mapping from typing import Any, cast from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator @@ -60,6 +61,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): @classmethod def invoke_app( cls, + session: Session, app_id: str, user_id: str, tenant_id: str, @@ -85,20 +87,21 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): if not query: raise ValueError("missing query") - return cls.invoke_chat_app(app, user, conversation_id, query, stream, inputs, files) + return cls.invoke_chat_app(session, app, user, conversation_id, query, stream, inputs, files) 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(app, user, stream, inputs, files) + return cls.invoke_completion_app(session, app, user, stream, inputs, files) case _: raise ValueError("unexpected app type") @classmethod def invoke_chat_app( cls, + session: Session, app: App, user: Account | EndUser, conversation_id: str, @@ -151,6 +154,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): ) case AppMode.CHAT: return ChatAppGenerator().generate( + session=session, app_model=app, user=user, args={ @@ -197,6 +201,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): @classmethod def invoke_completion_app( cls, + session: Session, app: App, user: EndUser | Account, stream: bool, @@ -207,6 +212,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation): invoke completion app """ return CompletionAppGenerator().generate( + session=session, app_model=app, user=user, args={"inputs": inputs, "files": files}, diff --git a/api/core/plugin/backwards_invocation/tool.py b/api/core/plugin/backwards_invocation/tool.py index 05854942691..0a296dfea10 100644 --- a/api/core/plugin/backwards_invocation/tool.py +++ b/api/core/plugin/backwards_invocation/tool.py @@ -1,6 +1,8 @@ from collections.abc import Generator from typing import Any +from sqlalchemy.orm import Session + from core.callback_handler.workflow_tool_callback_handler import DifyWorkflowCallbackHandler from core.plugin.backwards_invocation.base import BaseBackwardsInvocation from core.tools.entities.tool_entities import ToolInvokeMessage, ToolProviderType @@ -17,6 +19,7 @@ class PluginToolBackwardsInvocation(BaseBackwardsInvocation): @classmethod def invoke_tool( cls, + session: Session, tenant_id: str, user_id: str, tool_type: ToolProviderType, @@ -40,7 +43,7 @@ class PluginToolBackwardsInvocation(BaseBackwardsInvocation): credential_id=credential_id, ) response = ToolEngine.generic_invoke( - tool_runtime, tool_parameters, user_id, DifyWorkflowCallbackHandler(), workflow_call_depth=1 + session, tool_runtime, tool_parameters, user_id, DifyWorkflowCallbackHandler(), workflow_call_depth=1 ) response = ToolFileMessageTransformer.transform_tool_invoke_messages( diff --git a/api/core/rag/datasource/retrieval_service.py b/api/core/rag/datasource/retrieval_service.py index 85eb06045ac..50381f5e75c 100644 --- a/api/core/rag/datasource/retrieval_service.py +++ b/api/core/rag/datasource/retrieval_service.py @@ -192,6 +192,7 @@ class RetrievalService: @classmethod def external_retrieve( cls, + session: Session, dataset_id: str, query: str, external_retrieval_model: dict[str, Any] | None = None, @@ -207,6 +208,7 @@ class RetrievalService: else None ) all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( + session, dataset.tenant_id, dataset_id, query, diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index c8f1210ee8c..96c6e57e7b0 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -10,7 +10,7 @@ from typing import Any, Union, cast from flask import Flask, current_app from sqlalchemy import and_, func, literal, or_, select, update -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.app.app_config.entities import ( DatasetEntity, @@ -116,7 +116,7 @@ class DatasetRetrieval: else: self._llm_usage = self._llm_usage.plus(usage) - def knowledge_retrieval(self, request: KnowledgeRetrievalRequest) -> list[Source]: + def knowledge_retrieval(self, session: Session, request: KnowledgeRetrievalRequest) -> list[Source]: self._check_knowledge_rate_limit(request.tenant_id) available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids) available_datasets_ids = [i.id for i in available_datasets] @@ -145,6 +145,7 @@ class DatasetRetrieval: query = request.query if request.query is not None else "" metadata_filter_document_ids, metadata_condition = self.get_metadata_filter_condition( + session=session, dataset_ids=available_datasets_ids, query=query, tenant_id=request.tenant_id, @@ -214,6 +215,7 @@ class DatasetRetrieval: stop=stop, ) all_documents = self.single_retrieve( + session, request.app_id, request.tenant_id, request.user_id, @@ -350,6 +352,7 @@ class DatasetRetrieval: def retrieve( self, + session: Session, app_id: str, user_id: str, tenant_id: str, @@ -419,6 +422,7 @@ class DatasetRetrieval: inputs = {} available_datasets_ids = [dataset.id for dataset in available_datasets] metadata_filter_document_ids, metadata_condition = self.get_metadata_filter_condition( + session, available_datasets_ids, query, tenant_id, @@ -433,6 +437,7 @@ class DatasetRetrieval: user_from = "account" if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER} else "end_user" if retrieve_config.retrieve_strategy == DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE: all_documents = self.single_retrieve( + session, app_id, tenant_id, user_id, @@ -509,7 +514,7 @@ class DatasetRetrieval: ) ) if vision_enabled: - attachments_with_bindings = db.session.execute( + attachments_with_bindings = session.execute( select(SegmentAttachmentBinding, UploadFile) .join(UploadFile, UploadFile.id == SegmentAttachmentBinding.attachment_id) .where( @@ -545,11 +550,11 @@ class DatasetRetrieval: DatasetDocument.enabled == True, DatasetDocument.archived == False, ) - documents = db.session.execute(dataset_document_stmt).scalars().all() # type: ignore + documents = session.execute(dataset_document_stmt).scalars().all() # type: ignore dataset_stmt = select(Dataset).where( Dataset.id.in_(dataset_ids), ) - datasets = db.session.execute(dataset_stmt).scalars().all() # type: ignore + datasets = session.execute(dataset_stmt).scalars().all() # type: ignore dataset_map = {i.id: i for i in datasets} document_map = {i.id: i for i in documents} for record in records: @@ -589,13 +594,12 @@ class DatasetRetrieval: hit_callback.return_retriever_resource_info(retrieval_resource_list) if document_context_list: document_context_list = sorted(document_context_list, key=lambda x: x.score or 0.0, reverse=True) - return str( - "\n".join([document_context.content for document_context in document_context_list]) - ), context_files + return "\n".join([document_context.content for document_context in document_context_list]), context_files return "", context_files def single_retrieve( self, + session: Session, app_id: str, tenant_id: str, user_id: str, @@ -643,11 +647,12 @@ class DatasetRetrieval: if dataset_id: # get retrieval model config dataset_stmt = select(Dataset).where(Dataset.id == dataset_id) - selected_dataset = db.session.scalar(dataset_stmt) + selected_dataset = session.scalar(dataset_stmt) if selected_dataset: results = [] if selected_dataset.provider == "external": external_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( + session=session, tenant_id=selected_dataset.tenant_id, dataset_id=dataset_id, query=query, @@ -1076,6 +1081,7 @@ class DatasetRetrieval: def _retriever( self, flask_app: Flask, + session: Session, dataset_id: str, query: str, top_k: int, @@ -1086,13 +1092,14 @@ class DatasetRetrieval: ): with flask_app.app_context(): dataset_stmt = select(Dataset).where(Dataset.id == dataset_id) - dataset = db.session.scalar(dataset_stmt) + dataset = session.scalar(dataset_stmt) if not dataset: return [] if dataset.provider == "external" and query: external_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( + session=session, tenant_id=dataset.tenant_id, dataset_id=dataset_id, query=query, @@ -1154,6 +1161,7 @@ class DatasetRetrieval: def to_dataset_retriever_tool( self, + session: Session, tenant_id: str, dataset_ids: list[str], retrieve_config: DatasetRetrieveConfigEntity, @@ -1177,7 +1185,7 @@ class DatasetRetrieval: for dataset_id in dataset_ids: # get dataset from dataset id dataset_stmt = select(Dataset).where(Dataset.tenant_id == tenant_id, Dataset.id == dataset_id) - dataset = db.session.scalar(dataset_stmt) + dataset = session.scalar(dataset_stmt) # pass if dataset is not available if not dataset: @@ -1349,6 +1357,7 @@ class DatasetRetrieval: def get_metadata_filter_condition( self, + session: Session, dataset_ids: list[str], query: str, tenant_id: str, @@ -1370,7 +1379,7 @@ class DatasetRetrieval: return None, None elif metadata_filtering_mode == "automatic": automatic_metadata_filters = self._automatic_metadata_filter_func( - dataset_ids, query, tenant_id, user_id, metadata_model_config + session, dataset_ids, query, tenant_id, user_id, metadata_model_config ) if automatic_metadata_filters: conditions = [] @@ -1429,7 +1438,7 @@ class DatasetRetrieval: document_query = document_query.where(and_(*filters)) else: document_query = document_query.where(or_(*filters)) - documents = db.session.scalars(document_query).all() + documents = session.scalars(document_query).all() # group by dataset_id metadata_filter_document_ids = defaultdict(list) if documents else None # type: ignore for document in documents: @@ -1451,11 +1460,17 @@ class DatasetRetrieval: return output def _automatic_metadata_filter_func( - self, dataset_ids: list[str], query: str, tenant_id: str, user_id: str, metadata_model_config: ModelConfig + self, + session: Session, + dataset_ids: list[str], + query: str, + tenant_id: str, + user_id: str, + metadata_model_config: ModelConfig, ) -> list[dict[str, Any]] | None: # get all metadata field metadata_stmt = select(DatasetMetadata).where(DatasetMetadata.dataset_id.in_(dataset_ids)) - metadata_fields = db.session.scalars(metadata_stmt).all() + metadata_fields = session.scalars(metadata_stmt).all() all_metadata_fields = [metadata_field.name for metadata_field in metadata_fields] # get metadata model config if metadata_model_config is None: diff --git a/api/core/tools/__base/tool.py b/api/core/tools/__base/tool.py index b219ba49575..b16f80169fc 100644 --- a/api/core/tools/__base/tool.py +++ b/api/core/tools/__base/tool.py @@ -5,6 +5,8 @@ from collections.abc import Generator from copy import deepcopy from typing import TYPE_CHECKING, Any +from sqlalchemy.orm import Session + if TYPE_CHECKING: # pragma: no cover from models.model import File @@ -46,6 +48,7 @@ class Tool(ABC): def invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, @@ -59,6 +62,7 @@ class Tool(ABC): tool_parameters = self._transform_tool_parameters_type(tool_parameters) result = self._invoke( + session=session, user_id=user_id, tool_parameters=tool_parameters, conversation_id=conversation_id, @@ -97,6 +101,7 @@ class Tool(ABC): @abstractmethod def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/audio/tools/asr.py b/api/core/tools/builtin_tool/providers/audio/tools/asr.py index 0c3047244c9..8d5959d2d1b 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/asr.py +++ b/api/core/tools/builtin_tool/providers/audio/tools/asr.py @@ -2,6 +2,8 @@ import io from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.model_manager import ModelManager from core.plugin.entities.parameters import PluginParameterOption from core.tools.builtin_tool.tool import BuiltinTool @@ -17,6 +19,7 @@ class ASRTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/audio/tools/tts.py b/api/core/tools/builtin_tool/providers/audio/tools/tts.py index db653916101..37c78b61348 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/tts.py +++ b/api/core/tools/builtin_tool/providers/audio/tools/tts.py @@ -2,6 +2,8 @@ import io from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.model_manager import ModelManager from core.plugin.entities.parameters import PluginParameterOption from core.tools.builtin_tool.tool import BuiltinTool @@ -15,6 +17,7 @@ class TTSTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/code/tools/simple_code.py b/api/core/tools/builtin_tool/providers/code/tools/simple_code.py index dd041ef5eb9..78a01115b03 100644 --- a/api/core/tools/builtin_tool/providers/code/tools/simple_code.py +++ b/api/core/tools/builtin_tool/providers/code/tools/simple_code.py @@ -1,6 +1,8 @@ from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.helper.code_executor.code_executor import CodeExecutor, CodeLanguage from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -11,6 +13,7 @@ class SimpleCode(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/time/tools/current_time.py b/api/core/tools/builtin_tool/providers/time/tools/current_time.py index 1164cc11c62..e9530da9a92 100644 --- a/api/core/tools/builtin_tool/providers/time/tools/current_time.py +++ b/api/core/tools/builtin_tool/providers/time/tools/current_time.py @@ -3,6 +3,7 @@ from datetime import UTC, datetime from typing import Any, override from pytz import timezone as pytz_timezone # type: ignore[import-untyped] +from sqlalchemy.orm import Session from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -12,6 +13,7 @@ class CurrentTimeTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/time/tools/localtime_to_timestamp.py b/api/core/tools/builtin_tool/providers/time/tools/localtime_to_timestamp.py index 57363349458..0dfd4a03aa7 100644 --- a/api/core/tools/builtin_tool/providers/time/tools/localtime_to_timestamp.py +++ b/api/core/tools/builtin_tool/providers/time/tools/localtime_to_timestamp.py @@ -3,6 +3,7 @@ from datetime import datetime, tzinfo from typing import Any, cast, override import pytz # type: ignore[import-untyped] +from sqlalchemy.orm import Session from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -13,6 +14,7 @@ class LocaltimeToTimestampTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/time/tools/timestamp_to_localtime.py b/api/core/tools/builtin_tool/providers/time/tools/timestamp_to_localtime.py index f1bf1471444..489f450e6f4 100644 --- a/api/core/tools/builtin_tool/providers/time/tools/timestamp_to_localtime.py +++ b/api/core/tools/builtin_tool/providers/time/tools/timestamp_to_localtime.py @@ -3,6 +3,7 @@ from datetime import datetime from typing import Any, override import pytz # type: ignore[import-untyped] +from sqlalchemy.orm import Session from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -13,6 +14,7 @@ class TimestampToLocaltimeTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/time/tools/timezone_conversion.py b/api/core/tools/builtin_tool/providers/time/tools/timezone_conversion.py index 6f00e0009e7..6d9eb9e57f8 100644 --- a/api/core/tools/builtin_tool/providers/time/tools/timezone_conversion.py +++ b/api/core/tools/builtin_tool/providers/time/tools/timezone_conversion.py @@ -3,6 +3,7 @@ from datetime import datetime from typing import Any, override import pytz # type: ignore[import-untyped] +from sqlalchemy.orm import Session from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -13,6 +14,7 @@ class TimezoneConversionTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/time/tools/weekday.py b/api/core/tools/builtin_tool/providers/time/tools/weekday.py index 6a88538afe5..6920feae700 100644 --- a/api/core/tools/builtin_tool/providers/time/tools/weekday.py +++ b/api/core/tools/builtin_tool/providers/time/tools/weekday.py @@ -3,6 +3,8 @@ from collections.abc import Generator from datetime import datetime from typing import Any, override +from sqlalchemy.orm import Session + from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage @@ -11,6 +13,7 @@ class WeekdayTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/builtin_tool/providers/webscraper/tools/webscraper.py b/api/core/tools/builtin_tool/providers/webscraper/tools/webscraper.py index 44c6f720d0e..4c29ef3145b 100644 --- a/api/core/tools/builtin_tool/providers/webscraper/tools/webscraper.py +++ b/api/core/tools/builtin_tool/providers/webscraper/tools/webscraper.py @@ -1,6 +1,8 @@ from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.tools.builtin_tool.tool import BuiltinTool from core.tools.entities.tool_entities import ToolInvokeMessage from core.tools.errors import ToolInvokeError @@ -11,6 +13,7 @@ class WebscraperTool(BuiltinTool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/custom_tool/tool.py b/api/core/tools/custom_tool/tool.py index 3edb04d7b94..38a13609dbf 100644 --- a/api/core/tools/custom_tool/tool.py +++ b/api/core/tools/custom_tool/tool.py @@ -6,6 +6,7 @@ from typing import Any, Union, override from urllib.parse import urlencode import httpx +from sqlalchemy.orm import Session from core.helper import ssrf_proxy from core.tools.__base.tool import Tool @@ -379,6 +380,7 @@ class ApiTool(Tool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/mcp_tool/tool.py b/api/core/tools/mcp_tool/tool.py index 195acd6e1ad..07c9ff63b30 100644 --- a/api/core/tools/mcp_tool/tool.py +++ b/api/core/tools/mcp_tool/tool.py @@ -6,6 +6,8 @@ import logging from collections.abc import Generator, Mapping from typing import Any, cast, override +from sqlalchemy.orm import Session + from configs import dify_config from core.entities.mcp_provider import IdentityMode from core.mcp.auth_client import MCPClientWithAuthRetry @@ -65,6 +67,7 @@ class MCPTool(Tool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/plugin_tool/tool.py b/api/core/tools/plugin_tool/tool.py index ac17d542bc0..27362632c3a 100644 --- a/api/core/tools/plugin_tool/tool.py +++ b/api/core/tools/plugin_tool/tool.py @@ -3,6 +3,8 @@ from __future__ import annotations from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.plugin.impl.tool import PluginToolManager from core.plugin.utils.converter import convert_parameters_to_plugin_format from core.tools.__base.tool import Tool @@ -27,6 +29,7 @@ class PluginTool(Tool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/tools/tool_engine.py b/api/core/tools/tool_engine.py index b3cfd55265e..d9761eb4dbd 100644 --- a/api/core/tools/tool_engine.py +++ b/api/core/tools/tool_engine.py @@ -7,7 +7,7 @@ from datetime import UTC, datetime from mimetypes import guess_type from typing import Any, Union, cast -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from yarl import URL from core.app.entities.app_invoke_entities import InvokeFrom @@ -47,6 +47,7 @@ class ToolEngine: @staticmethod def agent_invoke( + session: Session, tool: Tool, tool_parameters: Union[str, dict[str, Any]], user_id: str, @@ -82,7 +83,7 @@ class ToolEngine: # hit the callback handler agent_tool_callback.on_tool_start(tool_name=tool.entity.identity.name, tool_inputs=tool_parameters) - messages = ToolEngine._invoke(tool, tool_parameters, user_id, conversation_id, app_id, message_id) + messages = ToolEngine._invoke(session, tool, tool_parameters, user_id, conversation_id, app_id, message_id) invocation_meta_dict: dict[str, ToolInvokeMeta] = {} def message_callback( @@ -157,6 +158,7 @@ class ToolEngine: @staticmethod def generic_invoke( + session: Session, tool: Tool, tool_parameters: dict[str, Any], user_id: str, @@ -180,6 +182,7 @@ class ToolEngine: tool_parameters = {**tool.runtime.runtime_parameters, **tool_parameters} response = tool.invoke( + session=session, user_id=user_id, tool_parameters=tool_parameters, conversation_id=conversation_id, @@ -201,6 +204,7 @@ class ToolEngine: @staticmethod def _invoke( + session: Session, tool: Tool, tool_parameters: dict[str, Any], user_id: str, @@ -224,7 +228,7 @@ class ToolEngine: }, ) try: - yield from tool.invoke(user_id, tool_parameters, conversation_id, app_id, message_id) + yield from tool.invoke(session, user_id, tool_parameters, conversation_id, app_id, message_id) except Exception as e: meta.error = str(e) raise ToolEngineInvokeError(meta) 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 55a406482a3..a3afe659563 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 @@ -4,6 +4,7 @@ from typing import override from flask import Flask, current_app from pydantic import BaseModel, Field from sqlalchemy import select +from sqlalchemy.orm import Session from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler from core.model_manager import ModelManager @@ -48,7 +49,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): ) @override - def _run(self, query: str) -> str: + def _run(self, session: Session, query: str) -> str: threads = [] all_documents: list[RagDocument] = [] for dataset_id in self.dataset_ids: @@ -147,7 +148,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool): for hit_callback in self.hit_callbacks: hit_callback.return_retriever_resource_info(context_list) - return str("\n".join(document_context_list)) + return "\n".join(document_context_list) return "" def _retriever( diff --git a/api/core/tools/utils/dataset_retriever/dataset_retriever_base_tool.py b/api/core/tools/utils/dataset_retriever/dataset_retriever_base_tool.py index dd0b4bedcff..2d92761bc90 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_base_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_base_tool.py @@ -1,6 +1,7 @@ from abc import ABC, abstractmethod from pydantic import BaseModel, ConfigDict +from sqlalchemy.orm import Session from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler @@ -18,12 +19,12 @@ class DatasetRetrieverBaseTool(BaseModel, ABC): retriever_from: str model_config = ConfigDict(arbitrary_types_allowed=True) - def run(self, query: str) -> str: + def run(self, session: Session, query: str) -> str: """Use the tool.""" - return self._run(query) + return self._run(session, query) @abstractmethod - def _run(self, query: str) -> str: + def _run(self, session: Session, query: str) -> str: """Use the tool. Add run_manager: Optional[CallbackManagerForToolRun] = None 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 d9f99884d0c..247bd0705fc 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py @@ -2,6 +2,7 @@ from typing import Any, cast, override from pydantic import BaseModel, Field from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import DatasetRetrieveConfigEntity, ModelConfig from core.rag.datasource.retrieval_service import DefaultRetrievalModelDict, RetrievalService @@ -57,7 +58,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): ) @override - def _run(self, query: str) -> str: + 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) @@ -67,6 +68,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): hit_callback.on_query(query, dataset.id, db.session) dataset_retrieval = DatasetRetrieval() metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition( + session, [dataset.id], query, self.tenant_id, @@ -83,6 +85,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): if dataset.provider == "external": results: list[RetrievalDocument] = [] external_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( + session=session, tenant_id=dataset.tenant_id, dataset_id=dataset.id, query=query, diff --git a/api/core/tools/utils/dataset_retriever_tool.py b/api/core/tools/utils/dataset_retriever_tool.py index d34bfe0aa1e..cd4f2849df9 100644 --- a/api/core/tools/utils/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever_tool.py @@ -1,6 +1,8 @@ from collections.abc import Generator from typing import Any, override +from sqlalchemy.orm import Session + from core.app.app_config.entities import DatasetRetrieveConfigEntity from core.app.entities.app_invoke_entities import InvokeFrom from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler @@ -26,6 +28,7 @@ class DatasetRetrieverTool(Tool): @staticmethod def get_dataset_tools( + session: Session, tenant_id: str, dataset_ids: list[str], retrieve_config: DatasetRetrieveConfigEntity | None, @@ -51,6 +54,7 @@ class DatasetRetrieverTool(Tool): original_retriever_mode = retrieve_config.retrieve_strategy retrieve_config.retrieve_strategy = DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE retrieval_tools = feature.to_dataset_retriever_tool( + session=session, tenant_id=tenant_id, dataset_ids=dataset_ids, retrieve_config=retrieve_config, @@ -113,6 +117,7 @@ class DatasetRetrieverTool(Tool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, @@ -127,7 +132,7 @@ class DatasetRetrieverTool(Tool): yield self.create_text_message(text="please input query") else: # invoke dataset retriever tool - result = self.retrieval_tool.run(query=query) + result = self.retrieval_tool.run(session=session, query=query) yield self.create_text_message(text=result) def validate_credentials( diff --git a/api/core/tools/workflow_as_tool/tool.py b/api/core/tools/workflow_as_tool/tool.py index 97222f3cfae..3aded386594 100644 --- a/api/core/tools/workflow_as_tool/tool.py +++ b/api/core/tools/workflow_as_tool/tool.py @@ -6,6 +6,7 @@ from collections.abc import Generator, Mapping, Sequence from typing import Any, cast, override from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.file_access import DatabaseFileAccessController from core.db.session_factory import session_factory @@ -79,6 +80,7 @@ class WorkflowTool(Tool): @override def _invoke( self, + session: Session, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, diff --git a/api/core/workflow/node_runtime.py b/api/core/workflow/node_runtime.py index 9964f65d0b6..7d9b9d854b2 100644 --- a/api/core/workflow/node_runtime.py +++ b/api/core/workflow/node_runtime.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, Literal, cast, overload, override from pydantic import JsonValue from sqlalchemy import select -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext from core.app.file_access import ( @@ -15,6 +15,7 @@ from core.app.file_access import ( is_retriever_segment_access_granted, ) from core.callback_handler.workflow_tool_callback_handler import DifyWorkflowCallbackHandler +from core.db.session_factory import session_factory from core.helper.trace_id_helper import ParentTraceContext from core.llm_generator.output_parser.errors import OutputParserError from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output @@ -456,9 +457,14 @@ class _WorkflowToolRuntimeBinding: class DifyToolNodeRuntime(ToolNodeRuntimeProtocol): - def __init__(self, run_context: Mapping[str, Any] | DifyRunContext) -> None: + def __init__( + self, + run_context: Mapping[str, Any] | DifyRunContext, + session_maker: sessionmaker[Session] | None = None, + ) -> None: self._run_context = resolve_dify_run_context(run_context) self._file_reference_factory = DifyFileReferenceFactory(self._run_context) + self._session_maker = session_maker @property def file_reference_factory(self) -> FileReferenceFactoryProtocol: @@ -556,27 +562,28 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol): tool.clear_trace_session_id() try: - messages = ToolEngine.generic_invoke( - tool=tool, - tool_parameters=dict(tool_parameters), - user_id=self._run_context.user_id, - workflow_tool_callback=callback, - workflow_call_depth=workflow_call_depth, - app_id=self._run_context.app_id, - conversation_id=runtime_binding.conversation_id, - ) + session_maker = self._session_maker or session_factory.get_session_maker() + with session_maker.begin() as session: + messages = ToolEngine.generic_invoke( + session=session, + tool=tool, + tool_parameters=dict(tool_parameters), + user_id=self._run_context.user_id, + workflow_tool_callback=callback, + workflow_call_depth=workflow_call_depth, + app_id=self._run_context.app_id, + conversation_id=runtime_binding.conversation_id, + ) + transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages( + messages=messages, + user_id=self._run_context.user_id, + tenant_id=self._run_context.tenant_id, + conversation_id=runtime_binding.conversation_id, + ) + yield from self._adapt_messages(transformed_messages, provider_name=provider_name) except Exception as exc: raise self._map_invocation_exception(exc, provider_name=provider_name) from exc - transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages( - messages=messages, - user_id=self._run_context.user_id, - tenant_id=self._run_context.tenant_id, - conversation_id=runtime_binding.conversation_id, - ) - - return self._adapt_messages(transformed_messages, provider_name=provider_name) - @override def get_usage( self, diff --git a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py index 6f1660390d3..11082d53fa2 100644 --- a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py +++ b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py @@ -8,8 +8,11 @@ import logging from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Literal, override +from sqlalchemy.orm import Session + from core.app.app_config.entities import DatasetRetrieveConfigEntity from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext +from core.db.session_factory import session_factory from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict from core.rag.retrieval.dataset_retrieval import DatasetRetrieval from core.workflow.file_reference import parse_file_reference @@ -75,6 +78,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD *, graph_init_params: "GraphInitParams", graph_runtime_state: "GraphRuntimeState", + session_maker=None, ) -> None: super().__init__( node_id=node_id, @@ -85,6 +89,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD # LLM file outputs, used for MultiModal outputs. self._file_outputs = [] self._rag_retrieval = DatasetRetrieval() + self._session_maker = session_maker or session_factory.get_session_maker() @classmethod @override @@ -130,20 +135,23 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD variables["attachments"] = [variable.value] try: - results, usage = self._fetch_dataset_retriever(node_data=self._node_data, variables=variables) - outputs = {"result": ArrayObjectSegment(value=[item.model_dump(by_alias=True) for item in results])} - return NodeRunResult( - status=WorkflowNodeExecutionStatus.SUCCEEDED, - inputs=variables, - process_data={"usage": jsonable_encoder(usage)}, - outputs=outputs, - metadata={ - WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS: usage.total_tokens, - WorkflowNodeExecutionMetadataKey.TOTAL_PRICE: usage.total_price, - WorkflowNodeExecutionMetadataKey.CURRENCY: usage.currency, - }, - llm_usage=usage, - ) + with self._session_maker() as session: + results, usage = self._fetch_dataset_retriever( + session=session, node_data=self._node_data, variables=variables + ) + outputs = {"result": ArrayObjectSegment(value=[item.model_dump(by_alias=True) for item in results])} + return NodeRunResult( + status=WorkflowNodeExecutionStatus.SUCCEEDED, + inputs=variables, + process_data={"usage": jsonable_encoder(usage)}, + outputs=outputs, + metadata={ + WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS: usage.total_tokens, + WorkflowNodeExecutionMetadataKey.TOTAL_PRICE: usage.total_price, + WorkflowNodeExecutionMetadataKey.CURRENCY: usage.currency, + }, + llm_usage=usage, + ) except RateLimitExceededError as e: logger.warning(e, exc_info=True) return NodeRunResult( @@ -174,7 +182,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD ) def _fetch_dataset_retriever( - self, node_data: KnowledgeRetrievalNodeData, variables: dict[str, Any] + self, session: Session, node_data: KnowledgeRetrievalNodeData, variables: dict[str, Any] ) -> tuple[list[Source], LLMUsage]: dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY)) dataset_ids = node_data.dataset_ids @@ -198,6 +206,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD raise ValueError("single_retrieval_config is required for single retrieval mode") model = node_data.single_retrieval_config.model retrieval_resource_list = self._rag_retrieval.knowledge_retrieval( + session=session, request=KnowledgeRetrievalRequest( tenant_id=dify_ctx.tenant_id, user_id=dify_ctx.user_id, @@ -213,7 +222,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD metadata_filtering_conditions=resolved_metadata_conditions, metadata_filtering_mode=metadata_filtering_mode, query=query, - ) + ), ) elif str(node_data.retrieval_mode) == DatasetRetrieveConfigEntity.RetrieveStrategy.MULTIPLE: if node_data.multiple_retrieval_config is None: @@ -251,6 +260,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD weights = None retrieval_resource_list = self._rag_retrieval.knowledge_retrieval( + session=session, request=KnowledgeRetrievalRequest( app_id=dify_ctx.app_id, tenant_id=dify_ctx.tenant_id, @@ -277,7 +287,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD ] if attachments else None, - ) + ), ) usage = self._rag_retrieval.llm_usage diff --git a/api/services/agent_tool_inner_service.py b/api/services/agent_tool_inner_service.py index aa3fbe1c701..4420f1b66b0 100644 --- a/api/services/agent_tool_inner_service.py +++ b/api/services/agent_tool_inner_service.py @@ -41,7 +41,7 @@ from services.errors.agent_tool_inner import AgentToolInnerServiceError class AgentToolInnerService: """Invoke one API-owned Agent tool declaration, including explicit plugin-via-core calls.""" - def invoke(self, request: AgentToolInvokeRequest, *, session: Session) -> AgentToolInvokeResponse: + def invoke(self, session: Session, request: AgentToolInvokeRequest) -> AgentToolInvokeResponse: app = session.get(App, request.caller.app_id) if app is None: raise AgentToolInnerServiceError( @@ -75,6 +75,7 @@ class AgentToolInnerService: use_default_for_missing_form_parameters=True, ) messages = ToolEngine.generic_invoke( + session=session, tool=tool_runtime, tool_parameters=dict(request.tool.tool_parameters), user_id=request.caller.user_id, diff --git a/api/services/app_generate_service.py b/api/services/app_generate_service.py index 79a46ab52c9..3e2c3c96403 100644 --- a/api/services/app_generate_service.py +++ b/api/services/app_generate_service.py @@ -6,6 +6,8 @@ import uuid from collections.abc import Callable, Generator, Mapping from typing import TYPE_CHECKING, Any +from sqlalchemy.orm import Session + from configs import dify_config from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator from core.app.apps.agent_app.app_generator import AgentAppGenerator @@ -118,6 +120,7 @@ class AppGenerateService: @trace_span(AppGenerateHandler) def generate( cls, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -138,6 +141,7 @@ class AppGenerateService: app_model=app_model, streaming=streaming, action=lambda rate_limit, request_id: cls._dispatch_generate( + session=session, app_model=app_model, user=user, args=args, @@ -185,6 +189,7 @@ class AppGenerateService: def _dispatch_generate( cls, *, + session: Session, app_model: App, user: Account | EndUser, args: Mapping[str, Any], @@ -202,7 +207,12 @@ class AppGenerateService: return rate_limit.generate( CompletionAppGenerator.convert_to_event_stream( CompletionAppGenerator().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=request_id, @@ -229,7 +239,12 @@ class AppGenerateService: return rate_limit.generate( ChatAppGenerator.convert_to_event_stream( ChatAppGenerator().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=request_id, @@ -441,6 +456,7 @@ class AppGenerateService: @classmethod def generate_more_like_this( cls, + session: Session, app_model: App, user: Account | EndUser, message_id: str, @@ -457,7 +473,12 @@ class AppGenerateService: :return: """ return CompletionAppGenerator().generate_more_like_this( - app_model=app_model, message_id=message_id, user=user, invoke_from=invoke_from, stream=streaming + session=session, + app_model=app_model, + message_id=message_id, + user=user, + invoke_from=invoke_from, + stream=streaming, ) @classmethod diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 17e1531db5f..30b04a620f8 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -405,6 +405,7 @@ class DatasetService: @staticmethod def create_empty_dataset( + session: Session, tenant_id: str, name: str, description: str | None, @@ -418,8 +419,6 @@ class DatasetService: embedding_model_name: str | None = None, retrieval_model: RetrievalModel | None = None, summary_index_setting: dict[str, Any] | None = None, - *, - session: scoped_session | Session, ): # check if dataset name already exists if session.scalar(select(Dataset).where(Dataset.name == name, Dataset.tenant_id == tenant_id).limit(1)): @@ -473,7 +472,7 @@ class DatasetService: if provider == "external" and external_knowledge_api_id: external_knowledge_api = ExternalDatasetService.get_external_knowledge_api( - external_knowledge_api_id, tenant_id + session, external_knowledge_api_id, tenant_id ) if not external_knowledge_api: raise ValueError("External API template not found.") @@ -632,7 +631,7 @@ class DatasetService: raise ValueError(ex.description) @staticmethod - def update_dataset(dataset_id, data, user, session: scoped_session | Session): + def update_dataset(session: Session, dataset_id, data, user): """ Update dataset configuration and settings. @@ -685,7 +684,7 @@ class DatasetService: return dataset is not None @staticmethod - def _update_external_dataset(dataset, data, user, session: scoped_session | Session): + def _update_external_dataset(dataset, data, user, session: Session): """ Update external dataset configuration. @@ -725,7 +724,7 @@ class DatasetService: if not external_knowledge_api_id: raise ValueError("External knowledge api id is required.") # Ensure the referenced external API template exists and belongs to the dataset tenant. - ExternalDatasetService.get_external_knowledge_api(external_knowledge_api_id, dataset.tenant_id) + ExternalDatasetService.get_external_knowledge_api(session, external_knowledge_api_id, dataset.tenant_id) # Update metadata fields dataset.updated_by = user.id if user else None dataset.updated_at = naive_utc_now() @@ -4315,9 +4314,7 @@ class DatasetPermissionService: raise e @classmethod - def check_permission( - cls, user, dataset, requested_permission, requested_partial_member_list, session: scoped_session | Session - ): + def check_permission(cls, session: Session, user, dataset, requested_permission, requested_partial_member_list): if not user.is_dataset_editor: raise NoPermissionError("User does not have permission to edit this dataset.") diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index c0aedfaceea..355c4844233 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -5,6 +5,7 @@ from urllib.parse import urlparse import httpx from sqlalchemy import func, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.helper import ssrf_proxy @@ -57,7 +58,9 @@ class ExternalDatasetService: raise ValueError("api_key is required") @staticmethod - def create_external_knowledge_api(tenant_id: str, user_id: str, args: dict[str, Any]) -> ExternalKnowledgeApis: + def create_external_knowledge_api( + tenant_id: str, user_id: str, args: dict[str, Any], session: Session + ) -> ExternalKnowledgeApis: settings = args.get("settings") if settings is None: raise ValueError("settings is required") @@ -71,8 +74,8 @@ class ExternalDatasetService: settings=json.dumps(args.get("settings"), ensure_ascii=False), ) - db.session.add(external_knowledge_api) - db.session.commit() + session.add(external_knowledge_api) + session.commit() return external_knowledge_api @staticmethod @@ -103,8 +106,10 @@ class ExternalDatasetService: raise ValueError(f"Forbidden: Authorization failed with api_key: {api_key}") @staticmethod - def get_external_knowledge_api(external_knowledge_api_id: str, tenant_id: str) -> ExternalKnowledgeApis: - external_knowledge_api: ExternalKnowledgeApis | None = db.session.scalar( + def get_external_knowledge_api( + session: Session, external_knowledge_api_id: str, tenant_id: str + ) -> ExternalKnowledgeApis: + external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) .limit(1) @@ -114,8 +119,10 @@ class ExternalDatasetService: return external_knowledge_api @staticmethod - def update_external_knowledge_api(tenant_id, user_id, external_knowledge_api_id, args) -> ExternalKnowledgeApis: - external_knowledge_api: ExternalKnowledgeApis | None = db.session.scalar( + def update_external_knowledge_api( + session: Session, tenant_id: str, user_id: str, external_knowledge_api_id: str, args + ) -> ExternalKnowledgeApis: + external_knowledge_api: ExternalKnowledgeApis | None = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) .limit(1) @@ -131,13 +138,13 @@ class ExternalDatasetService: external_knowledge_api.settings = json.dumps(args.get("settings"), ensure_ascii=False) external_knowledge_api.updated_by = user_id external_knowledge_api.updated_at = naive_utc_now() - db.session.commit() + session.commit() return external_knowledge_api @staticmethod - def delete_external_knowledge_api(tenant_id: str, external_knowledge_api_id: str): - external_knowledge_api = db.session.scalar( + def delete_external_knowledge_api(session: Session, tenant_id: str, external_knowledge_api_id: str): + external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) .limit(1) @@ -145,11 +152,13 @@ class ExternalDatasetService: if external_knowledge_api is None: raise ValueError("api template not found") - db.session.delete(external_knowledge_api) - db.session.commit() + session.delete(external_knowledge_api) + session.commit() @staticmethod - def external_knowledge_api_use_check(external_knowledge_api_id: str, tenant_id: str) -> tuple[bool, int]: + def external_knowledge_api_use_check( + session: Session, external_knowledge_api_id: str, tenant_id: str + ) -> tuple[bool, int]: """ Return usage for an external knowledge API within a single tenant. @@ -157,7 +166,7 @@ class ExternalDatasetService: same; otherwise the endpoint becomes a cross-tenant UUID oracle. """ count = ( - db.session.scalar( + session.scalar( select(func.count(ExternalKnowledgeBindings.id)).where( ExternalKnowledgeBindings.external_knowledge_api_id == external_knowledge_api_id, ExternalKnowledgeBindings.tenant_id == tenant_id, @@ -168,8 +177,10 @@ class ExternalDatasetService: return count > 0, count @staticmethod - def get_external_knowledge_binding_with_dataset_id(tenant_id: str, dataset_id: str) -> ExternalKnowledgeBindings: - external_knowledge_binding: ExternalKnowledgeBindings | None = db.session.scalar( + def get_external_knowledge_binding_with_dataset_id( + session: Session, tenant_id: str, dataset_id: str + ) -> ExternalKnowledgeBindings: + external_knowledge_binding: ExternalKnowledgeBindings | None = session.scalar( select(ExternalKnowledgeBindings) .where(ExternalKnowledgeBindings.dataset_id == dataset_id, ExternalKnowledgeBindings.tenant_id == tenant_id) .limit(1) @@ -180,9 +191,9 @@ class ExternalDatasetService: @staticmethod def document_create_args_validate( - tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any] + session: Session, tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any] ): - external_knowledge_api = db.session.scalar( + external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id) .limit(1) @@ -255,13 +266,13 @@ class ExternalDatasetService: return ExternalKnowledgeApiSetting.model_validate(settings) @staticmethod - def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any]) -> Dataset: + def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], session: Session) -> Dataset: # check if dataset name already exists - if db.session.scalar( + if session.scalar( select(Dataset).where(Dataset.name == args.get("name"), Dataset.tenant_id == tenant_id).limit(1) ): raise DatasetNameDuplicateError(f"Dataset with name {args.get('name')} already exists.") - external_knowledge_api = db.session.scalar( + external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where( ExternalKnowledgeApis.id == args.get("external_knowledge_api_id"), @@ -283,8 +294,8 @@ class ExternalDatasetService: maintainer=user_id, ) - db.session.add(dataset) - db.session.flush() + session.add(dataset) + session.flush() if args.get("external_knowledge_id") is None: raise ValueError("external_knowledge_id is required") if args.get("external_knowledge_api_id") is None: @@ -297,14 +308,15 @@ class ExternalDatasetService: external_knowledge_id=args.get("external_knowledge_id") or "", created_by=user_id, ) - db.session.add(external_knowledge_binding) + session.add(external_knowledge_binding) - db.session.commit() + session.commit() return dataset @staticmethod def fetch_external_knowledge_retrieval( + session: Session, tenant_id: str, dataset_id: str, query: str, @@ -320,7 +332,7 @@ class ExternalDatasetService: the inner knowledge retrieval API—can consistently expose ``502 external_knowledge_failed``. """ - external_knowledge_binding = db.session.scalar( + external_knowledge_binding = session.scalar( select(ExternalKnowledgeBindings) .where(ExternalKnowledgeBindings.dataset_id == dataset_id, ExternalKnowledgeBindings.tenant_id == tenant_id) .limit(1) @@ -328,7 +340,7 @@ class ExternalDatasetService: if not external_knowledge_binding: raise ExternalKnowledgeRetrievalError("external knowledge binding not found") - external_knowledge_api = db.session.scalar( + external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where( ExternalKnowledgeApis.id == external_knowledge_binding.external_knowledge_api_id, diff --git a/api/services/hit_testing_service.py b/api/services/hit_testing_service.py index 9a2843864d9..1b51a2d279b 100644 --- a/api/services/hit_testing_service.py +++ b/api/services/hit_testing_service.py @@ -105,7 +105,7 @@ class HitTestingService: @classmethod def retrieve( cls, - session: Session | scoped_session, + session: Session, dataset: Dataset, query: str, account: Account, @@ -131,6 +131,7 @@ class HitTestingService: metadata_filtering_conditions = MetadataFilteringCondition.model_validate(metadata_filtering_conditions_raw) metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition( + session=session, dataset_ids=[dataset.id], query=query, metadata_filtering_mode="manual", @@ -190,7 +191,7 @@ class HitTestingService: @classmethod def external_retrieve( cls, - session: Session | scoped_session, + session: Session, dataset: Dataset, query: str, account: Account, @@ -206,6 +207,7 @@ class HitTestingService: start = time.perf_counter() all_documents = RetrievalService.external_retrieve( + session=session, dataset_id=dataset.id, query=cls.escape_query_for_search(query), external_retrieval_model=external_retrieval_model, diff --git a/api/services/knowledge_retrieval_inner_service.py b/api/services/knowledge_retrieval_inner_service.py index 8759413f533..c10c413d4db 100644 --- a/api/services/knowledge_retrieval_inner_service.py +++ b/api/services/knowledge_retrieval_inner_service.py @@ -13,7 +13,7 @@ of a separate validation error. """ from sqlalchemy import select -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from core.rag.entities.metadata_entities import Condition, MetadataFilteringCondition from core.rag.retrieval.dataset_retrieval import DatasetRetrieval @@ -41,7 +41,7 @@ class InnerKnowledgeRetrievalService: def retrieve( self, request: InnerKnowledgeRetrieveRequest, - session: scoped_session, + session: Session, ) -> InnerKnowledgeRetrieveResponse: """Run tenant-scoped retrieval for a trusted internal caller. @@ -64,13 +64,13 @@ class InnerKnowledgeRetrievalService: self._validate_datasets(tenant_id=request.caller.tenant_id, dataset_ids=request.dataset_ids, session=session) rag = DatasetRetrieval() - results = rag.knowledge_retrieval(request=self._to_rag_request(request)) + results = rag.knowledge_retrieval(session=session, request=self._to_rag_request(request)) return InnerKnowledgeRetrieveResponse( results=results, usage=InnerKnowledgeRetrieveUsage.model_validate(jsonable_encoder(rag.llm_usage)), ) - def _validate_caller_app(self, *, tenant_id: str, app_id: str, session: scoped_session) -> None: + def _validate_caller_app(self, *, tenant_id: str, app_id: str, session: Session) -> None: app = session.scalar(select(App).where(App.id == app_id).limit(1)) if app is None: raise InnerKnowledgeRetrieveAppNotFoundError(f"App '{app_id}' not found") @@ -79,7 +79,7 @@ class InnerKnowledgeRetrievalService: f"App '{app_id}' does not belong to tenant '{tenant_id}'" ) - def _validate_datasets(self, *, tenant_id: str, dataset_ids: list[str], session: scoped_session) -> None: + def _validate_datasets(self, *, tenant_id: str, dataset_ids: list[str], session: Session) -> None: datasets = session.scalars(select(Dataset).where(Dataset.id.in_(dataset_ids))).all() found_ids = {dataset.id for dataset in datasets} diff --git a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py index a07fcf40fc2..ae37d305670 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py +++ b/api/tests/test_containers_integration_tests/controllers/console/app/test_app_apis.py @@ -117,7 +117,7 @@ class TestCompletionEndpoints: "/", json={"inputs": {}, "model_config": {}, "query": "hi"}, ): - resp = method(api, _make_account(), app_model=MagicMock(id="app-1")) + resp = method(api, MagicMock(spec=Session), _make_account(), app_model=MagicMock(id="app-1")) assert resp == {"result": {"text": "ok"}} @@ -138,7 +138,7 @@ class TestCompletionEndpoints: json={"inputs": {}, "model_config": {}, "query": "hi"}, ): with pytest.raises(NotFound): - method(api, _make_account(), app_model=MagicMock(id="app-1")) + method(api, MagicMock(spec=Session), _make_account(), app_model=MagicMock(id="app-1")) def test_completion_api_provider_not_initialized(self, app: Flask, monkeypatch: pytest.MonkeyPatch): api = completion_module.CompletionMessageApi() @@ -155,7 +155,7 @@ class TestCompletionEndpoints: json={"inputs": {}, "model_config": {}, "query": "hi"}, ): with pytest.raises(completion_module.ProviderNotInitializeError): - method(api, _make_account(), app_model=MagicMock(id="app-1")) + method(api, MagicMock(spec=Session), _make_account(), app_model=MagicMock(id="app-1")) def test_completion_api_quota_exceeded(self, app: Flask, monkeypatch: pytest.MonkeyPatch): api = completion_module.CompletionMessageApi() @@ -172,7 +172,7 @@ class TestCompletionEndpoints: json={"inputs": {}, "model_config": {}, "query": "hi"}, ): with pytest.raises(completion_module.ProviderQuotaExceededError): - method(api, _make_account(), app_model=MagicMock(id="app-1")) + method(api, MagicMock(spec=Session), _make_account(), app_model=MagicMock(id="app-1")) class TestAppEndpoints: 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 35e17a12b3a..372157813cc 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 @@ -501,7 +501,7 @@ class TestDatasetListApiPost: json={"name": "New Dataset"}, ): api = DatasetListApi() - response, status = unwrap(api.post)(api, tenant_id=mock_tenant.id) + response, status = unwrap(api.post)(api, Mock(spec=Session), tenant_id=mock_tenant.id) assert status == 200 assert_dataset_detail_shape(response) @@ -529,7 +529,7 @@ class TestDatasetListApiPost: ): api = DatasetListApi() with pytest.raises(DatasetNameDuplicateError): - unwrap(api.post)(api, tenant_id=mock_tenant.id) + unwrap(api.post)(api, Mock(spec=Session), tenant_id=mock_tenant.id) # --------------------------------------------------------------------------- @@ -722,14 +722,19 @@ class TestDatasetApiPatch: json=payload, ): api = DatasetApi() - response, status = unwrap(api.patch)(api, _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) + response, status = unwrap(api.patch)( + api, + Mock(spec=Session), + _=mock_dataset.tenant_id, + dataset_id=mock_dataset.id, + ) assert status == 200 assert_dataset_detail_shape(response, with_partial_members=True) assert response["name"] == "Updated Dataset" assert response["partial_member_list"] == ["user-1"] mock_dataset_svc.update_dataset.assert_called_once() - _, update_data, _, session = mock_dataset_svc.update_dataset.call_args.args + session, _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args assert isinstance(session, (Session, scoped_session)) assert update_data["name"] == "Updated Dataset" assert update_data["permission"] == "partial_members" diff --git a/api/tests/test_containers_integration_tests/core/rag/retrieval/test_dataset_retrieval_integration.py b/api/tests/test_containers_integration_tests/core/rag/retrieval/test_dataset_retrieval_integration.py index c0da09278e3..1b9a39ac891 100644 --- a/api/tests/test_containers_integration_tests/core/rag/retrieval/test_dataset_retrieval_integration.py +++ b/api/tests/test_containers_integration_tests/core/rag/retrieval/test_dataset_retrieval_integration.py @@ -518,7 +518,7 @@ class TestKnowledgeRetrievalIntegration: with patch.object(dataset_retrieval, "get_metadata_filter_condition", return_value=(None, None)): with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(db_session_with_containers, request) # Assert assert isinstance(result, list) @@ -567,7 +567,7 @@ class TestKnowledgeRetrievalIntegration: # Mock rate limit check with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(db_session_with_containers, request) # Assert assert result == [] @@ -620,7 +620,7 @@ class TestKnowledgeRetrievalIntegration: ): # Act & Assert with pytest.raises(Exception, match="Rate limit exceeded"): - dataset_retrieval.knowledge_retrieval(request) + dataset_retrieval.knowledge_retrieval(db_session_with_containers, request) @pytest.fixture diff --git a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py index fce0d26c484..473111f364a 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py @@ -234,7 +234,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -262,7 +267,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -288,7 +298,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -314,7 +329,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -342,7 +362,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -374,7 +399,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -401,7 +431,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.DEBUGGER, + streaming=True, ) # Verify the result @@ -426,7 +461,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=False + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=False, ) # Verify the result @@ -463,7 +503,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=end_user, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=end_user, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -490,7 +535,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -524,7 +574,12 @@ class TestAppGenerateService: # StatementError (from EnumText validation during autoflush) with pytest.raises((ValueError, sa.exc.StatementError)): AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) def test_generate_with_workflow_id_format_error( @@ -548,7 +603,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowIdFormatError with pytest.raises(WorkflowIdFormatError) as exc_info: AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify error message @@ -582,7 +642,12 @@ class TestAppGenerateService: # Execute the method under test and expect WorkflowNotFoundError with pytest.raises(WorkflowNotFoundError) as exc_info: AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify error message @@ -608,7 +673,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.DEBUGGER, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.DEBUGGER, + streaming=True, ) # Verify error message @@ -634,7 +704,12 @@ class TestAppGenerateService: # Execute the method under test and expect ValueError with pytest.raises(ValueError) as exc_info: AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify error message @@ -807,7 +882,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_more_like_this( - app_model=app, user=account, message_id=message_id, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + message_id=message_id, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -847,7 +927,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate_more_like_this( - app_model=app, user=end_user, message_id=message_id, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=end_user, + message_id=message_id, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -936,7 +1021,12 @@ class TestAppGenerateService: # Execute the method under test and expect exception with pytest.raises(Exception) as exc_info: AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify exception message @@ -964,7 +1054,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result @@ -999,7 +1094,12 @@ class TestAppGenerateService: # Execute the method under test result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=invoke_from, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=invoke_from, + streaming=True, ) # Verify the result @@ -1037,7 +1137,12 @@ class TestAppGenerateService: mock_exec_params.new.return_value = mock_payload result = AppGenerateService.generate( - app_model=app, user=account, args=args, invoke_from=InvokeFrom.SERVICE_API, streaming=True + session=db_session_with_containers, + app_model=app, + user=account, + args=args, + invoke_from=InvokeFrom.SERVICE_API, + streaming=True, ) # Verify the result diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_service.py index 912e00b0b7d..40c00267043 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service.py @@ -602,7 +602,10 @@ class TestDatasetServiceUpdateAndDeleteDataset: # Act / Assert with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( - source_dataset.id, {"name": "Existing Dataset"}, account, session=db_session_with_containers + db_session_with_containers, + source_dataset.id, + {"name": "Existing Dataset"}, + account, ) def test_delete_dataset_with_documents_success(self, db_session_with_containers: Session): @@ -725,7 +728,7 @@ class TestDatasetServiceRetrievalConfiguration: } # Act - result = DatasetService.update_dataset(dataset.id, update_data, account, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, account) # Assert db_session_with_containers.refresh(dataset) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index 6b32273624b..ced144e8d6e 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -574,7 +574,11 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="does not have permission"): DatasetPermissionService.check_permission( - user, dataset, DatasetPermissionEnum.ALL_TEAM, [], session=db_session_with_containers + db_session_with_containers, + user, + dataset, + DatasetPermissionEnum.ALL_TEAM, + [], ) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode( @@ -585,7 +589,11 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( - user, dataset, DatasetPermissionEnum.ONLY_ME, [], session=db_session_with_containers + db_session_with_containers, + user, + dataset, + DatasetPermissionEnum.ONLY_ME, + [], ) def test_check_permission_requires_partial_member_list_for_partial_members_mode( @@ -596,7 +604,11 @@ class TestDatasetPermissionServiceIntegration: with pytest.raises(ValueError, match="Partial member list is required"): DatasetPermissionService.check_permission( - user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [], session=db_session_with_containers + db_session_with_containers, + user, + dataset, + DatasetPermissionEnum.PARTIAL_TEAM, + [], ) def test_check_permission_rejects_dataset_operator_member_list_changes(self, db_session_with_containers: Session): @@ -606,11 +618,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( + db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-2"}], - session=db_session_with_containers, ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged( @@ -621,11 +633,11 @@ class TestDatasetPermissionServiceIntegration: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( + db_session_with_containers, user, dataset, DatasetPermissionEnum.PARTIAL_TEAM, [{"user_id": "user-1"}], - session=db_session_with_containers, ) def test_clear_partial_member_list_deletes_permissions_and_commits(self, db_session_with_containers: Session): diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py index d9fb23e8e33..f719a465dbd 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_update_dataset.py @@ -189,7 +189,7 @@ class TestDatasetServiceUpdateDataset: "external_knowledge_api_id": external_api.id, } - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) db_session_with_containers.refresh(dataset) updated_binding = db_session_with_containers.query(ExternalKnowledgeBindings).filter_by(id=binding_id).first() @@ -221,7 +221,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_api_id": str(uuid4())} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) assert "External knowledge id is required" in str(context.value) db_session_with_containers.rollback() @@ -245,7 +245,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name", "external_knowledge_id": "knowledge_id"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) assert "External knowledge api id is required" in str(context.value) db_session_with_containers.rollback() @@ -272,7 +272,7 @@ class TestDatasetServiceUpdateDataset: } with pytest.raises(ValueError) as context: - DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) assert "External knowledge binding not found" in str(context.value) db_session_with_containers.rollback() @@ -303,7 +303,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": "text-embedding-ada-002", } - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -338,7 +338,7 @@ class TestDatasetServiceUpdateDataset: "embedding_model": None, } - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -371,7 +371,7 @@ class TestDatasetServiceUpdateDataset: } with patch("services.dataset_service.deal_dataset_vector_index_task") as mock_task: - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) mock_task.delay.assert_called_once_with(dataset.id, "remove") db_session_with_containers.refresh(dataset) @@ -418,7 +418,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -462,7 +462,7 @@ class TestDatasetServiceUpdateDataset: "retrieval_model": "new_model", } - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) db_session_with_containers.refresh(dataset) assert dataset.name == "new_name" @@ -514,7 +514,7 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.return_value = embedding_model mock_get_binding.return_value = binding - result = DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + result = DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) mock_model_manager.return_value.get_model_instance.assert_called_once_with( tenant_id=tenant.id, @@ -545,7 +545,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(ValueError) as context: - DatasetService.update_dataset(str(uuid4()), update_data, user, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, str(uuid4()), update_data, user) assert "Dataset not found" in str(context.value) @@ -568,7 +568,7 @@ class TestDatasetServiceUpdateDataset: update_data = {"name": "new_name"} with pytest.raises(NoPermissionError): - DatasetService.update_dataset(dataset.id, update_data, outsider, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, outsider) def test_update_internal_dataset_embedding_model_error(self, db_session_with_containers: Session): """Test error when embedding model is not available.""" @@ -595,6 +595,6 @@ class TestDatasetServiceUpdateDataset: mock_model_manager.return_value.get_model_instance.side_effect = Exception("No Embedding Model available") with pytest.raises(Exception) as context: - DatasetService.update_dataset(dataset.id, update_data, user, session=db_session_with_containers) + DatasetService.update_dataset(db_session_with_containers, dataset.id, update_data, user) assert "No Embedding Model available".lower() in str(context.value).lower() 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 fbf993f7d69..4a73f98f50e 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 @@ -258,6 +258,7 @@ class TestHitTestingService: assert response.query.content == 'test "query"' assert response.records[0].content == "ext content" mock_ext_retrieve.assert_called_once_with( + session=db_session_with_containers, dataset_id=dataset.id, query='test \\"query\\"', external_retrieval_model={"model": "test"}, diff --git a/api/tests/unit_tests/commands/test_check_no_new_getattr.py b/api/tests/unit_tests/commands/test_check_no_new_getattr.py index c8fadfe982b..2fbaea62d64 100644 --- a/api/tests/unit_tests/commands/test_check_no_new_getattr.py +++ b/api/tests/unit_tests/commands/test_check_no_new_getattr.py @@ -144,6 +144,7 @@ def test_style_workflow_wires_no_new_getattr_guard() -> None: job_text, ) assert checkout_step is not None + assert "fetch-depth: 0" in checkout_step.group("step") changed_files_step = re.search( r"(?ms)^ - name: Check changed files\n.*?^ files: \|\n(?P(?:^ \S[^\n]*\n)+)", 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 e1351875fa7..e070a98c12a 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,6 +1,7 @@ from inspect import getsource, unwrap from types import SimpleNamespace from typing import Any, cast +from unittest.mock import Mock import pytest from flask import Flask @@ -1348,13 +1349,16 @@ def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( 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(), "tenant-1", SimpleNamespace(id=account_id), agent_id + 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} 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 @@ -1383,13 +1387,16 @@ def test_agent_build_chat_finalize_route_resolves_app_from_agent_id( monkeypatch.setattr(completion_controller, "resolve_agent_runtime_app_model", resolve_agent_app_model) monkeypatch.setattr(completion_controller, "_create_build_chat_finalization_message", create_finalization_message) + session = Mock() + with app.test_request_context(): assert unwrap(AgentBuildChatFinalizeApi.post)( - AgentBuildChatFinalizeApi(), "tenant-1", SimpleNamespace(id=account_id), agent_id + 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 assert finalize_call["app_model"] is app_model assert finalize_call["current_tenant_id"] == "tenant-1" assert finalize_call["agent_id"] == agent_id @@ -1429,6 +1436,7 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( current_user=SimpleNamespace(id=account_id), app_model=app_model, agent_id="agent-1", + session=Mock(), ) assert result == ({"result": "success"}, 200) @@ -1518,7 +1526,11 @@ def test_agent_chat_helper_forces_agent_streaming_and_external_trace( 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) + result = completion_controller._create_chat_message( + current_user=current_user, + app_model=app_model, + session=Mock(), + ) assert result == {"response": {"answer": "ok"}} assert captured["app_model"] is app_model @@ -1558,6 +1570,7 @@ def test_agent_chat_helper_rejects_foreign_debug_conversation( current_user=SimpleNamespace(id=account_id), app_model=app_model, agent_id="agent-1", + session=Mock(), ) @@ -1643,6 +1656,7 @@ def test_agent_chat_helper_maps_generation_errors( completion_controller._create_chat_message( current_user=SimpleNamespace(id="account-1"), app_model=app_model, + session=Mock(), ) 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 e05796d8853..2f971eaf74f 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -610,7 +610,7 @@ def test_advanced_chat_run_conversation_not_exists(app: Flask, monkeypatch: pyte json={"inputs": {}}, ): with pytest.raises(NotFound): - handler(api, "t1", app_model=SimpleNamespace(id="app")) + handler(api, Mock(), "t1", app_model=SimpleNamespace(id="app")) def test_workflow_online_users_filters_inaccessible_workflow(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: 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 76a09558987..53f4f139937 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py @@ -593,7 +593,7 @@ class TestDatasetListApiPost: return_value=dataset, ), ): - _, status = method(api, "tenant-1", user) + _, status = method(api, MagicMock(), "tenant-1", user) assert status == 201 @@ -610,7 +610,7 @@ class TestDatasetListApiPost: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user) + method(api, MagicMock(), "tenant-1", user) def test_post_duplicate_name(self, app: Flask): api = DatasetListApi() @@ -630,7 +630,7 @@ class TestDatasetListApiPost: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, "tenant-1", user) + method(api, MagicMock(), "tenant-1", user) def test_post_invalid_payload_missing_name(self, app: Flask): api = DatasetListApi() @@ -638,7 +638,7 @@ class TestDatasetListApiPost: with app.test_request_context("/datasets", json={}), patch.object(type(console_ns), "payload", {}): with pytest.raises(ValueError): - method(api, "tenant-1", make_account()) + method(api, MagicMock(), "tenant-1", make_account()) def test_post_invalid_indexing_technique(self, app: Flask): api = DatasetListApi() @@ -651,7 +651,7 @@ class TestDatasetListApiPost: 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, "tenant-1", make_account()) + method(api, MagicMock(), "tenant-1", make_account()) def test_post_invalid_provider(self, app: Flask): api = DatasetListApi() @@ -664,7 +664,7 @@ class TestDatasetListApiPost: with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError, match="Invalid provider"): - method(api, "tenant-1", make_account()) + method(api, MagicMock(), "tenant-1", make_account()) class TestDatasetApiGet: @@ -931,7 +931,7 @@ class TestDatasetApiPatch: return_value=[], ), ): - result, status = method(api, tenant_id, user, dataset_id) + result, status = method(api, MagicMock(), tenant_id, user, dataset_id) assert status == 200 assert result["partial_member_list"] == [] @@ -949,7 +949,7 @@ class TestDatasetApiPatch: ), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, "tenant-1", make_account(), "missing") + method(api, MagicMock(), "tenant-1", make_account(), "missing") def test_patch_permission_denied(self, app: Flask): api = DatasetApi() @@ -975,7 +975,7 @@ class TestDatasetApiPatch: ), ): with pytest.raises(Forbidden): - method(api, "tenant", make_account(), dataset_id) + method(api, MagicMock(), "tenant", make_account(), dataset_id) def test_patch_partial_members_update(self, app: Flask): api = DatasetApi() @@ -1019,7 +1019,7 @@ class TestDatasetApiPatch: return_value=["u1", "u2"], ), ): - result, _ = method(api, "tenant", make_account(), dataset_id) + result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) assert result["partial_member_list"] == ["u1", "u2"] @@ -1064,7 +1064,7 @@ class TestDatasetApiPatch: return_value=[], ), ): - result, _ = method(api, "tenant", make_account(), dataset_id) + result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) assert result["partial_member_list"] == [] 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 b7e16b91fb7..8ac40f03d3b 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -74,7 +74,7 @@ class TestExternalApiTemplateListApi: patch.object(ExternalDatasetService, "validate_api_list"), ): with pytest.raises(Forbidden): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) def test_post_duplicate_name(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() @@ -93,7 +93,7 @@ class TestExternalApiTemplateListApi: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) class TestExternalApiTemplateApi: @@ -110,7 +110,7 @@ class TestExternalApiTemplateApi: ), ): with pytest.raises(NotFound): - method(api, "tenant-1", "api-id") + method(api, MagicMock(), "tenant-1", "api-id") def test_delete_forbidden(self, app: Flask, current_user: Account): current_user.role = TenantAccountRole.NORMAL @@ -120,7 +120,7 @@ class TestExternalApiTemplateApi: with app.test_request_context("/"): with pytest.raises(Forbidden): - method(api, "tenant-1", current_user, "api-id") + method(api, MagicMock(), "tenant-1", current_user, "api-id") class TestExternalApiUseCheckApi: @@ -128,6 +128,8 @@ class TestExternalApiUseCheckApi: api = ExternalApiUseCheckApi() method = inspect.unwrap(api.get) + session = MagicMock() + with ( app.test_request_context("/"), patch.object( @@ -136,11 +138,11 @@ class TestExternalApiUseCheckApi: return_value=(True, 2), ) as mock_use_check, ): - response, status = method(api, "tenant-1", "api-id") + response, status = method(api, session, "tenant-1", "api-id") assert status == 200 assert response == {"is_using": True, "count": 2} - mock_use_check.assert_called_once_with("api-id", "tenant-1") + mock_use_check.assert_called_once_with(session, "api-id", "tenant-1") class TestExternalDatasetCreateApi: @@ -185,7 +187,7 @@ class TestExternalDatasetCreateApi: return_value=dataset, ), ): - _, status = method(api, "tenant-1", current_user) + _, status = method(api, MagicMock(), "tenant-1", current_user) assert status == 201 @@ -205,7 +207,7 @@ class TestExternalDatasetCreateApi: patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), ): with pytest.raises(Forbidden): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) class TestExternalKnowledgeHitTestingApi: @@ -222,7 +224,7 @@ class TestExternalKnowledgeHitTestingApi: ), ): with pytest.raises(NotFound): - method(api, current_user, "dataset-id") + method(api, MagicMock(), current_user, "dataset-id") def test_hit_testing_success(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -243,7 +245,7 @@ class TestExternalKnowledgeHitTestingApi: return_value={"ok": True}, ), ): - resp = method(api, current_user, "dataset-id") + resp = method(api, MagicMock(), current_user, "dataset-id") assert resp["ok"] is True @@ -291,7 +293,7 @@ class TestExternalApiTemplateListApiAdvanced: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) def test_get_with_pagination(self, app: Flask): api = ExternalApiTemplateListApi() @@ -331,7 +333,7 @@ class TestExternalDatasetCreateApiAdvanced: with app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(Forbidden): - method(api, "tenant-1", current_user) + method(api, MagicMock(), "tenant-1", current_user) class TestExternalKnowledgeHitTestingApiAdvanced: @@ -354,7 +356,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: ), ): with pytest.raises(NotFound): - method(api, current_user, "ds-1") + method(api, MagicMock(), current_user, "ds-1") def test_hit_testing_with_custom_retrieval_model(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -380,7 +382,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: return_value={"results": []}, ), ): - resp = method(api, current_user, "ds-1") + resp = method(api, MagicMock(), current_user, "ds-1") assert resp["results"] == [] diff --git a/api/tests/unit_tests/controllers/console/datasets/test_hit_testing.py b/api/tests/unit_tests/controllers/console/datasets/test_hit_testing.py index 3de780f3bbb..87b140715b9 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_hit_testing.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_hit_testing.py @@ -1,6 +1,6 @@ import uuid from inspect import unwrap -from unittest.mock import PropertyMock, patch +from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask @@ -135,7 +135,7 @@ class TestHitTestingApi: return_value={"query": {"content": "what is vector search"}, "records": []}, ), ): - result = method(api, account, "tenant-1", dataset_id) + result = method(api, MagicMock(), account, "tenant-1", dataset_id) assert "query" in result assert "records" in result @@ -173,7 +173,7 @@ class TestHitTestingApi: return_value={"query": {"content": payload["query"]}, "records": records}, ), ): - result = method(api, account, "tenant-1", dataset_id) + result = method(api, MagicMock(), account, "tenant-1", dataset_id) assert result["query"] == {"content": payload["query"]} assert result["records"][0]["segment"]["keywords"] == [] @@ -204,7 +204,7 @@ class TestHitTestingApi: ), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, account, "tenant-1", dataset_id) + method(api, MagicMock(), account, "tenant-1", dataset_id) def test_hit_testing_invalid_args(self, app: Flask, dataset, dataset_id, account: Account): api = HitTestingApi() @@ -234,4 +234,4 @@ class TestHitTestingApi: ), ): with pytest.raises(ValueError, match="Invalid parameters"): - method(api, account, "tenant-1", dataset_id) + method(api, MagicMock(), account, "tenant-1", dataset_id) 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 0fcf0df5262..b635fcf3efd 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 @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import Mock, patch import pytest from werkzeug.exceptions import Forbidden, InternalServerError, NotFound @@ -171,7 +171,9 @@ class TestPerformHitTesting: "retrieve", return_value=response, ): - result = DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + result = DatasetsHitTestingBase.perform_hit_testing( + Mock(), dataset, {"query": "hello"}, account, "tenant-1" + ) assert result["query"] == {"content": "hello"} assert result["records"] == [] @@ -187,7 +189,9 @@ class TestPerformHitTesting: "retrieve", return_value=response, ): - result = DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + result = DatasetsHitTestingBase.perform_hit_testing( + Mock(), dataset, {"query": "hello"}, account, "tenant-1" + ) assert result["query"] == {"content": "hello"} record = result["records"][0] @@ -208,7 +212,7 @@ class TestPerformHitTesting: ), pytest.raises(ValueError, match="Invalid hit testing query response"), ): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_invalid_records_response_raises_value_error(self): with pytest.raises(ValueError, match="Invalid hit testing records response"): @@ -225,7 +229,7 @@ class TestPerformHitTesting: side_effect=services.errors.index.IndexNotInitializedError(), ): with pytest.raises(DatasetNotInitializedError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_provider_token_not_init(self, dataset, account): with patch.object( @@ -234,7 +238,7 @@ class TestPerformHitTesting: side_effect=ProviderTokenNotInitError("token missing"), ): with pytest.raises(ProviderNotInitializeError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_quota_exceeded(self, dataset, account): with patch.object( @@ -243,7 +247,7 @@ class TestPerformHitTesting: side_effect=QuotaExceededError(), ): with pytest.raises(ProviderQuotaExceededError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_model_not_supported(self, dataset, account): with patch.object( @@ -252,7 +256,7 @@ class TestPerformHitTesting: side_effect=ModelCurrentlyNotSupportError(), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_llm_bad_request(self, dataset, account): with patch.object( @@ -261,7 +265,7 @@ class TestPerformHitTesting: side_effect=LLMBadRequestError("bad request"), ): with pytest.raises(ProviderNotInitializeError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_invoke_error(self, dataset, account): with patch.object( @@ -270,7 +274,7 @@ class TestPerformHitTesting: side_effect=InvokeError("invoke failed"), ): with pytest.raises(CompletionRequestError): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_value_error(self, dataset, account): with patch.object( @@ -279,7 +283,7 @@ class TestPerformHitTesting: side_effect=ValueError("bad args"), ): with pytest.raises(ValueError, match="bad args"): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") def test_unexpected_error(self, dataset, account): with patch.object( @@ -288,4 +292,4 @@ class TestPerformHitTesting: side_effect=Exception("boom"), ): with pytest.raises(InternalServerError, match="boom"): - DatasetsHitTestingBase.perform_hit_testing(dataset, {"query": "hello"}, account, "tenant-1") + DatasetsHitTestingBase.perform_hit_testing(Mock(), dataset, {"query": "hello"}, account, "tenant-1") 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 8b9121c4d7f..53d9badc5b8 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_completion.py +++ b/api/tests/unit_tests/controllers/console/explore/test_completion.py @@ -67,7 +67,7 @@ class TestCompletionApi: return_value=("ok", 200), ), ): - result = method(api, user, completion_app) + result = method(api, MagicMock(), user, completion_app) assert result == ("ok", 200) @@ -78,7 +78,7 @@ class TestCompletionApi: installed_app = MagicMock(app=MagicMock(mode=AppMode.CHAT)) with pytest.raises(NotCompletionAppError): - method(api, user, installed_app) + method(api, MagicMock(), user, installed_app) def test_conversation_completed(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -94,7 +94,7 @@ class TestCompletionApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_internal_error(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -110,7 +110,7 @@ class TestCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_conversation_not_exists(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -126,7 +126,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.NotFound): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_app_unavailable(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -142,7 +142,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.AppUnavailableError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_provider_not_initialized(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -158,7 +158,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderNotInitializeError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_quota_exceeded(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -174,7 +174,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderQuotaExceededError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_model_not_supported(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -190,7 +190,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderModelCurrentlyNotSupportError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) def test_invoke_error(self, app: Flask, completion_app, user, payload_patch): api = completion_module.CompletionApi() @@ -206,7 +206,7 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.CompletionRequestError): - method(api, user, completion_app) + method(api, MagicMock(), user, completion_app) class TestCompletionStopApi: @@ -249,7 +249,7 @@ class TestChatApi: return_value=("ok", 200), ), ): - result = method(api, user, chat_app) + result = method(api, MagicMock(), user, chat_app) assert result == ("ok", 200) @@ -260,7 +260,7 @@ class TestChatApi: installed_app = MagicMock(app=MagicMock(mode=AppMode.COMPLETION)) with pytest.raises(NotChatAppError): - method(api, user, installed_app) + method(api, MagicMock(), user, installed_app) def test_rate_limit_error(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -276,7 +276,7 @@ class TestChatApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_conversation_completed_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -292,7 +292,7 @@ class TestChatApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_conversation_not_exists_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -308,7 +308,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.NotFound): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_app_unavailable_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -324,7 +324,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.AppUnavailableError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_provider_not_initialized_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -340,7 +340,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderNotInitializeError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_quota_exceeded_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -356,7 +356,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderQuotaExceededError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_model_not_supported_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -372,7 +372,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderModelCurrentlyNotSupportError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_invoke_error_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -388,7 +388,7 @@ class TestChatApi: ), ): with pytest.raises(completion_module.CompletionRequestError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) def test_internal_error_chat(self, app: Flask, chat_app, user, payload_patch): api = completion_module.ChatApi() @@ -404,7 +404,7 @@ class TestChatApi: ), ): with pytest.raises(InternalServerError): - method(api, user, chat_app) + method(api, MagicMock(), user, chat_app) class TestChatStopApi: 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 cb63a52075d..f6182abb987 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_message.py @@ -200,7 +200,7 @@ class TestMessageMoreLikeThisApi: return_value=("ok", 200), ), ): - resp = method(MagicMock(), installed_app, "mid") + resp = method(MagicMock(), MagicMock(), installed_app, "mid") assert resp == ("ok", 200) @@ -212,7 +212,7 @@ class TestMessageMoreLikeThisApi: installed_app.app = MagicMock(mode="chat") with pytest.raises(NotCompletionAppError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_more_like_this_disabled(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -233,7 +233,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(AppMoreLikeThisDisabledError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_message_not_exists_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -254,7 +254,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(NotFound): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_provider_not_init_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -275,7 +275,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_quota_exceeded_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -296,7 +296,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_model_not_support_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -317,7 +317,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_invoke_error_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -338,7 +338,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(CompletionRequestError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") def test_unexpected_error_more_like_this(self, app: Flask): api = module.MessageMoreLikeThisApi() @@ -359,7 +359,7 @@ class TestMessageMoreLikeThisApi: ), ): with pytest.raises(InternalServerError): - method(MagicMock(), installed_app, "mid") + method(MagicMock(), MagicMock(), installed_app, "mid") class TestMessageSuggestedQuestionApi: 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 db800a23b84..6c80a93a558 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -105,7 +105,7 @@ class TestTrialAppWorkflowRunApi: with app.test_request_context("/"): with pytest.raises(NotWorkflowAppError): - method(api, account, MagicMock(mode=AppMode.CHAT)) + method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT)) def test_success(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -116,7 +116,7 @@ class TestTrialAppWorkflowRunApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, account, trial_app_workflow) + result = method(api, MagicMock(), account, trial_app_workflow) assert result is not None @@ -133,7 +133,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_quota_exceeded(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -148,7 +148,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_model_not_support(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -163,7 +163,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_invoke_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -178,7 +178,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_rate_limit_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -193,7 +193,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_value_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -208,7 +208,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ValueError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) def test_workflow_generic_exception(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -223,7 +223,7 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(InternalServerError): - method(api, account, trial_app_workflow) + method(api, MagicMock(), account, trial_app_workflow) class TestTrialChatApi: @@ -233,7 +233,7 @@ class TestTrialChatApi: with app.test_request_context("/", json={"inputs": {}, "query": "hi"}): with pytest.raises(NotChatAppError): - method(api, account, MagicMock(mode="completion")) + method(api, MagicMock(), account, MagicMock(mode="completion")) def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -244,7 +244,7 @@ class TestTrialChatApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, account, trial_app_chat) + result = method(api, MagicMock(), account, trial_app_chat) assert result is not None @@ -261,7 +261,7 @@ class TestTrialChatApi: ), ): with pytest.raises(NotFound): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_conversation_completed(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -276,7 +276,7 @@ class TestTrialChatApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -291,7 +291,7 @@ class TestTrialChatApi: ), ): with pytest.raises(AppUnavailableError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -306,7 +306,7 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -321,7 +321,7 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -336,7 +336,7 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -351,7 +351,7 @@ class TestTrialChatApi: ), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_rate_limit_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -366,7 +366,7 @@ class TestTrialChatApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_value_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -381,7 +381,7 @@ class TestTrialChatApi: ), ): with pytest.raises(ValueError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) def test_chat_generic_exception(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -396,7 +396,7 @@ class TestTrialChatApi: ), ): with pytest.raises(InternalServerError): - method(api, account, trial_app_chat) + method(api, MagicMock(), account, trial_app_chat) class TestTrialCompletionApi: @@ -406,7 +406,7 @@ class TestTrialCompletionApi: with app.test_request_context("/", json={"inputs": {}, "query": ""}): with pytest.raises(NotCompletionAppError): - method(api, account, MagicMock(mode=AppMode.CHAT)) + method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT)) def test_success(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -417,7 +417,7 @@ class TestTrialCompletionApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, account, trial_app_completion) + result = method(api, MagicMock(), account, trial_app_completion) assert result is not None @@ -434,7 +434,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(AppUnavailableError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_provider_not_init(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -449,7 +449,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_quota_exceeded(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -464,7 +464,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_model_not_support(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -479,7 +479,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_invoke_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -494,7 +494,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_rate_limit_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -509,7 +509,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_value_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -524,7 +524,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ValueError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) def test_completion_generic_exception(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -539,7 +539,7 @@ class TestTrialCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, account, trial_app_completion) + method(api, MagicMock(), account, trial_app_completion) class TestTrialMessageSuggestedQuestionApi: 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 83cfcadd093..70cd9fd5cc8 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/explore/test_workflow.py @@ -58,7 +58,7 @@ class TestInstalledAppWorkflowRunApi: with app.test_request_context("/"): with pytest.raises(NotWorkflowAppError): - method(api, MagicMock(), non_workflow_installed_app) + method(api, MagicMock(), MagicMock(), non_workflow_installed_app) def test_success(self, app: Flask, installed_workflow_app, user, payload): api = InstalledAppWorkflowRunApi() @@ -71,7 +71,7 @@ class TestInstalledAppWorkflowRunApi: return_value=MagicMock(), ) as generate_mock, ): - result = method(api, user, installed_workflow_app) + result = method(api, MagicMock(), user, installed_workflow_app) generate_mock.assert_called_once() assert generate_mock.call_args.kwargs["user"] is user @@ -89,7 +89,7 @@ class TestInstalledAppWorkflowRunApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, user, installed_workflow_app) + method(api, MagicMock(), user, installed_workflow_app) def test_unexpected_exception(self, app: Flask, installed_workflow_app, user, payload): api = InstalledAppWorkflowRunApi() @@ -103,7 +103,7 @@ class TestInstalledAppWorkflowRunApi: ), ): with pytest.raises(InternalServerError): - method(api, user, installed_workflow_app) + method(api, MagicMock(), user, installed_workflow_app) class TestInstalledAppWorkflowTaskStopApi: diff --git a/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py b/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py index aabed01262c..b82ab254d45 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_run_streaming.py @@ -77,6 +77,7 @@ def test_run_chat_always_calls_generate_with_streaming_true( _make_app(), _make_account(), AppRunRequest(inputs={}, query="hello"), + Mock(), ) _, kwargs = generate_mock.call_args assert kwargs["streaming"] is True 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 46ce1a85caa..393bdaf5eda 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 @@ -319,7 +319,7 @@ class TestCompletionControllerLogic: mock_compact.return_value = {"text": "compacted"} api = CompletionApi() - response = api.post.__wrapped__(api, mock_app_model, mock_end_user) + response = unwrap(api.post)(api, Mock(), mock_app_model, mock_end_user) assert response == {"text": "compacted"} mock_generate_service.generate.assert_called_once() @@ -335,7 +335,7 @@ class TestCompletionControllerLogic: with app.test_request_context(): with pytest.raises(AppUnavailableError): - CompletionApi().post.__wrapped__(CompletionApi(), mock_app_model, mock_end_user) + unwrap(CompletionApi().post)(CompletionApi(), Mock(), mock_app_model, mock_end_user) @patch("controllers.service_api.app.completion.service_api_ns") @patch("controllers.service_api.app.completion.AppGenerateService") @@ -356,7 +356,7 @@ class TestCompletionControllerLogic: mock_compact.return_value = {"text": "compacted"} api = ChatApi() - response = api.post.__wrapped__(api, mock_app_model, mock_end_user) + response = unwrap(api.post)(api, Mock(), mock_app_model, mock_end_user) assert response == {"text": "compacted"} @patch("controllers.service_api.app.completion.service_api_ns") @@ -370,7 +370,7 @@ class TestCompletionControllerLogic: with app.test_request_context(): with pytest.raises(NotChatAppError): - ChatApi().post.__wrapped__(ChatApi(), mock_app_model, mock_end_user) + unwrap(ChatApi().post)(ChatApi(), Mock(), mock_app_model, mock_end_user) @patch("controllers.service_api.app.completion.AppTaskService") def test_completion_stop_api_success(self, mock_task_service, app: Flask): @@ -427,7 +427,7 @@ class TestCompletionApiController: with app.test_request_context("/completion-messages", method="POST", json={"inputs": {}}): with pytest.raises(AppUnavailableError): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) def test_conversation_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -443,7 +443,7 @@ class TestCompletionApiController: with app.test_request_context("/completion-messages", method="POST", json={"inputs": {}}): with pytest.raises(NotFound): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) class TestCompletionStopApiController: @@ -482,7 +482,7 @@ class TestChatApiController: with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}): with pytest.raises(NotChatAppError): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) def test_workflow_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -498,7 +498,7 @@ class TestChatApiController: with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}): with pytest.raises(NotFound): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -514,7 +514,7 @@ class TestChatApiController: with app.test_request_context("/chat-messages", method="POST", json={"inputs": {}, "query": "hi"}): with pytest.raises(BadRequest): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) 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 576b95110a0..17f2ac3f17d 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 @@ -358,6 +358,7 @@ class TestHitlServiceApi: user.id = "user-id" result = AppGenerateService.generate( + session=Mock(), app_model=app_model, user=user, args={"workflow_id": None, "query": "hi", "inputs": {}}, diff --git a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py index 4f88ae69c2d..3cabfe43ddc 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py @@ -407,7 +407,7 @@ class TestWorkflowRunApi: with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}): with pytest.raises(NotWorkflowAppError): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) def test_rate_limit(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -423,7 +423,7 @@ class TestWorkflowRunApi: with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}): with pytest.raises(InvokeRateLimitHttpError): - handler(api, app_model=app_model, end_user=end_user) + handler(api, session=Mock(), app_model=app_model, end_user=end_user) class TestWorkflowRunByIdApi: @@ -441,7 +441,7 @@ class TestWorkflowRunByIdApi: with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}): with pytest.raises(NotFound): - handler(api, app_model=app_model, end_user=end_user, workflow_id="w1") + handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1") def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -457,7 +457,7 @@ class TestWorkflowRunByIdApi: with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}): with pytest.raises(BadRequest): - handler(api, app_model=app_model, end_user=end_user, workflow_id="w1") + handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1") class TestWorkflowTaskStopApi: 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 b983a945ad5..735cc437bd1 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 @@ -503,6 +503,7 @@ class TestBaseAgentRunnerInit: message = mocker.MagicMock(id="msg1", conversation_id="conv1") runner = BaseAgentRunner( + session=session, tenant_id="tenant", application_generate_entity=app_generate, conversation=mocker.MagicMock(), 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 a6cae351b17..8ccc3611b9d 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 @@ -80,6 +80,7 @@ def runner(mocker: MockerFixture): runner.agent_callback = None runner.memory = None runner.history_prompt_messages = [] + runner.session = MagicMock() return runner @@ -163,7 +164,7 @@ class TestFormatAssistantMessage: class TestHandleInvokeAction: def test_handle_invoke_action_tool_not_present(self, runner: DummyRunner): action = AgentScratchpadUnit.Action(action_name="missing", action_input={}) - response, meta = runner._handle_invoke_action(action, {}, []) + response, meta = runner._handle_invoke_action(runner.session, action, {}, []) assert "there is not a tool named" in response def test_tool_with_json_string_args(self, runner: DummyRunner, mocker: MockerFixture): @@ -176,7 +177,7 @@ class TestHandleInvokeAction: return_value=("result", [], MagicMock(to_dict=lambda: {})), ) - response, meta = runner._handle_invoke_action(action, tool_instances, []) + response, meta = runner._handle_invoke_action(runner.session, action, tool_instances, []) assert response == "result" @@ -200,7 +201,7 @@ class TestRun: return_value=[], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert isinstance(results, list) def test_run_with_action_and_tool_invocation(self, runner: DummyRunner, mocker: MockerFixture): @@ -222,7 +223,7 @@ class TestRun: runner.agent_callback = None with pytest.raises(AgentMaxIterationError): - list(runner.run(message, "query", {"tool": MagicMock()})) + list(runner.run(runner.session, message, "query", {"tool": MagicMock()})) def test_run_respects_max_iteration_boundary(self, runner: DummyRunner, mocker: MockerFixture): runner.app_config.agent.max_iteration = 1 @@ -244,7 +245,7 @@ class TestRun: runner.agent_callback = None with pytest.raises(AgentMaxIterationError): - list(runner.run(message, "query", {"tool": MagicMock()})) + list(runner.run(runner.session, message, "query", {"tool": MagicMock()})) def test_run_basic_flow(self, runner: DummyRunner, mocker: MockerFixture): message = MagicMock() @@ -255,7 +256,7 @@ class TestRun: return_value=[], ) - results = list(runner.run(message, "query", {"name": "John"})) + results = list(runner.run(runner.session, message, "query", {"name": "John"})) assert results def test_run_max_iteration_error(self, runner: DummyRunner, mocker: MockerFixture): @@ -271,7 +272,7 @@ class TestRun: ) with pytest.raises(AgentMaxIterationError): - list(runner.run(message, "query", {})) + list(runner.run(runner.session, message, "query", {})) def test_run_increase_usage_aggregation(self, runner: DummyRunner, mocker: MockerFixture): message = MagicMock() @@ -320,7 +321,7 @@ class TestRun: fake_prompt_tool.name = "tool" runner._init_prompt_tools = MagicMock(return_value=({"tool": MagicMock()}, [fake_prompt_tool])) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) final_usage = results[-1].delta.usage assert final_usage is not None assert final_usage.prompt_tokens == 2 @@ -339,7 +340,7 @@ class TestRun: return_value=[], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert results[-1].delta.message.content == "" def test_run_usage_missing_key_branch(self, runner: DummyRunner, mocker: MockerFixture): @@ -353,7 +354,7 @@ class TestRun: runner.model_instance.invoke_llm = MagicMock(return_value=[]) - list(runner.run(message, "query", {})) + list(runner.run(runner.session, message, "query", {})) def test_run_prompt_tool_update_branch(self, runner: DummyRunner, mocker: MockerFixture): message = MagicMock() @@ -383,7 +384,7 @@ class TestRun: runner.update_prompt_message_tool = MagicMock() runner.agent_callback = None - list(runner.run(message, "query", {})) + list(runner.run(runner.session, message, "query", {})) runner.update_prompt_message_tool.assert_called_once() @@ -435,7 +436,7 @@ class TestHandleInvokeActionExtended: ) message_file_ids = [] - response, meta = runner._handle_invoke_action(action, tool_instances, message_file_ids) + response, meta = runner._handle_invoke_action(runner.session, action, tool_instances, message_file_ids) assert response == "ok" assert message_file_ids == ["file1"] @@ -505,7 +506,7 @@ class TestRunAdditionalBranches: return_value=["thinking"], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert any(hasattr(r, "delta") for r in results) def test_run_with_final_answer_action_string(self, runner: DummyRunner, mocker: MockerFixture): @@ -519,7 +520,7 @@ class TestRunAdditionalBranches: return_value=[action], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert results[-1].delta.message.content == "done" def test_run_with_final_answer_action_dict(self, runner: DummyRunner, mocker: MockerFixture): @@ -533,7 +534,7 @@ class TestRunAdditionalBranches: return_value=[action], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert json.loads(results[-1].delta.message.content) == {"a": 1} def test_run_with_string_final_answer(self, runner: DummyRunner, mocker: MockerFixture): @@ -548,5 +549,5 @@ class TestRunAdditionalBranches: return_value=[action], ) - results = list(runner.run(message, "query", {})) + results = list(runner.run(runner.session, message, "query", {})) assert results[-1].delta.message.content == "12345" 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 9b2a1d70fdf..244dd8f6c62 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 @@ -130,6 +130,7 @@ def runner(mocker: MockerFixture): runner._current_thoughts = [] runner.files = [] runner.agent_callback = MagicMock() + runner.session = MagicMock() runner._init_prompt_tools = MagicMock(return_value=({}, [])) runner.create_agent_thought = MagicMock(return_value="thought1") @@ -296,7 +297,7 @@ class TestRunMethod: runner.model_instance.invoke_llm.return_value = result - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) == 1 runner.queue_manager.publish.assert_called() @@ -315,7 +316,7 @@ class TestRunMethod: runner.model_instance.invoke_llm.return_value = generator() - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) == 1 def test_run_streaming_tool_calls_list_content(self, runner: FunctionCallAgentRunner): @@ -338,7 +339,7 @@ class TestRunMethod: runner.model_instance.invoke_llm.side_effect = [generator(), final_result] - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) >= 1 def test_run_non_streaming_list_content(self, runner: FunctionCallAgentRunner): @@ -349,7 +350,7 @@ class TestRunMethod: runner.model_instance.invoke_llm.return_value = result - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) == 1 assert runner.save_agent_thought.call_args.kwargs["thought"] == "hi" @@ -378,7 +379,7 @@ class TestRunMethod: mocker.patch("core.agent.fc_agent_runner.json.dumps", side_effect=flaky_dumps) - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) == 1 def test_run_with_missing_tool_instance(self, runner: FunctionCallAgentRunner): @@ -396,7 +397,7 @@ class TestRunMethod: runner.model_instance.invoke_llm.side_effect = [result, final_result] - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) >= 1 def test_run_with_tool_instance_and_files(self, runner: FunctionCallAgentRunner, mocker: MockerFixture): @@ -425,7 +426,7 @@ class TestRunMethod: return_value=("ok", ["file1"], tool_invoke_meta), ) - outputs = list(runner.run(message, "query")) + outputs = list(runner.run(runner.session, message, "query")) assert len(outputs) >= 1 assert any( isinstance(call.args[0], QueueMessageFileEvent) @@ -450,4 +451,4 @@ class TestRunMethod: runner.model_instance.invoke_llm.return_value = result with pytest.raises(AgentMaxIterationError): - list(runner.run(message, "query")) + list(runner.run(runner.session, message, "query")) 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 467d404b7db..eb4bf9bc166 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 @@ -236,6 +236,7 @@ class TestAgentChatAppGeneratorWorker: generator._generate_worker( flask_app=mocker.MagicMock(), + session=mocker.MagicMock(), context=mocker.MagicMock(), application_generate_entity=mocker.MagicMock(), queue_manager=queue_manager, @@ -266,6 +267,7 @@ class TestAgentChatAppGeneratorWorker: generator._generate_worker( flask_app=mocker.MagicMock(), + session=mocker.MagicMock(), context=mocker.MagicMock(), application_generate_entity=mocker.MagicMock(), queue_manager=queue_manager, @@ -292,6 +294,7 @@ class TestAgentChatAppGeneratorWorker: 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 af2fb22ec78..a3a879bee3d 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(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(mocker.MagicMock(), generate_entity, 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(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) runner.direct_output.assert_called_once() @@ -82,7 +82,7 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "direct_output") queue_manager = mocker.MagicMock() - runner.run(generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock()) + runner.run(mocker.MagicMock(), generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock()) queue_manager.publish.assert_called_once() runner.direct_output.assert_called_once() @@ -109,7 +109,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(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(mocker.MagicMock(), generate_entity, 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 +144,7 @@ class TestAgentChatAppRunnerRun: mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) @pytest.mark.parametrize( ("mode", "expected_runner"), @@ -201,7 +201,7 @@ class TestAgentChatAppRunnerRun: runner_instance.run.return_value = [] mocker.patch.object(runner, "_handle_invoke_result") - runner.run(generate_entity, mocker.MagicMock(), conversation, message) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) runner_instance.run.assert_called_once() runner._handle_invoke_result.assert_called_once() @@ -247,7 +247,7 @@ class TestAgentChatAppRunnerRun: patch_create_session(mocker, side_effect=[app_record, conversation, message]) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), conversation, message) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) def test_run_function_calling_strategy_selected_by_features( self, runner: AgentChatAppRunner, mocker: MockerFixture @@ -299,7 +299,7 @@ class TestAgentChatAppRunnerRun: runner_instance.run.return_value = [] mocker.patch.object(runner, "_handle_invoke_result") - runner.run(generate_entity, mocker.MagicMock(), conversation, message) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) assert app_config.agent.strategy == AgentEntity.Strategy.FUNCTION_CALLING runner_instance.run.assert_called_once() @@ -333,7 +333,13 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "check_hosting_moderation", return_value=False) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg")) + runner.run( + mocker.MagicMock(), + generate_entity, + mocker.MagicMock(), + mocker.MagicMock(id="conv"), + mocker.MagicMock(id="msg"), + ) def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") @@ -364,7 +370,13 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "check_hosting_moderation", return_value=False) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg")) + runner.run( + mocker.MagicMock(), + generate_entity, + mocker.MagicMock(), + mocker.MagicMock(id="conv"), + mocker.MagicMock(id="msg"), + ) def test_run_invalid_agent_strategy_raises(self, runner: AgentChatAppRunner, mocker: MockerFixture): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") @@ -407,4 +419,4 @@ class TestAgentChatAppRunnerRun: patch_create_session(mocker, side_effect=[app_record, conversation, message]) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), conversation, message) + runner.run(mocker.MagicMock(), generate_entity, mocker.MagicMock(), conversation, message) 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 23334dbe67c..5b4e7bacc68 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 @@ -48,6 +48,7 @@ class TestChatAppGenerator: generator = ChatAppGenerator() with pytest.raises(ValueError): generator.generate( + session=MagicMock(), app_model=SimpleNamespace(), user=SimpleNamespace(), args={"inputs": {}}, @@ -59,6 +60,7 @@ class TestChatAppGenerator: generator = ChatAppGenerator() with pytest.raises(ValueError): generator.generate( + session=MagicMock(), app_model=SimpleNamespace(), user=SimpleNamespace(), args={"query": 1, "inputs": {}}, @@ -105,7 +107,7 @@ 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(app_model, user, args, InvokeFrom.DEBUGGER, streaming=False) + result = generator.generate(MagicMock(), app_model, user, args, InvokeFrom.DEBUGGER, streaming=False) assert result == {"ok": True} assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" @@ -119,6 +121,7 @@ class TestChatAppGenerator: ), ): generator.generate( + session=MagicMock(), app_model=SimpleNamespace(tenant_id="t1", id="a1", mode=AppMode.CHAT.value), user=SimpleNamespace(id="u1", session_id="s1"), args={"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}}, @@ -139,6 +142,7 @@ class TestChatAppGenerator: ): 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,6 +159,7 @@ class TestChatAppGenerator: ): 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", @@ -183,7 +188,9 @@ class TestChatAppRunner: with patched_create_session(return_value=None): with pytest.raises(ValueError): - runner.run(app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run( + MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + ) def test_run_moderation_error_direct_output(self): runner = ChatAppRunner() @@ -214,7 +221,9 @@ class TestChatAppRunner: patch.object(ChatAppRunner, "moderation_for_inputs", side_effect=ModerationError("blocked")), patch.object(ChatAppRunner, "direct_output") as mock_direct, ): - runner.run(app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run( + MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + ) mock_direct.assert_called_once() @@ -251,7 +260,7 @@ class TestChatAppRunner: patch.object(ChatAppRunner, "direct_output") as mock_direct, ): queue_manager = DummyQueueManager() - runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run(MagicMock(), app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) assert any(isinstance(item[0], QueueAnnotationReplyEvent) for item in queue_manager.published) mock_direct.assert_called_once() @@ -286,7 +295,9 @@ class TestChatAppRunner: patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None), patch.object(ChatAppRunner, "check_hosting_moderation", return_value=True), ): - runner.run(app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run( + MagicMock(), app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1") + ) def test_run_closes_scoped_session_before_stream_consumption(self): runner = ChatAppRunner() @@ -339,7 +350,7 @@ class TestChatAppRunner: 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(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) + runner.run(MagicMock(), app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1")) assert events == ["close", "invoke", "first-chunk"] mock_handle.assert_called_once_with( 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 2fdb197852a..69522b193bf 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(app_generate_entity, MagicMock(), MagicMock()) + runner.run(MagicMock(), app_generate_entity, 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(app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) 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(app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) runner._handle_invoke_result.assert_not_called() @@ -133,7 +133,7 @@ class TestCompletionAppRunner: mocker.patch.object(module, "ModelInstance", return_value=model_instance) with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant")) + runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant")) dataset_retrieval.retrieve.assert_called_once() assert dataset_retrieval.retrieve.call_args.kwargs["query"] == "query_from_input" @@ -167,7 +167,7 @@ class TestCompletionAppRunner: mocker.patch.object(module.db.session, "close", side_effect=lambda: events.append("close")) with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, queue_manager, MagicMock(id="msg")) + runner.run(MagicMock(), app_generate_entity, queue_manager, MagicMock(id="msg")) assert events == ["close", "invoke", "first-chunk"] runner._handle_invoke_result.assert_called_once_with( @@ -190,7 +190,7 @@ class TestCompletionAppRunner: runner.check_hosting_moderation = MagicMock(return_value=True) with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg")) + runner.run(MagicMock(), app_generate_entity, MagicMock(), MagicMock(id="msg")) assert ( runner.organize_prompt_messages.call_args.kwargs["image_detail_config"] 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 3acf0d46520..de0456851bd 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 @@ -56,6 +56,7 @@ class TestCompletionAppGenerator: def test_generate_invalid_query_type(self, generator): with pytest.raises(ValueError): generator.generate( + session=MagicMock(), app_model=_build_app_model(), user=_build_user(), args={"query": 123, "inputs": {}, "files": []}, @@ -66,6 +67,7 @@ class TestCompletionAppGenerator: def test_generate_override_not_debugger(self, generator): with pytest.raises(ValueError): generator.generate( + session=MagicMock(), app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {}, "files": [], "model_config": {}}, @@ -93,6 +95,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") result = generator.generate( + session=MagicMock(), app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {"a": 1}, "files": [], "trace_session_id": "session-1"}, @@ -126,6 +129,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") result = generator.generate( + session=MagicMock(), app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {"a": 1}, "files": [{"id": "f"}]}, @@ -161,6 +165,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") generator.generate( + session=MagicMock(), app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {}, "files": [], "model_config": override_config}, @@ -177,6 +182,7 @@ class TestCompletionAppGenerator: with pytest.raises(MessageNotExistsError): generator.generate_more_like_this( + session=session, app_model=_build_app_model(), message_id="msg", user=_build_user(), @@ -194,6 +200,7 @@ class TestCompletionAppGenerator: with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( + session=session, app_model=app_model, message_id="msg", user=_build_user(), @@ -211,6 +218,7 @@ class TestCompletionAppGenerator: with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( + session=session, app_model=app_model, message_id="msg", user=_build_user(), @@ -228,6 +236,7 @@ class TestCompletionAppGenerator: with pytest.raises(ValueError): generator.generate_more_like_this( + session=session, app_model=app_model, message_id="msg", user=_build_user(), @@ -275,6 +284,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") result = generator.generate_more_like_this( + session=session, app_model=app_model, message_id="msg", user=_build_user(), @@ -318,6 +328,7 @@ 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/mcp/server/test_streamable_http.py b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py index 57456085c34..3f95736ef92 100644 --- a/api/tests/unit_tests/core/mcp/server/test_streamable_http.py +++ b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py @@ -52,7 +52,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.type", request_type): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCResponse) @@ -68,7 +68,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.type", request_type): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCResponse) @@ -84,7 +84,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.type", request_type): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCResponse) @@ -109,7 +109,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.type", request_type): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCResponse) @@ -132,7 +132,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.type", request_type): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCError) @@ -151,7 +151,9 @@ class TestHandleMCPRequest: # Don't provide end_user to cause ValueError with patch("core.mcp.server.streamable_http.type", request_type): - result = handle_mcp_request(self.app, self.mock_request, self.user_input_form, self.mcp_server, None, 123) + result = handle_mcp_request( + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, None, 123 + ) assert isinstance(result, types.JSONRPCError) assert result.error.code == types.INVALID_PARAMS @@ -166,7 +168,7 @@ class TestHandleMCPRequest: with patch("core.mcp.server.streamable_http.handle_ping", side_effect=Exception("Test error")): with patch("core.mcp.server.streamable_http.type", return_value=types.PingRequest): result = handle_mcp_request( - self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 + Mock(), self.app, self.mock_request, self.user_input_form, self.mcp_server, self.end_user, 123 ) assert isinstance(result, types.JSONRPCError) @@ -228,7 +230,7 @@ class TestIndividualHandlers: mock_response = {"answer": "test answer"} mock_app_generate.generate.return_value = mock_response - result = handle_call_tool(app, mock_request, user_input_form, end_user) + result = handle_call_tool(Mock(), app, mock_request, user_input_form, end_user) assert isinstance(result, types.CallToolResult) assert len(result.content) == 1 @@ -244,7 +246,7 @@ class TestIndividualHandlers: user_input_form: list[VariableEntity] = [] with pytest.raises(ValueError, match="End user not found"): - handle_call_tool(app, mock_request, user_input_form, None) + handle_call_tool(Mock(), app, mock_request, user_input_form, None) class TestUtilityFunctions: 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 e95ca8a7bf8..d665dcd2e45 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,6 +116,7 @@ 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", @@ -142,6 +143,7 @@ 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", @@ -163,6 +165,7 @@ class TestPluginAppBackwardsInvocation: with pytest.raises(ValueError, match="missing query"): PluginAppBackwardsInvocation.invoke_app( + MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -179,6 +182,7 @@ class TestPluginAppBackwardsInvocation: with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_app( + MagicMock(), app_id="app", user_id="user", tenant_id="tenant", @@ -201,6 +205,7 @@ 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", @@ -231,6 +236,7 @@ class TestPluginAppBackwardsInvocation: ) result = PluginAppBackwardsInvocation.invoke_chat_app( + MagicMock(), app=app, user=MagicMock(), conversation_id="conv-1", @@ -251,6 +257,7 @@ 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", @@ -264,6 +271,7 @@ class TestPluginAppBackwardsInvocation: 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", @@ -311,6 +319,7 @@ 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", @@ -327,7 +336,7 @@ class TestPluginAppBackwardsInvocation: ) app = MagicMock(mode=AppMode.COMPLETION) - result = PluginAppBackwardsInvocation.invoke_completion_app(app, MagicMock(), False, {"x": 1}, []) + result = PluginAppBackwardsInvocation.invoke_completion_app(MagicMock(), app, MagicMock(), False, {"x": 1}, []) assert result == {"ok": 1} assert spy.call_count == 1 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 f72351ffa28..7c672570bfa 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 @@ -233,8 +233,10 @@ class TestRetrievalServiceInternals: mock_validate.return_value = "validated-condition" expected_documents = [create_mock_document("external-doc", "external-1", 0.8, provider="external")] mock_fetch.return_value = expected_documents + session = MagicMock() results = RetrievalService.external_retrieve( + session=session, dataset_id="dataset-1", query="test query", external_retrieval_model={"top_k": 3}, @@ -244,6 +246,7 @@ class TestRetrievalServiceInternals: assert results == expected_documents mock_validate.assert_called_once() mock_fetch.assert_called_once_with( + session, "tenant-1", "dataset-1", "test query", @@ -255,7 +258,7 @@ class TestRetrievalServiceInternals: def test_external_retrieve_returns_empty_when_dataset_not_found(self, mock_scalar): mock_scalar.return_value = None - results = RetrievalService.external_retrieve(dataset_id="missing", query="q") + results = RetrievalService.external_retrieve(session=MagicMock(), dataset_id="missing", query="q") assert results == [] 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 a0d047f6eda..ea7750eb6a6 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 @@ -1,6 +1,7 @@ from __future__ import annotations from types import SimpleNamespace +from unittest.mock import MagicMock from uuid import uuid4 import pytest @@ -148,6 +149,7 @@ def test_knowledge_retrieval_grants_returned_segments_to_current_scope(monkeypat with bind_file_access_scope(scope): results = retrieval.knowledge_retrieval( + MagicMock(), KnowledgeRetrievalRequest( tenant_id=tenant_id, user_id=str(uuid4()), @@ -156,7 +158,7 @@ def test_knowledge_retrieval_grants_returned_segments_to_current_scope(monkeypat dataset_ids=[dataset_id], query="desktop picture", retrieval_mode="multiple", - ) + ), ) current_scope = get_current_file_access_scope() 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 f9a58069ab3..b53520f1eaf 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 @@ -2589,7 +2589,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: mock_session.scalars.side_effect = [mock_datasets, mock_documents] # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -2638,7 +2638,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ) as mock_get_metadata: with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -2696,7 +2696,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ) with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[external_doc]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -2739,7 +2739,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: # Mock multiple_retrieve to return empty list with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert result == [] @@ -2779,7 +2779,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ): # Act & Assert with pytest.raises(exc.RateLimitExceededError): - dataset_retrieval.knowledge_retrieval(request) + dataset_retrieval.knowledge_retrieval(MagicMock(), request) def test_knowledge_retrieval_no_available_datasets(self): """ @@ -2813,7 +2813,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: # Mock _get_available_datasets to return empty list with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert result == [] @@ -4113,9 +4113,10 @@ class TestDatasetRetrievalAdditionalHelpers: usage = LLMUsage.from_metadata({"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}) session_scalars = Mock() session_scalars.all.return_value = [metadata_field] + session = MagicMock() + session.scalars.return_value = session_scalars with ( - patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=session_scalars), patch.object(retrieval, "_fetch_model_config", return_value=(model_instance, model_config)), patch.object(retrieval, "_get_prompt_template", return_value=(["prompt"], [])), patch.object(retrieval, "_handle_invoke_result", return_value=('{"metadata_map":[]}', usage)), @@ -4137,6 +4138,7 @@ class TestDatasetRetrievalAdditionalHelpers: ] } result = retrieval._automatic_metadata_filter_func( + session, dataset_ids=["d1"], query="python", tenant_id="tenant-1", @@ -4148,11 +4150,11 @@ class TestDatasetRetrievalAdditionalHelpers: mock_record_usage.assert_called_once_with(usage) with ( - patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=session_scalars), patch.object(retrieval, "_fetch_model_config", side_effect=RuntimeError("boom")), ): with pytest.raises(RuntimeError, match="boom"): retrieval._automatic_metadata_filter_func( + session, dataset_ids=["d1"], query="python", tenant_id="tenant-1", @@ -4163,27 +4165,29 @@ class TestDatasetRetrievalAdditionalHelpers: def test_get_metadata_filter_condition(self, retrieval: DatasetRetrieval) -> None: scalars_result = Mock() scalars_result.all.return_value = [SimpleNamespace(dataset_id="d1", id="doc-1")] + session = MagicMock() + session.scalars.return_value = scalars_result - with patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=scalars_result): - mapping, condition = retrieval.get_metadata_filter_condition( - dataset_ids=["d1"], - query="python", - tenant_id="tenant-1", - user_id="u1", - metadata_filtering_mode="disabled", - metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), - metadata_filtering_conditions=None, - inputs={}, - ) + mapping, condition = retrieval.get_metadata_filter_condition( + session, + dataset_ids=["d1"], + query="python", + tenant_id="tenant-1", + user_id="u1", + metadata_filtering_mode="disabled", + metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), + metadata_filtering_conditions=None, + inputs={}, + ) assert mapping is None assert condition is None automatic_filters = [{"condition": "contains", "metadata_name": "author", "value": "Alice"}] with ( - patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=scalars_result), patch.object(retrieval, "_automatic_metadata_filter_func", return_value=automatic_filters), ): mapping, condition = retrieval.get_metadata_filter_condition( + session, dataset_ids=["d1"], query="python", tenant_id="tenant-1", @@ -4201,35 +4205,35 @@ class TestDatasetRetrievalAdditionalHelpers: logical_operator="and", conditions=[AppCondition(name="author", comparison_operator="contains", value="{{name}}")], ) - with patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=scalars_result): - mapping, condition = retrieval.get_metadata_filter_condition( - dataset_ids=["d1"], - query="python", - tenant_id="tenant-1", - user_id="u1", - metadata_filtering_mode="manual", - metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), - metadata_filtering_conditions=manual_conditions, - inputs={"name": "Alice"}, - ) + mapping, condition = retrieval.get_metadata_filter_condition( + session, + dataset_ids=["d1"], + query="python", + tenant_id="tenant-1", + user_id="u1", + metadata_filtering_mode="manual", + metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), + metadata_filtering_conditions=manual_conditions, + inputs={"name": "Alice"}, + ) assert mapping == {"d1": ["doc-1"]} assert condition is not None assert condition.conditions first_condition = condition.conditions[0] assert first_condition.value == "Alice" - with patch("core.rag.retrieval.dataset_retrieval.db.session.scalars", return_value=scalars_result): - with pytest.raises(ValueError, match="Invalid metadata filtering mode"): - retrieval.get_metadata_filter_condition( - dataset_ids=["d1"], - query="python", - tenant_id="tenant-1", - user_id="u1", - metadata_filtering_mode="unsupported", - metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), - metadata_filtering_conditions=None, - inputs={}, - ) + with pytest.raises(ValueError, match="Invalid metadata filtering mode"): + retrieval.get_metadata_filter_condition( + session, + dataset_ids=["d1"], + query="python", + tenant_id="tenant-1", + user_id="u1", + metadata_filtering_mode="unsupported", + metadata_model_config=AppModelConfig(provider="openai", name="gpt", mode="chat"), + metadata_filtering_conditions=None, + inputs={}, + ) def test_get_available_datasets(self, retrieval: DatasetRetrieval) -> None: session = Mock() @@ -4362,7 +4366,7 @@ class TestKnowledgeRetrievalCoverage: patch.object(retrieval, "_check_knowledge_rate_limit"), patch.object(retrieval, "_get_available_datasets", return_value=[SimpleNamespace(id="d1")]), ): - assert retrieval.knowledge_retrieval(request) == [] + assert retrieval.knowledge_retrieval(MagicMock(), request) == [] def test_raises_when_metadata_model_config_missing(self, retrieval: DatasetRetrieval) -> None: request = KnowledgeRetrievalRequest( @@ -4381,7 +4385,7 @@ class TestKnowledgeRetrievalCoverage: patch.object(retrieval, "_get_available_datasets", return_value=[SimpleNamespace(id="d1")]), ): with pytest.raises(ValueError, match="metadata_model_config is required"): - retrieval.knowledge_retrieval(request) + retrieval.knowledge_retrieval(MagicMock(), request) @pytest.mark.parametrize( ("status", "error_cls"), @@ -4424,7 +4428,7 @@ class TestKnowledgeRetrievalCoverage: ): mock_model_manager.return_value.get_model_instance.return_value = model_instance with pytest.raises(Exception) as exc_info: - retrieval.knowledge_retrieval(request) + retrieval.knowledge_retrieval(MagicMock(), request) mock_model_manager.assert_called_once_with(tenant_id="tenant-1", user_id="user-1") assert error_cls in type(exc_info.value).__name__ @@ -4457,6 +4461,7 @@ class TestRetrieveCoverage: ), ) result = retrieval.retrieve( + MagicMock(), app_id="app-1", user_id="user-1", tenant_id="tenant-1", @@ -4486,6 +4491,7 @@ class TestRetrieveCoverage: with patch("core.rag.retrieval.dataset_retrieval.ModelManager.for_tenant") as mock_model_manager: mock_model_manager.return_value.get_model_instance.return_value = model_instance result = retrieval.retrieve( + MagicMock(), app_id="app-1", user_id="user-1", tenant_id="tenant-1", @@ -4528,6 +4534,7 @@ class TestRetrieveCoverage: ): mock_model_manager.return_value.get_model_instance.return_value = bound_model_instance context, files = retrieval.retrieve( + MagicMock(), app_id="app-1", user_id="user-1", tenant_id="tenant-1", @@ -4542,7 +4549,7 @@ class TestRetrieveCoverage: mock_model_manager.assert_called_once_with(tenant_id="tenant-1", user_id="user-1") mock_single_retrieve.assert_called_once() - assert mock_single_retrieve.call_args.args[8] == PlanningStrategy.ROUTER + assert mock_single_retrieve.call_args.args[9] == PlanningStrategy.ROUTER assert model_config.provider_model_bundle is bound_bundle assert model_config.credentials == {"api_key": "secret"} assert model_config.model_schema is bound_schema @@ -4577,6 +4584,7 @@ class TestRetrieveCoverage: bound_model_instance.model_type_instance.get_model_schema.return_value = SimpleNamespace(features=[]) mock_model_manager.return_value.get_model_instance.return_value = bound_model_instance context, files = retrieval.retrieve( + MagicMock(), app_id="app-1", user_id="user-1", tenant_id="tenant-1", @@ -4658,6 +4666,8 @@ class TestRetrieveCoverage: execute_docs = SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: [document_item])) execute_datasets = SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: [dataset_item])) hit_callback = Mock() + session = MagicMock() + session.execute.side_effect = [execute_attachments, execute_docs, execute_datasets] with ( patch("core.rag.retrieval.dataset_retrieval.ModelManager.for_tenant") as mock_model_manager, @@ -4669,7 +4679,6 @@ class TestRetrieveCoverage: return_value=[record], ), patch("core.rag.retrieval.dataset_retrieval.sign_upload_file_preview_url", return_value="https://signed"), - patch("core.rag.retrieval.dataset_retrieval.db.session.execute") as mock_execute, ): bound_model_instance = Mock() bound_model_instance.model_name = "gpt-4" @@ -4679,8 +4688,8 @@ class TestRetrieveCoverage: features=[ModelFeature.TOOL_CALL] ) mock_model_manager.return_value.get_model_instance.return_value = bound_model_instance - mock_execute.side_effect = [execute_attachments, execute_docs, execute_datasets] context, files = retrieval.retrieve( + session, app_id="app-1", user_id="user-1", tenant_id="tenant-1", @@ -4717,10 +4726,11 @@ class TestSingleAndMultipleRetrieveCoverage: ) app = Flask(__name__) usage = LLMUsage.from_metadata({"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}) + session = MagicMock() + session.scalar.return_value = dataset with app.app_context(): with ( patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls, - patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=dataset), patch( "core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval" ) as mock_external, @@ -4733,6 +4743,7 @@ class TestSingleAndMultipleRetrieveCoverage: {"content": "ext result", "metadata": {"k": "v"}, "score": 0.9, "title": "Ext Doc"} ] result = retrieval.single_retrieve( + session, app_id="app-1", tenant_id="tenant-1", user_id="user-1", @@ -4772,10 +4783,11 @@ class TestSingleAndMultipleRetrieveCoverage: app = Flask(__name__) usage = LLMUsage.from_metadata({"prompt_tokens": 1, "completion_tokens": 0, "total_tokens": 1}) result_doc = _doc(provider="dify", score=0.7, dataset_id="ds-1", document_id="doc-1", doc_id="node-1") + session = MagicMock() + session.scalar.return_value = dataset with app.app_context(): with ( patch("core.rag.retrieval.dataset_retrieval.FunctionCallMultiDatasetRouter") as mock_router_cls, - patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=dataset), patch( "core.rag.retrieval.dataset_retrieval.RetrievalService.retrieve", return_value=[result_doc] ) as mock_retrieve, @@ -4785,6 +4797,7 @@ class TestSingleAndMultipleRetrieveCoverage: ): mock_router_cls.return_value.invoke.return_value = ("ds-1", usage) results = retrieval.single_retrieve( + session, app_id="app-1", tenant_id="tenant-1", user_id="user-1", @@ -4806,6 +4819,7 @@ class TestSingleAndMultipleRetrieveCoverage: with patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls: mock_router_cls.return_value.invoke.return_value = (None, LLMUsage.empty_usage()) results = retrieval.single_retrieve( + MagicMock(), app_id="app-1", tenant_id="tenant-1", user_id="user-1", @@ -4832,10 +4846,12 @@ class TestSingleAndMultipleRetrieveCoverage: ) with ( patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls, - patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=dataset), ): + session = MagicMock() + session.scalar.return_value = dataset mock_router_cls.return_value.invoke.return_value = ("ds-1", LLMUsage.empty_usage()) no_filter = retrieval.single_retrieve( + session, app_id="app-1", tenant_id="tenant-1", user_id="user-1", @@ -4849,6 +4865,7 @@ class TestSingleAndMultipleRetrieveCoverage: metadata_condition=_metadata_condition(), ) missing_doc_ids = retrieval.single_retrieve( + session, app_id="app-1", tenant_id="tenant-1", user_id="user-1", @@ -5105,17 +5122,19 @@ class TestInternalHooksCoverage: flask_app = SimpleNamespace(app_context=lambda: nullcontext()) all_documents: list[Document] = [] - with patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=None): - assert ( - retrieval._retriever( - flask_app=flask_app, # type: ignore[arg-type] - dataset_id="d1", - query="python", - top_k=1, - all_documents=all_documents, - ) - == [] + session = MagicMock() + session.scalar.return_value = None + assert ( + retrieval._retriever( + flask_app=flask_app, # type: ignore[arg-type] + session=session, + dataset_id="d1", + query="python", + top_k=1, + all_documents=all_documents, ) + == [] + ) external_dataset = SimpleNamespace( id="ext-ds", @@ -5126,14 +5145,16 @@ class TestInternalHooksCoverage: indexing_technique="high_quality", ) with ( - patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=external_dataset), patch( "core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval" ) as mock_external, ): + session = MagicMock() + session.scalar.return_value = external_dataset mock_external.return_value = [{"content": "e", "metadata": {}, "score": 0.8, "title": "Ext"}] retrieval._retriever( flask_app=flask_app, # type: ignore[arg-type] + session=session, dataset_id="ext-ds", query="python", top_k=1, @@ -5161,16 +5182,16 @@ class TestInternalHooksCoverage: }, indexing_technique="high_quality", ) + session = MagicMock() + session.scalar.side_effect = [economy_dataset, high_dataset] with ( - patch( - "core.rag.retrieval.dataset_retrieval.db.session.scalar", side_effect=[economy_dataset, high_dataset] - ), patch( "core.rag.retrieval.dataset_retrieval.RetrievalService.retrieve", return_value=[_doc(provider="dify")] ) as mock_retrieve, ): retrieval._retriever( flask_app=flask_app, # type: ignore[arg-type] + session=session, dataset_id="eco-ds", query="python", top_k=2, @@ -5178,6 +5199,7 @@ class TestInternalHooksCoverage: ) retrieval._retriever( flask_app=flask_app, # type: ignore[arg-type] + session=session, dataset_id="hq-ds", query="python", top_k=2, @@ -5199,17 +5221,16 @@ class TestInternalHooksCoverage: retrieve_strategy=DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE, metadata_filtering_mode="disabled", ) + session = MagicMock() + session.scalar.side_effect = [None, dataset_skip_zero, dataset_ok_single] with ( - patch( - "core.rag.retrieval.dataset_retrieval.db.session.scalar", - side_effect=[None, dataset_skip_zero, dataset_ok_single], - ), patch( "core.tools.utils.dataset_retriever.dataset_retriever_tool.DatasetRetrieverTool.from_dataset", return_value="single-tool", ) as mock_single_tool, ): single_tools = retrieval.to_dataset_retriever_tool( + session=session, tenant_id="tenant-1", dataset_ids=["missing", "d1", "d2"], retrieve_config=single_config, @@ -5228,18 +5249,20 @@ class TestInternalHooksCoverage: metadata_filtering_mode="disabled", reranking_model=None, ) - with patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=dataset_ok_single): - with pytest.raises(ValueError, match="Reranking model is required"): - retrieval.to_dataset_retriever_tool( - tenant_id="tenant-1", - dataset_ids=["d2"], - retrieve_config=multiple_config_missing, - return_resource=True, - invoke_from=InvokeFrom.WEB_APP, - hit_callback=Mock(), - user_id="user-1", - inputs={}, - ) + session = MagicMock() + session.scalar.return_value = dataset_ok_single + with pytest.raises(ValueError, match="Reranking model is required"): + retrieval.to_dataset_retriever_tool( + session=session, + tenant_id="tenant-1", + dataset_ids=["d2"], + retrieve_config=multiple_config_missing, + return_resource=True, + invoke_from=InvokeFrom.WEB_APP, + hit_callback=Mock(), + user_id="user-1", + inputs={}, + ) multiple_config = DatasetRetrieveConfigEntity( retrieve_strategy=DatasetRetrieveConfigEntity.RetrieveStrategy.MULTIPLE, @@ -5248,14 +5271,16 @@ class TestInternalHooksCoverage: score_threshold=0.2, reranking_model={"reranking_provider_name": "cohere", "reranking_model_name": "rerank-v3"}, ) + session = MagicMock() + session.scalar.return_value = dataset_ok_single with ( - patch("core.rag.retrieval.dataset_retrieval.db.session.scalar", return_value=dataset_ok_single), patch( "core.tools.utils.dataset_retriever.dataset_multi_retriever_tool.DatasetMultiRetrieverTool.from_dataset", return_value="multi-tool", ) as mock_multi_tool, ): multi_tools = retrieval.to_dataset_retriever_tool( + session=session, tenant_id="tenant-1", dataset_ids=["d2"], retrieve_config=multiple_config, @@ -5281,6 +5306,7 @@ class TestInternalHooksCoverage: mock_scalars.return_value.all.return_value = [] with pytest.raises(ValueError): retrieval._automatic_metadata_filter_func( + MagicMock(), dataset_ids=["d1"], query="python", tenant_id="tenant-1", @@ -5301,6 +5327,7 @@ class TestInternalHooksCoverage: with patch.object(retrieval, "_fetch_model_config", return_value=(model_instance, Mock())): assert ( retrieval._automatic_metadata_filter_func( + MagicMock(), dataset_ids=["d1"], query="python", tenant_id="tenant-1", diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_attachment_entry.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_attachment_entry.py index adcf5585d39..1e41827d316 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_attachment_entry.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_attachment_entry.py @@ -30,7 +30,7 @@ def test_knowledge_retrieval_allows_attachment_only_requests() -> None: patch.object(retrieval, "_get_available_datasets", return_value=[available_dataset]), patch.object(retrieval, "multiple_retrieve", return_value=[]) as mock_multiple, ): - result = retrieval.knowledge_retrieval(request) + result = retrieval.knowledge_retrieval(MagicMock(), request) assert result == [] mock_multiple.assert_called_once() diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py index aace419d154..98413840d00 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py @@ -458,7 +458,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: mock_session.scalars.side_effect = [mock_datasets, mock_documents] # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -507,7 +507,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ) as mock_get_metadata: with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -565,7 +565,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ) with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[external_doc]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert isinstance(result, list) @@ -608,7 +608,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: # Mock multiple_retrieve to return empty list with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert result == [] @@ -648,7 +648,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: ): # Act & Assert with pytest.raises(exc.RateLimitExceededError): - dataset_retrieval.knowledge_retrieval(request) + dataset_retrieval.knowledge_retrieval(MagicMock(), request) def test_knowledge_retrieval_no_available_datasets(self): """ @@ -682,7 +682,7 @@ class TestDatasetRetrievalKnowledgeRetrieval: # Mock _get_available_datasets to return empty list with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[]): # Act - result = dataset_retrieval.knowledge_retrieval(request) + result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) # Assert assert result == [] diff --git a/api/tests/unit_tests/core/tools/test_base_tool.py b/api/tests/unit_tests/core/tools/test_base_tool.py index 9486144e980..9e80e086472 100644 --- a/api/tests/unit_tests/core/tools/test_base_tool.py +++ b/api/tests/unit_tests/core/tools/test_base_tool.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Generator from dataclasses import dataclass from typing import Any, cast +from unittest.mock import MagicMock from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool import Tool @@ -48,6 +49,7 @@ class DummyTool(Tool): def _invoke( self, + session: Any, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, @@ -107,6 +109,7 @@ def test_invoke_supports_single_message_and_parameter_casting(): messages = list( tool.invoke( + session=MagicMock(), user_id="user-1", tool_parameters={"age": "18", "raw": "keep"}, conversation_id="conv-1", @@ -129,7 +132,7 @@ def test_invoke_supports_single_message_and_parameter_casting(): def test_invoke_supports_list_and_generator_results(): tool = _build_tool() tool.result = [tool.create_text_message("a"), tool.create_text_message("b")] - list_messages = list(tool.invoke(user_id="user-1", tool_parameters={})) + list_messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={})) assert [msg.message.text for msg in list_messages] == ["a", "b"] def _message_generator() -> Generator[ToolInvokeMessage, None, None]: @@ -137,7 +140,7 @@ def test_invoke_supports_list_and_generator_results(): yield tool.create_text_message("g2") tool.result = _message_generator() - generated_messages = list(tool.invoke(user_id="user-2", tool_parameters={})) + generated_messages = list(tool.invoke(session=MagicMock(), user_id="user-2", tool_parameters={})) assert [msg.message.text for msg in generated_messages] == ["g1", "g2"] @@ -315,4 +318,4 @@ def test_message_factory_helpers(): def test_base_abstract_invoke_placeholder_returns_none(): tool = _build_tool() - assert Tool._invoke(tool, user_id="u", tool_parameters={}) is None + assert Tool._invoke(tool, session=MagicMock(), user_id="u", tool_parameters={}) is None diff --git a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py index a99fd1f248f..4dac9b7260d 100644 --- a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py +++ b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py @@ -4,6 +4,7 @@ import calendar import math from datetime import date from types import SimpleNamespace +from unittest.mock import MagicMock from zoneinfo import ZoneInfo import pytest @@ -52,17 +53,23 @@ def _raise_runtime_error(*_args: object, **_kwargs: object) -> None: def test_current_time_tool(): current_tool = _build_builtin_tool(CurrentTimeTool) - utc_text = list(current_tool.invoke(user_id="u", tool_parameters={"timezone": "UTC"}))[0].message.text + utc_text = list(current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "UTC"}))[ + 0 + ].message.text assert utc_text - invalid_tz = list(current_tool.invoke(user_id="u", tool_parameters={"timezone": "Invalid/TZ"}))[0].message.text + invalid_tz = list( + current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "Invalid/TZ"}) + )[0].message.text assert "Invalid timezone" in invalid_tz def test_localtime_to_timestamp_tool(): localtime_tool = _build_builtin_tool(LocaltimeToTimestampTool) ts_message = list( - localtime_tool.invoke(user_id="u", tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"}) + localtime_tool.invoke( + session=MagicMock(), user_id="u", tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"} + ) )[0].message.text ts_value = float(ts_message.strip()) assert math.isfinite(ts_value) @@ -87,9 +94,11 @@ def test_localtime_to_timestamp_tool(): def test_timestamp_to_localtime_tool(): to_local_tool = _build_builtin_tool(TimestampToLocaltimeTool) - local_text = list(to_local_tool.invoke(user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"}))[ - 0 - ].message.text + local_text = list( + to_local_tool.invoke( + session=MagicMock(), user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"} + ) + )[0].message.text assert "2024" in local_text with pytest.raises(ToolInvokeError): TimestampToLocaltimeTool.timestamp_to_localtime("bad", "UTC") # type: ignore[arg-type] @@ -99,6 +108,7 @@ def test_timezone_conversion_tool(): timezone_tool = _build_builtin_tool(TimezoneConversionTool) converted = list( timezone_tool.invoke( + session=MagicMock(), user_id="u", tool_parameters={ "current_time": "2024-01-01 08:00:00", @@ -114,7 +124,9 @@ def test_timezone_conversion_tool(): def test_weekday_tool(): weekday_tool = _build_builtin_tool(WeekdayTool) - valid = list(weekday_tool.invoke(user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}))[0].message.text + valid = list( + weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}) + )[0].message.text expected_date = date(2024, 1, 1) expected_message = ( f"{calendar.month_name[expected_date.month]} " @@ -122,12 +134,12 @@ def test_weekday_tool(): f"is {calendar.day_name[expected_date.weekday()]}." ) assert valid == expected_message - invalid = list(weekday_tool.invoke(user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}))[ - 0 - ].message.text + invalid = list( + weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}) + )[0].message.text assert "Invalid date" in invalid with pytest.raises(ValueError, match="Month is required"): - list(weekday_tool.invoke(user_id="u", tool_parameters={"year": 2024, "day": 1})) + list(weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "day": 1})) def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): @@ -139,6 +151,7 @@ def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): ) result = list( simple_code.invoke( + session=MagicMock(), user_id="u", tool_parameters={"language": "python3", "code": "print(1)"}, ) @@ -150,7 +163,11 @@ def test_simple_code_invalid_language(): simple_code = _build_builtin_tool(SimpleCode) with pytest.raises(ValueError, match="Only python3 and javascript"): - list(simple_code.invoke(user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"})) + list( + simple_code.invoke( + session=MagicMock(), user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"} + ) + ) def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): @@ -161,19 +178,25 @@ def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): _raise_runtime_error, ) with pytest.raises(ToolInvokeError, match="boom"): - list(simple_code.invoke(user_id="u", tool_parameters={"language": "python3", "code": "print(1)"})) + list( + simple_code.invoke( + session=MagicMock(), user_id="u", tool_parameters={"language": "python3", "code": "print(1)"} + ) + ) def test_webscraper_empty_url(): webscraper = _build_builtin_tool(WebscraperTool) - empty = list(webscraper.invoke(user_id="u", tool_parameters={"url": ""}))[0].message.text + empty = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": ""}))[0].message.text assert empty == "Please input url" def test_webscraper_fetch(monkeypatch: pytest.MonkeyPatch): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr("core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", lambda *a, **k: "page") - full = list(webscraper.invoke(user_id="u", tool_parameters={"url": "https://example.com"}))[0].message.text + full = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"}))[ + 0 + ].message.text assert full == "page" @@ -183,6 +206,7 @@ def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(webscraper, "summary", lambda user_id, content: "summary") summarized = list( webscraper.invoke( + session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com", "generate_summary": True}, ) @@ -197,13 +221,15 @@ def test_webscraper_fetch_error(monkeypatch: pytest.MonkeyPatch): _raise_runtime_error, ) with pytest.raises(ToolInvokeError, match="boom"): - list(webscraper.invoke(user_id="u", tool_parameters={"url": "https://example.com"})) + list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"})) def test_asr_invalid_file(): asr = _build_builtin_tool(ASRTool) file_obj = SimpleNamespace(type=FileType.DOCUMENT) - invalid_file = list(asr.invoke(user_id="u", tool_parameters={"audio_file": file_obj}))[0].message.text + invalid_file = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": file_obj}))[ + 0 + ].message.text assert "not a valid audio file" in invalid_file @@ -219,7 +245,9 @@ def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch): lambda **kwargs: captured_manager_kwargs.update(kwargs) or model_manager, ) audio_file = SimpleNamespace(type=FileType.AUDIO) - ok = list(asr.invoke(user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}))[0].message.text + ok = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}))[ + 0 + ].message.text assert ok == "transcript" assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -253,7 +281,7 @@ def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch): or type("M", (), {"get_model_instance": lambda *a, **k: voices_model_instance})() ), ) - messages = list(tts.invoke(user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + messages = list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.TEXT, ToolInvokeMessage.MessageType.BLOB] assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -269,7 +297,7 @@ def test_tts_tool_raises_when_runtime_missing(): tts = _build_builtin_tool(TTSTool) tts.runtime = None with pytest.raises(ValueError, match="Runtime is required"): - list(tts.invoke(user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) @pytest.mark.parametrize( @@ -292,7 +320,7 @@ def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices): lambda **_: type("Manager", (), {"get_model_instance": lambda *args, **kwargs: model_without_voice})(), ) with pytest.raises(ValueError, match="no voice available"): - list(tts.invoke(user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) def test_tts_tool_get_available_models_and_runtime_parameters(monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/core/tools/test_custom_tool.py b/api/tests/unit_tests/core/tools/test_custom_tool.py index f525baeaf23..3f664456bd2 100644 --- a/api/tests/unit_tests/core/tools/test_custom_tool.py +++ b/api/tests/unit_tests/core/tools/test_custom_tool.py @@ -2,6 +2,7 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any +from unittest.mock import MagicMock import httpx import pytest @@ -276,11 +277,11 @@ def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytes monkeypatch.setattr(tool, "assembling_request", lambda parameters: {}) monkeypatch.setattr(tool, "do_http_request", lambda *args, **kwargs: httpx.Response(200, text='{"a":1}')) monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse({"a": 1}, True)) - messages = list(tool.invoke(user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.JSON, ToolInvokeMessage.MessageType.TEXT] # _invoke text path monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse("plain", False)) - messages = list(tool.invoke(user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) assert len(messages) == 1 assert messages[0].message.text == "plain" diff --git a/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py b/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py index 23c0be9487d..9dece8cfe07 100644 --- a/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py +++ b/api/tests/unit_tests/core/tools/test_dataset_retriever_tool.py @@ -3,7 +3,7 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch from core.app.app_config.entities import DatasetRetrieveConfigEntity from core.app.entities.app_invoke_entities import InvokeFrom @@ -20,6 +20,7 @@ def test_get_dataset_tools_returns_empty_for_empty_dataset_ids() -> None: # Act tools = DatasetRetrieverTool.get_dataset_tools( + session=MagicMock(), tenant_id="tenant", dataset_ids=[], retrieve_config=retrieve_config, @@ -40,6 +41,7 @@ def test_get_dataset_tools_returns_empty_for_missing_retrieve_config() -> None: # Act tools = DatasetRetrieverTool.get_dataset_tools( + session=MagicMock(), tenant_id="tenant", dataset_ids=dataset_ids, retrieve_config=None, # type: ignore[arg-type] @@ -64,6 +66,7 @@ def test_get_dataset_tools_builds_tool_and_restores_strategy() -> None: # Act with patch("core.tools.utils.dataset_retriever_tool.DatasetRetrieval", return_value=feature): tools = DatasetRetrieverTool.get_dataset_tools( + session=MagicMock(), tenant_id="tenant", dataset_ids=["d1"], retrieve_config=retrieve_config, @@ -81,11 +84,16 @@ def test_get_dataset_tools_builds_tool_and_restores_strategy() -> None: def _build_dataset_tool() -> tuple[DatasetRetrieverTool, SimpleNamespace]: - retrieval_tool = SimpleNamespace(name="dataset_tool", description="desc", run=lambda query: f"result:{query}") + retrieval_tool = SimpleNamespace( + name="dataset_tool", + description="desc", + run=lambda session, query: f"result:{query}", + ) feature = Mock() feature.to_dataset_retriever_tool.return_value = [retrieval_tool] with patch("core.tools.utils.dataset_retriever_tool.DatasetRetrieval", return_value=feature): tools = DatasetRetrieverTool.get_dataset_tools( + session=MagicMock(), tenant_id="tenant", dataset_ids=["d1"], retrieve_config=_retrieve_config(), @@ -115,7 +123,7 @@ def test_empty_query_behavior() -> None: tool, _ = _build_dataset_tool() # Act - empty_query = list(tool.invoke(user_id="u", tool_parameters={})) + empty_query = list(tool.invoke(session=MagicMock(), user_id="u", tool_parameters={})) # Assert assert len(empty_query) == 1 @@ -127,7 +135,7 @@ def test_query_invocation_result() -> None: tool, _ = _build_dataset_tool() # Act - result = list(tool.invoke(user_id="u", tool_parameters={"query": "hello"})) + result = list(tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"query": "hello"})) # Assert assert len(result) == 1 diff --git a/api/tests/unit_tests/core/tools/test_mcp_tool.py b/api/tests/unit_tests/core/tools/test_mcp_tool.py index 2e2b961bf2a..984794ae967 100644 --- a/api/tests/unit_tests/core/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/core/tools/test_mcp_tool.py @@ -1,7 +1,7 @@ from __future__ import annotations import base64 -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest @@ -123,7 +123,7 @@ def test_mcp_tool_invoke_handles_content_types_and_structured_output(): ) with patch.object(MCPTool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool.invoke(user_id="user-1", tool_parameters={"a": 1})) + messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"a": 1})) types = [m.type for m in messages] assert ToolInvokeMessage.MessageType.JSON in types @@ -141,7 +141,7 @@ def test_mcp_tool_invoke_raises_for_unsupported_embedded_resource(): with patch.object(MCPTool, "invoke_remote_mcp_tool", return_value=result): with pytest.raises(ToolInvokeError, match="Unsupported embedded resource type"): - list(tool.invoke(user_id="user-1", tool_parameters={})) + list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={})) def test_mcp_tool_handle_none_parameter_filters_empty_values(): diff --git a/api/tests/unit_tests/core/tools/test_plugin_tool.py b/api/tests/unit_tests/core/tools/test_plugin_tool.py index 4378432a0f0..d974b004e43 100644 --- a/api/tests/unit_tests/core/tools/test_plugin_tool.py +++ b/api/tests/unit_tests/core/tools/test_plugin_tool.py @@ -1,6 +1,6 @@ from __future__ import annotations -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -47,7 +47,7 @@ def test_plugin_tool_invoke_and_fork_runtime(): "core.tools.plugin_tool.tool.convert_parameters_to_plugin_format", return_value={"converted": 1}, ): - messages = list(tool.invoke(user_id="user-1", tool_parameters={"raw": 1})) + messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"raw": 1})) assert [m.message.text for m in messages] == ["ok"] manager.invoke.assert_called_once() diff --git a/api/tests/unit_tests/core/tools/test_tool_engine.py b/api/tests/unit_tests/core/tools/test_tool_engine.py index ec1da418228..f38ab2a2fab 100644 --- a/api/tests/unit_tests/core/tools/test_tool_engine.py +++ b/api/tests/unit_tests/core/tools/test_tool_engine.py @@ -42,6 +42,7 @@ class _DummyTool(Tool): def _invoke( self, + session: Any, user_id: str, tool_parameters: dict[str, Any], conversation_id: str | None = None, @@ -152,7 +153,7 @@ def test_create_message_files_and_invoke_generator(): mock_db.session.close.assert_not_called() tool = _build_tool() - invoked = list(ToolEngine._invoke(tool, {"a": 1}, user_id="u")) + invoked = list(ToolEngine._invoke(MagicMock(), tool, {"a": 1}, user_id="u")) assert invoked[0].type == ToolInvokeMessage.MessageType.TEXT assert isinstance(invoked[-1], ToolInvokeMeta) assert invoked[-1].error is None @@ -164,6 +165,7 @@ def test_generic_invoke_success_and_error_paths(): callback.on_tool_execution.side_effect = lambda **kwargs: kwargs["tool_outputs"] response = list( ToolEngine.generic_invoke( + session=MagicMock(), tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -184,6 +186,7 @@ def test_generic_invoke_success_and_error_paths(): with pytest.raises(RuntimeError, match="boom"): list( ToolEngine.generic_invoke( + session=MagicMock(), tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -208,6 +211,7 @@ def test_agent_invoke_success(): with patch.object(ToolEngine, "_extract_tool_response_binary_and_text", return_value=iter([])): with patch.object(ToolEngine, "_create_message_files", return_value=[]): result_text, message_files, result_meta = ToolEngine.agent_invoke( + session=MagicMock(), tool=tool, tool_parameters="hello", user_id="u1", @@ -231,6 +235,7 @@ def test_agent_invoke_param_validation_error(): with patch.object(ToolEngine, "_invoke", side_effect=ToolParameterValidationError("bad-param")): error_text, files, error_meta = ToolEngine.agent_invoke( + session=MagicMock(), tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -253,6 +258,7 @@ def test_agent_invoke_engine_meta_error(): with patch.object(ToolEngine, "_invoke", side_effect=engine_error): error_text, files, error_meta = ToolEngine.agent_invoke( + session=MagicMock(), tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -296,6 +302,7 @@ def test_agent_invoke_tool_invoke_error(): with patch.object(ToolEngine, "_invoke", side_effect=ToolInvokeError("invoke boom")): error_text, files, _ = ToolEngine.agent_invoke( + session=MagicMock(), tool=tool, tool_parameters={"a": 1}, user_id="u1", 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 cb3be81ab61..5829098f6b4 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 Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from yaml import YAMLError @@ -159,7 +159,7 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources(): "fetch_external_knowledge_retrieval", return_value=external_documents, ) as fetch_mock: - result = tool.run(query="hello") + result = tool.run(session=MagicMock(), query="hello") assert result == "first\nsecond" assert callback.queries == [("hello", "dataset-1")] @@ -197,7 +197,7 @@ def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_do 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(query="hello") + result = tool.run(session=MagicMock(), query="hello") assert result == "" retrieve_mock.assert_not_called() @@ -284,7 +284,7 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): "format_retrieval_documents", return_value=records, ): - result = tool.run(query="hello") + result = tool.run(session=MagicMock(), query="hello") assert result == "signed high\nsummary low\nquestion:signed low answer:low answer" assert callback.documents == documents @@ -464,7 +464,7 @@ def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(): 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(query="hello") + result = tool.run(session=MagicMock(), query="hello") assert result == "signed one\nquestion:signed two answer:answer two" assert retriever_mock.call_count == 2 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 c9e8a1aaf15..5cb6acf50a7 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 @@ -122,7 +122,7 @@ def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_fiel with pytest.raises(ToolInvokeError) as exc_info: # WorkflowTool always returns a generator, so we need to iterate to # actually `run` the tool. - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) assert exc_info.value.args == ("oops",) @@ -140,7 +140,7 @@ def test_workflow_tool_does_not_use_pause_state_config(monkeypatch: pytest.Monke monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "pause_state_config" in call_kwargs @@ -165,7 +165,7 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["parent_trace_context"].model_dump() == { @@ -197,7 +197,7 @@ def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.Monkey monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {"trace_session_id": "user-input-session"})) + list(tool.invoke(MagicMock(), "test_user", {"trace_session_id": "user-input-session"})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["inputs"]["trace_session_id"] == "user-input-session" @@ -238,6 +238,7 @@ def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypat list( tool.invoke( + MagicMock(), "test_user", { "outer_workflow_run_id": "user-workflow-input", @@ -274,7 +275,7 @@ def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.Monkey monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "parent_trace_context" not in call_kwargs["args"] @@ -296,7 +297,7 @@ def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatc monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "trace_session_id" not in call_kwargs["args"] @@ -329,7 +330,7 @@ def test_workflow_tool_omits_parent_trace_context_when_runtime_is_incomplete( monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke("test_user", {})) + list(tool.invoke(MagicMock(), "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "parent_trace_context" not in call_kwargs["args"] @@ -358,7 +359,7 @@ def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) # Execute tool invocation - messages = list(tool.invoke("test_user", {})) + messages = list(tool.invoke(MagicMock(), "test_user", {})) # Verify variable messages variable_messages = [msg for msg in messages if msg.type == ToolInvokeMessage.MessageType.VARIABLE] @@ -401,7 +402,7 @@ def test_workflow_tool_should_handle_empty_outputs(monkeypatch: pytest.MonkeyPat monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) # Execute tool invocation - messages = list(tool.invoke("test_user", {})) + messages = list(tool.invoke(MagicMock(), "test_user", {})) # Verify generated messages # Should contain: 0 variable messages + 1 text message + 1 JSON message = 2 messages @@ -551,7 +552,7 @@ def test_invoke_raises_when_user_not_found(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(tool, "_resolve_user", lambda *args, **kwargs: None) with pytest.raises(ToolInvokeError, match="User not found"): - list(tool.invoke("missing", {})) + list(tool.invoke(MagicMock(), "missing", {})) def test_resolve_user_from_database_returns_account(monkeypatch: pytest.MonkeyPatch): @@ -737,7 +738,7 @@ def test_workflow_tool_invocation_normalizes_optional_files_parameter(monkeypatc generate_mock = MagicMock(return_value={"data": {}}) monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) - list(tool.invoke("test_user", {"images": None})) + list(tool.invoke(MagicMock(), "test_user", {"images": None})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["inputs"]["images"] == [] diff --git a/api/tests/unit_tests/core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py b/api/tests/unit_tests/core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py index 3c821e75ba3..ce7984ea26f 100644 --- a/api/tests/unit_tests/core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py @@ -488,7 +488,7 @@ class TestFetchDatasetRetriever: ) # Act - results, usage = node._fetch_dataset_retriever(node_data=node_data, variables=variables) + results, usage = node._fetch_dataset_retriever(Mock(), node_data=node_data, variables=variables) # Assert assert len(results) == 1 @@ -525,7 +525,7 @@ class TestFetchDatasetRetriever: ) # Act - results, usage = node._fetch_dataset_retriever(node_data=sample_node_data, variables=variables) + results, usage = node._fetch_dataset_retriever(Mock(), node_data=sample_node_data, variables=variables) # Assert assert isinstance(results, list) @@ -580,7 +580,7 @@ class TestFetchDatasetRetriever: ) # Act - results, usage = node._fetch_dataset_retriever(node_data=node_data, variables=variables) + results, usage = node._fetch_dataset_retriever(Mock(), node_data=node_data, variables=variables) # Assert assert isinstance(results, list) @@ -692,7 +692,7 @@ class TestFetchDatasetRetriever: mock_rag_retrieval.llm_usage = LLMUsage.empty_usage() # Act - node._fetch_dataset_retriever(node_data=node_data, variables=variables) + node._fetch_dataset_retriever(Mock(), node_data=node_data, variables=variables) # Assert the passed request has resolved value call_args = mock_rag_retrieval.knowledge_retrieval.call_args diff --git a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py index f0438376437..501225fdbab 100644 --- a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node_runtime.py @@ -44,7 +44,10 @@ def runtime(monkeypatch) -> DifyToolNodeRuntime: invoke_from="debugger", call_depth=0, ) - return DifyToolNodeRuntime(init_params.run_context) + session_maker = MagicMock() + session_maker.begin.return_value.__enter__.return_value = MagicMock(name="session") + session_maker.begin.return_value.__exit__.return_value = None + return DifyToolNodeRuntime(init_params.run_context, session_maker=session_maker) def _build_tool_node_data() -> ToolNodeData: @@ -106,6 +109,7 @@ def test_invoke_creates_callback_and_converts_messages(runtime: DifyToolNodeRunt callback = generic_invoke_mock.call_args.kwargs["workflow_tool_callback"] assert isinstance(callback, DifyWorkflowCallbackHandler) + assert generic_invoke_mock.call_args.kwargs["session"] is not None assert generic_invoke_mock.call_args.kwargs["conversation_id"] == "conversation-id" transform_kwargs = transform_tool_messages.call_args.kwargs @@ -117,11 +121,13 @@ def test_invoke_maps_plugin_errors_to_graph_errors(runtime: DifyToolNodeRuntime) with patch.object(ToolEngine, "generic_invoke", side_effect=invoke_error): with pytest.raises(ToolRuntimeInvocationError, match="An error occurred in the provider"): - runtime.invoke( - tool_runtime=ToolRuntimeHandle(raw=MagicMock()), - tool_parameters={}, - workflow_call_depth=0, - provider_name="provider", + list( + runtime.invoke( + tool_runtime=ToolRuntimeHandle(raw=MagicMock()), + tool_parameters={}, + workflow_call_depth=0, + provider_name="provider", + ) ) diff --git a/api/tests/unit_tests/extensions/otel/decorators/handlers/test_generate_handler.py b/api/tests/unit_tests/extensions/otel/decorators/handlers/test_generate_handler.py index 12e91f190f1..b83602e98cf 100644 --- a/api/tests/unit_tests/extensions/otel/decorators/handlers/test_generate_handler.py +++ b/api/tests/unit_tests/extensions/otel/decorators/handlers/test_generate_handler.py @@ -6,7 +6,7 @@ Test objectives: 2. Verify span attribute mapping correctness """ -from unittest.mock import patch +from unittest.mock import MagicMock, patch from core.app.entities.app_invoke_entities import InvokeFrom from extensions.otel.decorators.handlers.generate_handler import AppGenerateHandler @@ -31,6 +31,7 @@ class TestAppGenerateHandler: handler = AppGenerateHandler() kwargs = { + "session": MagicMock(), "app_model": mock_app_model, "user": mock_account_user, "args": {"workflow_id": "test-wf-123"}, diff --git a/api/tests/unit_tests/services/test_agent_tool_inner_service.py b/api/tests/unit_tests/services/test_agent_tool_inner_service.py index 61049d29e9e..93e20d267da 100644 --- a/api/tests/unit_tests/services/test_agent_tool_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_tool_inner_service.py @@ -70,7 +70,7 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None: side_effect=lambda messages, **_kwargs: messages, ), ): - response = AgentToolInnerService().invoke(_request(), session=session) + response = AgentToolInnerService().invoke(session, _request()) assert response.observation == "ok" assert response.metadata == { @@ -89,7 +89,7 @@ def test_invoke_raises_app_not_found_when_session_has_no_app() -> None: session.get.return_value = None with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(_request(), session=session) + AgentToolInnerService().invoke(session, _request()) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -102,7 +102,7 @@ def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant() -> session.get.return_value = fake_app with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(_request(), session=session) + AgentToolInnerService().invoke(session, _request()) assert exc_info.value.error_code == "app_tenant_mismatch" assert exc_info.value.status_code == 403 @@ -120,7 +120,7 @@ def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_co patch("services.agent_tool_inner_service.ToolEngine.generic_invoke", side_effect=ValueError("app not found")), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(_request(), session=session) + AgentToolInnerService().invoke(session, _request()) assert exc_info.value.error_code == "app_not_found" assert exc_info.value.status_code == 404 @@ -141,7 +141,7 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N ), ): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(_request(), session=session) + AgentToolInnerService().invoke(session, _request()) assert exc_info.value.error_code == "agent_tool_invoke_failed" @@ -161,6 +161,6 @@ def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(error: Excepti with patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", side_effect=error): with pytest.raises(AgentToolInnerServiceError) as exc_info: - AgentToolInnerService().invoke(_request(), session=session) + AgentToolInnerService().invoke(session, _request()) assert exc_info.value.error_code == expected_code 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 f34afb243eb..865410fddf1 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -235,6 +235,7 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -255,6 +256,7 @@ class TestGenerate: side_effect=lambda x: x, ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.AGENT_CHAT), user=_make_user(), args={"inputs": {}}, @@ -276,6 +278,7 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=True) result = AppGenerateService.generate( + MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, @@ -297,6 +300,7 @@ class TestGenerate: ) app = _make_app(AppMode.CHAT, is_agent=False) result = AppGenerateService.generate( + MagicMock(), app_model=app, user=_make_user(), args={"inputs": {}}, @@ -338,6 +342,7 @@ class TestGenerate: ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, @@ -370,6 +375,7 @@ class TestGenerate: ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.ADVANCED_CHAT), user=_make_user(), args={"workflow_id": None, "query": "hi", "inputs": {}}, @@ -395,6 +401,7 @@ class TestGenerate: ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, @@ -428,6 +435,7 @@ class TestGenerate: ) result = AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.WORKFLOW), user=_make_user(), args={"inputs": {}}, @@ -443,6 +451,7 @@ class TestGenerate: app = _make_app("invalid-mode", is_agent=False) with pytest.raises(ValueError, match="Invalid app mode"): AppGenerateService.generate( + MagicMock(), app_model=app, user=_make_user(), args={}, @@ -480,6 +489,7 @@ class TestGenerateBilling: ) AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -503,6 +513,7 @@ class TestGenerateBilling: with pytest.raises(InvokeRateLimitError): AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -528,6 +539,7 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -559,6 +571,7 @@ class TestGenerateBilling: ) AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -656,6 +669,7 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -687,6 +701,7 @@ class TestGenerateBilling: with pytest.raises(RuntimeError, match="boom"): AppGenerateService.generate( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), args={"inputs": {}}, @@ -874,6 +889,7 @@ class TestGenerateMoreLikeThis: return_value={"result": "similar"}, ) result = AppGenerateService.generate_more_like_this( + MagicMock(), app_model=_make_app(AppMode.COMPLETION), user=_make_user(), message_id="msg-1", 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 6560e5a1c0b..3d623e8abf8 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -344,9 +344,7 @@ class TestDatasetServiceCreationAndUpdate: mock_db.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 - ) + DatasetService.create_empty_dataset(mock_db.session, "tenant-1", "Dataset", None, "economy", account) def test_create_empty_dataset_uses_default_embedding_model_for_high_quality_dataset(self): account = SimpleNamespace(id="user-1") @@ -514,7 +512,7 @@ class TestDatasetServiceCreationAndUpdate: session = MagicMock() with patch.object(DatasetService, "get_dataset", return_value=None): with pytest.raises(ValueError, match="Dataset not found"): - DatasetService.update_dataset("dataset-1", {}, SimpleNamespace(id="user-1"), session) + DatasetService.update_dataset(session, "dataset-1", {}, SimpleNamespace(id="user-1")) def test_update_dataset_raises_when_new_name_conflicts(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1", tenant_id="tenant-1") @@ -526,10 +524,10 @@ class TestDatasetServiceCreationAndUpdate: ): with pytest.raises(ValueError, match="Dataset name already exists"): DatasetService.update_dataset( + MagicMock(), "dataset-1", {"name": "New Dataset"}, SimpleNamespace(id="user-1"), - MagicMock(), ) def test_update_dataset_routes_external_datasets_to_external_helper(self): @@ -543,7 +541,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_external_dataset", return_value="updated") as update_external, ): session = MagicMock() - result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session) + result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) assert result == "updated" check_permission.assert_called_once() @@ -562,7 +560,7 @@ class TestDatasetServiceCreationAndUpdate: patch.object(DatasetService, "_update_internal_dataset", return_value="updated") as update_internal, ): session = MagicMock() - result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user, session) + result = DatasetService.update_dataset(session, "dataset-1", {"name": dataset.name}, user) assert result == "updated" check_permission.assert_called_once() @@ -614,7 +612,7 @@ class TestDatasetServiceCreationAndUpdate: assert dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM assert dataset.updated_by == "user-1" assert dataset.updated_at is now - get_external_knowledge_api.assert_called_once_with("api-1", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with(mock_db.session, "api-1", dataset.tenant_id) 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() @@ -654,7 +652,7 @@ class TestDatasetServiceCreationAndUpdate: mock_db.session, ) - get_external_knowledge_api.assert_called_once_with("foreign-api", dataset.tenant_id) + get_external_knowledge_api.assert_called_once_with(mock_db.session, "foreign-api", dataset.tenant_id) update_binding.assert_not_called() mock_db.session.commit.assert_not_called() @@ -1459,7 +1457,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="does not have permission"): - DatasetPermissionService.check_permission(user, dataset, "all_team", [], session) + DatasetPermissionService.check_permission(session, user, dataset, "all_team", []) def test_check_permission_prevents_dataset_operator_from_changing_permission_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1467,7 +1465,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(NoPermissionError, match="cannot change the dataset permissions"): - DatasetPermissionService.check_permission(user, dataset, "only_me", [], session) + DatasetPermissionService.check_permission(session, user, dataset, "only_me", []) def test_check_permission_requires_partial_member_list_for_partial_members_mode(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1475,7 +1473,7 @@ class TestDatasetPermissionService: session = MagicMock() with pytest.raises(ValueError, match="Partial member list is required"): - DatasetPermissionService.check_permission(user, dataset, "partial_members", [], session) + DatasetPermissionService.check_permission(session, user, dataset, "partial_members", []) def test_check_permission_rejects_dataset_operator_member_list_changes(self): user = SimpleNamespace(is_dataset_editor=True, is_dataset_operator=True) @@ -1487,11 +1485,11 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): with pytest.raises(ValueError, match="cannot change the dataset permissions"): DatasetPermissionService.check_permission( + session, user, dataset, "partial_members", [{"user_id": "user-2"}], - session, ) def test_check_permission_allows_dataset_operator_when_member_list_is_unchanged(self): @@ -1503,11 +1501,11 @@ class TestDatasetPermissionService: with patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["user-1"]): DatasetPermissionService.check_permission( + session, user, dataset, "partial_members", [{"user_id": "user-1"}], - session, ) def test_clear_partial_member_list_rolls_back_on_exception(self): 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 248a2b1f68a..453ed7560c5 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -449,7 +449,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args) + result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, mock_db.session) # Assert assert result.name == "Test API" @@ -474,7 +474,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) # Assert assert result.name == "Minimal API" @@ -490,7 +490,7 @@ class TestExternalDatasetServiceCreateAPI: # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) @patch("services.external_knowledge_service.db") def test_create_external_knowledge_api_none_settings(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): @@ -500,7 +500,7 @@ class TestExternalDatasetServiceCreateAPI: # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") @@ -517,7 +517,7 @@ class TestExternalDatasetServiceCreateAPI: args = {"name": "Test API", "settings": settings} # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) # Assert assert isinstance(result.settings, str) @@ -538,7 +538,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) # Assert assert result.name == "测试API" @@ -559,7 +559,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) # Assert assert result.description == long_description @@ -833,7 +833,7 @@ class TestExternalDatasetServiceGetAPI: # Act tenant_id = "tenant-123" - result = ExternalDatasetService.get_external_knowledge_api(api_id, tenant_id) + result = ExternalDatasetService.get_external_knowledge_api(mock_db.session, api_id, tenant_id) # Assert assert result.id == api_id @@ -846,7 +846,7 @@ class TestExternalDatasetServiceGetAPI: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.get_external_knowledge_api("nonexistent-id", "tenant-123") + ExternalDatasetService.get_external_knowledge_api(mock_db.session, "nonexistent-id", "tenant-123") class TestExternalDatasetServiceUpdateAPI: @@ -876,7 +876,7 @@ class TestExternalDatasetServiceUpdateAPI: mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(tenant_id, user_id, api_id, args) + result = ExternalDatasetService.update_external_knowledge_api(mock_db.session, tenant_id, user_id, api_id, args) # Assert assert result.name == "Updated API" @@ -908,7 +908,9 @@ class TestExternalDatasetServiceUpdateAPI: mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(tenant_id, "user-123", api_id, args) + result = ExternalDatasetService.update_external_knowledge_api( + mock_db.session, tenant_id, "user-123", api_id, args + ) # Assert settings = json.loads(result.settings) @@ -924,7 +926,9 @@ class TestExternalDatasetServiceUpdateAPI: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.update_external_knowledge_api("tenant-123", "user-123", "api-123", args) + ExternalDatasetService.update_external_knowledge_api( + mock_db.session, "tenant-123", "user-123", "api-123", args + ) @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_tenant_mismatch( @@ -938,7 +942,9 @@ class TestExternalDatasetServiceUpdateAPI: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.update_external_knowledge_api("wrong-tenant", "user-123", "api-123", args) + ExternalDatasetService.update_external_knowledge_api( + mock_db.session, "wrong-tenant", "user-123", "api-123", args + ) @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_name_only(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): @@ -954,7 +960,9 @@ class TestExternalDatasetServiceUpdateAPI: mock_db.session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api("tenant-123", "user-123", "api-123", args) + result = ExternalDatasetService.update_external_knowledge_api( + mock_db.session, "tenant-123", "user-123", "api-123", args + ) # Assert assert result.name == "New Name Only" @@ -975,7 +983,7 @@ class TestExternalDatasetServiceDeleteAPI: mock_db.session.scalar.return_value = existing_api # Act - ExternalDatasetService.delete_external_knowledge_api(tenant_id, api_id) + ExternalDatasetService.delete_external_knowledge_api(mock_db.session, tenant_id, api_id) # Assert mock_db.session.delete.assert_called_once_with(existing_api) @@ -989,7 +997,7 @@ class TestExternalDatasetServiceDeleteAPI: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api("tenant-123", "api-123") + ExternalDatasetService.delete_external_knowledge_api(mock_db.session, "tenant-123", "api-123") @patch("services.external_knowledge_service.db") def test_delete_external_knowledge_api_tenant_mismatch( @@ -1001,7 +1009,7 @@ class TestExternalDatasetServiceDeleteAPI: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api("wrong-tenant", "api-123") + ExternalDatasetService.delete_external_knowledge_api(mock_db.session, "wrong-tenant", "api-123") class TestExternalDatasetServiceAPIUseCheck: @@ -1019,7 +1027,7 @@ class TestExternalDatasetServiceAPIUseCheck: mock_db.session.scalar.return_value = 1 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) # Assert assert in_use is True @@ -1038,7 +1046,7 @@ class TestExternalDatasetServiceAPIUseCheck: mock_db.session.scalar.return_value = 10 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) # Assert assert in_use is True @@ -1054,7 +1062,7 @@ class TestExternalDatasetServiceAPIUseCheck: mock_db.session.scalar.return_value = 0 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) # Assert assert in_use is False @@ -1076,7 +1084,9 @@ class TestExternalDatasetServiceGetBinding: mock_db.session.scalar.return_value = expected_binding # Act - result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id(tenant_id, dataset_id) + result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( + mock_db.session, tenant_id, dataset_id + ) # Assert assert result.dataset_id == dataset_id @@ -1090,7 +1100,9 @@ class TestExternalDatasetServiceGetBinding: # Act & Assert with pytest.raises(ValueError, match="external knowledge binding not found"): - ExternalDatasetService.get_external_knowledge_binding_with_dataset_id("tenant-123", "dataset-123") + ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( + mock_db.session, "tenant-123", "dataset-123" + ) class TestExternalDatasetServiceDocumentValidate: @@ -1120,7 +1132,7 @@ class TestExternalDatasetServiceDocumentValidate: process_parameter = {"param1": "value1", "param2": "value2"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate(mock_db.session, tenant_id, api_id, process_parameter) @patch("services.external_knowledge_service.db") def test_document_create_args_validate_missing_required_param( @@ -1141,7 +1153,7 @@ class TestExternalDatasetServiceDocumentValidate: # Act & Assert with pytest.raises(ValueError, match="required_param is required"): - ExternalDatasetService.document_create_args_validate(tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate(mock_db.session, tenant_id, api_id, process_parameter) @patch("services.external_knowledge_service.db") def test_document_create_args_validate_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): @@ -1151,7 +1163,7 @@ class TestExternalDatasetServiceDocumentValidate: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate(mock_db.session, "tenant-123", "api-123", {}) @patch("services.external_knowledge_service.db") def test_document_create_args_validate_no_custom_parameters( @@ -1165,7 +1177,7 @@ class TestExternalDatasetServiceDocumentValidate: mock_db.session.scalar.return_value = api # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate(mock_db.session, "tenant-123", "api-123", {}) @patch("services.external_knowledge_service.db") def test_document_create_args_validate_optional_params_not_required( @@ -1187,7 +1199,9 @@ class TestExternalDatasetServiceDocumentValidate: process_parameter = {"required_param": "value"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", process_parameter) + ExternalDatasetService.document_create_args_validate( + mock_db.session, "tenant-123", "api-123", process_parameter + ) class TestExternalDatasetServiceProcessAPI: @@ -1496,7 +1510,7 @@ class TestExternalDatasetServiceCreateDataset: mock_db.session.scalar.side_effect = [None, api] # Act - result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args) + result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, mock_db.session) # Assert assert result.name == "Test External Dataset" @@ -1520,7 +1534,7 @@ class TestExternalDatasetServiceCreateDataset: # Act & Assert with pytest.raises(DatasetNameDuplicateError): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) @patch("services.external_knowledge_service.db") def test_create_external_dataset_api_not_found_error(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): @@ -1532,7 +1546,7 @@ class TestExternalDatasetServiceCreateDataset: # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) @patch("services.external_knowledge_service.db") def test_create_external_dataset_missing_knowledge_id_error( @@ -1548,7 +1562,7 @@ class TestExternalDatasetServiceCreateDataset: # Act & Assert with pytest.raises(ValueError, match="external_knowledge_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) @patch("services.external_knowledge_service.db") def test_create_external_dataset_missing_api_id_error( @@ -1564,7 +1578,7 @@ class TestExternalDatasetServiceCreateDataset: # Act & Assert with pytest.raises(ValueError, match="external_knowledge_api_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) class TestExternalDatasetServiceFetchRetrieval: @@ -1602,7 +1616,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - tenant_id, dataset_id, query, external_retrieval_parameters + mock_db.session, tenant_id, dataset_id, query, external_retrieval_parameters ) # Assert @@ -1620,7 +1634,9 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external knowledge binding not found"): - ExternalDatasetService.fetch_external_knowledge_retrieval("tenant-123", "dataset-123", "query", {}) + ExternalDatasetService.fetch_external_knowledge_retrieval( + mock_db.session, "tenant-123", "dataset-123", "query", {} + ) @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_cross_tenant_api_template_error( @@ -1633,7 +1649,9 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external api template not found"): - ExternalDatasetService.fetch_external_knowledge_retrieval("tenant-123", "dataset-123", "query", {}) + ExternalDatasetService.fetch_external_knowledge_retrieval( + mock_db.session, "tenant-123", "dataset-123", "query", {} + ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @patch("services.external_knowledge_service.db") @@ -1654,7 +1672,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) # Assert @@ -1685,7 +1703,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", external_retrieval_parameters + mock_db.session, "tenant-123", "dataset-123", "query", external_retrieval_parameters ) # Assert @@ -1714,7 +1732,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="Internal Server Error: Database connection failed"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @pytest.mark.parametrize( @@ -1754,7 +1772,9 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match=re.escape(error_message)): - ExternalDatasetService.fetch_external_knowledge_retrieval(tenant_id, dataset_id, "query", {"top_k": 5}) + ExternalDatasetService.fetch_external_knowledge_retrieval( + mock_db.session, tenant_id, dataset_id, "query", {"top_k": 5} + ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @patch("services.external_knowledge_service.db") @@ -1776,7 +1796,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @@ -1795,7 +1815,7 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @@ -1814,7 +1834,7 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @@ -1833,7 +1853,7 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") @@ -1848,5 +1868,5 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="connection reset by peer"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) diff --git a/api/tests/unit_tests/test_pytest_dify.py b/api/tests/unit_tests/test_pytest_dify.py index 8110a11d801..c902df5df4d 100644 --- a/api/tests/unit_tests/test_pytest_dify.py +++ b/api/tests/unit_tests/test_pytest_dify.py @@ -28,6 +28,8 @@ def test_ensure_backend_test_environment_uses_example_env_and_stable_logging( monkeypatch.delenv("LOG_OUTPUT_FORMAT", raising=False) monkeypatch.delenv("DIFY_TEST_ENV_FILE", raising=False) monkeypatch.delenv("DIFY_VDB_TEST_ENV_FILE", raising=False) + monkeypatch.delenv("STORAGE_TYPE", raising=False) + monkeypatch.delenv("OPENDAL_SCHEME", raising=False) monkeypatch.setenv("OPENDAL_FS_ROOT", str(storage_root)) ensure_backend_test_environment(repo_root) diff --git a/api/tests/unit_tests/tools/test_api_tool.py b/api/tests/unit_tests/tools/test_api_tool.py index 2a8c6686d7d..e62f74b8019 100644 --- a/api/tests/unit_tests/tools/test_api_tool.py +++ b/api/tests/unit_tests/tools/test_api_tool.py @@ -74,7 +74,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -149,7 +149,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -186,7 +186,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -212,7 +212,7 @@ class TestApiToolInvoke: mock_get.return_value = mock_response # Invoke the tool - result_generator = self.api_tool._invoke(user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) # Get the result from the generator result = list(result_generator) @@ -236,7 +236,7 @@ class TestApiToolInvoke: mock_response.text = "Not Found" mock_get.return_value = mock_response - result_generator = self.api_tool._invoke(user_id="test_user", tool_parameters={}) + result_generator = self.api_tool._invoke(session=Mock(), user_id="test_user", tool_parameters={}) # Invoke the tool and expect an error with pytest.raises(Exception) as exc_info: diff --git a/api/tests/unit_tests/tools/test_mcp_tool.py b/api/tests/unit_tests/tools/test_mcp_tool.py index 689b9730971..5984b3b6744 100644 --- a/api/tests/unit_tests/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/tools/test_mcp_tool.py @@ -64,7 +64,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -80,7 +80,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -101,7 +101,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[content]) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) assert len(messages) == 1 msg = messages[0] @@ -115,7 +115,7 @@ class TestMCPToolInvoke: result = CallToolResult(content=[], structuredContent={"a": 1, "b": "x"}) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - messages = list(tool._invoke(user_id="test_user", tool_parameters={})) + messages = list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) # Expect two variable messages corresponding to keys a and b assert len(messages) == 2 @@ -281,7 +281,7 @@ class TestMCPToolUsageExtraction: result = CallToolResult(content=[TextContent(type="text", text="test")], _meta=meta) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - list(tool._invoke(user_id="test_user", tool_parameters={})) + list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) # Verify latest_usage was set correctly assert tool.latest_usage.prompt_tokens == 200 @@ -295,7 +295,7 @@ class TestMCPToolUsageExtraction: result = CallToolResult(content=[TextContent(type="text", text="test")], _meta=None) with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): - list(tool._invoke(user_id="test_user", tool_parameters={})) + list(tool._invoke(session=Mock(), user_id="test_user", tool_parameters={})) # Verify latest_usage is empty assert tool.latest_usage.total_tokens == 0