From 5cb76f5effffcced285e23fa5c6386f4f5344cbe Mon Sep 17 00:00:00 2001 From: "Byron.wang" Date: Fri, 3 Jul 2026 16:15:37 +0800 Subject: [PATCH] refactor: manage rag pipeline sessions explicitly (#38274) --- .../datasets/rag_pipeline/rag_pipeline.py | 15 +- .../rag_pipeline/rag_pipeline_workflow.py | 16 +- .../rag_pipeline/rag_pipeline_workflow.py | 6 +- .../rag_pipeline/pipeline_generate_service.py | 13 +- .../built_in/built_in_retrieval.py | 7 +- .../customized/customized_retrieval.py | 29 +- .../database/database_retrieval.py | 22 +- .../pipeline_template_base.py | 8 +- .../remote/remote_retrieval.py | 11 +- api/services/rag_pipeline/rag_pipeline.py | 661 ++++++++++-------- .../rag_pipeline_transform_service.py | 38 +- .../rag_pipeline/test_rag_pipeline.py | 8 +- .../test_rag_pipeline_workflow.py | 10 +- .../test_rag_pipeline_service_db.py | 22 +- .../rag_pipeline/test_rag_pipeline.py | 24 +- .../pipeline_template/conftest.py | 18 + .../test_built_in_retrieval.py | 8 +- .../test_customized_retrieval.py | 18 +- .../test_database_retrieval.py | 18 +- .../test_pipeline_template_base.py | 13 +- .../test_remote_retrieval.py | 10 +- .../test_pipeline_generate_service.py | 71 +- .../rag_pipeline/test_rag_pipeline_service.py | 639 +++++++++-------- .../test_rag_pipeline_task_proxy.py | 2 +- .../test_rag_pipeline_transform_service.py | 78 +-- 25 files changed, 972 insertions(+), 793 deletions(-) create mode 100644 api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py index 3187124f121..4027fa487a2 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py @@ -5,7 +5,7 @@ from flask import request from flask_restx import Resource from pydantic import BaseModel, Field from sqlalchemy import select -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import NotFound from controllers.common.fields import SimpleDataResponse @@ -16,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.wraps import ( account_initialization_required, enterprise_license_required, @@ -102,10 +103,13 @@ class PipelineTemplateListApi(Resource): @account_initialization_required @enterprise_license_required @with_current_tenant_id - def get(self, current_tenant_id: str) -> JsonResponseWithStatus: + @with_session + def get(self, session: Session, current_tenant_id: str) -> JsonResponseWithStatus: query = PipelineTemplateListQuery.model_validate(request.args.to_dict(flat=True)) # get pipeline templates - pipeline_templates = RagPipelineService.get_pipeline_templates(query.type, query.language, current_tenant_id) + pipeline_templates = RagPipelineService.get_pipeline_templates( + session, query.type, query.language, current_tenant_id + ) return dump_response(PipelineTemplateListResponse, pipeline_templates), 200 @@ -117,10 +121,11 @@ class PipelineTemplateDetailApi(Resource): @login_required @account_initialization_required @enterprise_license_required - def get(self, template_id: str) -> JsonResponseWithStatus: + @with_session + def get(self, session: Session, template_id: str) -> JsonResponseWithStatus: query = PipelineTemplateDetailQuery.model_validate(request.args.to_dict(flat=True)) rag_pipeline_service = RagPipelineService() - pipeline_template = rag_pipeline_service.get_pipeline_template_detail(template_id, query.type) + pipeline_template = rag_pipeline_service.get_pipeline_template_detail(session, template_id, query.type) if pipeline_template is None: raise NotFound("Pipeline template not found from upstream service.") return dump_response(PipelineTemplateDetailResponse, pipeline_template), 200 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index e0a6ee0a83e..c52385f6cf2 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -6,7 +6,7 @@ from uuid import UUID from flask import abort, request from flask_restx import Resource from pydantic import BaseModel, Field, RootModel, ValidationError -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, Forbidden, InternalServerError, NotFound import services @@ -26,6 +26,7 @@ from controllers.console.app.workflow import ( WorkflowPaginationResponse, WorkflowResponse, ) +from controllers.console.app.wraps import with_session from controllers.console.datasets.wraps import get_rag_pipeline from controllers.console.wraps import ( RBACPermission, @@ -343,7 +344,8 @@ class DraftRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline): + @with_session + def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run draft workflow """ @@ -352,6 +354,7 @@ class DraftRagPipelineRunApi(Resource): try: response = PipelineGenerateService.generate( + session=session, pipeline=pipeline, user=current_user, args=args, @@ -375,7 +378,8 @@ class PublishedRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline): + @with_session + def post(self, session: Session, current_user: Account, pipeline: Pipeline): """ Run published workflow """ @@ -385,6 +389,7 @@ class PublishedRagPipelineRunApi(Resource): try: response = PipelineGenerateService.generate( + session=session, pipeline=pipeline, user=current_user, args=args, @@ -1014,13 +1019,14 @@ class RagPipelineTransformApi(Resource): @login_required @account_initialization_required @with_current_user - def post(self, current_user: Account, dataset_id: UUID): + @with_session + def post(self, session: Session, current_user: Account, dataset_id: UUID): if not (current_user.has_edit_permission or current_user.is_dataset_operator): raise Forbidden() dataset_id_str = str(dataset_id) rag_pipeline_transform_service = RagPipelineTransformService() - result = rag_pipeline_transform_service.transform_dataset(dataset_id_str, db.session) + result = rag_pipeline_transform_service.transform_dataset(dataset_id_str, session) return result diff --git a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py index a6a61262cdc..8c4063398bc 100644 --- a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py @@ -5,6 +5,7 @@ from uuid import UUID from flask import request from pydantic import BaseModel, Field, RootModel from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -16,6 +17,7 @@ from controllers.common.schema import ( register_schema_model, register_schema_models, ) +from controllers.console.app.wraps import with_session from controllers.service_api import service_api_ns from controllers.service_api.dataset.error import PipelineRunError from controllers.service_api.dataset.rag_pipeline.serializers import serialize_upload_file @@ -264,7 +266,8 @@ class PipelineRunApi(DatasetApiResource): "Pipeline run successfully", service_api_ns.models[GeneratedAppResponse.__name__], ) - def post(self, tenant_id: str, dataset_id: UUID): + @with_session + def post(self, session: Session, tenant_id: str, dataset_id: UUID): """Resource for running a rag pipeline.""" dataset_id_str = str(dataset_id) # Verify dataset ownership @@ -282,6 +285,7 @@ class PipelineRunApi(DatasetApiResource): pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str) try: response: dict[Any, Any] | Generator[str, Any, None] = PipelineGenerateService.generate( + session=session, pipeline=pipeline, user=current_user, args=payload.model_dump(), diff --git a/api/services/rag_pipeline/pipeline_generate_service.py b/api/services/rag_pipeline/pipeline_generate_service.py index 56bc7859589..e77ff9687ed 100644 --- a/api/services/rag_pipeline/pipeline_generate_service.py +++ b/api/services/rag_pipeline/pipeline_generate_service.py @@ -1,10 +1,11 @@ from collections.abc import Mapping from typing import Any +from sqlalchemy.orm import Session + from configs import dify_config from core.app.apps.pipeline.pipeline_generator import PipelineGenerator from core.app.entities.app_invoke_entities import InvokeFrom -from extensions.ext_database import db from models.dataset import Document, Pipeline from models.enums import IndexingStatus from models.model import Account, App, EndUser @@ -16,6 +17,7 @@ class PipelineGenerateService: @classmethod def generate( cls, + session: Session, pipeline: Pipeline, user: Account | EndUser, args: Mapping[str, Any], @@ -35,7 +37,7 @@ class PipelineGenerateService: workflow = cls._get_workflow(pipeline, invoke_from) if original_document_id := args.get("original_document_id"): # update document status to waiting - cls.update_document_status(original_document_id) + cls.update_document_status(original_document_id, session) return PipelineGenerator.convert_to_event_stream( PipelineGenerator().generate( pipeline=pipeline, @@ -105,13 +107,12 @@ class PipelineGenerateService: return workflow @classmethod - def update_document_status(cls, document_id: str): + def update_document_status(cls, document_id: str, session: Session): """ Update document status to waiting :param document_id: document id """ - document = db.session.get(Document, document_id) + document = session.get(Document, document_id) if document: document.indexing_status = IndexingStatus.WAITING - db.session.add(document) - db.session.commit() + session.add(document) diff --git a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py index 4e4cf2d19f5..6de0be33a4a 100644 --- a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import Any, override from flask import current_app +from sqlalchemy.orm import Session from services.rag_pipeline.pipeline_template.pipeline_template_base import PipelineTemplateRetrievalBase from services.rag_pipeline.pipeline_template.pipeline_template_type import PipelineTemplateType @@ -21,13 +22,15 @@ class BuiltInPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return PipelineTemplateType.BUILTIN @override - def get_pipeline_templates(self, language: str, current_tenant_id: str | None = None) -> dict[str, Any]: + def get_pipeline_templates( + self, session: Session, language: str, current_tenant_id: str | None = None + ) -> dict[str, Any]: del current_tenant_id result = self.fetch_pipeline_templates_from_builtin(language) return result @override - def get_pipeline_template_detail(self, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: result = self.fetch_pipeline_template_detail_from_builtin(template_id) return result diff --git a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py index 57dfefed2e0..4faaf342f66 100644 --- a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py @@ -2,8 +2,8 @@ from typing import Any, TypedDict, override import yaml from sqlalchemy import select +from sqlalchemy.orm import Session -from extensions.ext_database import db from libs.login import resolve_tenant_id_fallback from models.dataset import PipelineCustomizedTemplate from services.rag_pipeline.pipeline_template.pipeline_template_base import PipelineTemplateRetrievalBase @@ -40,29 +40,38 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_templates(self, language: str, current_tenant_id: str | None = None) -> dict[str, Any]: + def get_pipeline_templates( + self, session: Session, language: str, current_tenant_id: str | None = None + ) -> dict[str, Any]: current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) - return self.fetch_pipeline_templates_from_customized(tenant_id=current_tenant_id, language=language) + return self.fetch_pipeline_templates_from_customized( + session=session, tenant_id=current_tenant_id, language=language + ) @override - def get_pipeline_template_detail(self, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(template_id) + def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(session, template_id) @override def get_type(self) -> str: return PipelineTemplateType.CUSTOMIZED @classmethod - def fetch_pipeline_templates_from_customized(cls, tenant_id: str, language: str) -> dict[str, Any]: + def fetch_pipeline_templates_from_customized( + cls, session: Session, tenant_id: str, language: str + ) -> dict[str, Any]: """ Fetch pipeline templates from db. :param tenant_id: tenant id :param language: language :return: """ - pipeline_customized_templates = db.session.scalars( + pipeline_customized_templates = session.scalars( select(PipelineCustomizedTemplate) - .where(PipelineCustomizedTemplate.tenant_id == tenant_id, PipelineCustomizedTemplate.language == language) + .where( + PipelineCustomizedTemplate.tenant_id == tenant_id, + PipelineCustomizedTemplate.language == language, + ) .order_by(PipelineCustomizedTemplate.position.asc(), PipelineCustomizedTemplate.created_at.desc()) ).all() recommended_pipelines_results: list[CustomizedTemplateItemDict] = [] @@ -80,13 +89,13 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param template_id: Template ID :return: """ - pipeline_template = db.session.get(PipelineCustomizedTemplate, template_id) + pipeline_template = session.get(PipelineCustomizedTemplate, template_id) if not pipeline_template: return None diff --git a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py index 0f6d0727c76..f6d2731e21a 100644 --- a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py @@ -2,8 +2,8 @@ from typing import Any, TypedDict, override import yaml from sqlalchemy import select +from sqlalchemy.orm import Session -from extensions.ext_database import db from models.dataset import PipelineBuiltInTemplate from services.rag_pipeline.pipeline_template.pipeline_template_base import PipelineTemplateRetrievalBase from services.rag_pipeline.pipeline_template.pipeline_template_type import PipelineTemplateType @@ -40,20 +40,22 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_templates(self, language: str, current_tenant_id: str | None = None) -> dict[str, Any]: + def get_pipeline_templates( + self, session: Session, language: str, current_tenant_id: str | None = None + ) -> dict[str, Any]: del current_tenant_id - return self.fetch_pipeline_templates_from_db(language) + return self.fetch_pipeline_templates_from_db(session, language) @override - def get_pipeline_template_detail(self, template_id: str) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(template_id) + def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(session, template_id) @override def get_type(self) -> str: return PipelineTemplateType.DATABASE @classmethod - def fetch_pipeline_templates_from_db(cls, language: str) -> dict[str, Any]: + def fetch_pipeline_templates_from_db(cls, session: Session, language: str) -> dict[str, Any]: """ Fetch pipeline templates from db. :param language: language @@ -61,9 +63,7 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ pipeline_built_in_templates = list( - db.session.scalars( - select(PipelineBuiltInTemplate).where(PipelineBuiltInTemplate.language == language) - ).all() + session.scalars(select(PipelineBuiltInTemplate).where(PipelineBuiltInTemplate.language == language)).all() ) recommended_pipelines_results: list[PipelineTemplateItemDict] = [] @@ -83,14 +83,14 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, template_id: str) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db(cls, session: Session, template_id: str) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param pipeline_id: Pipeline ID :return: """ # is in public recommended list - pipeline_template = db.session.get(PipelineBuiltInTemplate, template_id) + pipeline_template = session.get(PipelineBuiltInTemplate, template_id) if not pipeline_template: return None diff --git a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py index 84d8f5674bb..ff53dc1f79e 100644 --- a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py +++ b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py @@ -1,11 +1,15 @@ from typing import Any, Protocol +from sqlalchemy.orm import Session + class PipelineTemplateRetrievalBase(Protocol): """Interface for pipeline template retrieval.""" - def get_pipeline_templates(self, language: str, current_tenant_id: str | None = None) -> dict[str, Any]: ... + def get_pipeline_templates( + self, session: Session, language: str, current_tenant_id: str | None = None + ) -> dict[str, Any]: ... - def get_pipeline_template_detail(self, template_id: str) -> dict[str, Any] | None: ... + def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: ... def get_type(self) -> str: ... diff --git a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py index 5cf46915ab0..7f9fe1b56ea 100644 --- a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py @@ -2,6 +2,7 @@ import logging from typing import Any, override import httpx +from sqlalchemy.orm import Session from configs import dify_config from services.rag_pipeline.pipeline_template.database.database_retrieval import DatabasePipelineTemplateRetrieval @@ -17,21 +18,23 @@ class RemotePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_template_detail(self, template_id: str) -> dict[str, Any] | None: + def get_pipeline_template_detail(self, session: Session, template_id: str) -> dict[str, Any] | None: try: return self.fetch_pipeline_template_detail_from_dify_official(template_id) except Exception as e: logger.warning("fetch recommended app detail from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db(template_id) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_template_detail_from_db(session, template_id) @override - def get_pipeline_templates(self, language: str, current_tenant_id: str | None = None) -> dict[str, Any]: + def get_pipeline_templates( + self, session: Session, language: str, current_tenant_id: str | None = None + ) -> dict[str, Any]: del current_tenant_id try: return self.fetch_pipeline_templates_from_dify_official(language) except Exception as e: logger.warning("fetch pipeline templates from dify official failed: %r, switch to database.", e) - return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(language) + return DatabasePipelineTemplateRetrieval.fetch_pipeline_templates_from_db(session, language) @override def get_type(self) -> str: diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index 10112ef19a0..9e17a05be16 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -27,6 +27,7 @@ from core.datasource.entities.datasource_entities import ( from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin from core.datasource.online_drive.online_drive_plugin import OnlineDriveDatasourcePlugin from core.datasource.website_crawl.website_crawl_plugin import WebsiteCrawlDatasourcePlugin +from core.db.session_factory import session_factory from core.helper import marketplace from core.rag.entities import DatasourceCompletedEvent, DatasourceErrorEvent, DatasourceProcessingEvent from core.repositories.factory import DifyCoreRepositoryFactory, OrderConfig @@ -98,7 +99,8 @@ class RagPipelineService: def __init__(self, session_maker: sessionmaker | None = None): """Initialize RagPipelineService with repository dependencies.""" if session_maker is None: - session_maker = sessionmaker(bind=db.engine, expire_on_commit=False) + session_maker = session_factory.get_session_maker() + self._session_maker = session_maker self._node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( session_maker ) @@ -107,6 +109,7 @@ class RagPipelineService: @classmethod def get_pipeline_templates( cls, + session: Session, type: str = "built-in", language: str = "en-US", current_tenant_id: str | None = None, @@ -114,7 +117,7 @@ class RagPipelineService: if type == "built-in": mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) if not result.get("pipeline_templates") and language != "en-US": template_retrieval = PipelineTemplateRetrievalFactory.get_built_in_pipeline_template_retrieval() result = template_retrieval.fetch_pipeline_templates_from_builtin("en-US") @@ -122,11 +125,13 @@ class RagPipelineService: else: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - result = retrieval_instance.get_pipeline_templates(language, current_tenant_id) + result = retrieval_instance.get_pipeline_templates(session, language, current_tenant_id) return result @classmethod - def get_pipeline_template_detail(cls, template_id: str, type: str = "built-in") -> dict[str, Any] | None: + def get_pipeline_template_detail( + cls, session: Session, template_id: str, type: str = "built-in" + ) -> dict[str, Any] | None: """ Get pipeline template detail. @@ -137,7 +142,9 @@ class RagPipelineService: if type == "built-in": mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - built_in_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail(template_id) + built_in_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( + session, template_id + ) if built_in_result is None: logger.warning( "pipeline template retrieval returned empty result, template_id: %s, mode: %s", @@ -148,7 +155,9 @@ class RagPipelineService: else: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() - customized_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail(template_id) + customized_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( + session, template_id + ) return customized_result @classmethod @@ -158,6 +167,7 @@ class RagPipelineService: template_info: PipelineTemplateInfoEntity, current_user: Account | None = None, current_tenant_id: str | None = None, + session: Session | None = None, ): """ Update pipeline template. @@ -165,7 +175,17 @@ class RagPipelineService: :param template_info: template info """ current_user, current_tenant_id = resolve_account_fallback(current_user, current_tenant_id) - customized_template: PipelineCustomizedTemplate | None = db.session.scalar( + if session is None: + with session_factory.get_session_maker().begin() as new_session: + return cls.update_customized_pipeline_template( + template_id, + template_info, + current_user, + current_tenant_id, + session=new_session, + ) + + customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( PipelineCustomizedTemplate.id == template_id, @@ -178,7 +198,7 @@ class RagPipelineService: # check template name is exist template_name = template_info.name if template_name: - template = db.session.scalar( + template = session.scalar( select(PipelineCustomizedTemplate) .where( PipelineCustomizedTemplate.name == template_name, @@ -193,16 +213,22 @@ class RagPipelineService: customized_template.description = template_info.description customized_template.icon = template_info.icon_info.model_dump() customized_template.updated_by = current_user.id - db.session.commit() return customized_template @classmethod - def delete_customized_pipeline_template(cls, template_id: str, current_tenant_id: str | None = None): + def delete_customized_pipeline_template( + cls, template_id: str, current_tenant_id: str | None = None, session: Session | None = None + ): """ Delete customized pipeline template. """ current_tenant_id = resolve_tenant_id_fallback(current_tenant_id) - customized_template: PipelineCustomizedTemplate | None = db.session.scalar( + if session is None: + with session_factory.get_session_maker().begin() as new_session: + cls.delete_customized_pipeline_template(template_id, current_tenant_id, session=new_session) + return + + customized_template: PipelineCustomizedTemplate | None = session.scalar( select(PipelineCustomizedTemplate) .where( PipelineCustomizedTemplate.id == template_id, @@ -212,23 +238,23 @@ class RagPipelineService: ) if not customized_template: raise ValueError("Customized pipeline template not found.") - db.session.delete(customized_template) - db.session.commit() + session.delete(customized_template) def get_draft_workflow(self, pipeline: Pipeline) -> Workflow | None: """ Get draft workflow """ # fetch draft workflow by rag pipeline - workflow = db.session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == "draft", + with self._session_maker() as session: + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == "draft", + ) + .limit(1) ) - .limit(1) - ) # return draft workflow return workflow @@ -242,29 +268,31 @@ class RagPipelineService: return None # fetch published workflow by workflow_id - workflow = db.session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == pipeline.workflow_id, + with self._session_maker() as session: + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == pipeline.workflow_id, + ) + .limit(1) ) - .limit(1) - ) return workflow def get_published_workflow_by_id(self, pipeline: Pipeline, workflow_id: str) -> Workflow | None: """Fetch a published workflow snapshot by ID for restore operations.""" - workflow = db.session.scalar( - select(Workflow) - .where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.id == workflow_id, + with self._session_maker() as session: + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == workflow_id, + ) + .limit(1) ) - .limit(1) - ) if workflow and workflow.version == Workflow.VERSION_DRAFT: raise IsDraftWorkflowError("source workflow must be published") return workflow @@ -322,39 +350,51 @@ class RagPipelineService: Sync draft workflow :raises WorkflowHashNotEqualError """ - # fetch draft workflow by app_model - workflow = self.get_draft_workflow(pipeline=pipeline) + with self._session_maker.begin() as session: + managed_pipeline = session.get(Pipeline, pipeline.id) + if not managed_pipeline: + raise ValueError("Pipeline not found") - if workflow and workflow.unique_hash != unique_hash: - raise WorkflowHashNotEqualError() - - # create draft workflow if not found - if not workflow: - workflow = Workflow( - tenant_id=pipeline.tenant_id, - app_id=pipeline.id, - features="{}", - type=WorkflowType.RAG_PIPELINE.value, - version="draft", - graph=json.dumps(graph), - created_by=account.id, - environment_variables=environment_variables, - conversation_variables=conversation_variables, - rag_pipeline_variables=rag_pipeline_variables, + # fetch draft workflow by app_model + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == managed_pipeline.tenant_id, + Workflow.app_id == managed_pipeline.id, + Workflow.version == "draft", + ) + .limit(1) ) - db.session.add(workflow) - db.session.flush() - pipeline.workflow_id = workflow.id - # update draft workflow if found - else: - workflow.graph = json.dumps(graph) - workflow.updated_by = account.id - workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) - workflow.environment_variables = environment_variables - workflow.conversation_variables = conversation_variables - workflow.rag_pipeline_variables = rag_pipeline_variables - # commit db session changes - db.session.commit() + + if workflow and workflow.unique_hash != unique_hash: + raise WorkflowHashNotEqualError() + + # create draft workflow if not found + if not workflow: + workflow = Workflow( + tenant_id=managed_pipeline.tenant_id, + app_id=managed_pipeline.id, + features="{}", + type=WorkflowType.RAG_PIPELINE.value, + version="draft", + graph=json.dumps(graph), + created_by=account.id, + environment_variables=environment_variables, + conversation_variables=conversation_variables, + rag_pipeline_variables=rag_pipeline_variables, + ) + session.add(workflow) + session.flush() + managed_pipeline.workflow_id = workflow.id + pipeline.workflow_id = workflow.id + # update draft workflow if found + else: + workflow.graph = json.dumps(graph) + workflow.updated_by = account.id + workflow.updated_at = datetime.now(UTC).replace(tzinfo=None) + workflow.environment_variables = environment_variables + workflow.conversation_variables = conversation_variables + workflow.rag_pipeline_variables = rag_pipeline_variables # trigger workflow events TODO # app_draft_workflow_was_synced.send(pipeline, synced_draft_workflow=workflow) @@ -375,26 +415,48 @@ class RagPipelineService: the pipeline-specific flush/link step that wires a newly created draft back onto ``pipeline.workflow_id``. """ - source_workflow = self.get_published_workflow_by_id(pipeline=pipeline, workflow_id=workflow_id) - if not source_workflow: - raise WorkflowNotFoundError("Workflow not found.") + with self._session_maker.begin() as session: + managed_pipeline = session.get(Pipeline, pipeline.id) + if not managed_pipeline: + raise ValueError("Pipeline not found") - draft_workflow = self.get_draft_workflow(pipeline=pipeline) - draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( - tenant_id=pipeline.tenant_id, - app_id=pipeline.id, - source_workflow=source_workflow, - draft_workflow=draft_workflow, - account=account, - updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), - ) + source_workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == managed_pipeline.tenant_id, + Workflow.app_id == managed_pipeline.id, + Workflow.id == workflow_id, + ) + .limit(1) + ) + if source_workflow and source_workflow.version == Workflow.VERSION_DRAFT: + raise IsDraftWorkflowError("source workflow must be published") + if not source_workflow: + raise WorkflowNotFoundError("Workflow not found.") - if is_new_draft: - db.session.add(draft_workflow) - db.session.flush() - pipeline.workflow_id = draft_workflow.id + draft_workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == managed_pipeline.tenant_id, + Workflow.app_id == managed_pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, + ) + .limit(1) + ) + draft_workflow, is_new_draft = apply_published_workflow_snapshot_to_draft( + tenant_id=managed_pipeline.tenant_id, + app_id=managed_pipeline.id, + source_workflow=source_workflow, + draft_workflow=draft_workflow, + account=account, + updated_at_factory=lambda: datetime.now(UTC).replace(tzinfo=None), + ) - db.session.commit() + if is_new_draft: + session.add(draft_workflow) + session.flush() + managed_pipeline.workflow_id = draft_workflow.id + pipeline.workflow_id = draft_workflow.id return draft_workflow @@ -571,7 +633,7 @@ class RagPipelineService: workflow_node_execution.id ) - with sessionmaker(bind=db.engine).begin() as session: + with self._session_maker.begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -988,23 +1050,22 @@ class RagPipelineService: dataset_id = get_system_segment(variable_pool, SystemVariableKey.DATASET_ID) pipeline_id = get_system_segment(variable_pool, SystemVariableKey.APP_ID) if document_id and dataset_id and pipeline_id: - document = db.session.scalar( - select(Document) - .join(Dataset, Dataset.id == Document.dataset_id) - .where( - Document.id == document_id.value, - Document.tenant_id == tenant_id, - Document.dataset_id == dataset_id.value, - Dataset.tenant_id == tenant_id, - Dataset.pipeline_id == pipeline_id.value, + with self._session_maker.begin() as session: + document = session.scalar( + select(Document) + .join(Dataset, Dataset.id == Document.dataset_id) + .where( + Document.id == document_id.value, + Document.tenant_id == tenant_id, + Document.dataset_id == dataset_id.value, + Dataset.tenant_id == tenant_id, + Dataset.pipeline_id == pipeline_id.value, + ) + .limit(1) ) - .limit(1) - ) - if document: - document.indexing_status = IndexingStatus.ERROR - document.error = error - db.session.add(document) - db.session.commit() + if document: + document.indexing_status = IndexingStatus.ERROR + document.error = error return workflow_node_execution @@ -1220,82 +1281,81 @@ class RagPipelineService: Publish customized pipeline template """ current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - pipeline = db.session.get(Pipeline, pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - if not pipeline.workflow_id: - raise ValueError("Pipeline workflow not found") - workflow = db.session.get(Workflow, pipeline.workflow_id) - if not workflow: - raise ValueError("Workflow not found") - with sessionmaker(db.engine).begin() as session: + with session_factory.get_session_maker().begin() as session: + pipeline = session.get(Pipeline, pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + if not pipeline.workflow_id: + raise ValueError("Pipeline workflow not found") + workflow = session.get(Workflow, pipeline.workflow_id) + if not workflow: + raise ValueError("Workflow not found") dataset = pipeline.retrieve_dataset(session=session) if not dataset: raise ValueError("Dataset not found") - # check template name is exist - template_name = args.get("name") - if template_name: - template = db.session.scalar( - select(PipelineCustomizedTemplate) - .where( - PipelineCustomizedTemplate.name == template_name, - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, + # check template name is exist + template_name = args.get("name") + if template_name: + template = session.scalar( + select(PipelineCustomizedTemplate) + .where( + PipelineCustomizedTemplate.name == template_name, + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, + ) + .limit(1) + ) + if template: + raise ValueError("Template name is already exists") + + max_position = session.scalar( + select(func.max(PipelineCustomizedTemplate.position)).where( + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id ) - .limit(1) ) - if template: - raise ValueError("Template name is already exists") - max_position = db.session.scalar( - select(func.max(PipelineCustomizedTemplate.position)).where( - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id - ) - ) + from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineDslService - - with sessionmaker(db.engine).begin() as session: rag_pipeline_dsl_service = RagPipelineDslService(session) dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) - if args.get("icon_info") is None: - args["icon_info"] = {} - if args.get("description") is None: - raise ValueError("Description is required") - if args.get("name") is None: - raise ValueError("Name is required") - pipeline_customized_template = PipelineCustomizedTemplate( - name=args.get("name") or "", - description=args.get("description") or "", - icon=args.get("icon_info") or {}, - tenant_id=pipeline.tenant_id, - yaml_content=dsl, - install_count=0, - position=max_position + 1 if max_position else 1, - chunk_structure=dataset.chunk_structure, - language="en-US", - created_by=current_user.id, - ) - db.session.add(pipeline_customized_template) - db.session.commit() + if args.get("icon_info") is None: + args["icon_info"] = {} + if args.get("description") is None: + raise ValueError("Description is required") + if args.get("name") is None: + raise ValueError("Name is required") + pipeline_customized_template = PipelineCustomizedTemplate( + name=args.get("name") or "", + description=args.get("description") or "", + icon=args.get("icon_info") or {}, + tenant_id=pipeline.tenant_id, + yaml_content=dsl, + install_count=0, + position=max_position + 1 if max_position else 1, + chunk_structure=dataset.chunk_structure, + language="en-US", + created_by=current_user.id, + ) + session.add(pipeline_customized_template) def is_workflow_exist(self, pipeline: Pipeline) -> bool: - return ( - db.session.scalar( - select(func.count(Workflow.id)).where( - Workflow.tenant_id == pipeline.tenant_id, - Workflow.app_id == pipeline.id, - Workflow.version == Workflow.VERSION_DRAFT, + with self._session_maker() as session: + return ( + session.scalar( + select(func.count(Workflow.id)).where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, + ) ) - ) - or 0 - ) > 0 + or 0 + ) > 0 def get_node_last_run( self, pipeline: Pipeline, workflow: Workflow, node_id: str ) -> WorkflowNodeExecutionModel | None: node_execution_service_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository( - sessionmaker(db.engine) + self._session_maker ) node_exec = node_execution_service_repo.get_node_last_execution( @@ -1371,7 +1431,7 @@ class RagPipelineService: # Convert node_execution to WorkflowNodeExecution after save workflow_node_execution_db_model = repository._to_db_model(workflow_node_execution) # type: ignore - with sessionmaker(bind=db.engine).begin() as session: + with self._session_maker.begin() as session: draft_var_saver = DraftVariableSaver( session=session, app_id=pipeline.id, @@ -1405,7 +1465,10 @@ class RagPipelineService: if type and type != "all": stmt = stmt.where(PipelineRecommendedPlugin.type == type) - pipeline_recommended_plugins = db.session.scalars(stmt.order_by(PipelineRecommendedPlugin.position.asc())).all() + with self._session_maker() as session: + pipeline_recommended_plugins = session.scalars( + stmt.order_by(PipelineRecommendedPlugin.position.asc()) + ).all() if not pipeline_recommended_plugins: return { @@ -1444,139 +1507,173 @@ class RagPipelineService: """ Retry error document """ - document_pipeline_execution_log = db.session.scalar( - select(DocumentPipelineExecutionLog).where(DocumentPipelineExecutionLog.document_id == document.id).limit(1) - ) - if not document_pipeline_execution_log: - raise ValueError("Document pipeline execution log not found") - pipeline = db.session.get(Pipeline, document_pipeline_execution_log.pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") - # convert to app config - workflow = self.get_published_workflow(pipeline) - if not workflow: - raise ValueError("Workflow not found") - PipelineGenerator().generate( - pipeline=pipeline, - workflow=workflow, - user=user, - args={ - "inputs": document_pipeline_execution_log.input_data, - "start_node_id": document_pipeline_execution_log.datasource_node_id, - "datasource_type": document_pipeline_execution_log.datasource_type, - "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], - "original_document_id": document.id, - }, - invoke_from=InvokeFrom.PUBLISHED_PIPELINE, - streaming=False, - call_depth=0, - workflow_thread_pool_id=None, - is_retry=True, - ) + with self._session_maker() as session: + document_pipeline_execution_log = session.scalar( + select(DocumentPipelineExecutionLog) + .where(DocumentPipelineExecutionLog.document_id == document.id) + .limit(1) + ) + if not document_pipeline_execution_log: + raise ValueError("Document pipeline execution log not found") + pipeline = session.get(Pipeline, document_pipeline_execution_log.pipeline_id) + if not pipeline: + raise ValueError("Pipeline not found") + # convert to app config + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == pipeline.workflow_id, + ) + .limit(1) + ) + if not workflow: + raise ValueError("Workflow not found") + PipelineGenerator().generate( + pipeline=pipeline, + workflow=workflow, + user=user, + args={ + "inputs": document_pipeline_execution_log.input_data, + "start_node_id": document_pipeline_execution_log.datasource_node_id, + "datasource_type": document_pipeline_execution_log.datasource_type, + "datasource_info_list": [json.loads(document_pipeline_execution_log.datasource_info)], + "original_document_id": document.id, + }, + invoke_from=InvokeFrom.PUBLISHED_PIPELINE, + streaming=False, + call_depth=0, + workflow_thread_pool_id=None, + is_retry=True, + ) def get_datasource_plugins(self, tenant_id: str, dataset_id: str, is_published: bool) -> list[dict]: """ Get datasource plugins """ - dataset: Dataset | None = db.session.scalar( - select(Dataset) - .where( - Dataset.id == dataset_id, - Dataset.tenant_id == tenant_id, - ) - .limit(1) - ) - if not dataset: - raise ValueError("Dataset not found") - pipeline: Pipeline | None = db.session.scalar( - select(Pipeline) - .where( - Pipeline.id == dataset.pipeline_id, - Pipeline.tenant_id == tenant_id, - ) - .limit(1) - ) - if not pipeline: - raise ValueError("Pipeline not found") - - workflow: Workflow | None = None - if is_published: - workflow = self.get_published_workflow(pipeline=pipeline) - else: - workflow = self.get_draft_workflow(pipeline=pipeline) - if not pipeline or not workflow: - raise ValueError("Pipeline or workflow not found") - - datasource_nodes = workflow.graph_dict.get("nodes", []) - datasource_plugins = [] - for datasource_node in datasource_nodes: - if datasource_node.get("data", {}).get("type") == "datasource": - datasource_node_data = datasource_node["data"] - if not datasource_node_data: - continue - - variables = workflow.rag_pipeline_variables - if variables: - variables_map = {item["variable"]: item for item in variables} - else: - variables_map = {} - - datasource_parameters = datasource_node_data.get("datasource_parameters", {}) - user_input_variables_keys = [] - user_input_variables = [] - - for _, value in datasource_parameters.items(): - if value.get("value") and isinstance(value.get("value"), str): - pattern = r"\{\{#([a-zA-Z0-9_]{1,50}(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" - match = re.match(pattern, value["value"]) - if match: - full_path = match.group(1) - last_part = full_path.split(".")[-1] - user_input_variables_keys.append(last_part) - elif value.get("value") and isinstance(value.get("value"), list): - last_part = value.get("value")[-1] - user_input_variables_keys.append(last_part) - for key, value in variables_map.items(): - if key in user_input_variables_keys: - user_input_variables.append(value) - - # get credentials - datasource_provider_service: DatasourceProviderService = DatasourceProviderService() - credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( - tenant_id=tenant_id, - provider=datasource_node_data.get("provider_name"), - plugin_id=datasource_node_data.get("plugin_id"), + with self._session_maker() as session: + dataset: Dataset | None = session.scalar( + select(Dataset) + .where( + Dataset.id == dataset_id, + Dataset.tenant_id == tenant_id, ) - credential_info_list: list[Any] = [] - for credential in credentials: - credential_info_list.append( + .limit(1) + ) + if not dataset: + raise ValueError("Dataset not found") + pipeline: Pipeline | None = session.scalar( + select(Pipeline) + .where( + Pipeline.id == dataset.pipeline_id, + Pipeline.tenant_id == tenant_id, + ) + .limit(1) + ) + if not pipeline: + raise ValueError("Pipeline not found") + + if is_published: + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.id == pipeline.workflow_id, + ) + .limit(1) + ) + else: + workflow = session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, + ) + .limit(1) + ) + if not pipeline or not workflow: + raise ValueError("Pipeline or workflow not found") + + datasource_nodes = workflow.graph_dict.get("nodes", []) + datasource_plugins = [] + for datasource_node in datasource_nodes: + if datasource_node.get("data", {}).get("type") == "datasource": + datasource_node_data = datasource_node["data"] + if not datasource_node_data: + continue + + variables = workflow.rag_pipeline_variables + if variables: + variables_map = {item["variable"]: item for item in variables} + else: + variables_map = {} + + datasource_parameters = datasource_node_data.get("datasource_parameters", {}) + user_input_variables_keys = [] + user_input_variables = [] + + for _, value in datasource_parameters.items(): + if value.get("value") and isinstance(value.get("value"), str): + pattern = ( + r"\{\{#([a-zA-Z0-9_]{1,50}" + r"(?:\.[a-zA-Z0-9_][a-zA-Z0-9_]{0,29}){1,10})#\}\}" + ) + match = re.match(pattern, value["value"]) + if match: + full_path = match.group(1) + last_part = full_path.split(".")[-1] + user_input_variables_keys.append(last_part) + elif value.get("value") and isinstance(value.get("value"), list): + last_part = value.get("value")[-1] + user_input_variables_keys.append(last_part) + for key, value in variables_map.items(): + if key in user_input_variables_keys: + user_input_variables.append(value) + + # get credentials + datasource_provider_service: DatasourceProviderService = DatasourceProviderService() + credentials: list[dict[Any, Any]] = datasource_provider_service.list_datasource_credentials( + tenant_id=tenant_id, + provider=datasource_node_data.get("provider_name"), + plugin_id=datasource_node_data.get("plugin_id"), + ) + credential_info_list: list[Any] = [] + for credential in credentials: + credential_info_list.append( + { + "id": credential.get("id"), + "name": credential.get("name"), + "type": credential.get("type"), + "is_default": credential.get("is_default"), + } + ) + + datasource_plugins.append( { - "id": credential.get("id"), - "name": credential.get("name"), - "type": credential.get("type"), - "is_default": credential.get("is_default"), + "node_id": datasource_node.get("id"), + "plugin_id": datasource_node_data.get("plugin_id"), + "provider_name": datasource_node_data.get("provider_name"), + "datasource_type": datasource_node_data.get("provider_type"), + "title": datasource_node_data.get("title"), + "user_input_variables": user_input_variables, + "credentials": credential_info_list, } ) - datasource_plugins.append( - { - "node_id": datasource_node.get("id"), - "plugin_id": datasource_node_data.get("plugin_id"), - "provider_name": datasource_node_data.get("provider_name"), - "datasource_type": datasource_node_data.get("provider_type"), - "title": datasource_node_data.get("title"), - "user_input_variables": user_input_variables, - "credentials": credential_info_list, - } - ) + return datasource_plugins - return datasource_plugins - - def get_pipeline(self, tenant_id: str, dataset_id: str) -> Pipeline: + def get_pipeline(self, tenant_id: str, dataset_id: str, session: Session | None = None) -> Pipeline: """ Get pipeline """ - dataset: Dataset | None = db.session.scalar( + if session is None: + with self._session_maker() as new_session: + return self.get_pipeline(tenant_id, dataset_id, session=new_session) + + dataset: Dataset | None = session.scalar( select(Dataset) .where( Dataset.id == dataset_id, @@ -1586,7 +1683,7 @@ class RagPipelineService: ) if not dataset: raise ValueError("Dataset not found") - pipeline: Pipeline | None = db.session.scalar( + pipeline: Pipeline | None = session.scalar( select(Pipeline) .where( Pipeline.id == dataset.pipeline_id, diff --git a/api/services/rag_pipeline/rag_pipeline_transform_service.py b/api/services/rag_pipeline/rag_pipeline_transform_service.py index dc3eeae201e..daefaa9e30e 100644 --- a/api/services/rag_pipeline/rag_pipeline_transform_service.py +++ b/api/services/rag_pipeline/rag_pipeline_transform_service.py @@ -8,7 +8,7 @@ from uuid import uuid4 import yaml from flask_login import current_user from sqlalchemy import select -from sqlalchemy.orm import scoped_session +from sqlalchemy.orm import Session from configs import dify_config from constants import DOCUMENT_EXTENSIONS @@ -16,7 +16,6 @@ from core.plugin.impl.plugin import PluginInstaller from core.plugin.plugin_service import PluginService from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.retrieval.retrieval_methods import RetrievalMethod -from extensions.ext_database import db from factories import variable_factory from models.dataset import Dataset, Document, DocumentPipelineExecutionLog, Pipeline from models.enums import DatasetRuntimeMode, DataSourceType @@ -29,7 +28,7 @@ logger = logging.getLogger(__name__) class RagPipelineTransformService: - def transform_dataset(self, dataset_id: str, session: scoped_session): + def transform_dataset(self, dataset_id: str, session: Session): dataset = session.get(Dataset, dataset_id) if not dataset: raise ValueError("Dataset not found") @@ -45,11 +44,11 @@ class RagPipelineTransformService: indexing_technique = dataset.indexing_technique if not datasource_type and not indexing_technique: - return self._transform_to_empty_pipeline(dataset) + return self._transform_to_empty_pipeline(dataset, session=session) doc_form = dataset.doc_form if not doc_form: - return self._transform_to_empty_pipeline(dataset) + return self._transform_to_empty_pipeline(dataset, session=session) retrieval_model = RetrievalSetting.model_validate(dataset.retrieval_model) if dataset.retrieval_model else None pipeline_yaml = self._get_transform_yaml(doc_form, datasource_type, indexing_technique) # deal dependencies @@ -81,7 +80,7 @@ class RagPipelineTransformService: workflow_data["graph"] = graph pipeline_yaml["workflow"] = workflow_data # create pipeline - pipeline = self._create_pipeline(pipeline_yaml) + pipeline = self._create_pipeline(pipeline_yaml, session=session) # save chunk structure to dataset if doc_form == IndexStructureType.PARENT_CHILD_INDEX: @@ -97,7 +96,7 @@ class RagPipelineTransformService: # deal document data self._deal_document_data(dataset, session) - session.commit() + session.flush() return { "pipeline_id": pipeline.id, "dataset_id": dataset_id, @@ -195,6 +194,7 @@ class RagPipelineTransformService: def _create_pipeline( self, data: dict[str, Any], + session: Session, ) -> Pipeline: """Create a new app or update an existing one.""" pipeline_data = data.get("rag_pipeline", {}) @@ -227,8 +227,8 @@ class RagPipelineTransformService: ) pipeline.id = str(uuid4()) - db.session.add(pipeline) - db.session.flush() + session.add(pipeline) + session.flush() # create draft workflow draft_workflow = Workflow( tenant_id=pipeline.tenant_id, @@ -254,11 +254,11 @@ class RagPipelineTransformService: conversation_variables=conversation_variables, rag_pipeline_variables=rag_pipeline_variables_list, ) - db.session.add(draft_workflow) - db.session.add(published_workflow) - db.session.flush() + session.add(draft_workflow) + session.add(published_workflow) + session.flush() pipeline.workflow_id = published_workflow.id - db.session.add(pipeline) + session.add(pipeline) return pipeline def _deal_dependencies(self, pipeline_yaml: dict[str, Any], tenant_id: str): @@ -289,29 +289,29 @@ class RagPipelineTransformService: logger.debug("Installing missing pipeline plugins %s", need_install_plugin_unique_identifiers) PluginService.install_from_marketplace_pkg(tenant_id, need_install_plugin_unique_identifiers) - def _transform_to_empty_pipeline(self, dataset: Dataset): + def _transform_to_empty_pipeline(self, dataset: Dataset, session: Session): pipeline = Pipeline( tenant_id=dataset.tenant_id, name=dataset.name, description=dataset.description, created_by=current_user.id, ) - db.session.add(pipeline) - db.session.flush() + session.add(pipeline) + session.flush() dataset.pipeline_id = pipeline.id dataset.runtime_mode = DatasetRuntimeMode.RAG_PIPELINE dataset.updated_by = current_user.id dataset.updated_at = datetime.now(UTC).replace(tzinfo=None) - db.session.add(dataset) - db.session.commit() + session.add(dataset) + session.flush() return { "pipeline_id": pipeline.id, "dataset_id": dataset.id, "status": "success", } - def _deal_document_data(self, dataset: Dataset, session: scoped_session): + def _deal_document_data(self, dataset: Dataset, session: Session): file_node_id = "1752479895761" notion_node_id = "1752489759475" jina_node_id = "1752491761974" diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index 75b0d3c5002..c34810c97d0 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -54,7 +54,7 @@ class TestPipelineTemplateListApi: return_value=templates, ), ): - response, status = method(api, str(uuid4())) + response, status = method(api, MagicMock(), str(uuid4())) assert status == 200 assert response == { @@ -100,7 +100,7 @@ class TestPipelineTemplateDetailApi: return_value=service, ), ): - response, status = method(api, "tpl-1") + response, status = method(api, MagicMock(), "tpl-1") assert status == 200 assert response == {**template, "created_by": None} @@ -120,7 +120,7 @@ class TestPipelineTemplateDetailApi: ), ): with pytest.raises(NotFound): - method(api, "non-existent-id") + method(api, MagicMock(), "non-existent-id") def test_get_returns_404_for_customized_type_not_found(self, app: Flask) -> None: api = PipelineTemplateDetailApi() @@ -137,7 +137,7 @@ class TestPipelineTemplateDetailApi: ), ): with pytest.raises(NotFound): - method(api, "non-existent-id") + method(api, MagicMock(), "non-existent-id") class TestCustomizedPipelineTemplateApi: diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py index bdec903ef33..e1bdff4d23c 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py @@ -470,7 +470,7 @@ class TestPipelineRunApis: return_value={"ok": True}, ), ): - assert method(api, user, pipeline) == {"ok": True} + assert method(api, MagicMock(), user, pipeline) == {"ok": True} def test_draft_run_rate_limit(self, app: Flask) -> None: api = DraftRagPipelineRunApi() @@ -498,7 +498,7 @@ class TestPipelineRunApis: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, user, pipeline) + method(api, MagicMock(), user, pipeline) class TestDraftNodeRun: @@ -600,7 +600,7 @@ class TestMiscApis: app.test_request_context("/"), ): with pytest.raises(Forbidden): - method(api, user, "ds1") + method(api, MagicMock(spec=Session), user, "ds1") def test_recommended_plugins(self, app: Flask) -> None: api = RagPipelineRecommendedPluginApi() @@ -655,7 +655,7 @@ class TestPublishedRagPipelineRunApi: return_value={"ok": True}, ), ): - result = method(api, user, pipeline) + result = method(api, MagicMock(), user, pipeline) assert result == {"ok": True} def test_published_run_rate_limit(self, app: Flask) -> None: @@ -681,7 +681,7 @@ class TestPublishedRagPipelineRunApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, user, pipeline) + method(api, MagicMock(), user, pipeline) class TestDefaultBlockConfigApi: diff --git a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py index 8f126e1cff0..2e7df67d266 100644 --- a/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py +++ b/api/tests/test_containers_integration_tests/services/rag_pipeline/test_rag_pipeline_service_db.py @@ -101,8 +101,8 @@ class TestRagPipelineServiceGetPipeline: service = self._make_service(flask_app_with_containers) - with pytest.raises(ValueError, match="(Dataset not found|Pipeline not found)"): - service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) + with pytest.raises(ValueError, match="Pipeline not found"): + service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) def test_get_pipeline_returns_pipeline_when_found( self, db_session_with_containers: Session, flask_app_with_containers: Flask @@ -117,7 +117,7 @@ class TestRagPipelineServiceGetPipeline: service = self._make_service(flask_app_with_containers) - result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id) + result = service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset.id, session=db_session_with_containers) assert result.id == pipeline.id @@ -165,7 +165,9 @@ class TestUpdateCustomizedPipelineTemplate: description="Updated description", icon_info=IconInfo(icon="🔥"), ) - result = RagPipelineService.update_customized_pipeline_template(template.id, info, account, tenant_id) + result = RagPipelineService.update_customized_pipeline_template( + template.id, info, account, tenant_id, session=db_session_with_containers + ) assert result.name == "Updated Name" assert result.description == "Updated description" @@ -203,7 +205,9 @@ class TestUpdateCustomizedPipelineTemplate: icon_info=IconInfo(icon="📄"), ) with pytest.raises(ValueError, match="Template name is already exists"): - RagPipelineService.update_customized_pipeline_template(template1.id, info, account, tenant_id) + RagPipelineService.update_customized_pipeline_template( + template1.id, info, account, tenant_id, session=db_session_with_containers + ) class TestDeleteCustomizedPipelineTemplate: @@ -241,14 +245,14 @@ class TestDeleteCustomizedPipelineTemplate: template_id = template.id db_session_with_containers.flush() - RagPipelineService.delete_customized_pipeline_template(template_id, tenant_id) + RagPipelineService.delete_customized_pipeline_template( + template_id, tenant_id, session=db_session_with_containers + ) # Verify the record is deleted within the same context from sqlalchemy import select - from extensions.ext_database import db as ext_db - - remaining = ext_db.session.scalar( + remaining = db_session_with_containers.scalar( select(PipelineCustomizedTemplate).where(PipelineCustomizedTemplate.id == template_id) ) assert remaining is None diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index bca2d73ad9f..2a1970d3837 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -1,7 +1,7 @@ from __future__ import annotations from inspect import unwrap -from unittest.mock import PropertyMock, patch +from unittest.mock import Mock, PropertyMock, patch import pytest from flask import Flask @@ -64,7 +64,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates(template_type: str, language: str, current_tenant_id: str) -> dict[str, object]: + def get_pipeline_templates( + session: Mock, template_type: str, language: str, current_tenant_id: str + ) -> dict[str, object]: service_calls.append((template_type, language, current_tenant_id)) return {"pipeline_templates": [_template_item()]} @@ -72,7 +74,7 @@ class TestPipelineTemplateListApi: app.test_request_context("/rag/pipeline/templates"), patch.object(module.RagPipelineService, "get_pipeline_templates", side_effect=get_pipeline_templates), ): - response, status = method(api, tenant_id) + response, status = method(api, Mock(), tenant_id) assert status == 200 assert service_calls == [("built-in", "en-US", tenant_id)] @@ -92,7 +94,9 @@ class TestPipelineTemplateListApi: tenant_id = "tenant-1" service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_templates(template_type: str, language: str, current_tenant_id: str) -> dict[str, object]: + def get_pipeline_templates( + session: Mock, template_type: str, language: str, current_tenant_id: str + ) -> dict[str, object]: service_calls.append((template_type, language, current_tenant_id)) return {"pipeline_templates": []} @@ -100,7 +104,7 @@ class TestPipelineTemplateListApi: app.test_request_context("/rag/pipeline/templates?type=customized&language=ja-JP"), patch.object(module.RagPipelineService, "get_pipeline_templates", side_effect=get_pipeline_templates), ): - response, status = method(api, tenant_id) + response, status = method(api, Mock(), tenant_id) assert status == 200 assert response == {"pipeline_templates": []} @@ -114,7 +118,9 @@ class TestPipelineTemplateDetailApi: service_calls: list[tuple[str, str]] = [] class Service: - def get_pipeline_template_detail(self, template_id: str, template_type: str) -> dict[str, object]: + def get_pipeline_template_detail( + self, session: Mock, template_id: str, template_type: str + ) -> dict[str, object]: service_calls.append((template_id, template_type)) return _template_detail() @@ -122,7 +128,7 @@ class TestPipelineTemplateDetailApi: app.test_request_context("/rag/pipeline/templates/template-1?type=customized"), patch.object(module, "RagPipelineService", Service), ): - response, status = method(api, "template-1") + response, status = method(api, Mock(), "template-1") assert status == 200 assert response == {**_template_detail(), "created_by": None} @@ -133,7 +139,7 @@ class TestPipelineTemplateDetailApi: method = unwrap(api.get) class Service: - def get_pipeline_template_detail(self, template_id: str, template_type: str) -> None: + def get_pipeline_template_detail(self, session: Mock, template_id: str, template_type: str) -> None: return None with ( @@ -141,7 +147,7 @@ class TestPipelineTemplateDetailApi: patch.object(module, "RagPipelineService", Service), ): with pytest.raises(NotFound): - method(api, "missing") + method(api, Mock(), "missing") class TestCustomizedPipelineTemplateApi: diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py new file mode 100644 index 00000000000..2e65be48086 --- /dev/null +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py @@ -0,0 +1,18 @@ +from collections.abc import Callable +from unittest.mock import Mock + +import pytest +from pytest_mock import MockerFixture + + +@pytest.fixture +def patch_session_factory(mocker: MockerFixture) -> Callable[[str, Mock], Mock]: + def _patch(module_path: str, session_mock: Mock) -> Mock: + session_context = mocker.MagicMock() + session_context.__enter__.return_value = session_mock + session_context.__exit__.return_value = None + session_maker = mocker.Mock(return_value=session_context) + mocker.patch(f"{module_path}.session_factory.get_session_maker", return_value=session_maker) + return session_maker + + return _patch diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py index 441a914ee62..5bc41fdb5bd 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py @@ -23,7 +23,7 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: ) retrieval = BuiltInPipelineTemplateRetrieval() - templates = retrieval.get_pipeline_templates("en-US") + templates = retrieval.get_pipeline_templates(mocker.Mock(), "en-US") assert templates == {"pipeline_templates": [{"id": "tpl-1"}]} @@ -40,7 +40,7 @@ def test_get_pipeline_template_detail(mocker: MockerFixture) -> None: ) retrieval = BuiltInPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail("tpl-1") + detail = retrieval.get_pipeline_template_detail(mocker.Mock(), "tpl-1") assert detail == {"id": "tpl-1", "name": "Template 1"} @@ -53,7 +53,7 @@ def test_get_pipeline_templates_missing_language_returns_empty_dict(mocker: Mock ) retrieval = BuiltInPipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates("fr-FR") + result = retrieval.get_pipeline_templates(mocker.Mock(), "fr-FR") assert result == {} @@ -66,7 +66,7 @@ def test_get_pipeline_template_detail_returns_none_for_unknown_id(mocker: Mocker ) retrieval = BuiltInPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail("nonexistent-id") + result = retrieval.get_pipeline_template_detail(mocker.Mock(), "nonexistent-id") assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py index 168ec8fce3c..b3ef79961d3 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py @@ -19,13 +19,9 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: scalars_mock.all.return_value = [customized_template] session_mock = mocker.Mock() session_mock.scalars.return_value = scalars_mock - mocker.patch( - "services.rag_pipeline.pipeline_template.customized.customized_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates("en-US", "tenant-id") + result = retrieval.get_pipeline_templates(session_mock, "en-US", "tenant-id") assert retrieval.get_type() == PipelineTemplateType.CUSTOMIZED assert result == { @@ -53,13 +49,9 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N yaml_content="workflow:\n graph:\n edges: []", created_user_name="creator", ) - mocker.patch( - "services.rag_pipeline.pipeline_template.customized.customized_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = CustomizedPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail("tpl-1") + detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") assert detail == { "id": "tpl-1", @@ -76,12 +68,8 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: MockerFixture) -> None: session_mock = mocker.Mock() session_mock.get.return_value = None - mocker.patch( - "services.rag_pipeline.pipeline_template.customized.customized_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail("missing") + result = retrieval.get_pipeline_template_detail(session_mock, "missing") assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py index 41f60e45755..cae79175b1c 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py @@ -21,13 +21,9 @@ def test_get_pipeline_templates(mocker: MockerFixture) -> None: scalars_mock.all.return_value = [built_in_template] session_mock = mocker.Mock() session_mock.scalars.return_value = scalars_mock - mocker.patch( - "services.rag_pipeline.pipeline_template.database.database_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_templates("en-US") + result = retrieval.get_pipeline_templates(session_mock, "en-US") assert retrieval.get_type() == PipelineTemplateType.DATABASE assert result == { @@ -56,13 +52,9 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N chunk_structure="general", yaml_content="workflow:\n graph:\n nodes: []", ) - mocker.patch( - "services.rag_pipeline.pipeline_template.database.database_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = DatabasePipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail("tpl-1") + detail = retrieval.get_pipeline_template_detail(session_mock, "tpl-1") assert detail == { "id": "tpl-1", @@ -78,12 +70,8 @@ def test_get_pipeline_template_detail_returns_detail(mocker: MockerFixture) -> N def test_get_pipeline_template_detail_returns_none_when_not_found(mocker: MockerFixture) -> None: session_mock = mocker.Mock() session_mock.get.return_value = None - mocker.patch( - "services.rag_pipeline.pipeline_template.database.database_retrieval.db", - new=SimpleNamespace(session=session_mock), - ) retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail("missing") + result = retrieval.get_pipeline_template_detail(session_mock, "missing") assert result is None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py index 5918d74f891..c8af1869732 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py @@ -1,11 +1,15 @@ +from unittest.mock import Mock + from services.rag_pipeline.pipeline_template.pipeline_template_base import PipelineTemplateRetrievalBase class DummyRetrieval(PipelineTemplateRetrievalBase): - def get_pipeline_templates(self, language: str) -> dict: + def get_pipeline_templates(self, session: Mock, language: str, current_tenant_id: str | None = None) -> dict: + del session, current_tenant_id return {"language": language} - def get_pipeline_template_detail(self, template_id: str) -> dict | None: + def get_pipeline_template_detail(self, session: Mock, template_id: str) -> dict | None: + del session return {"id": template_id} def get_type(self) -> str: @@ -14,7 +18,8 @@ class DummyRetrieval(PipelineTemplateRetrievalBase): def test_pipeline_template_retrieval_base_concrete_implementation() -> None: retrieval = DummyRetrieval() + session = Mock() - assert retrieval.get_pipeline_templates("en-US") == {"language": "en-US"} - assert retrieval.get_pipeline_template_detail("tpl-1") == {"id": "tpl-1"} + assert retrieval.get_pipeline_templates(session, "en-US") == {"language": "en-US"} + assert retrieval.get_pipeline_template_detail(session, "tpl-1") == {"id": "tpl-1"} assert retrieval.get_type() == "dummy" diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py index 5da6684926c..8f55b4b1c2f 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py @@ -18,13 +18,14 @@ def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFix return_value={"pipeline_templates": [{"id": "db-1"}]}, ) retrieval = RemotePipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_templates("en-US") + result = retrieval.get_pipeline_templates(session, "en-US") assert retrieval.get_type() == PipelineTemplateType.REMOTE assert result == {"pipeline_templates": [{"id": "db-1"}]} fetch_mock.assert_called_once_with("en-US") - fallback_mock.assert_called_once_with("en-US") + fallback_mock.assert_called_once_with(session, "en-US") def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: MockerFixture) -> None: @@ -39,12 +40,13 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: Moc return_value={"id": "db-1"}, ) retrieval = RemotePipelineTemplateRetrieval() + session = mocker.Mock() - result = retrieval.get_pipeline_template_detail("tpl-1") + result = retrieval.get_pipeline_template_detail(session, "tpl-1") assert result == {"id": "db-1"} fetch_mock.assert_called_once_with("tpl-1") - fallback_mock.assert_called_once_with("tpl-1") + fallback_mock.assert_called_once_with(session, "tpl-1") def test_fetch_pipeline_templates_from_dify_official(mocker: MockerFixture) -> None: diff --git a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py index 178f4595359..0ae2ba97f1a 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py @@ -1,15 +1,30 @@ +from collections.abc import Iterator from types import SimpleNamespace from typing import cast +from uuid import uuid4 import pytest from pytest_mock import MockerFixture +from sqlalchemy import create_engine, func, select +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom -from models.dataset import Pipeline +from models.dataset import Document, Pipeline +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus from models.model import Account, App, EndUser from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService +@pytest.fixture +def document_session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + Document.__table__.create(engine) + session_factory = sessionmaker(bind=engine, expire_on_commit=False) + with session_factory() as session: + yield session + engine.dispose() + + def test_get_max_active_requests_uses_smallest_non_zero_limit(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_DEFAULT_ACTIVE_REQUESTS", 5) mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_MAX_ACTIVE_REQUESTS", 3) @@ -63,6 +78,7 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) update_status_mock = mocker.patch.object(PipelineGenerateService, "update_document_status") + session = mocker.Mock() generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") generator_instance = generator_cls.return_value @@ -70,6 +86,7 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke generator_cls.convert_to_event_stream.return_value = "stream-events" result = PipelineGenerateService.generate( + session=session, pipeline=pipeline, user=user, args=args, @@ -78,42 +95,40 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke ) assert result == "stream-events" - update_status_mock.assert_called_once_with("doc-1") + update_status_mock.assert_called_once_with("doc-1", session) -def test_update_document_status_updates_existing_document(mocker: MockerFixture) -> None: - document = SimpleNamespace(indexing_status="completed") - - session_mock = mocker.Mock() - session_mock.get.return_value = document - add_mock = session_mock.add - commit_mock = session_mock.commit - mocker.patch( - "services.rag_pipeline.pipeline_generate_service.db", - new=SimpleNamespace(session=session_mock), +def test_update_document_status_updates_existing_document(document_session: Session) -> None: + session = document_session + document_id = str(uuid4()) + document = Document( + id=document_id, + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="Doc", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + indexing_status=IndexingStatus.COMPLETED, ) + session.add(document) + session.commit() - PipelineGenerateService.update_document_status("doc-1") + PipelineGenerateService.update_document_status(document_id, session) - assert document.indexing_status == "waiting" - add_mock.assert_called_once_with(document) - commit_mock.assert_called_once() + updated_document = session.get(Document, document_id) + assert updated_document is not None + assert updated_document.indexing_status == IndexingStatus.WAITING -def test_update_document_status_skips_when_document_missing(mocker: MockerFixture) -> None: - session_mock = mocker.Mock() - session_mock.get.return_value = None - add_mock = session_mock.add - commit_mock = session_mock.commit - mocker.patch( - "services.rag_pipeline.pipeline_generate_service.db", - new=SimpleNamespace(session=session_mock), - ) +def test_update_document_status_skips_when_document_missing(document_session: Session) -> None: + session = document_session - PipelineGenerateService.update_document_status("missing") + PipelineGenerateService.update_document_status(str(uuid4()), session) - add_mock.assert_not_called() - commit_mock.assert_not_called() + assert session.scalar(select(func.count()).select_from(Document)) == 0 # --- generate_single_iteration --- diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py index b2daaa42b66..37141c97c83 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py @@ -1,11 +1,12 @@ import json import time +from dataclasses import dataclass from datetime import datetime from types import SimpleNamespace +from unittest.mock import Mock import pytest from pytest_mock import MockerFixture -from sqlalchemy.orm import sessionmaker from models import Account, Tenant from models.dataset import Dataset, Pipeline, PipelineCustomizedTemplate, PipelineRecommendedPlugin @@ -15,8 +16,29 @@ from services.rag_pipeline.rag_pipeline import RagPipelineService from services.workflow_ref_service import WorkflowRef +@dataclass +class RagPipelineServiceTestContext: + service: RagPipelineService + session: Mock + session_maker: Mock + + +def _make_mock_session_maker(mocker: MockerFixture, session: Mock) -> Mock: + session_context = mocker.MagicMock() + session_context.__enter__.return_value = session + session_context.__exit__.return_value = None + + transaction_context = mocker.MagicMock() + transaction_context.__enter__.return_value = session + transaction_context.__exit__.return_value = None + + session_maker = mocker.Mock(return_value=session_context) + session_maker.begin.return_value = transaction_context + return session_maker + + @pytest.fixture -def rag_pipeline_service(mocker) -> RagPipelineService: +def rag_pipeline_service(mocker: MockerFixture) -> RagPipelineServiceTestContext: mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", return_value=MockRepo(), @@ -25,7 +47,12 @@ def rag_pipeline_service(mocker) -> RagPipelineService: "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_run_repository", return_value=MockRepo(), ) - return RagPipelineService(session_maker=sessionmaker()) + session = mocker.Mock() + session_maker = _make_mock_session_maker(mocker, session) + mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) + mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock())) + service = RagPipelineService(session_maker=session_maker) + return RagPipelineServiceTestContext(service=service, session=session, session_maker=session_maker) class MockRepo: @@ -117,6 +144,7 @@ def _make_recommended_plugin(plugin_id: str) -> PipelineRecommendedPlugin: def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_result(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") + session = mocker.Mock() remote_retrieval = mocker.Mock() remote_retrieval.get_pipeline_templates.return_value = {"pipeline_templates": []} @@ -128,58 +156,63 @@ def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_resul builtin_retrieval.fetch_pipeline_templates_from_builtin.return_value = {"pipeline_templates": [{"id": "builtin-1"}]} factory_mock.get_built_in_pipeline_template_retrieval.return_value = builtin_retrieval - result = RagPipelineService.get_pipeline_templates(type="built-in", language="ja-JP") + result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="ja-JP") assert result == {"pipeline_templates": [{"id": "builtin-1"}]} + remote_retrieval.get_pipeline_templates.assert_called_once_with(session, "ja-JP", None) builtin_retrieval.fetch_pipeline_templates_from_builtin.assert_called_once_with("en-US") def test_get_pipeline_templates_customized_mode_uses_customized_factory(mocker: MockerFixture) -> None: + session = mocker.Mock() retrieval = mocker.Mock() retrieval.get_pipeline_templates.return_value = {"pipeline_templates": [{"id": "custom-1"}]} factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_templates(type="customized", language="en-US") + result = RagPipelineService.get_pipeline_templates(session, type="customized", language="en-US") assert result == {"pipeline_templates": [{"id": "custom-1"}]} factory_mock.get_pipeline_template_factory.assert_called_with("customized") + retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) @pytest.mark.parametrize("template_type", ["built-in", "customized"]) def test_get_pipeline_template_detail_uses_expected_mode(mocker: MockerFixture, template_type: str) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") + session = mocker.Mock() retrieval = mocker.Mock() retrieval.get_pipeline_template_detail.return_value = {"id": "tpl-1"} factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_template_detail("tpl-1", type=template_type) + result = RagPipelineService.get_pipeline_template_detail(session, "tpl-1", type=template_type) assert result == {"id": "tpl-1"} expected_mode = "remote" if template_type == "built-in" else "customized" factory_mock.get_pipeline_template_factory.assert_called_with(expected_mode) + retrieval.get_pipeline_template_detail.assert_called_once_with(session, "tpl-1") def test_get_published_workflow_returns_none_when_pipeline_has_no_workflow_id( - rag_pipeline_service: RagPipelineService, + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: pipeline = _make_pipeline(workflow_id=None) - result = rag_pipeline_service.get_published_workflow(pipeline) + result = rag_pipeline_service.service.get_published_workflow(pipeline) assert result is None def test_get_all_published_workflow_returns_empty_for_unpublished_pipeline( - rag_pipeline_service: RagPipelineService, + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: pipeline = _make_pipeline(workflow_id=None) session = SimpleNamespace() - workflows, has_more = rag_pipeline_service.get_all_published_workflow( + workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( session=session, pipeline=pipeline, page=1, @@ -192,12 +225,14 @@ def test_get_all_published_workflow_returns_empty_for_unpublished_pipeline( assert has_more is False -def test_get_all_published_workflow_applies_limit_and_has_more(rag_pipeline_service: RagPipelineService) -> None: +def test_get_all_published_workflow_applies_limit_and_has_more( + rag_pipeline_service: RagPipelineServiceTestContext, +) -> None: scalars_result = SimpleNamespace(all=lambda: ["wf1", "wf2", "wf3"]) session = SimpleNamespace(scalars=lambda stmt: scalars_result) pipeline = _make_pipeline(pipeline_id="pipeline-1", workflow_id="wf-live") - workflows, has_more = rag_pipeline_service.get_all_published_workflow( + workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( session=session, pipeline=pipeline, page=1, @@ -214,25 +249,16 @@ def test_get_all_published_workflow_applies_limit_and_has_more(rag_pipeline_serv def test_sync_draft_workflow_creates_new_when_none_exists( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=None) - - class FakeWorkflow: - def __init__(self, **kwargs): - for k, v in kwargs.items(): - setattr(self, k, v) - self.id = "wf-new" - - mocker.patch("services.rag_pipeline.rag_pipeline.Workflow", FakeWorkflow) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.add") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.flush") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.commit") + session = rag_pipeline_service.session + session.get.return_value = _make_pipeline(workflow_id=None) + session.scalar.return_value = None pipeline = _make_pipeline(workflow_id=None) account = _make_account() - result = rag_pipeline_service.sync_draft_workflow( + result = rag_pipeline_service.service.sync_draft_workflow( pipeline=pipeline, graph={"nodes": []}, unique_hash=None, @@ -242,23 +268,25 @@ def test_sync_draft_workflow_creates_new_when_none_exists( rag_pipeline_variables=[], ) - assert result.id == "wf-new" - assert pipeline.workflow_id == "wf-new" + assert result.app_id == "p1" + assert pipeline.workflow_id == result.id def test_sync_draft_workflow_raises_on_hash_mismatch( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from services.errors.app import WorkflowHashNotEqualError existing_wf = _make_workflow(graph={"nodes": [{"id": "old"}]}) - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=existing_wf) + session = rag_pipeline_service.session + session.get.return_value = _make_pipeline() + session.scalar.return_value = existing_wf pipeline = _make_pipeline() account = _make_account() with pytest.raises(WorkflowHashNotEqualError): - rag_pipeline_service.sync_draft_workflow( + rag_pipeline_service.service.sync_draft_workflow( pipeline=pipeline, graph={"nodes": []}, unique_hash="hash-different", @@ -269,7 +297,9 @@ def test_sync_draft_workflow_raises_on_hash_mismatch( ) -def test_sync_draft_workflow_updates_existing(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_sync_draft_workflow_updates_existing( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: existing_wf = SimpleNamespace( unique_hash="hash-1", graph=None, @@ -279,13 +309,14 @@ def test_sync_draft_workflow_updates_existing(mocker: MockerFixture, rag_pipelin conversation_variables=None, rag_pipeline_variables=None, ) - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=existing_wf) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.commit") + session = rag_pipeline_service.session + session.get.return_value = _make_pipeline() + session.scalar.return_value = existing_wf pipeline = _make_pipeline() account = _make_account() - result = rag_pipeline_service.sync_draft_workflow( + result = rag_pipeline_service.service.sync_draft_workflow( pipeline=pipeline, graph={"nodes": [{"id": "n1"}]}, unique_hash="hash-1", @@ -304,7 +335,7 @@ def test_sync_draft_workflow_updates_existing(mocker: MockerFixture, rag_pipelin def test_get_default_block_config_returns_config_for_valid_type( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: fake_node_class = mocker.Mock() fake_node_class.get_default_config.return_value = {"type": "start", "config": {}} @@ -318,20 +349,22 @@ def test_get_default_block_config_returns_config_for_valid_type( ) mocker.patch("services.rag_pipeline.rag_pipeline.LATEST_VERSION", "1") - result = rag_pipeline_service.get_default_block_config("start") + result = rag_pipeline_service.service.get_default_block_config("start") assert result == {"type": "start", "config": {}} -def test_get_default_block_config_returns_none_for_unmapped_type(rag_pipeline_service: RagPipelineService) -> None: - assert rag_pipeline_service.get_default_block_config("nonexistent-type") is None +def test_get_default_block_config_returns_none_for_unmapped_type( + rag_pipeline_service: RagPipelineServiceTestContext, +) -> None: + assert rag_pipeline_service.service.get_default_block_config("nonexistent-type") is None # --- update_workflow --- def test_update_workflow_updates_allowed_fields( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" @@ -340,7 +373,7 @@ def test_update_workflow_updates_allowed_fields( session = mocker.Mock() session.scalar.return_value = workflow - result = rag_pipeline_service.update_workflow( + result = rag_pipeline_service.service.update_workflow( session=session, account_id="u1", data={"marked_name": "v1", "marked_comment": "release", "disallowed": "hacked"}, @@ -354,12 +387,12 @@ def test_update_workflow_updates_allowed_fields( def test_update_workflow_returns_none_when_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: session = mocker.Mock() session.scalar.return_value = None - result = rag_pipeline_service.update_workflow( + result = rag_pipeline_service.service.update_workflow( session=session, account_id="u1", data={"marked_name": "v1"}, @@ -370,7 +403,7 @@ def test_update_workflow_returns_none_when_not_found( def test_update_workflow_with_ref_scopes_lookup_to_pipeline( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" @@ -379,7 +412,7 @@ def test_update_workflow_with_ref_scopes_lookup_to_pipeline( session = mocker.Mock() session.scalar.return_value = workflow - result = rag_pipeline_service.update_workflow( + result = rag_pipeline_service.service.update_workflow( session=session, account_id="u1", data={"marked_name": "v1"}, @@ -402,15 +435,17 @@ def test_update_workflow_with_ref_scopes_lookup_to_pipeline( def test_get_rag_pipeline_paginate_workflow_runs_delegates( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: expected = mocker.Mock() repo_mock = mocker.Mock() repo_mock.get_paginated_workflow_runs.return_value = expected - rag_pipeline_service._workflow_run_repo = repo_mock + rag_pipeline_service.service._workflow_run_repo = repo_mock pipeline = _make_pipeline() - result = rag_pipeline_service.get_rag_pipeline_paginate_workflow_runs(pipeline, {"limit": 10, "last_id": "abc"}) + result = rag_pipeline_service.service.get_rag_pipeline_paginate_workflow_runs( + pipeline, {"limit": 10, "last_id": "abc"} + ) assert result is expected repo_mock.get_paginated_workflow_runs.assert_called_once_with( @@ -426,15 +461,15 @@ def test_get_rag_pipeline_paginate_workflow_runs_delegates( def test_get_rag_pipeline_workflow_run_delegates( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: expected = mocker.Mock() repo_mock = mocker.Mock() repo_mock.get_workflow_run_by_id.return_value = expected - rag_pipeline_service._workflow_run_repo = repo_mock + rag_pipeline_service.service._workflow_run_repo = repo_mock pipeline = _make_pipeline() - result = rag_pipeline_service.get_rag_pipeline_workflow_run(pipeline, "run-1") + result = rag_pipeline_service.service.get_rag_pipeline_workflow_run(pipeline, "run-1") assert result is expected repo_mock.get_workflow_run_by_id.assert_called_once_with(tenant_id="t1", app_id="p1", run_id="run-1") @@ -444,27 +479,27 @@ def test_get_rag_pipeline_workflow_run_delegates( def test_is_workflow_exist_returns_true_when_draft_exists( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=1) + rag_pipeline_service.session.scalar.return_value = 1 pipeline = _make_pipeline() - assert rag_pipeline_service.is_workflow_exist(pipeline) is True + assert rag_pipeline_service.service.is_workflow_exist(pipeline) is True def test_is_workflow_exist_returns_false_when_no_draft( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=0) + rag_pipeline_service.session.scalar.return_value = 0 pipeline = _make_pipeline() - assert rag_pipeline_service.is_workflow_exist(pipeline) is False + assert rag_pipeline_service.service.is_workflow_exist(pipeline) is False # --- publish_workflow --- -def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext) -> None: # Don't import Workflow from rag_pipeline to avoid confusion during patching # 1. Mock select to bypass SQLAlchemy validation @@ -510,9 +545,7 @@ def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: R new_wf.graph_dict = draft_wf.graph mock_workflow_class.new.return_value = new_wf - # 5. Mock entire db object and DatasetService - mock_db = mocker.Mock() - mocker.patch("services.rag_pipeline.rag_pipeline.db", mock_db) + # 5. Mock DatasetService mock_dataset_service_class = mocker.patch("services.dataset_service.DatasetService") # 6. Mock session and dataset lookup @@ -524,7 +557,7 @@ def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: R pipeline.retrieve_dataset.return_value = dataset # 7. Run test - result = rag_pipeline_service.publish_workflow(session=mock_session, pipeline=pipeline, account=account) + result = rag_pipeline_service.service.publish_workflow(session=mock_session, pipeline=pipeline, account=account) # 8. Assertions assert result == new_wf @@ -536,7 +569,7 @@ def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: R def test_run_datasource_workflow_node_website_crawl( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceProviderType @@ -560,7 +593,7 @@ def test_run_datasource_workflow_node_website_crawl( } ] } - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) # 2. Mock DatasourceManager and Runtime mock_runtime = mocker.Mock() @@ -590,7 +623,7 @@ def test_run_datasource_workflow_node_website_crawl( mocker.patch("services.rag_pipeline.rag_pipeline.DatasourceProviderType", DatasourceProviderType) # 5. Run test - gen = rag_pipeline_service.run_datasource_workflow_node( + gen = rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="node-1", user_inputs={"url": "https://example.com"}, @@ -614,7 +647,7 @@ def test_run_datasource_workflow_node_website_crawl( def test_run_datasource_node_preview_online_document( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceMessage, DatasourceProviderType @@ -642,7 +675,7 @@ def test_run_datasource_node_preview_online_document( } ] } - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) # 2. Mock Runtime and results mock_runtime = mocker.Mock() @@ -672,7 +705,7 @@ def test_run_datasource_node_preview_online_document( mocker.patch("services.rag_pipeline.rag_pipeline.DatasourceProviderType", DatasourceProviderType) # 3. Run test - result = rag_pipeline_service.run_datasource_node_preview( + result = rag_pipeline_service.service.run_datasource_node_preview( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -688,7 +721,9 @@ def test_run_datasource_node_preview_online_document( # --- _handle_node_run_result --- -def test_handle_node_run_result_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_handle_node_run_result_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunSucceededEvent from graphon.node_events.base import NodeRunResult @@ -718,7 +753,7 @@ def test_handle_node_run_result_success(mocker: MockerFixture, rag_pipeline_serv yield event # 2. Run test - result = rag_pipeline_service._handle_node_run_result( + result = rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, mock_getter()), start_at=time.perf_counter(), tenant_id="t1", node_id="node-1" ) @@ -732,7 +767,9 @@ def test_handle_node_run_result_success(mocker: MockerFixture, rag_pipeline_serv # --- get_first_step_parameters / get_second_step_parameters --- -def test_get_first_step_parameters_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_get_first_step_parameters_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: # 1. Setup mock workflow pipeline = mocker.Mock() workflow = mocker.Mock() @@ -740,17 +777,19 @@ def test_get_first_step_parameters_success(mocker: MockerFixture, rag_pipeline_s "nodes": [{"id": "node-1", "data": {"datasource_parameters": {"url": {"value": "{{#start.url#}}"}}}}] } workflow.rag_pipeline_variables = [{"variable": "url", "label": "URL", "type": "string"}] - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) # 2. Run test - result = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id="node-1", is_draft=False) + result = rag_pipeline_service.service.get_first_step_parameters(pipeline=pipeline, node_id="node-1", is_draft=False) # 3. Assertions assert len(result) == 1 assert result[0]["variable"] == "url" -def test_get_second_step_parameters_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_get_second_step_parameters_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: # 1. Setup mock workflow pipeline = mocker.Mock() workflow = mocker.Mock() @@ -763,10 +802,12 @@ def test_get_second_step_parameters_success(mocker: MockerFixture, rag_pipeline_ ] } workflow.rag_pipeline_variables = [{"variable": "var1", "label": "Var 1"}] - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) # 2. Run test - result = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id="node-1", is_draft=False) + result = rag_pipeline_service.service.get_second_step_parameters( + pipeline=pipeline, node_id="node-1", is_draft=False + ) # 3. Assertions # Note: get_second_step_parameters also filters by variable names found in node data @@ -779,39 +820,32 @@ def test_get_second_step_parameters_success(mocker: MockerFixture, rag_pipeline_ def test_publish_customized_pipeline_template_success( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: # 1. Setup mocks pipeline = _make_pipeline(workflow_id="wf-1", is_published=True) workflow = _make_workflow(workflow_id="wf-1") - # Mock db itself to avoid app context errors - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - - # Mock get() for Pipeline and Workflow PK lookups - mock_db.session.get.side_effect = [pipeline, workflow] - # Mock scalar() for template name check (None) and max position (5) - mock_db.session.scalar.side_effect = [None, 5] + session = rag_pipeline_service.session + session.get.side_effect = [pipeline, workflow] + session.scalar.side_effect = [None, 5] # Mock retrieve_dataset dataset = _make_dataset() dataset.chunk_structure = "paragraph" + pipeline.retrieve_dataset = mocker.Mock(return_value=dataset) # Mock RagPipelineDslService mock_dsl_service = mocker.Mock() mock_dsl_service.export_rag_pipeline_dsl.return_value = {"dsl": "content"} mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService", return_value=mock_dsl_service) - # Mock Session and commit - session_factory = mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker") - session_factory.return_value.begin.return_value.__enter__.return_value.scalar.return_value = dataset - account = _make_account(account_id="user-123") # 2. Run test args = {"name": "New Template", "description": "Desc", "icon_info": {"icon": "star"}, "tags": ["tag1"]} - rag_pipeline_service.publish_customized_pipeline_template("p1", args, account, "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1") # 3. Assertions # Verify a new template was added to session or similar? @@ -823,7 +857,9 @@ def test_publish_customized_pipeline_template_success( # --- get_datasource_plugins --- -def test_get_datasource_plugins_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_get_datasource_plugins_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: # 1. Setup mocks dataset = _make_dataset() @@ -846,10 +882,7 @@ def test_get_datasource_plugins_success(mocker: MockerFixture, rag_pipeline_serv } workflow.rag_pipeline_variables = [] - # Mock queries - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) - - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] # Mock DatasourceProviderService mock_provider_service = mocker.Mock() @@ -859,7 +892,7 @@ def test_get_datasource_plugins_success(mocker: MockerFixture, rag_pipeline_serv mocker.patch("services.rag_pipeline.rag_pipeline.DatasourceProviderService", return_value=mock_provider_service) # 2. Run test - result = rag_pipeline_service.get_datasource_plugins("t1", "d1", True) + result = rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) # 3. Assertions assert len(result) == 1 @@ -870,7 +903,9 @@ def test_get_datasource_plugins_success(mocker: MockerFixture, rag_pipeline_serv # --- retry_error_document --- -def test_retry_error_document_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_retry_error_document_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: from models.dataset import Document, DocumentPipelineExecutionLog, Pipeline # 1. Setup mocks @@ -887,11 +922,8 @@ def test_retry_error_document_success(mocker: MockerFixture, rag_pipeline_servic workflow = mocker.Mock() - # Mock queries: Log lookup via scalar, Pipeline lookup via get - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=log) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=pipeline) - - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + rag_pipeline_service.session.scalar.side_effect = [log, workflow] + rag_pipeline_service.session.get.return_value = pipeline # Mock PipelineGenerator mock_gen_instance = mocker.Mock() @@ -899,7 +931,7 @@ def test_retry_error_document_success(mocker: MockerFixture, rag_pipeline_servic # 2. Run test user = mocker.Mock() - rag_pipeline_service.retry_error_document(dataset, document, user) + rag_pipeline_service.service.retry_error_document(dataset, document, user) # 3. Assertions mock_gen_instance.generate.assert_called_once() @@ -908,16 +940,13 @@ def test_retry_error_document_success(mocker: MockerFixture, rag_pipeline_servic # --- set_datasource_variables --- -def test_set_datasource_variables_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_set_datasource_variables_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: from graphon.entities.workflow_node_execution import WorkflowNodeExecution from models.dataset import Pipeline # 1. Setup mocks - # Mock db aggressively - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.engine = mocker.Mock() - mock_db.session.scalar.return_value = mocker.Mock() - pipeline = mocker.Mock(spec=Pipeline) pipeline.id = "p-1" pipeline.tenant_id = "t1" @@ -925,14 +954,14 @@ def test_set_datasource_variables_success(mocker: MockerFixture, rag_pipeline_se draft_wf = mocker.Mock() draft_wf.id = "wf-1" draft_wf.get_enclosing_node_type_and_id.return_value = None # Avoid unpacking error - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=draft_wf) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=draft_wf) execution = mocker.Mock(spec=WorkflowNodeExecution) execution.id = "exec-1" execution.process_data = {} execution.inputs = {} execution.outputs = {} - mocker.patch.object(rag_pipeline_service, "_handle_node_run_result", return_value=execution) + mocker.patch.object(rag_pipeline_service.service, "_handle_node_run_result", return_value=execution) # Mock Repository mock_repo_instance = mocker.Mock() @@ -957,7 +986,7 @@ def test_set_datasource_variables_success(mocker: MockerFixture, rag_pipeline_se args = {"start_node_id": "node-1"} user = mocker.Mock() user.id = "user-1" - rag_pipeline_service.set_datasource_variables(pipeline, args, user) + rag_pipeline_service.service.set_datasource_variables(pipeline, args, user) # 3. Assertions mock_repo_instance.save.assert_called_once() @@ -967,56 +996,56 @@ def test_set_datasource_variables_success(mocker: MockerFixture, rag_pipeline_se # --- Utility Methods --- -def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_get_draft_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext) -> None: # 1. Setup mocks pipeline = _make_pipeline() workflow = _make_workflow() - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.session.scalar.return_value = workflow + rag_pipeline_service.session.scalar.return_value = workflow # 2. Run test - result = rag_pipeline_service.get_draft_workflow(pipeline) + result = rag_pipeline_service.service.get_draft_workflow(pipeline) # 3. Assertions assert result == workflow -def test_get_published_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineService) -> None: +def test_get_published_workflow_success( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: # 1. Setup mocks pipeline = _make_pipeline(workflow_id="wf-pub") workflow = _make_workflow(workflow_id="wf-pub") - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.session.scalar.return_value = workflow + rag_pipeline_service.session.scalar.return_value = workflow # 2. Run test - result = rag_pipeline_service.get_published_workflow(pipeline) + result = rag_pipeline_service.service.get_published_workflow(pipeline) # 3. Assertions assert result == workflow -def test_get_default_block_configs_success(rag_pipeline_service: RagPipelineService) -> None: +def test_get_default_block_configs_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: # This calls static methods on node classes, should be safe with default mocks or as-is # unless they access db. - result = rag_pipeline_service.get_default_block_configs() + result = rag_pipeline_service.service.get_default_block_configs() assert isinstance(result, list) assert len(result) > 0 -def test_get_default_block_config_success(rag_pipeline_service: RagPipelineService) -> None: +def test_get_default_block_config_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: from graphon.enums import BuiltinNodeTypes - result = rag_pipeline_service.get_default_block_config(BuiltinNodeTypes.LLM) + result = rag_pipeline_service.service.get_default_block_config(BuiltinNodeTypes.LLM) assert result is not None assert result["type"] == "llm" def test_publish_workflow_raises_when_draft_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: session = mocker.Mock() session.scalar.return_value = None @@ -1024,21 +1053,21 @@ def test_publish_workflow_raises_when_draft_workflow_missing( account = _make_account() with pytest.raises(ValueError, match="No valid workflow found"): - rag_pipeline_service.publish_workflow(session=session, pipeline=pipeline, account=account) + rag_pipeline_service.service.publish_workflow(session=session, pipeline=pipeline, account=account) def test_get_default_block_config_returns_none_when_mapped_type_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from graphon.enums import BuiltinNodeTypes mocker.patch("services.rag_pipeline.rag_pipeline.get_node_type_classes_mapping", return_value={}) - assert rag_pipeline_service.get_default_block_config(BuiltinNodeTypes.START) is None + assert rag_pipeline_service.service.get_default_block_config(BuiltinNodeTypes.START) is None def test_get_default_block_config_injects_http_request_filter( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from graphon.enums import BuiltinNodeTypes @@ -1050,43 +1079,44 @@ def test_get_default_block_config_injects_http_request_filter( ) mocker.patch("services.rag_pipeline.rag_pipeline.LATEST_VERSION", "1") - rag_pipeline_service.get_default_block_config(BuiltinNodeTypes.HTTP_REQUEST) + rag_pipeline_service.service.get_default_block_config(BuiltinNodeTypes.HTTP_REQUEST) called_filters = fake_node_cls.get_default_config.call_args.kwargs["filters"] assert "http_request_config" in called_filters def test_run_draft_workflow_node_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline() account = _make_account() - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=None) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=None) with pytest.raises(ValueError, match="Workflow not initialized"): - rag_pipeline_service.run_draft_workflow_node(pipeline, "node-1", {}, account) + rag_pipeline_service.service.run_draft_workflow_node(pipeline, "node-1", {}, account) def test_run_draft_workflow_node_saves_execution_and_variables( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db", mocker.Mock(engine=mocker.Mock())) pipeline = _make_pipeline() account = _make_account() draft_workflow = mocker.Mock(id="wf-1") draft_workflow.get_node_config_by_id.return_value = {"id": "node-1"} draft_workflow.get_enclosing_node_type_and_id.return_value = ("loop", "enclosing-node") - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=draft_workflow) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=draft_workflow) execution = SimpleNamespace(id="exec-1", node_id="node-1", node_type="llm", process_data={}, outputs={}) - mocker.patch.object(rag_pipeline_service, "_handle_node_run_result", return_value=execution) + mocker.patch.object(rag_pipeline_service.service, "_handle_node_run_result", return_value=execution) repo = mocker.Mock() mocker.patch( "services.rag_pipeline.rag_pipeline.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", return_value=repo, ) - rag_pipeline_service._node_execution_service_repo = mocker.Mock(get_execution_by_id=mocker.Mock(return_value="db")) + rag_pipeline_service.service._node_execution_service_repo = mocker.Mock( + get_execution_by_id=mocker.Mock(return_value="db") + ) saver = mocker.Mock() mocker.patch("services.rag_pipeline.rag_pipeline.DraftVariableSaver", return_value=saver) @@ -1095,7 +1125,7 @@ def test_run_draft_workflow_node_saves_execution_and_variables( session_ctx.begin.return_value = begin_ctx mocker.patch("services.rag_pipeline.rag_pipeline.Session", return_value=session_ctx) - result = rag_pipeline_service.run_draft_workflow_node(pipeline, "node-1", {"q": "x"}, account) + result = rag_pipeline_service.service.run_draft_workflow_node(pipeline, "node-1", {"q": "x"}, account) assert result == "db" assert execution.workflow_id == "wf-1" @@ -1104,13 +1134,13 @@ def test_run_draft_workflow_node_saves_execution_and_variables( def test_run_datasource_workflow_node_returns_error_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = SimpleNamespace(id="p1", tenant_id="t1") - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=None) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=None) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -1124,7 +1154,7 @@ def test_run_datasource_workflow_node_returns_error_when_workflow_missing( def test_run_datasource_workflow_node_online_document_success( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceProviderType @@ -1144,7 +1174,7 @@ def test_run_datasource_workflow_node_online_document_success( } ] } - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() runtime.runtime = SimpleNamespace(credentials=None) @@ -1157,7 +1187,7 @@ def test_run_datasource_workflow_node_online_document_success( ) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -1172,7 +1202,7 @@ def test_run_datasource_workflow_node_online_document_success( def test_run_datasource_workflow_node_online_drive_success( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceProviderType @@ -1192,7 +1222,7 @@ def test_run_datasource_workflow_node_online_drive_success( } ] } - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() runtime.runtime = SimpleNamespace(credentials=None) @@ -1205,7 +1235,7 @@ def test_run_datasource_workflow_node_online_drive_success( ) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="node-1", user_inputs={"bucket": "bucket-1"}, @@ -1220,7 +1250,7 @@ def test_run_datasource_workflow_node_online_drive_success( def test_handle_node_run_result_default_value_strategy( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from datetime import datetime @@ -1254,7 +1284,7 @@ def test_handle_node_run_result_default_value_strategy( node_run_result=failed_result, ) - result = rag_pipeline_service._handle_node_run_result( + result = rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, _events()), start_at=time.perf_counter(), tenant_id="t1", @@ -1267,17 +1297,17 @@ def test_handle_node_run_result_default_value_strategy( def test_get_first_step_parameters_raises_when_datasource_node_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace(graph_dict={"nodes": []}, rag_pipeline_variables=[{"variable": "url"}]) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) with pytest.raises(ValueError, match="Datasource node data not found"): - rag_pipeline_service.get_first_step_parameters(SimpleNamespace(), "missing-node") + rag_pipeline_service.service.get_first_step_parameters(SimpleNamespace(), "missing-node") def test_get_second_step_parameters_handles_string_and_list_variable_references( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( rag_pipeline_variables=[ @@ -1299,20 +1329,20 @@ def test_get_second_step_parameters_handles_string_and_list_variable_references( ] }, ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) - result = rag_pipeline_service.get_second_step_parameters(SimpleNamespace(), "node-1") + result = rag_pipeline_service.service.get_second_step_parameters(SimpleNamespace(), "node-1") assert result == [{"variable": "keep", "belong_to_node_id": "node-1"}] def test_get_rag_pipeline_workflow_run_node_executions_empty_when_run_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline() - mocker.patch.object(rag_pipeline_service, "get_rag_pipeline_workflow_run", return_value=None) + mocker.patch.object(rag_pipeline_service.service, "get_rag_pipeline_workflow_run", return_value=None) - result = rag_pipeline_service.get_rag_pipeline_workflow_run_node_executions( + result = rag_pipeline_service.service.get_rag_pipeline_workflow_run_node_executions( pipeline=pipeline, run_id="run-1", user=_make_account() ) @@ -1320,16 +1350,17 @@ def test_get_rag_pipeline_workflow_run_node_executions_empty_when_run_missing( def test_get_rag_pipeline_workflow_run_node_executions_returns_sorted_executions( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db", mocker.Mock(engine=mocker.Mock())) pipeline = _make_pipeline() - mocker.patch.object(rag_pipeline_service, "get_rag_pipeline_workflow_run", return_value=SimpleNamespace(id="run-1")) + mocker.patch.object( + rag_pipeline_service.service, "get_rag_pipeline_workflow_run", return_value=SimpleNamespace(id="run-1") + ) repo = mocker.Mock() repo.get_db_models_by_workflow_run.return_value = ["n1", "n2"] mocker.patch("services.rag_pipeline.rag_pipeline.SQLAlchemyWorkflowNodeExecutionRepository", return_value=repo) - result = rag_pipeline_service.get_rag_pipeline_workflow_run_node_executions( + result = rag_pipeline_service.service.get_rag_pipeline_workflow_run_node_executions( pipeline=pipeline, run_id="run-1", user=_make_account() ) @@ -1337,12 +1368,11 @@ def test_get_rag_pipeline_workflow_run_node_executions_returns_sorted_executions def test_get_recommended_plugins_returns_empty_when_no_active_plugins( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.session.scalars.return_value.all.return_value = [] + rag_pipeline_service.session.scalars.return_value.all.return_value = [] - result = rag_pipeline_service.get_recommended_plugins("all", _make_account(), "t1") + result = rag_pipeline_service.service.get_recommended_plugins("all", _make_account(), "t1") assert result == { "installed_recommended_plugins": [], @@ -1351,12 +1381,11 @@ def test_get_recommended_plugins_returns_empty_when_no_active_plugins( def test_get_recommended_plugins_returns_installed_and_uninstalled( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: plugin_a = _make_recommended_plugin("plugin-a") plugin_b = _make_recommended_plugin("plugin-b") - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.session.scalars.return_value.all.return_value = [plugin_a, plugin_b] + rag_pipeline_service.session.scalars.return_value.all.return_value = [plugin_a, plugin_b] mocker.patch( "services.rag_pipeline.rag_pipeline.BuiltinToolManageService.list_builtin_tools", return_value=[SimpleNamespace(plugin_id="plugin-a", to_dict=lambda: {"plugin_id": "plugin-a"})], @@ -1366,16 +1395,15 @@ def test_get_recommended_plugins_returns_installed_and_uninstalled( return_value=[{"plugin_id": "plugin-b", "name": "Plugin B"}], ) - result = rag_pipeline_service.get_recommended_plugins("custom", _make_account(), "t1") + result = rag_pipeline_service.service.get_recommended_plugins("custom", _make_account(), "t1") assert result["installed_recommended_plugins"] == [{"plugin_id": "plugin-a"}] assert result["uninstalled_recommended_plugins"] == [{"plugin_id": "plugin-b", "name": "Plugin B"}] def test_get_node_last_run_delegates_to_repository( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db", mocker.Mock(engine=mocker.Mock())) repo = mocker.Mock() repo.get_node_last_execution.return_value = "node-exec" mocker.patch( @@ -1385,24 +1413,24 @@ def test_get_node_last_run_delegates_to_repository( pipeline = _make_pipeline() workflow = _make_workflow(workflow_id="wf1") - result = rag_pipeline_service.get_node_last_run(pipeline, workflow, "node-1") + result = rag_pipeline_service.service.get_node_last_run(pipeline, workflow, "node-1") assert result == "node-exec" def test_set_datasource_variables_raises_when_node_id_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = SimpleNamespace(id="p1", tenant_id="t1") workflow = mocker.Mock() - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=workflow) with pytest.raises(ValueError, match="Node id is required"): - rag_pipeline_service.set_datasource_variables(pipeline, {"start_node_id": ""}, SimpleNamespace(id="u1")) + rag_pipeline_service.service.set_datasource_variables(pipeline, {"start_node_id": ""}, SimpleNamespace(id="u1")) def test_get_default_block_configs_skips_empty_configs( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from graphon.enums import BuiltinNodeTypes @@ -1420,7 +1448,7 @@ def test_get_default_block_configs_skips_empty_configs( ) mocker.patch("services.rag_pipeline.rag_pipeline.LATEST_VERSION", "1") - result = rag_pipeline_service.get_default_block_configs() + result = rag_pipeline_service.service.get_default_block_configs() assert result == [{"type": "http-request"}] http_node.get_default_config.assert_called_once() @@ -1428,14 +1456,14 @@ def test_get_default_block_configs_skips_empty_configs( def test_run_datasource_workflow_node_returns_error_when_node_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = SimpleNamespace(id="p1", tenant_id="t1") workflow = SimpleNamespace(graph_dict={"nodes": []}) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="missing-node", user_inputs={}, @@ -1450,7 +1478,7 @@ def test_run_datasource_workflow_node_returns_error_when_node_missing( def test_run_datasource_workflow_node_online_document_exception( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = SimpleNamespace(id="p1", tenant_id="t1") workflow = SimpleNamespace( @@ -1468,7 +1496,7 @@ def test_run_datasource_workflow_node_online_document_exception( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() @@ -1488,7 +1516,7 @@ def test_run_datasource_workflow_node_online_document_exception( ) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -1504,7 +1532,7 @@ def test_run_datasource_workflow_node_online_document_exception( def test_run_datasource_node_preview_raises_for_stream_non_string( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceMessage @@ -1524,7 +1552,7 @@ def test_run_datasource_node_preview_raises_for_stream_non_string( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() @@ -1543,7 +1571,7 @@ def test_run_datasource_node_preview_raises_for_stream_non_string( ) with pytest.raises(RuntimeError, match="must be a string"): - rag_pipeline_service.run_datasource_node_preview( + rag_pipeline_service.service.run_datasource_node_preview( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -1554,21 +1582,21 @@ def test_run_datasource_node_preview_raises_for_stream_non_string( def test_get_first_step_parameters_returns_empty_when_no_rag_variables( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( graph_dict={"nodes": [{"id": "node-1", "data": {"datasource_parameters": {"url": {"value": "literal"}}}}]}, rag_pipeline_variables=[], ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) - result = rag_pipeline_service.get_first_step_parameters(SimpleNamespace(), "node-1") + result = rag_pipeline_service.service.get_first_step_parameters(SimpleNamespace(), "node-1") assert result == [] def test_get_second_step_parameters_filters_first_step_variables( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( graph_dict={ @@ -1591,38 +1619,37 @@ def test_get_second_step_parameters_filters_first_step_variables( {"variable": "other-node", "belong_to_node_id": "node-x"}, ], ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) - result = rag_pipeline_service.get_second_step_parameters(SimpleNamespace(), "node-1") + result = rag_pipeline_service.service.get_second_step_parameters(SimpleNamespace(), "node-1") assert result == [{"variable": "keep", "belong_to_node_id": "shared"}] def test_retry_error_document_raises_when_execution_log_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=None) + rag_pipeline_service.session.scalar.return_value = None with pytest.raises(ValueError, match="Document pipeline execution log not found"): - rag_pipeline_service.retry_error_document( + rag_pipeline_service.service.retry_error_document( SimpleNamespace(), SimpleNamespace(id="doc-1"), SimpleNamespace(id="u1") ) def test_get_datasource_plugins_raises_when_workflow_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = SimpleNamespace(pipeline_id="p1") - pipeline = SimpleNamespace(id="p1", tenant_id="t1") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=None) + pipeline = SimpleNamespace(id="p1", tenant_id="t1", workflow_id="wf-1") + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, None] with pytest.raises(ValueError, match="Pipeline or workflow not found"): - rag_pipeline_service.get_datasource_plugins("t1", "d1", True) + rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) def test_handle_node_run_result_raises_when_no_terminal_event( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: node_instance = SimpleNamespace( workflow_id="wf-1", @@ -1636,7 +1663,7 @@ def test_handle_node_run_result_raises_when_no_terminal_event( yield object() with pytest.raises(ValueError, match="Node run failed with no run result"): - rag_pipeline_service._handle_node_run_result( + rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, _event_generator()), start_at=time.perf_counter(), tenant_id="t1", @@ -1645,7 +1672,7 @@ def test_handle_node_run_result_raises_when_no_terminal_event( def test_handle_node_run_result_marks_document_error_for_published_invoke( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.app.entities.app_invoke_entities import InvokeFrom from graphon.enums import WorkflowNodeExecutionStatus @@ -1691,12 +1718,9 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( ) document = SimpleNamespace(indexing_status="waiting", error=None) - scalar_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=document) - get_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get") - add_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.add") - commit_mock = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.commit") + rag_pipeline_service.session.scalar.return_value = document - result = rag_pipeline_service._handle_node_run_result( + result = rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, _event_generator()), start_at=time.perf_counter(), tenant_id="t1", @@ -1704,7 +1728,7 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( ) assert result.status == WorkflowNodeExecutionStatus.FAILED - stmt = scalar_mock.call_args.args[0] + stmt = rag_pipeline_service.session.scalar.call_args.args[0] compiled = stmt.compile() statement = str(compiled) assert "documents.id" in statement @@ -1716,15 +1740,12 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( assert "t1" in compiled.params.values() assert "dataset-1" in compiled.params.values() assert "pipeline-1" in compiled.params.values() - get_mock.assert_not_called() assert document.indexing_status == "error" assert document.error == "boom" - add_mock.assert_called_once_with(document) - commit_mock.assert_called_once() def test_run_datasource_node_preview_raises_for_unsupported_provider( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = SimpleNamespace(id="p1", tenant_id="t1") workflow = SimpleNamespace( @@ -1742,7 +1763,7 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() runtime.datasource_provider_type.return_value = "unsupported" mocker.patch("core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", return_value=runtime) @@ -1751,7 +1772,7 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( ) with pytest.raises(RuntimeError, match="Unsupported datasource provider"): - rag_pipeline_service.run_datasource_node_preview( + rag_pipeline_service.service.run_datasource_node_preview( pipeline=pipeline, node_id="node-1", user_inputs={}, @@ -1762,49 +1783,51 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( def test_publish_customized_pipeline_template_raises_for_missing_pipeline( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=None) + rag_pipeline_service.session.get.return_value = None with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") def test_publish_customized_pipeline_template_raises_for_missing_workflow_id( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id=None) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=pipeline) + rag_pipeline_service.session.get.return_value = pipeline with pytest.raises(ValueError, match="Pipeline workflow not found"): - rag_pipeline_service.publish_customized_pipeline_template( + rag_pipeline_service.service.publish_customized_pipeline_template( "p1", {"name": "template-name"}, _make_account(), "t1" ) def test_get_pipeline_raises_when_dataset_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=None) + rag_pipeline_service.session.scalar.return_value = None with pytest.raises(ValueError, match="Dataset not found"): - rag_pipeline_service.get_pipeline("t1", "d1") + rag_pipeline_service.service.get_pipeline("t1", "d1") def test_get_pipeline_raises_when_pipeline_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = SimpleNamespace(pipeline_id="p1") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, None]) + rag_pipeline_service.session.scalar.side_effect = [dataset, None] with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.get_pipeline("t1", "d1") + rag_pipeline_service.service.get_pipeline("t1", "d1") def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None: default_session_maker = mocker.Mock() - mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker", return_value=default_session_maker) - mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock())) + mocker.patch( + "services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", + return_value=default_session_maker, + ) create_exec_repo = mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository" ) @@ -1820,35 +1843,41 @@ def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") + session = mocker.Mock() retrieval = mocker.Mock() retrieval.get_pipeline_templates.return_value = {"pipeline_templates": []} factory = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory.get_pipeline_template_factory.return_value.return_value = retrieval builtin = factory.get_built_in_pipeline_template_retrieval.return_value - result = RagPipelineService.get_pipeline_templates(type="built-in", language="en-US") + result = RagPipelineService.get_pipeline_templates(session, type="built-in", language="en-US") assert result == {"pipeline_templates": []} + retrieval.get_pipeline_templates.assert_called_once_with(session, "en-US", None) builtin.fetch_pipeline_templates_from_builtin.assert_not_called() def test_update_customized_pipeline_template_commits_when_name_empty(mocker: MockerFixture) -> None: template = _make_customized_template() - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=template) - commit = mocker.patch("services.rag_pipeline.rag_pipeline.db.session.commit") + session = mocker.Mock() + session.scalar.return_value = template + session_maker = _make_mock_session_maker(mocker, session) + mocker.patch("services.rag_pipeline.rag_pipeline.session_factory.get_session_maker", return_value=session_maker) info = PipelineTemplateInfoEntity(name="", description="updated", icon_info=IconInfo(icon="i")) result = RagPipelineService.update_customized_pipeline_template("tpl-1", info, _make_account(), "t1") assert result.description == "updated" - commit.assert_called_once() + session_maker.begin.assert_called_once() -def test_get_all_published_workflow_without_filters_has_no_more(rag_pipeline_service: RagPipelineService) -> None: +def test_get_all_published_workflow_without_filters_has_no_more( + rag_pipeline_service: RagPipelineServiceTestContext, +) -> None: session = SimpleNamespace(scalars=lambda stmt: SimpleNamespace(all=lambda: ["wf1"])) pipeline = _make_pipeline(workflow_id="wf-live") - workflows, has_more = rag_pipeline_service.get_all_published_workflow( + workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( session=session, pipeline=pipeline, page=1, @@ -1862,7 +1891,7 @@ def test_get_all_published_workflow_without_filters_has_no_more(rag_pipeline_ser def test_publish_workflow_skips_dataset_update_for_non_knowledge_nodes( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: draft = SimpleNamespace( type="workflow", @@ -1878,7 +1907,7 @@ def test_publish_workflow_skips_dataset_update_for_non_knowledge_nodes( mocker.patch("services.rag_pipeline.rag_pipeline.select") mocker.patch("services.rag_pipeline.rag_pipeline.Workflow.new", return_value=published) - result = rag_pipeline_service.publish_workflow( + result = rag_pipeline_service.service.publish_workflow( session=session, pipeline=SimpleNamespace(id="p1", tenant_id="t1", is_published=False, retrieve_dataset=lambda session: None), account=SimpleNamespace(id="u1"), @@ -1888,7 +1917,7 @@ def test_publish_workflow_skips_dataset_update_for_non_knowledge_nodes( def test_get_default_block_config_returns_none_when_default_empty( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from graphon.enums import BuiltinNodeTypes @@ -1900,11 +1929,11 @@ def test_get_default_block_config_returns_none_when_default_empty( ) mocker.patch("services.rag_pipeline.rag_pipeline.LATEST_VERSION", "1") - assert rag_pipeline_service.get_default_block_config("start") is None + assert rag_pipeline_service.service.get_default_block_config("start") is None def test_run_datasource_workflow_node_handles_variable_parameter_types( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceProviderType @@ -1927,7 +1956,7 @@ def test_run_datasource_workflow_node_handles_variable_parameter_types( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() def crawl_gen(**kwargs): @@ -1941,7 +1970,7 @@ def test_run_datasource_workflow_node_handles_variable_parameter_types( ) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="node-1", user_inputs={"k": "mapped"}, @@ -1956,7 +1985,7 @@ def test_run_datasource_workflow_node_handles_variable_parameter_types( def test_run_datasource_workflow_node_online_drive_branch( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceProviderType @@ -1975,7 +2004,7 @@ def test_run_datasource_workflow_node_online_drive_branch( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() def drive_gen(**kwargs): @@ -1989,7 +2018,7 @@ def test_run_datasource_workflow_node_online_drive_branch( ) events = list( - rag_pipeline_service.run_datasource_workflow_node( + rag_pipeline_service.service.run_datasource_workflow_node( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="node-1", user_inputs={}, @@ -2004,7 +2033,7 @@ def test_run_datasource_workflow_node_online_drive_branch( def test_run_datasource_node_preview_not_published_uses_draft( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceMessage @@ -2023,7 +2052,7 @@ def test_run_datasource_node_preview_not_published_uses_draft( ] } ) - get_draft = mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=workflow) + get_draft = mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=workflow) runtime = mocker.Mock() def doc_gen(**kwargs): @@ -2038,7 +2067,7 @@ def test_run_datasource_node_preview_not_published_uses_draft( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.get_datasource_credentials", return_value=None ) - result = rag_pipeline_service.run_datasource_node_preview( + result = rag_pipeline_service.service.run_datasource_node_preview( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="n1", user_inputs={}, @@ -2052,12 +2081,12 @@ def test_run_datasource_node_preview_not_published_uses_draft( def test_run_free_workflow_node_delegates_to_handle_result( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: expected = SimpleNamespace(id="exec-1") - handle = mocker.patch.object(rag_pipeline_service, "_handle_node_run_result", return_value=expected) + handle = mocker.patch.object(rag_pipeline_service.service, "_handle_node_run_result", return_value=expected) - result = rag_pipeline_service.run_free_workflow_node( + result = rag_pipeline_service.service.run_free_workflow_node( node_data={"type": "start"}, tenant_id="t1", user_id="u1", @@ -2070,89 +2099,83 @@ def test_run_free_workflow_node_delegates_to_handle_result( def test_publish_customized_pipeline_template_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", side_effect=[pipeline, None]) + rag_pipeline_service.session.get.side_effect = [pipeline, None] with pytest.raises(ValueError, match="Workflow not found"): - rag_pipeline_service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") def test_publish_customized_pipeline_template_raises_when_dataset_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") workflow = _make_workflow(workflow_id="wf-1") - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.engine = mocker.Mock() - mock_db.session.get.side_effect = [pipeline, workflow] - session_factory = mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker") - session_factory.return_value.begin.return_value.__enter__.return_value.scalar.return_value = None + pipeline.retrieve_dataset = mocker.Mock(return_value=None) + rag_pipeline_service.session.get.side_effect = [pipeline, workflow] with pytest.raises(ValueError, match="Dataset not found"): - rag_pipeline_service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") + rag_pipeline_service.service.publish_customized_pipeline_template("p1", {}, _make_account(), "t1") def test_get_recommended_plugins_skips_manifest_when_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: plugin = _make_recommended_plugin("plugin-a") - mock_db = mocker.patch("services.rag_pipeline.rag_pipeline.db") - mock_db.session.scalars.return_value.all.return_value = [plugin] + rag_pipeline_service.session.scalars.return_value.all.return_value = [plugin] mocker.patch("services.rag_pipeline.rag_pipeline.BuiltinToolManageService.list_builtin_tools", return_value=[]) mocker.patch("services.rag_pipeline.rag_pipeline.marketplace.batch_fetch_plugin_by_ids", return_value=[]) - result = rag_pipeline_service.get_recommended_plugins("all", _make_account(), "t1") + result = rag_pipeline_service.service.get_recommended_plugins("all", _make_account(), "t1") assert result["installed_recommended_plugins"] == [] assert result["uninstalled_recommended_plugins"] == [] def test_retry_error_document_raises_when_pipeline_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: exec_log = SimpleNamespace(pipeline_id="p1") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=exec_log) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=None) + rag_pipeline_service.session.scalar.return_value = exec_log + rag_pipeline_service.session.get.return_value = None with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.retry_error_document( + rag_pipeline_service.service.retry_error_document( SimpleNamespace(), SimpleNamespace(id="doc-1"), SimpleNamespace(id="u1") ) def test_retry_error_document_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: exec_log = SimpleNamespace(pipeline_id="p1") - pipeline = SimpleNamespace(id="p1") - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", return_value=exec_log) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.get", return_value=pipeline) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=None) + pipeline = SimpleNamespace(id="p1", tenant_id="t1", workflow_id="wf-1") + rag_pipeline_service.session.scalar.side_effect = [exec_log, None] + rag_pipeline_service.session.get.return_value = pipeline with pytest.raises(ValueError, match="Workflow not found"): - rag_pipeline_service.retry_error_document( + rag_pipeline_service.service.retry_error_document( SimpleNamespace(), SimpleNamespace(id="doc-1"), SimpleNamespace(id="u1") ) def test_get_datasource_plugins_returns_empty_for_non_datasource_nodes( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() workflow = SimpleNamespace( graph_dict={"nodes": [{"id": "n1", "data": {"type": "start"}}]}, rag_pipeline_variables=[] ) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] - assert rag_pipeline_service.get_datasource_plugins("t1", "d1", True) == [] + assert rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) == [] def test_publish_workflow_raises_when_knowledge_index_dataset_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: draft = SimpleNamespace( type="workflow", @@ -2173,16 +2196,18 @@ def test_publish_workflow_raises_when_knowledge_index_dataset_missing( pipeline = SimpleNamespace(id="p1", tenant_id="t1", is_published=False, retrieve_dataset=lambda session: None) with pytest.raises(ValueError, match="Dataset not found"): - rag_pipeline_service.publish_workflow(session=session, pipeline=pipeline, account=SimpleNamespace(id="u1")) + rag_pipeline_service.service.publish_workflow( + session=session, pipeline=pipeline, account=SimpleNamespace(id="u1") + ) def test_run_datasource_node_preview_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=None) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=None) with pytest.raises(RuntimeError, match="Workflow not initialized"): - rag_pipeline_service.run_datasource_node_preview( + rag_pipeline_service.service.run_datasource_node_preview( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="n1", user_inputs={}, @@ -2193,14 +2218,14 @@ def test_run_datasource_node_preview_raises_when_workflow_missing( def test_run_datasource_node_preview_raises_when_node_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: mocker.patch.object( - rag_pipeline_service, "get_published_workflow", return_value=SimpleNamespace(graph_dict={"nodes": []}) + rag_pipeline_service.service, "get_published_workflow", return_value=SimpleNamespace(graph_dict={"nodes": []}) ) with pytest.raises(RuntimeError, match="Datasource node data not found"): - rag_pipeline_service.run_datasource_node_preview( + rag_pipeline_service.service.run_datasource_node_preview( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="missing", user_inputs={}, @@ -2211,7 +2236,7 @@ def test_run_datasource_node_preview_raises_when_node_missing( def test_run_datasource_node_preview_keeps_existing_user_input( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: from core.datasource.entities.datasource_entities import DatasourceMessage @@ -2230,7 +2255,7 @@ def test_run_datasource_node_preview_keeps_existing_user_input( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() def gen(**kwargs): @@ -2247,7 +2272,7 @@ def test_run_datasource_node_preview_keeps_existing_user_input( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.get_datasource_credentials", return_value=None ) - result = rag_pipeline_service.run_datasource_node_preview( + result = rag_pipeline_service.service.run_datasource_node_preview( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="n1", user_inputs={"workspace_id": "existing"}, @@ -2259,7 +2284,7 @@ def test_run_datasource_node_preview_keeps_existing_user_input( def test_run_datasource_node_preview_ignores_non_variable_messages( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: workflow = SimpleNamespace( graph_dict={ @@ -2276,7 +2301,7 @@ def test_run_datasource_node_preview_ignores_non_variable_messages( ] } ) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + mocker.patch.object(rag_pipeline_service.service, "get_published_workflow", return_value=workflow) runtime = mocker.Mock() def gen(**kwargs): @@ -2288,7 +2313,7 @@ def test_run_datasource_node_preview_ignores_non_variable_messages( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.get_datasource_credentials", return_value=None ) - result = rag_pipeline_service.run_datasource_node_preview( + result = rag_pipeline_service.service.run_datasource_node_preview( pipeline=SimpleNamespace(id="p1", tenant_id="t1"), node_id="n1", user_inputs={}, @@ -2300,12 +2325,12 @@ def test_run_datasource_node_preview_ignores_non_variable_messages( def test_set_datasource_variables_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=None) + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=None) with pytest.raises(ValueError, match="Workflow not initialized"): - rag_pipeline_service.set_datasource_variables( + rag_pipeline_service.service.set_datasource_variables( SimpleNamespace(id="p1", tenant_id="t1"), {"start_node_id": "n1"}, SimpleNamespace(id="u1"), @@ -2313,7 +2338,7 @@ def test_set_datasource_variables_raises_when_workflow_missing( def test_get_datasource_plugins_handles_empty_datasource_data_and_non_published( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() @@ -2321,19 +2346,18 @@ def test_get_datasource_plugins_handles_empty_datasource_data_and_non_published( graph_dict={"nodes": [{"id": "n1", "data": {"type": "datasource", "datasource_parameters": {}}}]}, rag_pipeline_variables=[{"variable": "v1", "belong_to_node_id": "shared"}], ) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) - mocker.patch.object(rag_pipeline_service, "get_draft_workflow", return_value=workflow) + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] mocker.patch( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.list_datasource_credentials", return_value=[] ) - result = rag_pipeline_service.get_datasource_plugins("t1", "d1", False) + result = rag_pipeline_service.service.get_datasource_plugins("t1", "d1", False) assert len(result) == 1 def test_get_datasource_plugins_extracts_user_inputs_and_credentials( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() @@ -2362,14 +2386,13 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( {"variable": "v3", "belong_to_node_id": "shared"}, ], ) - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) - mocker.patch.object(rag_pipeline_service, "get_published_workflow", return_value=workflow) + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] mocker.patch( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.list_datasource_credentials", return_value=[{"id": "c1", "name": "Cred", "type": "api", "is_default": True}], ) - result = rag_pipeline_service.get_datasource_plugins("t1", "d1", True) + result = rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) assert len(result) == 1 assert len(result[0]["user_input_variables"]) == 2 @@ -2377,12 +2400,12 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( def test_get_pipeline_returns_pipeline_when_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineService + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() - mocker.patch("services.rag_pipeline.rag_pipeline.db.session.scalar", side_effect=[dataset, pipeline]) + rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline] - result = rag_pipeline_service.get_pipeline("t1", "d1") + result = rag_pipeline_service.service.get_pipeline("t1", "d1") assert result is pipeline diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py index a05930c73ce..281de58abad 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py @@ -152,7 +152,7 @@ def test_upload_invoke_entities_returns_file_id(mocker: MockerFixture, proxy) -> upload_file = SimpleNamespace(id="uploaded-file-1") file_service_cls = mocker.patch("services.rag_pipeline.rag_pipeline_task_proxy.FileService") file_service_cls.return_value.upload_text.return_value = upload_file - mocker.patch("services.rag_pipeline.rag_pipeline_task_proxy.db", mocker.Mock(engine="fake-engine")) + mocker.patch("services.rag_pipeline.rag_pipeline_task_proxy.db", SimpleNamespace(engine="fake-engine")) result = proxy._upload_invoke_entities() diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py index f6a3f524fef..21df81d3ea8 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py @@ -1,16 +1,31 @@ import logging +from collections.abc import Iterator from datetime import UTC, datetime from types import SimpleNamespace from typing import cast import pytest from pytest_mock import MockerFixture +from sqlalchemy import create_engine, select +from sqlalchemy.orm import Session, sessionmaker -from models.dataset import Dataset +from models.dataset import Dataset, Pipeline +from models.enums import DatasetRuntimeMode from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService +@pytest.fixture +def pipeline_session() -> Iterator[Session]: + engine = create_engine("sqlite:///:memory:") + Dataset.__table__.create(engine) + Pipeline.__table__.create(engine) + session_factory = sessionmaker(bind=engine, expire_on_commit=False) + with session_factory() as session: + yield session + engine.dispose() + + @pytest.mark.parametrize( ("doc_form", "datasource_type", "indexing_technique"), [ @@ -92,51 +107,37 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi install_mock.assert_called_once_with("tenant-1", ["missing-plugin:1.0.0"]) -def test_transform_to_empty_pipeline_updates_dataset_and_commits(mocker: MockerFixture) -> None: +def test_transform_to_empty_pipeline_updates_dataset_and_flushes( + mocker: MockerFixture, pipeline_session: Session +) -> None: service = RagPipelineTransformService() mocker.patch( "services.rag_pipeline.rag_pipeline_transform_service.current_user", SimpleNamespace(id="user-1"), ) - class FakePipeline: - def __init__(self, **kwargs): - self.id = "pipeline-1" - self.tenant_id = kwargs["tenant_id"] - self.name = kwargs["name"] - self.description = kwargs["description"] - self.created_by = kwargs["created_by"] - - mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.Pipeline", FakePipeline) - session_mock = mocker.Mock() - add_mock = session_mock.add - flush_mock = session_mock.flush - commit_mock = session_mock.commit - mocker.patch( - "services.rag_pipeline.rag_pipeline_transform_service.db", - new=SimpleNamespace(session=session_mock), - ) - - dataset = SimpleNamespace( - id="dataset-1", + session = pipeline_session + dataset = Dataset( tenant_id="tenant-1", name="Dataset", description="desc", - pipeline_id=None, - runtime_mode="general", - updated_by=None, - updated_at=None, + created_by="user-1", ) + session.add(dataset) + session.commit() + flush_spy = mocker.spy(session, "flush") + commit_spy = mocker.spy(session, "commit") - result = service._transform_to_empty_pipeline(cast(Dataset, dataset)) + result = service._transform_to_empty_pipeline(dataset, session) - assert result == {"pipeline_id": "pipeline-1", "dataset_id": "dataset-1", "status": "success"} - assert dataset.pipeline_id == "pipeline-1" - assert dataset.runtime_mode == "rag_pipeline" + assert flush_spy.call_count == 2 + commit_spy.assert_not_called() + pipeline = session.scalar(select(Pipeline).where(Pipeline.id == dataset.pipeline_id)) + assert pipeline is not None + assert result == {"pipeline_id": pipeline.id, "dataset_id": dataset.id, "status": "success"} + assert dataset.pipeline_id == pipeline.id + assert dataset.runtime_mode == DatasetRuntimeMode.RAG_PIPELINE assert dataset.updated_by == "user-1" - add_mock.assert_called() - flush_mock.assert_called_once() - commit_mock.assert_called_once() # --- transform_dataset --- @@ -372,7 +373,6 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_deal_document_data") - session_mock.commit = mocker.Mock() # Mock current_user to have the same tenant_id as dataset mock_current_user = SimpleNamespace(current_tenant_id="t1") @@ -386,6 +386,8 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: assert result["pipeline_id"] == "p-new" assert dataset.runtime_mode == "rag_pipeline" assert dataset.chunk_structure == "text_model" + session_mock.flush.assert_called_once_with() + session_mock.commit.assert_not_called() def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(mocker: MockerFixture) -> None: @@ -405,10 +407,6 @@ def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create ) session_mock = mocker.Mock() session_mock.get.return_value = dataset - mocker.patch( - "services.rag_pipeline.rag_pipeline_transform_service.db", - new=SimpleNamespace(session=session_mock), - ) mocker.patch.object(service, "_get_transform_yaml", return_value={"workflow": {"graph": {"nodes": []}}}) mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_create_pipeline", return_value=SimpleNamespace(id="p-new")) @@ -441,11 +439,11 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: M service.transform_dataset("d1", session_mock) -def test_create_pipeline_raises_when_workflow_data_missing() -> None: +def test_create_pipeline_raises_when_workflow_data_missing(pipeline_session: Session) -> None: service = RagPipelineTransformService() with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service._create_pipeline({"rag_pipeline": {"name": "N"}}) + service._create_pipeline({"rag_pipeline": {"name": "N"}}, pipeline_session) def test_deal_document_data_upload_file_with_existing_file(mocker: MockerFixture) -> None: