diff --git a/api/commands/vector.py b/api/commands/vector.py index 095d0dea4ed..39884418b23 100644 --- a/api/commands/vector.py +++ b/api/commands/vector.py @@ -13,6 +13,7 @@ from core.rag.index_processor.constant.built_in_field import BuiltInField from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.models.document import ChildDocument, Document from extensions.ext_database import db +from libs.pagination import paginate_query from models.dataset import Dataset, DatasetCollectionBinding, DatasetMetadata, DatasetMetadataBinding, DocumentSegment from models.dataset import Document as DatasetDocument from models.enums import DatasetMetadataType, IndexingStatus, SegmentStatus @@ -183,7 +184,7 @@ def migrate_knowledge_vector_database(): .order_by(Dataset.created_at.desc()) ) - datasets = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False) + datasets = paginate_query(stmt, page=page, per_page=50, max_per_page=50) if not datasets.items: break except SQLAlchemyError: @@ -409,7 +410,7 @@ def old_metadata_migration(): .where(DatasetDocument.doc_metadata.is_not(None)) .order_by(DatasetDocument.created_at.desc()) ) - documents = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False) + documents = paginate_query(stmt, page=page, per_page=50, max_per_page=50) except SQLAlchemyError: raise if not documents: diff --git a/api/controllers/console/app/conversation.py b/api/controllers/console/app/conversation.py index 6387ea441e5..a80935e5e33 100644 --- a/api/controllers/console/app/conversation.py +++ b/api/controllers/console/app/conversation.py @@ -41,6 +41,7 @@ from fields.conversation_fields import ( from libs.datetime_utils import naive_utc_now, parse_time_range from libs.helper import dump_response from libs.login import login_required +from libs.pagination import paginate_query from models import Conversation, EndUser, Message, MessageAnnotation from models.account import Account from models.model import App, AppMode @@ -156,7 +157,7 @@ class CompletionConversationApi(Resource): query = query.order_by(Conversation.created_at.desc()) - conversations = db.paginate(query, page=args.page, per_page=args.limit, error_out=False) + conversations = paginate_query(query, page=args.page, per_page=args.limit) return dump_response(ConversationPaginationResponse, conversations) @@ -310,7 +311,7 @@ class ChatConversationApi(Resource): case _: query = query.order_by(Conversation.created_at.desc()) - conversations = db.paginate(query, page=args.page, per_page=args.limit, error_out=False) + conversations = paginate_query(query, page=args.page, per_page=args.limit) return dump_response(ConversationWithSummaryPaginationResponse, conversations) diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index afa617535e1..a6263c8e2e3 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -46,6 +46,7 @@ from graphon.model_runtime.errors.invoke import InvokeAuthorizationError from libs.datetime_utils import naive_utc_now from libs.helper import dump_response, to_timestamp from libs.login import login_required +from libs.pagination import paginate_query from models import Account, DatasetProcessRule, Document, DocumentSegment, UploadFile from models.dataset import DocumentPipelineExecutionLog from models.enums import IndexingStatus, SegmentStatus @@ -368,7 +369,7 @@ class DatasetDocumentListApi(Resource): desc(Document.position), ) - paginated_documents = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False) + paginated_documents = paginate_query(query, page=page, per_page=limit, max_per_page=100) documents = paginated_documents.items DocumentService.enrich_documents_with_summary_index_status( diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 40ae9b207fd..5cccd2453dc 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -57,6 +57,7 @@ from fields.segment_fields import ( from graphon.model_runtime.entities.model_entities import ModelType from libs.helper import dump_response, escape_like_pattern from libs.login import login_required +from libs.pagination import paginate_query from models import Account from models.dataset import Dataset, Document, DocumentSegment from models.model import UploadFile @@ -270,7 +271,7 @@ class DatasetDocumentSegmentListApi(Resource): elif args.enabled.lower() == "false": query = query.where(DocumentSegment.enabled == False) - segments = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False) + segments = paginate_query(query, page=page, per_page=limit, max_per_page=100) segment_list = list(segments.items) segment_ids = [segment.id for segment in segment_list] diff --git a/api/controllers/console/workspace/workspace.py b/api/controllers/console/workspace/workspace.py index 746e971bce8..418c3eb66e1 100644 --- a/api/controllers/console/workspace/workspace.py +++ b/api/controllers/console/workspace/workspace.py @@ -33,6 +33,7 @@ from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import OptionalTimestampField, TimestampField, dump_response, to_timestamp from libs.login import login_required +from libs.pagination import paginate_query from models.account import Account, Tenant, TenantAccountJoin, TenantCustomConfigDict, TenantStatus from services.account_service import TenantService from services.billing_service import BillingService, SubscriptionPlan @@ -294,7 +295,7 @@ class WorkspaceListApi(Resource): args = WorkspaceListQuery.model_validate(payload) stmt = select(Tenant).order_by(Tenant.created_at.desc()) - tenants = db.paginate(select=stmt, page=args.page, per_page=args.limit, error_out=False) + tenants = paginate_query(stmt, page=args.page, per_page=args.limit) has_more = False if tenants.has_next: diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 49ccb1bd55c..4c083d3d50f 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -59,6 +59,7 @@ from fields.document_fields import ( ) from libs.helper import dump_response from libs.login import current_user +from libs.pagination import paginate_query from models.dataset import Dataset, Document, DocumentSegment from models.enums import SegmentStatus from services.dataset_service import DatasetService, DocumentService @@ -945,8 +946,8 @@ class DocumentListApi(DatasetApiResource): query = query.order_by(desc(Document.created_at), desc(Document.position)) - paginated_documents = db.paginate( - select=query, page=query_params.page, per_page=query_params.limit, max_per_page=100, error_out=False + paginated_documents = paginate_query( + query, page=query_params.page, per_page=query_params.limit, max_per_page=100 ) documents = paginated_documents.items diff --git a/api/libs/pagination.py b/api/libs/pagination.py new file mode 100644 index 00000000000..c38297efc2b --- /dev/null +++ b/api/libs/pagination.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import math +from dataclasses import dataclass + +from sqlalchemy import Select, func, select +from sqlalchemy.orm import Session, scoped_session + + +@dataclass +class PaginatedResult[T]: + """Minimal pagination container backed by plain SQLAlchemy queries. + + Drop-in replacement for Flask-SQLAlchemy's ``db.paginate`` return value. + Only the attributes actually consumed across the codebase are exposed: + ``items``, ``total``, ``page``, ``per_page``, ``pages``, ``has_next``. + """ + + items: list[T] + total: int + page: int + per_page: int + + @property + def pages(self) -> int: + if self.per_page == 0: + return 0 + return max(1, math.ceil(self.total / self.per_page)) + + @property + def has_next(self) -> bool: + return self.page < self.pages + + def __iter__(self): + return iter(self.items) + + +def paginate_query( + stmt: Select, + *, + page: int = 1, + per_page: int = 20, + max_per_page: int | None = None, + session: Session | scoped_session | None = None, +) -> PaginatedResult: + """Execute *stmt* as a paginated query using plain SQLAlchemy. + + Parameters + ---------- + stmt: + A SQLAlchemy ``select()`` statement. + page: + 1-based page number. + per_page: + Number of items per page. + max_per_page: + Hard ceiling for *per_page*; ``None`` means no cap. + session: + The session to use. Falls back to ``db.session`` when omitted. + """ + if session is None: + from extensions.ext_database import db + + session = db.session + + if max_per_page is not None: + per_page = min(per_page, max_per_page) + + page = max(1, page) + per_page = max(1, per_page) + + # total count — wrap in a scalar subquery so arbitrary selects work + count_stmt = select(func.count()).select_from(stmt.subquery()) + total: int = session.scalar(count_stmt) or 0 # type: ignore[assignment] + + # fetch the page + offset = (page - 1) * per_page + page_stmt = stmt.limit(per_page).offset(offset) + items = list(session.scalars(page_stmt).all()) + + return PaginatedResult( + items=items, + total=total, + page=page, + per_page=per_page, + ) diff --git a/api/schedule/clean_unused_datasets_task.py b/api/schedule/clean_unused_datasets_task.py index 849274311a2..03417647724 100644 --- a/api/schedule/clean_unused_datasets_task.py +++ b/api/schedule/clean_unused_datasets_task.py @@ -12,6 +12,7 @@ from core.rag.index_processor.index_processor_factory import IndexProcessorFacto from enums.cloud_plan import CloudPlan from extensions.ext_database import db from extensions.ext_redis import redis_client +from libs.pagination import paginate_query from models.dataset import Dataset, DatasetAutoDisableLog, DatasetQuery, Document from services.feature_service import FeatureService @@ -88,7 +89,7 @@ def clean_unused_datasets_task(): .order_by(Dataset.created_at.desc()) ) - datasets = db.paginate(stmt, page=page, per_page=50, error_out=False) + datasets = paginate_query(stmt, page=page, per_page=50) except SQLAlchemyError: raise diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 080527d9769..03e445a938b 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -13,6 +13,7 @@ from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now from libs.login import current_account_with_tenant +from libs.pagination import paginate_query from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation from services.app_ref_service import AnnotationRef, AppRef from services.feature_service import FeatureService @@ -242,7 +243,7 @@ class AppAnnotationService: .where(MessageAnnotation.app_id == app_id) .order_by(MessageAnnotation.created_at.desc(), MessageAnnotation.id.desc()) ) - annotations = db.paginate(select=stmt, page=page, per_page=limit, max_per_page=100, error_out=False) + annotations = paginate_query(stmt, page=page, per_page=limit, max_per_page=100) return annotations.items, annotations.total or 0 @classmethod @@ -573,9 +574,7 @@ class AppAnnotationService: ) .order_by(AppAnnotationHitHistory.created_at.desc()) ) - annotation_hit_histories = db.paginate( - select=stmt, page=page, per_page=limit, max_per_page=100, error_out=False - ) + annotation_hit_histories = paginate_query(stmt, page=page, per_page=limit, max_per_page=100) return annotation_hit_histories.items, annotation_hit_histories.total or 0 @classmethod diff --git a/api/services/app_service.py b/api/services/app_service.py index bd0e3fd08e0..08cd30974e3 100644 --- a/api/services/app_service.py +++ b/api/services/app_service.py @@ -5,7 +5,6 @@ from datetime import datetime from typing import Any, Literal, NotRequired, TypedDict, cast, override import sqlalchemy as sa -from flask_sqlalchemy.pagination import Pagination from pydantic import BaseModel, Field from sqlalchemy import ColumnElement, select from sqlalchemy.exc import IntegrityError @@ -24,6 +23,7 @@ from graphon.model_runtime.entities.model_entities import ModelPropertyKey, Mode from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel from libs.datetime_utils import naive_utc_now from libs.login import current_user +from libs.pagination import PaginatedResult, paginate_query from models import Account, AppStar from models.agent import Agent, AgentIconType, AgentScope, AgentSource, AgentStatus from models.model import App, AppMode, AppModelConfig, IconType, Site @@ -220,7 +220,7 @@ class AppService: def get_paginate_apps( self, user_id: str, tenant_id: str, params: AppListParams, session: scoped_session - ) -> Pagination | None: + ) -> PaginatedResult | None: """ Get app list with pagination, filters, and explicit sort order. :param user_id: user id @@ -234,11 +234,10 @@ class AppService: order_by = self._build_app_list_order_by(params.sort_by) - app_models = db.paginate( + app_models = paginate_query( sa.select(App).where(*filters).order_by(order_by), page=params.page, per_page=params.limit, - error_out=False, ) app_ids = [str(app.id) for app in app_models.items] @@ -255,7 +254,7 @@ class AppService: def get_paginate_starred_apps( self, user_id: str, tenant_id: str, params: StarredAppListParams, session: scoped_session - ) -> Pagination | None: + ) -> PaginatedResult | None: """ Get apps starred by the current account with pagination, filters, and explicit sort order. """ @@ -264,7 +263,7 @@ class AppService: return None order_by = self._build_app_list_order_by(params.sort_by) - app_models = db.paginate( + app_models = paginate_query( sa.select(App) .join( AppStar, @@ -278,7 +277,6 @@ class AppService: .order_by(order_by), page=params.page, per_page=params.limit, - error_out=False, ) for app in app_models.items: diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 30b04a620f8..b36926a32c5 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -34,6 +34,7 @@ from graphon.model_runtime.model_providers.base.text_embedding_model import Text from libs import helper from libs.datetime_utils import naive_utc_now from libs.login import current_user +from libs.pagination import paginate_query from models import Account, TenantAccountRole from models.dataset import ( AppDatasetJoin, @@ -355,7 +356,7 @@ class DatasetService: else: return [], 0 - datasets = db.paginate(select=query, page=page, per_page=per_page, max_per_page=100, error_out=False) + datasets = paginate_query(query, page=page, per_page=per_page, max_per_page=100) return datasets.items, datasets.total @@ -399,7 +400,7 @@ class DatasetService: accessible_filter = sa.or_(Dataset.maintainer == user.id, accessible_filter) stmt = stmt.where(accessible_filter) - datasets = db.paginate(select=stmt, page=1, per_page=len(ids), max_per_page=len(ids), error_out=False) + datasets = paginate_query(stmt, page=1, per_page=len(ids), max_per_page=len(ids)) return datasets.items, datasets.total @@ -1403,7 +1404,7 @@ class DatasetService: def get_dataset_queries(dataset_id: str, page: int, per_page: int): stmt = select(DatasetQuery).filter_by(dataset_id=dataset_id).order_by(db.desc(DatasetQuery.created_at)) - dataset_queries = db.paginate(select=stmt, page=page, per_page=per_page, max_per_page=100, error_out=False) + dataset_queries = paginate_query(stmt, page=page, per_page=per_page, max_per_page=100) return dataset_queries.items, dataset_queries.total @@ -4120,7 +4121,7 @@ class SegmentService: if keyword: escaped_keyword = helper.escape_like_pattern(keyword) query = query.where(ChildChunk.content.ilike(f"%{escaped_keyword}%", escape="\\")) - return db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False) + return paginate_query(query, page=page, per_page=limit, max_per_page=100) @classmethod def get_child_chunk_by_id( @@ -4172,7 +4173,7 @@ class SegmentService: query = query.where(DocumentSegment.content.ilike(f"%{escaped_keyword}%", escape="\\")) query = query.order_by(DocumentSegment.position.asc(), DocumentSegment.id.asc()) - paginated_segments = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False) + paginated_segments = paginate_query(query, page=page, per_page=limit, max_per_page=100) return paginated_segments.items, paginated_segments.total diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index 355c4844233..42e7eca29d7 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -10,9 +10,9 @@ from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.helper import ssrf_proxy from core.rag.entities import MetadataFilteringCondition -from extensions.ext_database import db from graphon.nodes.http_request.exc import InvalidHttpMethodError from libs.datetime_utils import naive_utc_now +from libs.pagination import paginate_query from models.dataset import ( Dataset, ExternalKnowledgeApis, @@ -42,9 +42,7 @@ class ExternalDatasetService: escaped_search = escape_like_pattern(search) query = query.where(ExternalKnowledgeApis.name.ilike(f"%{escaped_search}%", escape="\\")) - external_knowledge_apis = db.paginate( - select=query, page=page, per_page=per_page, max_per_page=100, error_out=False - ) + external_knowledge_apis = paginate_query(query, page=page, per_page=per_page, max_per_page=100) return external_knowledge_apis.items, external_knowledge_apis.total diff --git a/api/services/plugin/plugin_migration.py b/api/services/plugin/plugin_migration.py index 8239186bbcf..82eeb5a7261 100644 --- a/api/services/plugin/plugin_migration.py +++ b/api/services/plugin/plugin_migration.py @@ -25,6 +25,7 @@ from core.plugin.impl.plugin import PluginInstaller from core.plugin.plugin_service import PluginService from core.tools.entities.tool_entities import ToolProviderType from extensions.ext_database import db +from libs.pagination import paginate_query from models.account import Tenant from models.model import App, AppMode, AppModelConfig from models.provider_ids import ModelProviderID, ToolProviderID @@ -499,7 +500,7 @@ class PluginMigration: total_failed_tenant = 0 while True: # paginate - tenants = db.paginate(sa.select(Tenant).order_by(Tenant.created_at.desc()), page=page, per_page=100) + tenants = paginate_query(sa.select(Tenant).order_by(Tenant.created_at.desc()), page=page, per_page=100) if tenants.items is None or len(tenants.items) == 0: break diff --git a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py index 5de07ff14ea..2d2d5b4f361 100644 --- a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py @@ -30,7 +30,7 @@ def test_completion_conversation_list_returns_paginated_result(app: Flask, monke paginate_result.total = 0 paginate_result.has_next = False paginate_result.items = [] - monkeypatch.setattr(conversation_module.db, "paginate", lambda *_args, **_kwargs: paginate_result) + monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) with app.test_request_context("/console/api/apps/app-1/completion-conversations", method="GET"): response = method(api, account, app_model=SimpleNamespace(id="app-1")) @@ -71,7 +71,7 @@ def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeyp paginate_result.total = 0 paginate_result.has_next = False paginate_result.items = [] - monkeypatch.setattr(conversation_module.db, "paginate", lambda *_args, **_kwargs: paginate_result) + monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) with app.test_request_context("/console/api/apps/app-1/chat-conversations", method="GET"): response = method(api, account, app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT)) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py index 99d7cf626f0..be0a4eea40a 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py @@ -207,7 +207,7 @@ class TestDatasetDocumentListApi: with ( app.test_request_context("/?fetch=true"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( @@ -237,7 +237,7 @@ class TestDatasetDocumentListApi: with ( app.test_request_context("/?keyword=test&status=enabled&sort=created_at"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( @@ -263,7 +263,7 @@ class TestDatasetDocumentListApi: with ( app.test_request_context("/"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( @@ -341,7 +341,7 @@ class TestDatasetDocumentListApi: with ( app.test_request_context("/?fetch=maybe"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( @@ -363,7 +363,7 @@ class TestDatasetDocumentListApi: with ( app.test_request_context("/?sort=hit_count"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( @@ -1537,7 +1537,7 @@ class TestDocumentListAdvancedCases: with ( app.test_request_context("/?sort=updated_at"), patch( - "controllers.console.datasets.datasets_document.db.paginate", + "controllers.console.datasets.datasets_document.paginate_query", return_value=pagination, ), patch( diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py index 2e54c977260..93fd25a610e 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py @@ -161,7 +161,7 @@ class TestDatasetDocumentSegmentListApi: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.paginate", + "controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination, ), patch( @@ -1207,7 +1207,7 @@ class TestSegmentListAdvancedCases: return_value=document, ), patch( - "controllers.console.datasets.datasets_segments.db.paginate", + "controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination, ), patch( @@ -1255,7 +1255,7 @@ class TestSegmentListAdvancedCases: SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME="postgresql"), ), patch( - "controllers.console.datasets.datasets_segments.db.paginate", + "controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination, ) as paginate_mock, ): @@ -1267,7 +1267,7 @@ class TestSegmentListAdvancedCases: "33333333-3333-3333-3333-333333333333", ) - query = paginate_mock.call_args.kwargs["select"] + query = paginate_mock.call_args.args[0] sql = str(query.compile(compile_kwargs={"literal_binds": True})) assert "jsonb_array_elements_text(CASE" in sql assert "ELSE CAST('[]' AS JSONB)" in sql diff --git a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py index 47e9f51fb27..b4fb60910e5 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py @@ -291,7 +291,7 @@ class TestWorkspaceListApi: with ( app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 20}), - patch("controllers.console.workspace.workspace.db.paginate", return_value=paginate_result), + patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result), ): result, status = method(api) @@ -308,7 +308,7 @@ class TestWorkspaceListApi: with ( app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 1}), - patch("controllers.console.workspace.workspace.db.paginate", return_value=paginate_result), + patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result), ): result, status = method(api) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index e68eb647063..dd2caf4f3fc 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -851,9 +851,10 @@ class TestDocumentApiDelete: class TestDocumentListApi: """Test suite for DocumentListApi endpoint.""" + @patch("controllers.service_api.dataset.document.paginate_query") @patch("controllers.service_api.dataset.document.DocumentService") @patch("controllers.service_api.dataset.document.db") - def test_list_documents_success(self, mock_db, mock_doc_svc, app: Flask, mock_tenant, mock_dataset): + def test_list_documents_success(self, mock_db, mock_doc_svc, mock_paginate, app: Flask, mock_tenant, mock_dataset): """Test successful document list retrieval.""" # Arrange mock_db.session.scalar.return_value = mock_dataset @@ -868,7 +869,7 @@ class TestDocumentListApi: make_serializable_document(id="doc-2", name="Document 2"), ] mock_pagination.total = 2 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination mock_doc_svc.enrich_documents_with_summary_index_status.return_value = None diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index c483440dd2c..2975c4df14c 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -395,10 +395,11 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.paginate_query") as mock_paginate, patch("libs.helper.escape_like_pattern", return_value="safe"), ): mock_db.session.scalar.return_value = app - mock_db.paginate.return_value = pagination + mock_paginate.return_value = pagination # Act items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "keyword") @@ -417,9 +418,10 @@ class TestAppAnnotationServiceListAndExport: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.paginate_query") as mock_paginate, ): mock_db.session.scalar.return_value = app - mock_db.paginate.return_value = pagination + mock_paginate.return_value = pagination # Act items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "") @@ -1101,9 +1103,11 @@ class TestAppAnnotationServiceHitHistoryAndSettings: with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), patch("services.annotation_service.db") as mock_db, + patch("services.annotation_service.paginate_query") as mock_paginate, ): - mock_db.session.scalar.return_value = annotation - mock_db.paginate.return_value = pagination + mock_db.session.scalar.return_value = app + mock_db.session.get.return_value = annotation + mock_paginate.return_value = pagination # Act items, total = AppAnnotationService.get_annotation_hit_histories( diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 3d623e8abf8..02d965f4bd2 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -174,18 +174,18 @@ class TestDatasetServiceRetrievalPermissions: mock_db = MagicMock() explicit_session = MagicMock() explicit_session.scalars.return_value.all.return_value = [] - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( "services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get", return_value=SimpleNamespace(workspace=SimpleNamespace(permission_keys=[])), ), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets( page=1, per_page=20, @@ -198,7 +198,7 @@ class TestDatasetServiceRetrievalPermissions: explicit_session.scalars.assert_called_once() mock_db.session.scalars.assert_not_called() - select_stmt = mock_db.paginate.call_args.kwargs["select"] + select_stmt = mock_paginate.call_args.args[0] visibility_clause = str(select_stmt._where_criteria[1]) assert "maintainer" in visibility_clause assert "IN" in visibility_clause @@ -206,18 +206,18 @@ class TestDatasetServiceRetrievalPermissions: def test_get_datasets_filters_only_by_rbac_overrides_without_manage_own_permission(self): mock_db = MagicMock() mock_db.session.scalars.return_value.all.return_value = [] - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( "services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get", return_value=SimpleNamespace(workspace=SimpleNamespace(permission_keys=[])), ), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets( page=1, per_page=20, @@ -227,21 +227,21 @@ class TestDatasetServiceRetrievalPermissions: accessible_dataset_ids=["dataset-shared"], ) - select_stmt = mock_db.paginate.call_args.kwargs["select"] + select_stmt = mock_paginate.call_args.args[0] visibility_clause = str(select_stmt._where_criteria[1]) assert "maintainer" not in visibility_clause assert "IN" in visibility_clause def test_get_datasets_by_ids_applies_rbac_visibility(self): mock_db = MagicMock() - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets_by_ids( ["dataset-requested", "dataset-shared"], "tenant-1", @@ -250,7 +250,7 @@ class TestDatasetServiceRetrievalPermissions: include_own_datasets=True, ) - select_stmt = mock_db.paginate.call_args.kwargs["select"] + select_stmt = mock_paginate.call_args.args[0] visibility_clause = str(select_stmt._where_criteria[-1]) assert "maintainer" in visibility_clause assert "IN" in visibility_clause @@ -262,20 +262,20 @@ class TestDatasetServiceRetrievalPermissions: def test_get_datasets_rbac_include_all_uses_workspace_permission(self): mock_db = MagicMock() mock_db.session.scalars.return_value.all.return_value = [] - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL) mock_permissions = SimpleNamespace(workspace=SimpleNamespace(permission_keys=["dataset.create_and_management"])) with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), patch( "services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get", return_value=mock_permissions, ), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets( page=1, per_page=20, @@ -286,39 +286,39 @@ class TestDatasetServiceRetrievalPermissions: ) mock_db.session.scalars.assert_called_once() - mock_db.paginate.assert_called_once() - select_stmt = mock_db.paginate.call_args.kwargs["select"] + mock_paginate.assert_called_once() + select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 1 def test_get_datasets_rbac_without_user_returns_empty_result(self): mock_db = MagicMock() mock_db.session.scalars.return_value.all.return_value = [] - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", True), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets(page=1, per_page=20, session=mock_db.session, tenant_id="tenant-1", user=None) mock_db.session.scalars.assert_not_called() - mock_db.paginate.assert_called_once() - select_stmt = mock_db.paginate.call_args.kwargs["select"] + mock_paginate.assert_called_once() + select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 2 def test_get_datasets_legacy_owner_include_all_keeps_full_access(self): mock_db = MagicMock() mock_db.session.scalars.return_value.all.return_value = [] - mock_db.paginate.return_value.items = [] - mock_db.paginate.return_value.total = 0 user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.OWNER) with ( patch("services.dataset_service.db", mock_db), + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.dify_config.RBAC_ENABLED", False), ): + mock_paginate.return_value = SimpleNamespace(items=[], total=0) DatasetService.get_datasets( page=1, per_page=20, @@ -329,8 +329,8 @@ class TestDatasetServiceRetrievalPermissions: ) mock_db.session.scalars.assert_called_once() - mock_db.paginate.assert_called_once() - select_stmt = mock_db.paginate.call_args.kwargs["select"] + mock_paginate.assert_called_once() + select_stmt = mock_paginate.call_args.args[0] assert len(select_stmt._where_criteria) == 1 diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index f2c08324774..34f3f947f96 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -270,9 +270,10 @@ class TestSegmentServiceQueries: with ( patch("services.dataset_service.db") as mock_db, + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like, ): - mock_db.paginate.return_value = paginated + mock_paginate.return_value = paginated result = SegmentService.get_child_chunks( segment_id="segment-1", @@ -285,7 +286,7 @@ class TestSegmentServiceQueries: assert result is paginated escape_like.assert_called_once_with("needle") - mock_db.paginate.assert_called_once() + mock_paginate.assert_called_once() def test_get_child_chunk_by_id_returns_only_child_chunk_instances(self): child_chunk = _make_child_chunk() @@ -324,9 +325,10 @@ class TestSegmentServiceQueries: with ( patch("services.dataset_service.db") as mock_db, + patch("services.dataset_service.paginate_query") as mock_paginate, patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like, ): - mock_db.paginate.return_value = paginated + mock_paginate.return_value = paginated items, total = SegmentService.get_segments( document_id="doc-1", @@ -340,7 +342,7 @@ class TestSegmentServiceQueries: assert items == ["segment"] assert total == 1 escape_like.assert_called_once_with("needle") - mock_db.paginate.assert_called_once() + mock_paginate.assert_called_once() def test_get_segment_by_id_returns_only_document_segment_instances(self): segment = DocumentSegment( diff --git a/api/tests/unit_tests/services/test_external_dataset_service.py b/api/tests/unit_tests/services/test_external_dataset_service.py index 453ed7560c5..dbb4627759c 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -143,8 +143,10 @@ def factory(): class TestExternalDatasetServiceGetAPIs: """Test get_external_knowledge_apis operations - comprehensive coverage.""" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_apis_success_basic(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.paginate_query") + def test_get_external_knowledge_apis_success_basic( + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + ): """Test successful retrieval of external knowledge APIs with pagination.""" # Arrange tenant_id = "tenant-123" @@ -156,7 +158,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 5 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -168,11 +170,11 @@ class TestExternalDatasetServiceGetAPIs: assert result_total == 5 assert result_items[0].id == "api-0" assert result_items[4].id == "api-4" - mock_db.paginate.assert_called_once() + mock_paginate.assert_called_once() - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_with_search_filter( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with search filter.""" # Arrange @@ -184,7 +186,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -196,14 +198,16 @@ class TestExternalDatasetServiceGetAPIs: assert result_total == 1 assert result_items[0].name == "Production API" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_apis_empty_results(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + @patch("services.external_knowledge_service.paginate_query") + def test_get_external_knowledge_apis_empty_results( + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory + ): """Test retrieval with no results.""" # Arrange mock_pagination = MagicMock() mock_pagination.items = [] mock_pagination.total = 0 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -214,9 +218,9 @@ class TestExternalDatasetServiceGetAPIs: assert len(result_items) == 0 assert result_total == 0 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_large_result_set( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with large result set.""" # Arrange @@ -225,7 +229,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[:10] mock_pagination.total = 100 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -236,9 +240,9 @@ class TestExternalDatasetServiceGetAPIs: assert len(result_items) == 10 assert result_total == 100 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_pagination_last_page( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test last page pagination with partial results.""" # Arrange @@ -247,7 +251,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 100 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -258,9 +262,9 @@ class TestExternalDatasetServiceGetAPIs: assert len(result_items) == 5 assert result_total == 100 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_case_insensitive_search( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test case-insensitive search functionality.""" # Arrange @@ -272,7 +276,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 2 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -283,9 +287,9 @@ class TestExternalDatasetServiceGetAPIs: assert len(result_items) == 2 assert result_total == 2 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_special_characters_search( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test search with special characters.""" # Arrange @@ -294,7 +298,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -304,9 +308,9 @@ class TestExternalDatasetServiceGetAPIs: # Assert assert len(result_items) == 1 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_max_per_page_limit( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test that max_per_page limit is enforced.""" # Arrange @@ -315,7 +319,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis mock_pagination.total = 1000 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -323,12 +327,12 @@ class TestExternalDatasetServiceGetAPIs: ) # Assert - call_args = mock_db.paginate.call_args + call_args = mock_paginate.call_args assert call_args.kwargs["max_per_page"] == 100 - @patch("services.external_knowledge_service.db") + @patch("services.external_knowledge_service.paginate_query") def test_get_external_knowledge_apis_ordered_by_created_at_desc( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_paginate, factory: ExternalDatasetServiceTestDataFactory ): """Test that results are ordered by created_at descending.""" # Arrange @@ -340,7 +344,7 @@ class TestExternalDatasetServiceGetAPIs: mock_pagination = MagicMock() mock_pagination.items = apis[::-1] # Reversed to simulate DESC order mock_pagination.total = 5 - mock_db.paginate.return_value = mock_pagination + mock_paginate.return_value = mock_pagination # Act result_items, result_total = ExternalDatasetService.get_external_knowledge_apis( @@ -433,13 +437,13 @@ class TestExternalDatasetServiceValidateAPIList: class TestExternalDatasetServiceCreateAPI: """Test create_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_success_full( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory ): """Test successful creation with all fields.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -449,7 +453,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, mock_session) # Assert assert result.name == "Test API" @@ -458,57 +462,55 @@ class TestExternalDatasetServiceCreateAPI: assert result.created_by == user_id assert result.updated_by == user_id mock_check.assert_called_once_with(args["settings"]) - mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + mock_session.add.assert_called_once() + mock_session.commit.assert_called_once() - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_minimal_fields( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with minimal required fields.""" # Arrange + mock_session = MagicMock() args = { "name": "Minimal API", "settings": {"endpoint": "https://api.example.com", "api_key": "key"}, } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) # Assert assert result.name == "Minimal API" assert result.description == "" - @patch("services.external_knowledge_service.db") - def test_create_external_knowledge_api_missing_settings( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_create_external_knowledge_api_missing_settings(self, factory: ExternalDatasetServiceTestDataFactory): """Test creation fails when settings are missing.""" # Arrange + mock_session = MagicMock() args = {"name": "Test API", "description": "Test"} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) - @patch("services.external_knowledge_service.db") - def test_create_external_knowledge_api_none_settings(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_knowledge_api_none_settings(self, factory: ExternalDatasetServiceTestDataFactory): """Test creation fails when settings are explicitly None.""" # Arrange + mock_session = MagicMock() args = {"name": "Test API", "settings": None} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_settings_json_serialization( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory ): """Test that settings are properly JSON serialized.""" # Arrange + mock_session = MagicMock() settings = { "endpoint": "https://api.example.com", "api_key": "test-key", @@ -517,20 +519,20 @@ class TestExternalDatasetServiceCreateAPI: args = {"name": "Test API", "settings": settings} # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) # Assert assert isinstance(result.settings, str) parsed_settings = json.loads(result.settings) assert parsed_settings == settings - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_unicode_handling( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory ): """Test proper handling of Unicode characters in name and description.""" # Arrange + mock_session = MagicMock() args = { "name": "测试API", "description": "テストの説明", @@ -538,19 +540,19 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) # Assert assert result.name == "测试API" assert result.description == "テストの説明" - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_long_description( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory ): """Test creation with very long description.""" # Arrange + mock_session = MagicMock() long_description = "A" * 1000 args = { "name": "Test API", @@ -559,7 +561,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, mock_session) # Assert assert result.description == long_description @@ -822,43 +824,43 @@ class TestExternalDatasetServiceCheckEndpoint: class TestExternalDatasetServiceGetAPI: """Test get_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge API.""" # Arrange + mock_session = MagicMock() api_id = "api-123" expected_api = factory.create_external_knowledge_api_mock(api_id=api_id) - mock_db.session.scalar.return_value = expected_api + mock_session.scalar.return_value = expected_api # Act tenant_id = "tenant-123" - result = ExternalDatasetService.get_external_knowledge_api(mock_db.session, api_id, tenant_id) + result = ExternalDatasetService.get_external_knowledge_api(mock_session, api_id, tenant_id) # Assert assert result.id == api_id - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.get_external_knowledge_api(mock_db.session, "nonexistent-id", "tenant-123") + ExternalDatasetService.get_external_knowledge_api(mock_session, "nonexistent-id", "tenant-123") class TestExternalDatasetServiceUpdateAPI: """Test update_external_knowledge_api operations.""" @patch("services.external_knowledge_service.naive_utc_now") - @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_success_all_fields( - self, mock_db, mock_now, factory: ExternalDatasetServiceTestDataFactory + self, mock_now, factory: ExternalDatasetServiceTestDataFactory ): """Test successful update with all fields.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" user_id = "user-456" @@ -873,24 +875,24 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://new.example.com", "api_key": "new-key"}, } - mock_db.session.scalar.return_value = existing_api + mock_session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api(mock_db.session, tenant_id, user_id, api_id, args) + result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, user_id, api_id, args) # Assert assert result.name == "Updated API" assert result.description == "Updated description" assert result.updated_by == user_id assert result.updated_at == current_time - mock_db.session.commit.assert_called_once() + mock_session.commit.assert_called_once() - @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_preserve_hidden_api_key( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory ): """Test that hidden API key is preserved from existing settings.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" @@ -905,51 +907,47 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://api.example.com", "api_key": HIDDEN_VALUE}, } - mock_db.session.scalar.return_value = existing_api + mock_session.scalar.return_value = existing_api # Act - result = ExternalDatasetService.update_external_knowledge_api( - mock_db.session, tenant_id, "user-123", api_id, args - ) + result = ExternalDatasetService.update_external_knowledge_api(mock_session, tenant_id, "user-123", api_id, args) # Assert settings = json.loads(result.settings) assert settings["api_key"] == "original-secret-key" - @patch("services.external_knowledge_service.db") - def test_update_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_update_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_db.session, "tenant-123", "user-123", "api-123", args + mock_session, "tenant-123", "user-123", "api-123", args ) - @patch("services.external_knowledge_service.db") - def test_update_external_knowledge_api_tenant_mismatch( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_update_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when tenant ID doesn't match.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - mock_db.session, "wrong-tenant", "user-123", "api-123", args + mock_session, "wrong-tenant", "user-123", "api-123", args ) - @patch("services.external_knowledge_service.db") - def test_update_external_knowledge_api_name_only(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_update_external_knowledge_api_name_only(self, factory: ExternalDatasetServiceTestDataFactory): """Test updating only the name field.""" # Arrange + mock_session = MagicMock() existing_api = factory.create_external_knowledge_api_mock( description="Original description", settings={"endpoint": "https://api.example.com", "api_key": "key"}, @@ -957,11 +955,11 @@ class TestExternalDatasetServiceUpdateAPI: args = {"name": "New Name Only"} - mock_db.session.scalar.return_value = existing_api + mock_session.scalar.return_value = existing_api # Act result = ExternalDatasetService.update_external_knowledge_api( - mock_db.session, "tenant-123", "user-123", "api-123", args + mock_session, "tenant-123", "user-123", "api-123", args ) # Assert @@ -971,98 +969,92 @@ class TestExternalDatasetServiceUpdateAPI: class TestExternalDatasetServiceDeleteAPI: """Test delete_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") - def test_delete_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_delete_external_knowledge_api_success(self, factory: ExternalDatasetServiceTestDataFactory): """Test successful deletion of external knowledge API.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" existing_api = factory.create_external_knowledge_api_mock(api_id=api_id, tenant_id=tenant_id) - mock_db.session.scalar.return_value = existing_api + mock_session.scalar.return_value = existing_api # Act - ExternalDatasetService.delete_external_knowledge_api(mock_db.session, tenant_id, api_id) + ExternalDatasetService.delete_external_knowledge_api(mock_session, tenant_id, api_id) # Assert - mock_db.session.delete.assert_called_once_with(existing_api) - mock_db.session.commit.assert_called_once() + mock_session.delete.assert_called_once_with(existing_api) + mock_session.commit.assert_called_once() - @patch("services.external_knowledge_service.db") - def test_delete_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_delete_external_knowledge_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when API is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_db.session, "tenant-123", "api-123") + ExternalDatasetService.delete_external_knowledge_api(mock_session, "tenant-123", "api-123") - @patch("services.external_knowledge_service.db") - def test_delete_external_knowledge_api_tenant_mismatch( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_delete_external_knowledge_api_tenant_mismatch(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when tenant ID doesn't match.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api(mock_db.session, "wrong-tenant", "api-123") + ExternalDatasetService.delete_external_knowledge_api(mock_session, "wrong-tenant", "api-123") class TestExternalDatasetServiceAPIUseCheck: """Test external_knowledge_api_use_check operations.""" - @patch("services.external_knowledge_service.db") - def test_external_knowledge_api_use_check_in_use_single( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_external_knowledge_api_use_check_in_use_single(self, factory: ExternalDatasetServiceTestDataFactory): """Test API use check when API has one binding.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 1 + mock_session.scalar.return_value = 1 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) # Assert assert in_use is True assert count == 1 - assert "tenant_id" in str(mock_db.session.scalar.call_args.args[0]) + assert "tenant_id" in str(mock_session.scalar.call_args.args[0]) - @patch("services.external_knowledge_service.db") - def test_external_knowledge_api_use_check_in_use_multiple( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_external_knowledge_api_use_check_in_use_multiple(self, factory: ExternalDatasetServiceTestDataFactory): """Test API use check with multiple bindings.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 10 + mock_session.scalar.return_value = 10 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) # Assert assert in_use is True assert count == 10 - @patch("services.external_knowledge_service.db") - def test_external_knowledge_api_use_check_not_in_use(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_external_knowledge_api_use_check_not_in_use(self, factory: ExternalDatasetServiceTestDataFactory): """Test API use check when API is not in use.""" # Arrange + mock_session = MagicMock() api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 0 + mock_session.scalar.return_value = 0 # Act - in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_db.session, api_id, tenant_id) + in_use, count = ExternalDatasetService.external_knowledge_api_use_check(mock_session, api_id, tenant_id) # Assert assert in_use is False @@ -1072,48 +1064,46 @@ class TestExternalDatasetServiceAPIUseCheck: class TestExternalDatasetServiceGetBinding: """Test get_external_knowledge_binding_with_dataset_id operations.""" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_binding_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_binding_success(self, factory: ExternalDatasetServiceTestDataFactory): """Test successful retrieval of external knowledge binding.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" expected_binding = factory.create_external_knowledge_binding_mock(tenant_id=tenant_id, dataset_id=dataset_id) - mock_db.session.scalar.return_value = expected_binding + mock_session.scalar.return_value = expected_binding # Act result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_db.session, tenant_id, dataset_id + mock_session, tenant_id, dataset_id ) # Assert assert result.dataset_id == dataset_id assert result.tenant_id == tenant_id - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_binding_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_binding_not_found(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when binding is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="external knowledge binding not found"): ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - mock_db.session, "tenant-123", "dataset-123" + mock_session, "tenant-123", "dataset-123" ) class TestExternalDatasetServiceDocumentValidate: """Test document_create_args_validate operations.""" - @patch("services.external_knowledge_service.db") - def test_document_create_args_validate_success_all_params( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_document_create_args_validate_success_all_params(self, factory: ExternalDatasetServiceTestDataFactory): """Test successful validation with all required parameters.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1127,19 +1117,17 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_db.session.scalar.return_value = api + mock_session.scalar.return_value = api process_parameter = {"param1": "value1", "param2": "value2"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_db.session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) - @patch("services.external_knowledge_service.db") - def test_document_create_args_validate_missing_required_param( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_document_create_args_validate_missing_required_param(self, factory: ExternalDatasetServiceTestDataFactory): """Test validation fails when required parameter is missing.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" api_id = "api-123" @@ -1147,44 +1135,42 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - mock_db.session.scalar.return_value = api + mock_session.scalar.return_value = api process_parameter = {} # Act & Assert with pytest.raises(ValueError, match="required_param is required"): - ExternalDatasetService.document_create_args_validate(mock_db.session, tenant_id, api_id, process_parameter) + ExternalDatasetService.document_create_args_validate(mock_session, tenant_id, api_id, process_parameter) - @patch("services.external_knowledge_service.db") - def test_document_create_args_validate_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_document_create_args_validate_api_not_found(self, factory: ExternalDatasetServiceTestDataFactory): """Test validation fails when API is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.document_create_args_validate(mock_db.session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) - @patch("services.external_knowledge_service.db") - def test_document_create_args_validate_no_custom_parameters( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_document_create_args_validate_no_custom_parameters(self, factory: ExternalDatasetServiceTestDataFactory): """Test validation succeeds when no custom parameters defined.""" # Arrange + mock_session = MagicMock() settings = {} api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_db.session.scalar.return_value = api + mock_session.scalar.return_value = api # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate(mock_db.session, "tenant-123", "api-123", {}) + ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", {}) - @patch("services.external_knowledge_service.db") def test_document_create_args_validate_optional_params_not_required( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory ): """Test that optional parameters don't cause validation failure.""" # Arrange + mock_session = MagicMock() settings = { "document_process_setting": [ {"name": "required_param", "required": True}, @@ -1194,14 +1180,12 @@ class TestExternalDatasetServiceDocumentValidate: api = factory.create_external_knowledge_api_mock(settings=[settings]) - mock_db.session.scalar.return_value = api + mock_session.scalar.return_value = api process_parameter = {"required_param": "value"} # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate( - mock_db.session, "tenant-123", "api-123", process_parameter - ) + ExternalDatasetService.document_create_args_validate(mock_session, "tenant-123", "api-123", process_parameter) class TestExternalDatasetServiceProcessAPI: @@ -1491,10 +1475,10 @@ class TestExternalDatasetServiceGetSettings: class TestExternalDatasetServiceCreateDataset: """Test create_external_dataset operations.""" - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_success_full(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_dataset_success_full(self, factory: ExternalDatasetServiceTestDataFactory): """Test successful creation of external dataset with all fields.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" user_id = "user-123" args = { @@ -1507,90 +1491,84 @@ class TestExternalDatasetServiceCreateDataset: api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_db.session.scalar.side_effect = [None, api] + mock_session.scalar.side_effect = [None, api] # Act - result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, mock_db.session) + result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, mock_session) # Assert assert result.name == "Test External Dataset" assert result.description == "Comprehensive test description" assert result.provider == "external" assert result.created_by == user_id - mock_db.session.add.assert_called() - mock_db.session.commit.assert_called_once() + mock_session.add.assert_called() + mock_session.commit.assert_called_once() - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_duplicate_name_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_create_external_dataset_duplicate_name_error(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when dataset name already exists.""" # Arrange + mock_session = MagicMock() existing_dataset = factory.create_dataset_mock(name="Duplicate Dataset") - mock_db.session.scalar.return_value = existing_dataset + mock_session.scalar.return_value = existing_dataset args = {"name": "Duplicate Dataset"} # Act & Assert with pytest.raises(DatasetNameDuplicateError): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_api_not_found_error(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_dataset_api_not_found_error(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when external knowledge API is not found.""" # Arrange - mock_db.session.scalar.side_effect = [None, None] + mock_session = MagicMock() + mock_session.scalar.side_effect = [None, None] args = {"name": "Test Dataset", "external_knowledge_api_id": "nonexistent-api"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_missing_knowledge_id_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_create_external_dataset_missing_knowledge_id_error(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when external_knowledge_id is missing.""" # Arrange + mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [None, api] + mock_session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_api_id": "api-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_missing_api_id_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory - ): + def test_create_external_dataset_missing_api_id_error(self, factory: ExternalDatasetServiceTestDataFactory): """Test error when external_knowledge_api_id is missing.""" # Arrange + mock_session = MagicMock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [None, api] + mock_session.scalar.side_effect = [None, api] args = {"name": "Test Dataset", "external_knowledge_id": "knowledge-123"} # Act & Assert with pytest.raises(ValueError, match="external_knowledge_api_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, mock_session) class TestExternalDatasetServiceFetchRetrieval: """Test fetch_external_knowledge_retrieval operations.""" @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_success_with_results( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test successful external knowledge retrieval with results.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" query = "test query for retrieval" @@ -1600,7 +1578,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1616,7 +1594,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, tenant_id, dataset_id, query, external_retrieval_parameters + mock_session, tenant_id, dataset_id, query, external_retrieval_parameters ) # Assert @@ -1624,46 +1602,46 @@ class TestExternalDatasetServiceFetchRetrieval: assert result[0]["content"] == "result 1" assert result[1]["score"] == 0.8 - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_binding_not_found_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory ): """Test error when external knowledge binding is not found.""" # Arrange - mock_db.session.scalar.return_value = None + mock_session = MagicMock() + mock_session.scalar.return_value = None # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external knowledge binding not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {} + mock_session, "tenant-123", "dataset-123", "query", {} ) - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_cross_tenant_api_template_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory ): """Test error when a binding points to an API template outside the dataset tenant.""" # Arrange + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() - mock_db.session.scalar.side_effect = [binding, None] + mock_session.scalar.side_effect = [binding, None] # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external api template not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {} + mock_session, "tenant-123", "dataset-123", "query", {} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_results( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with empty results.""" # Arrange + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1672,23 +1650,23 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) # Assert assert len(result) == 0 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_with_score_threshold( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test retrieval with score threshold enabled.""" # Arrange + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1703,7 +1681,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", external_retrieval_parameters + mock_session, "tenant-123", "dataset-123", "query", external_retrieval_parameters ) # Assert @@ -1713,16 +1691,16 @@ class TestExternalDatasetServiceFetchRetrieval: assert call_args.params["retrieval_setting"]["score_threshold"] == 0.75 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_non_200_status_raises_exception( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test that non-200 status code raises Exception with response text.""" # Arrange + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 500 @@ -1732,7 +1710,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="Internal Server Error: Database connection failed"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @pytest.mark.parametrize( @@ -1749,12 +1727,12 @@ class TestExternalDatasetServiceFetchRetrieval: ], ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_various_error_status_codes( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message ): """Test that various error status codes raise exceptions with response text.""" # Arrange + mock_session = MagicMock() tenant_id = "tenant-123" dataset_id = "dataset-123" @@ -1763,7 +1741,7 @@ class TestExternalDatasetServiceFetchRetrieval: ) api = factory.create_external_knowledge_api_mock(api_id="api-123") - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = status_code @@ -1773,20 +1751,20 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match=re.escape(error_message)): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, tenant_id, dataset_id, "query", {"top_k": 5} + mock_session, tenant_id, dataset_id, "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_response_text( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory ): """Test exception with empty response text.""" # Arrange + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 503 @@ -1796,17 +1774,17 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_process, factory): """Test malformed JSON success responses are normalized to external retrieval errors.""" + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1815,17 +1793,17 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_process, factory): """Test malformed success payload shapes are normalized to external retrieval errors.""" + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1834,17 +1812,17 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_process, factory): """Test non-list records payloads are normalized to external retrieval errors.""" + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_response = MagicMock() mock_response.status_code = 200 @@ -1853,20 +1831,20 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_process, factory): """Test transport/runtime failures are normalized to external retrieval errors.""" + mock_session = MagicMock() binding = factory.create_external_knowledge_binding_mock() api = factory.create_external_knowledge_api_mock() - mock_db.session.scalar.side_effect = [binding, api] + mock_session.scalar.side_effect = [binding, api] mock_process.side_effect = RuntimeError("connection reset by peer") with pytest.raises(ExternalKnowledgeRetrievalError, match="connection reset by peer"): ExternalDatasetService.fetch_external_knowledge_retrieval( - mock_db.session, "tenant-123", "dataset-123", "query", {"top_k": 5} + mock_session, "tenant-123", "dataset-123", "query", {"top_k": 5} )