mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 02:28:30 +08:00
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:
parent
dcc06dee20
commit
2f0c6c10bd
1
.github/workflows/style.yml
vendored
1
.github/workflows/style.yml
vendored
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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],
|
||||
)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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),
|
||||
)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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.")
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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},
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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.")
|
||||
|
||||
|
||||
@ -5,6 +5,7 @@ from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from constants import HIDDEN_VALUE
|
||||
from core.helper import ssrf_proxy
|
||||
@ -57,7 +58,9 @@ class ExternalDatasetService:
|
||||
raise ValueError("api_key is required")
|
||||
|
||||
@staticmethod
|
||||
def create_external_knowledge_api(tenant_id: str, user_id: str, args: dict[str, Any]) -> ExternalKnowledgeApis:
|
||||
def create_external_knowledge_api(
|
||||
tenant_id: str, user_id: str, args: dict[str, Any], session: Session
|
||||
) -> ExternalKnowledgeApis:
|
||||
settings = args.get("settings")
|
||||
if settings is None:
|
||||
raise ValueError("settings is required")
|
||||
@ -71,8 +74,8 @@ class ExternalDatasetService:
|
||||
settings=json.dumps(args.get("settings"), ensure_ascii=False),
|
||||
)
|
||||
|
||||
db.session.add(external_knowledge_api)
|
||||
db.session.commit()
|
||||
session.add(external_knowledge_api)
|
||||
session.commit()
|
||||
return external_knowledge_api
|
||||
|
||||
@staticmethod
|
||||
@ -103,8 +106,10 @@ class ExternalDatasetService:
|
||||
raise ValueError(f"Forbidden: Authorization failed with api_key: {api_key}")
|
||||
|
||||
@staticmethod
|
||||
def get_external_knowledge_api(external_knowledge_api_id: str, tenant_id: str) -> ExternalKnowledgeApis:
|
||||
external_knowledge_api: ExternalKnowledgeApis | None = db.session.scalar(
|
||||
def get_external_knowledge_api(
|
||||
session: Session, external_knowledge_api_id: str, tenant_id: str
|
||||
) -> ExternalKnowledgeApis:
|
||||
external_knowledge_api: ExternalKnowledgeApis | None = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -114,8 +119,10 @@ class ExternalDatasetService:
|
||||
return external_knowledge_api
|
||||
|
||||
@staticmethod
|
||||
def update_external_knowledge_api(tenant_id, user_id, external_knowledge_api_id, args) -> ExternalKnowledgeApis:
|
||||
external_knowledge_api: ExternalKnowledgeApis | None = db.session.scalar(
|
||||
def update_external_knowledge_api(
|
||||
session: Session, tenant_id: str, user_id: str, external_knowledge_api_id: str, args
|
||||
) -> ExternalKnowledgeApis:
|
||||
external_knowledge_api: ExternalKnowledgeApis | None = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -131,13 +138,13 @@ class ExternalDatasetService:
|
||||
external_knowledge_api.settings = json.dumps(args.get("settings"), ensure_ascii=False)
|
||||
external_knowledge_api.updated_by = user_id
|
||||
external_knowledge_api.updated_at = naive_utc_now()
|
||||
db.session.commit()
|
||||
session.commit()
|
||||
|
||||
return external_knowledge_api
|
||||
|
||||
@staticmethod
|
||||
def delete_external_knowledge_api(tenant_id: str, external_knowledge_api_id: str):
|
||||
external_knowledge_api = db.session.scalar(
|
||||
def delete_external_knowledge_api(session: Session, tenant_id: str, external_knowledge_api_id: str):
|
||||
external_knowledge_api = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -145,11 +152,13 @@ class ExternalDatasetService:
|
||||
if external_knowledge_api is None:
|
||||
raise ValueError("api template not found")
|
||||
|
||||
db.session.delete(external_knowledge_api)
|
||||
db.session.commit()
|
||||
session.delete(external_knowledge_api)
|
||||
session.commit()
|
||||
|
||||
@staticmethod
|
||||
def external_knowledge_api_use_check(external_knowledge_api_id: str, tenant_id: str) -> tuple[bool, int]:
|
||||
def external_knowledge_api_use_check(
|
||||
session: Session, external_knowledge_api_id: str, tenant_id: str
|
||||
) -> tuple[bool, int]:
|
||||
"""
|
||||
Return usage for an external knowledge API within a single tenant.
|
||||
|
||||
@ -157,7 +166,7 @@ class ExternalDatasetService:
|
||||
same; otherwise the endpoint becomes a cross-tenant UUID oracle.
|
||||
"""
|
||||
count = (
|
||||
db.session.scalar(
|
||||
session.scalar(
|
||||
select(func.count(ExternalKnowledgeBindings.id)).where(
|
||||
ExternalKnowledgeBindings.external_knowledge_api_id == external_knowledge_api_id,
|
||||
ExternalKnowledgeBindings.tenant_id == tenant_id,
|
||||
@ -168,8 +177,10 @@ class ExternalDatasetService:
|
||||
return count > 0, count
|
||||
|
||||
@staticmethod
|
||||
def get_external_knowledge_binding_with_dataset_id(tenant_id: str, dataset_id: str) -> ExternalKnowledgeBindings:
|
||||
external_knowledge_binding: ExternalKnowledgeBindings | None = db.session.scalar(
|
||||
def get_external_knowledge_binding_with_dataset_id(
|
||||
session: Session, tenant_id: str, dataset_id: str
|
||||
) -> ExternalKnowledgeBindings:
|
||||
external_knowledge_binding: ExternalKnowledgeBindings | None = session.scalar(
|
||||
select(ExternalKnowledgeBindings)
|
||||
.where(ExternalKnowledgeBindings.dataset_id == dataset_id, ExternalKnowledgeBindings.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -180,9 +191,9 @@ class ExternalDatasetService:
|
||||
|
||||
@staticmethod
|
||||
def document_create_args_validate(
|
||||
tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any]
|
||||
session: Session, tenant_id: str, external_knowledge_api_id: str, process_parameter: dict[str, Any]
|
||||
):
|
||||
external_knowledge_api = db.session.scalar(
|
||||
external_knowledge_api = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(ExternalKnowledgeApis.id == external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -255,13 +266,13 @@ class ExternalDatasetService:
|
||||
return ExternalKnowledgeApiSetting.model_validate(settings)
|
||||
|
||||
@staticmethod
|
||||
def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any]) -> Dataset:
|
||||
def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], session: Session) -> Dataset:
|
||||
# check if dataset name already exists
|
||||
if db.session.scalar(
|
||||
if session.scalar(
|
||||
select(Dataset).where(Dataset.name == args.get("name"), Dataset.tenant_id == tenant_id).limit(1)
|
||||
):
|
||||
raise DatasetNameDuplicateError(f"Dataset with name {args.get('name')} already exists.")
|
||||
external_knowledge_api = db.session.scalar(
|
||||
external_knowledge_api = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(
|
||||
ExternalKnowledgeApis.id == args.get("external_knowledge_api_id"),
|
||||
@ -283,8 +294,8 @@ class ExternalDatasetService:
|
||||
maintainer=user_id,
|
||||
)
|
||||
|
||||
db.session.add(dataset)
|
||||
db.session.flush()
|
||||
session.add(dataset)
|
||||
session.flush()
|
||||
if args.get("external_knowledge_id") is None:
|
||||
raise ValueError("external_knowledge_id is required")
|
||||
if args.get("external_knowledge_api_id") is None:
|
||||
@ -297,14 +308,15 @@ class ExternalDatasetService:
|
||||
external_knowledge_id=args.get("external_knowledge_id") or "",
|
||||
created_by=user_id,
|
||||
)
|
||||
db.session.add(external_knowledge_binding)
|
||||
session.add(external_knowledge_binding)
|
||||
|
||||
db.session.commit()
|
||||
session.commit()
|
||||
|
||||
return dataset
|
||||
|
||||
@staticmethod
|
||||
def fetch_external_knowledge_retrieval(
|
||||
session: Session,
|
||||
tenant_id: str,
|
||||
dataset_id: str,
|
||||
query: str,
|
||||
@ -320,7 +332,7 @@ class ExternalDatasetService:
|
||||
the inner knowledge retrieval API—can consistently expose
|
||||
``502 external_knowledge_failed``.
|
||||
"""
|
||||
external_knowledge_binding = db.session.scalar(
|
||||
external_knowledge_binding = session.scalar(
|
||||
select(ExternalKnowledgeBindings)
|
||||
.where(ExternalKnowledgeBindings.dataset_id == dataset_id, ExternalKnowledgeBindings.tenant_id == tenant_id)
|
||||
.limit(1)
|
||||
@ -328,7 +340,7 @@ class ExternalDatasetService:
|
||||
if not external_knowledge_binding:
|
||||
raise ExternalKnowledgeRetrievalError("external knowledge binding not found")
|
||||
|
||||
external_knowledge_api = db.session.scalar(
|
||||
external_knowledge_api = session.scalar(
|
||||
select(ExternalKnowledgeApis)
|
||||
.where(
|
||||
ExternalKnowledgeApis.id == external_knowledge_binding.external_knowledge_api_id,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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}
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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"},
|
||||
|
||||
@ -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)+)",
|
||||
|
||||
@ -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(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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"] == []
|
||||
|
||||
|
||||
@ -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"] == []
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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": {}},
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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(),
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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"))
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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"]
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 == []
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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
Loading…
Reference in New Issue
Block a user