refactor: thread explicit sessions through app retrieval paths (#38309)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Byron.wang 2026-07-03 01:00:47 +08:00 committed by GitHub
parent dcc06dee20
commit 2f0c6c10bd
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
121 changed files with 1356 additions and 677 deletions

View File

@ -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

View File

@ -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/<uuid:agent_id>/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,

View File

@ -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,

View File

@ -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],
)

View File

@ -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,

View File

@ -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),
)

View File

@ -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,

View File

@ -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)

View File

@ -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,

View File

@ -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)

View File

@ -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)

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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)

View File

@ -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:

View File

@ -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)

View File

@ -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)

View File

@ -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.")

View File

@ -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))

View File

@ -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)

View File

@ -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,

View File

@ -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)

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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},

View File

@ -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(

View File

@ -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,

View File

@ -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:

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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)

View File

@ -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(

View File

@ -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

View File

@ -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,

View File

@ -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(

View File

@ -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,

View File

@ -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,

View File

@ -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

View File

@ -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,

View File

@ -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

View File

@ -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.")

View File

@ -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 APIcan 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,

View File

@ -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,

View File

@ -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}

View File

@ -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:

View File

@ -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"

View File

@ -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

View File

@ -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

View File

@ -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)

View File

@ -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):

View File

@ -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()

View File

@ -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"},

View File

@ -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<files>(?:^ \S[^\n]*\n)+)",

View File

@ -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(),
)

View File

@ -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:

View File

@ -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"] == []

View File

@ -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"] == []

View File

@ -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)

View File

@ -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")

View File

@ -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:

View File

@ -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:

View File

@ -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:

View File

@ -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:

View File

@ -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

View File

@ -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:

View File

@ -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": {}},

View File

@ -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:

View File

@ -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(),

View File

@ -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"

View File

@ -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"))

View File

@ -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,

View File

@ -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)

View File

@ -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(

View File

@ -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"]

View File

@ -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",

View File

@ -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:

View File

@ -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

View File

@ -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 == []

View File

@ -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()

View File

@ -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",

Some files were not shown because too many files have changed in this diff Show More