test: move RAG pipeline workflow coverage to unit tests (#38937)

This commit is contained in:
Asuka Minato
2026-08-08 12:00:32 +00:00
committed by GitHub
parent 4e18ab0ad4
commit 5a84cde198
2 changed files with 238 additions and 343 deletions
@@ -1,40 +1,51 @@
"""RAG pipeline workflow controller serialization tests.
Handlers that own transactions run against real SQLite sessions so response
DTOs must be materialized before those transaction contexts close.
"""
"""Unit coverage for RAG workflow controllers using real models and disposable SQLite state."""
from __future__ import annotations
import json
from collections.abc import Iterator
from datetime import datetime
from inspect import unwrap as unwrap_all
from types import SimpleNamespace
from unittest.mock import PropertyMock, patch
from unittest.mock import MagicMock, patch
from uuid import UUID
import pytest
from flask import Flask
from sqlalchemy.engine import Engine
from sqlalchemy import Engine
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden
from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as module
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
from models.account import Account, TenantAccountRole
from models.dataset import Pipeline
from models.engine import db
from models.tools import WorkflowToolProvider
from models.workflow import Workflow, WorkflowType
from services.errors.llm import InvokeRateLimitError
from services.rag_pipeline.rag_pipeline import RagPipelineService
DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001"
DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002"
DEFAULT_WORKFLOW_CREATED_BY = "00000000-0000-0000-0000-000000000003"
DEFAULT_WORKFLOW_ID = "00000000-0000-0000-0000-000000000004"
def _make_workflow(**overrides):
workflow = SimpleNamespace(
id="workflow-1",
graph_dict={"nodes": [], "edges": []},
features_dict={"file_upload": {"enabled": False}},
unique_hash="hash-1",
version="1",
def _make_workflow(**overrides: object) -> Workflow:
workflow = Workflow(
id=DEFAULT_WORKFLOW_ID,
tenant_id=DEFAULT_WORKFLOW_TENANT_ID,
app_id=DEFAULT_WORKFLOW_APP_ID,
type=WorkflowType.WORKFLOW,
version=Workflow.VERSION_DRAFT,
marked_name="Release 1",
marked_comment="Initial release",
created_by_account=SimpleNamespace(id="user-1", name="Alice", email="alice@example.com"),
graph=json.dumps({"nodes": [], "edges": []}),
features=json.dumps({"file_upload": {"enabled": False}}),
created_by=DEFAULT_WORKFLOW_CREATED_BY,
created_at=datetime(2024, 1, 1, 12, 0, 0),
updated_by_account=None,
updated_by=None,
updated_at=datetime(2024, 1, 1, 12, 1, 0),
tool_published=False,
environment_variables=[],
conversation_variables=[],
rag_pipeline_variables=[],
@@ -46,137 +57,130 @@ def _make_workflow(**overrides):
def _account() -> Account:
account = Account(name="Alice", email="alice@example.com")
account.id = "user-1"
account.id = DEFAULT_WORKFLOW_CREATED_BY
account.role = TenantAccountRole.EDITOR
return account
def _pipeline() -> Pipeline:
pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline", description="desc")
pipeline.id = "pipeline-1"
pipeline = Pipeline(tenant_id=DEFAULT_WORKFLOW_TENANT_ID, name="Pipeline", description="desc")
pipeline.id = DEFAULT_WORKFLOW_APP_ID
return pipeline
def test_draft_rag_pipeline_workflow_get_serializes_response_model(monkeypatch: pytest.MonkeyPatch) -> None:
def _persist_workflow(workflow: Workflow) -> None:
db.session.add(workflow)
db.session.commit()
db.session.expunge(workflow)
@pytest.fixture
def database_app() -> Iterator[Flask]:
app = Flask(__name__)
app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:"
db.init_app(app)
with app.app_context():
Account.__table__.create(db.engine)
WorkflowToolProvider.__table__.create(db.engine)
Workflow.__table__.create(db.engine)
db.session.add(_account())
db.session.commit()
try:
yield app
finally:
db.session.remove()
@pytest.mark.usefixtures("database_app")
def test_draft_rag_pipeline_workflow_get_serializes_response_model() -> None:
workflow = _make_workflow()
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow),
)
expected_hash = workflow.unique_hash
_persist_workflow(workflow)
api = module.DraftRagPipelineApi()
handler = unwrap_all(api.get)
response = handler(api, _pipeline())
assert response["id"] == "workflow-1"
assert response["id"] == DEFAULT_WORKFLOW_ID
assert response["graph"] == {"nodes": [], "edges": []}
assert response["features"] == {"file_upload": {"enabled": False}}
assert response["hash"] == "hash-1"
assert response["created_by"] == {"id": "user-1", "name": "Alice", "email": "alice@example.com"}
assert response["hash"] == expected_hash
assert response["created_by"] == {
"id": DEFAULT_WORKFLOW_CREATED_BY,
"name": "Alice",
"email": "alice@example.com",
}
assert response["updated_by"] is None
assert response["created_at"] == int(datetime(2024, 1, 1, 12, 0, 0).timestamp())
assert response["updated_at"] == int(datetime(2024, 1, 1, 12, 1, 0).timestamp())
def test_published_rag_pipeline_workflows_serialize_items_before_session_closes(
app, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
database_app: Flask,
) -> None:
api = module.PublishedAllRagPipelineApi()
handler = unwrap_all(api.get)
session_state: dict[str, Session] = {}
workflow = _make_workflow(version="1")
_persist_workflow(workflow)
pipeline = _pipeline()
pipeline.workflow_id = DEFAULT_WORKFLOW_ID
base_workflow = _make_workflow()
with database_app.test_request_context(
"/rag/pipelines/pipeline-1/workflows",
method="GET",
query_string={"page": 1, "limit": 10, "user_id": "", "named_only": "false"},
):
response = handler(api, _account(), pipeline=pipeline)
class _Workflow:
def __getattr__(self, name: str):
assert session_state["session"].in_transaction() is True
return getattr(base_workflow, name)
def _get_all_published_workflow(**kwargs):
session_state["session"] = kwargs["session"]
return [_Workflow()], False
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_all_published_workflow=_get_all_published_workflow),
)
with Session(sqlite_engine) as request_session:
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=lambda: request_session))
with app.test_request_context(
"/rag/pipelines/pipeline-1/workflows",
method="GET",
query_string={"page": 1, "limit": 10, "user_id": "", "named_only": "false"},
):
response = handler(api, _account(), pipeline=_pipeline())
assert session_state["session"].in_transaction() is False
assert response["items"][0]["id"] == "workflow-1"
assert response["items"][0]["id"] == DEFAULT_WORKFLOW_ID
assert response["page"] == 1
assert response["limit"] == 10
assert response["has_more"] is False
def test_rag_pipeline_workflow_patch_serializes_response_model(
app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine
database_app: Flask,
) -> None:
workflow = _make_workflow(marked_name="Updated release")
captured_session: dict[str, Session] = {}
def _update_workflow(**kwargs):
captured_session["session"] = kwargs["session"]
assert kwargs["session"].in_transaction() is True
return workflow
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(update_workflow=_update_workflow),
)
expected_hash = workflow.unique_hash
_persist_workflow(workflow)
payload: dict[str, object] = {"marked_name": "Updated release"}
api = module.RagPipelineByIdApi()
handler = unwrap_all(api.patch)
with Session(sqlite_engine) as request_session:
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=lambda: request_session))
with (
app.test_request_context("/rag/pipelines/pipeline-1/workflows/workflow-1", method="PATCH", json=payload),
patch.object(type(module.console_ns), "payload", new_callable=PropertyMock, return_value=payload),
):
response = handler(
api,
_account(),
pipeline=_pipeline(),
workflow_id="workflow-1",
)
with database_app.test_request_context(
f"/rag/pipelines/{DEFAULT_WORKFLOW_APP_ID}/workflows/{DEFAULT_WORKFLOW_ID}", method="PATCH", json=payload
):
response = handler(
api,
_account(),
pipeline=_pipeline(),
workflow_id=DEFAULT_WORKFLOW_ID,
)
assert captured_session["session"].in_transaction() is False
assert response["id"] == "workflow-1"
assert response["id"] == DEFAULT_WORKFLOW_ID
assert response["marked_name"] == "Updated release"
assert response["hash"] == "hash-1"
assert response["hash"] == expected_hash
def test_default_rag_pipeline_block_configs_serializes_root_response(monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.usefixtures("database_app")
def test_default_rag_pipeline_block_configs_serializes_root_response() -> None:
block_configs = [{"type": "start", "config": {"title": "Start"}}]
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_default_block_configs=lambda: block_configs),
)
api = module.DefaultRagPipelineBlockConfigsApi()
handler = unwrap_all(api.get)
response = handler(api, _pipeline())
with patch.object(RagPipelineService, "get_default_block_configs", return_value=block_configs):
response = handler(api, _pipeline())
assert response == block_configs
def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, monkeypatch: pytest.MonkeyPatch) -> None:
def test_draft_rag_pipeline_second_step_parameters_serializes_variables(database_app: Flask) -> None:
variables = [
{
"belong_to_node_id": "shared",
@@ -187,36 +191,114 @@ def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, mon
"required": True,
}
]
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables),
)
api = module.DraftRagPipelineSecondStepApi()
handler = unwrap_all(api.get)
with app.test_request_context("/?node_id=node-1"):
with (
database_app.test_request_context("/?node_id=node-1"),
patch.object(RagPipelineService, "get_second_step_parameters", return_value=variables),
):
response = handler(api, _pipeline())
assert response["variables"] == variables
def test_rag_pipeline_recommended_plugins_serializes_known_envelope(app, monkeypatch: pytest.MonkeyPatch) -> None:
def test_rag_pipeline_recommended_plugins_serializes_known_envelope(database_app: Flask) -> None:
recommended_plugins = {
"installed_recommended_plugins": [{"name": "Dify Extractor", "meta": {"version": "1.0.0"}}],
"uninstalled_recommended_plugins": [{"plugin_id": "langgenius/notion_datasource"}],
}
monkeypatch.setattr(
module,
"RagPipelineService",
lambda *_args, **_kwargs: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins),
)
api = module.RagPipelineRecommendedPluginApi()
handler = unwrap_all(api.get)
with app.test_request_context("/?type=tool"):
response = handler(api, "tenant-1", _account())
with (
database_app.test_request_context("/?type=tool"),
patch.object(RagPipelineService, "get_recommended_plugins", return_value=recommended_plugins),
):
response = handler(api, DEFAULT_WORKFLOW_TENANT_ID, _account())
assert response == recommended_plugins
def test_rag_pipeline_transform_rejects_read_only_member(app: Flask, sqlite_engine: Engine) -> None:
account = _account()
account.role = TenantAccountRole.NORMAL
api = module.RagPipelineTransformApi()
handler = unwrap_all(api.post)
with (
Session(sqlite_engine) as session,
app.test_request_context("/"),
pytest.raises(Forbidden),
):
handler(api, session, account, UUID("44444444-4444-4444-4444-444444444444"))
@pytest.mark.parametrize(
("api_type", "payload"),
[
(
module.DraftRagPipelineRunApi,
{"inputs": {}, "datasource_type": "x", "datasource_info_list": [], "start_node_id": "node-1"},
),
(
module.PublishedRagPipelineRunApi,
{
"inputs": {},
"datasource_type": "x",
"datasource_info_list": [],
"start_node_id": "node-1",
"response_mode": "blocking",
},
),
],
)
def test_rag_pipeline_run_uses_sqlite_session(
app: Flask,
sqlite_engine: Engine,
api_type: type,
payload: dict[str, object],
) -> None:
api = api_type()
handler = unwrap_all(api.post)
pipeline = _pipeline()
with (
Session(sqlite_engine) as session,
app.test_request_context("/", json=payload),
patch.object(module, "load_rag_pipeline", return_value=pipeline) as load_pipeline,
patch.object(module.PipelineGenerateService, "generate", return_value=MagicMock()) as generate,
patch.object(module.helper, "compact_generate_response", return_value={"ok": True}),
):
response = handler(api, session, _account(), pipeline.id)
assert response == {"ok": True}
load_pipeline.assert_called_once_with(session, pipeline.id)
assert generate.call_args.kwargs["session"] is session
assert session.get_bind() is sqlite_engine
@pytest.mark.parametrize("api_type", [module.DraftRagPipelineRunApi, module.PublishedRagPipelineRunApi])
def test_rag_pipeline_run_translates_rate_limit(
app: Flask,
sqlite_engine: Engine,
api_type: type,
) -> None:
payload = {
"inputs": {},
"datasource_type": "x",
"datasource_info_list": [],
"start_node_id": "node-1",
}
api = api_type()
handler = unwrap_all(api.post)
pipeline = _pipeline()
with (
Session(sqlite_engine) as session,
app.test_request_context("/", json=payload),
patch.object(module, "load_rag_pipeline", return_value=pipeline),
patch.object(module.PipelineGenerateService, "generate", side_effect=InvokeRateLimitError("limit")),
pytest.raises(InvokeRateLimitHttpError),
):
handler(api, session, _account(), pipeline.id)
@@ -0,0 +1,802 @@
"""Unit tests for rag_pipeline_workflow controller endpoints."""
from __future__ import annotations
import json
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import datetime
from inspect import unwrap
from typing import TypedDict, Unpack
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from flask import Flask
from sqlalchemy import Engine
from sqlalchemy.orm import Session, scoped_session, sessionmaker
from werkzeug.exceptions import BadRequest, Forbidden, HTTPException, NotFound
import models.workflow as workflow_models
import services
from controllers.console.app.error import DraftWorkflowNotExist, DraftWorkflowNotSync
from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as workflow_controller
from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import (
DefaultRagPipelineBlockConfigApi,
DraftRagPipelineApi,
PublishedAllRagPipelineApi,
PublishedRagPipelineApi,
RagPipelineByIdApi,
RagPipelineDatasourceVariableApi,
RagPipelineDraftNodeRunApi,
RagPipelineDraftRunIterationNodeApi,
RagPipelineDraftRunLoopNodeApi,
RagPipelineDraftWorkflowRestoreApi,
RagPipelineRecommendedPluginApi,
RagPipelineTaskStopApi,
RagPipelineWorkflowLastRunApi,
RagPipelineWorkflowRunNodeExecutionListApi,
)
from graphon.enums import WorkflowNodeExecutionStatus
from libs.datetime_utils import naive_utc_now
from models.account import Account, TenantAccountRole
from models.dataset import Pipeline
from models.enums import CreatorUserRole
from models.workflow import Workflow, WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom
from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError
DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001"
DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002"
DEFAULT_WORKFLOW_CREATED_BY = "00000000-0000-0000-0000-000000000003"
type WorkflowVariablePayload = dict[str, object]
@dataclass(frozen=True)
class SQLiteDatabase:
"""Expose the concrete SQLite engine and scoped session interface used by controller code."""
engine: Engine
session: scoped_session[Session]
@pytest.fixture(autouse=True)
def sqlite_database(
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
) -> Iterator[scoped_session[Session]]:
"""Route controller transactions and model author lookups through SQLite."""
database_session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False))
database = SQLiteDatabase(engine=sqlite_engine, session=database_session)
monkeypatch.setattr(workflow_controller, "db", database)
monkeypatch.setattr(workflow_models, "db", database)
with database_session() as session:
default_author = Account(name="Default Author", email="default-author@example.com")
default_author.id = DEFAULT_WORKFLOW_CREATED_BY
session.add(default_author)
session.commit()
try:
yield database_session
finally:
database_session.remove()
def empty_mapping() -> dict[str, object]:
return {}
def empty_list() -> list[object]:
return []
class WorkflowFactoryPayload(TypedDict):
id: str
tenant_id: str
app_id: str
type: str
version: str
marked_name: str
marked_comment: str
graph: str
features: str
created_by: str
created_at: datetime
updated_by: str | None
updated_at: datetime | None
environment_variables: list[WorkflowVariablePayload]
conversation_variables: list[WorkflowVariablePayload]
rag_pipeline_variables: list[WorkflowVariablePayload]
class WorkflowFactoryOverrides(TypedDict, total=False):
id: str
tenant_id: str
app_id: str
type: str
version: str
marked_name: str
marked_comment: str
graph: str
features: str
created_by: str
created_at: datetime
updated_by: str | None
updated_at: datetime | None
environment_variables: list[WorkflowVariablePayload]
conversation_variables: list[WorkflowVariablePayload]
rag_pipeline_variables: list[WorkflowVariablePayload]
class NodeExecutionOverrides(TypedDict, total=False):
id: str
tenant_id: str
app_id: str
workflow_id: str
workflow_run_id: str | None
index: int
predecessor_node_id: str | None
node_execution_id: str | None
node_id: str
node_type: str
title: str
inputs: str | None
process_data: str | None
outputs: str | None
status: WorkflowNodeExecutionStatus
error: str | None
elapsed_time: float
execution_metadata: str | None
created_at: datetime
created_by_role: CreatorUserRole
created_by: str
finished_at: datetime | None
def make_node_execution(**overrides: Unpack[NodeExecutionOverrides]) -> WorkflowNodeExecutionModel:
payload: NodeExecutionOverrides = {
"id": "node-exec-1",
"tenant_id": DEFAULT_WORKFLOW_TENANT_ID,
"app_id": DEFAULT_WORKFLOW_APP_ID,
"workflow_id": "workflow-1",
"workflow_run_id": None,
"index": 1,
"predecessor_node_id": None,
"node_execution_id": None,
"node_id": "node1",
"node_type": "start",
"title": "Start",
"inputs": json.dumps({"query": "hello"}),
"process_data": json.dumps({}),
"outputs": json.dumps({"answer": "world"}),
"status": WorkflowNodeExecutionStatus.SUCCEEDED,
"error": None,
"elapsed_time": 1.0,
"execution_metadata": json.dumps({}),
"created_at": datetime(2026, 1, 1, 0, 0, 0),
"created_by_role": CreatorUserRole.ACCOUNT,
"created_by": DEFAULT_WORKFLOW_CREATED_BY,
"finished_at": datetime(2026, 1, 1, 0, 0, 1),
}
payload.update(overrides)
execution = WorkflowNodeExecutionModel(
triggered_from=WorkflowNodeExecutionTriggeredFrom.RAG_PIPELINE_RUN,
**payload,
)
execution.offload_data = []
return execution
def default_workflow_payload() -> WorkflowFactoryPayload:
return {
"id": "workflow-1",
"tenant_id": DEFAULT_WORKFLOW_TENANT_ID,
"app_id": DEFAULT_WORKFLOW_APP_ID,
"type": "workflow",
"version": "1",
"marked_name": "Release 1",
"marked_comment": "Initial release",
"graph": json.dumps({"nodes": [], "edges": []}),
"features": json.dumps({"file_upload": {"enabled": False}}),
"created_by": DEFAULT_WORKFLOW_CREATED_BY,
"created_at": datetime(2024, 1, 1, 12, 0, 0),
"updated_by": None,
"updated_at": datetime(2024, 1, 1, 12, 1, 0),
"environment_variables": [],
"conversation_variables": [],
"rag_pipeline_variables": [],
}
def make_workflow(**overrides: Unpack[WorkflowFactoryOverrides]) -> Workflow:
payload = default_workflow_payload()
payload.update(overrides)
return Workflow(**payload)
def make_account(*, id: str = "account-1", role: TenantAccountRole = TenantAccountRole.EDITOR) -> Account:
account = Account(name="Alice", email=f"{id}@example.com")
account.id = id
account.role = role
return account
def make_pipeline(
*,
id: str = "pipeline-1",
tenant_id: str = "tenant-1",
workflow_id: str | None = None,
is_published: bool = False,
) -> Pipeline:
pipeline = Pipeline(tenant_id=tenant_id, name="test-pipeline", description="test")
pipeline.id = id
pipeline.workflow_id = workflow_id
pipeline.is_published = is_published
return pipeline
@pytest.fixture
def workflow_author(sqlite_database: scoped_session[Session]) -> Account:
account = Account(name="Alice", email=f"alice-{uuid4()}@example.com")
account.id = str(uuid4())
sqlite_database.add(account)
sqlite_database.commit()
return account
class TestDraftWorkflowApi:
def test_get_draft_success(self, app: Flask, workflow_author: Account) -> None:
api = DraftRagPipelineApi()
method = unwrap(api.get)
pipeline = make_pipeline()
workflow = make_workflow(created_by=workflow_author.id)
service = MagicMock()
service.get_draft_workflow.return_value = workflow
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, pipeline)
assert result["id"] == "workflow-1"
assert result["graph"] == {"nodes": [], "edges": []}
assert result["features"] == {"file_upload": {"enabled": False}}
assert result["hash"] == workflow.unique_hash
assert result["created_by"] == {
"id": workflow_author.id,
"name": workflow_author.name,
"email": workflow_author.email,
}
assert result["updated_by"] is None
def test_get_draft_not_exist(self, app: Flask) -> None:
api = DraftRagPipelineApi()
method = unwrap(api.get)
pipeline = make_pipeline()
service = MagicMock()
service.get_draft_workflow.return_value = None
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(DraftWorkflowNotExist):
method(api, pipeline)
def test_sync_hash_not_match(self, app: Flask) -> None:
api = DraftRagPipelineApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
service = MagicMock()
service.sync_draft_workflow.side_effect = WorkflowHashNotEqualError()
with (
app.test_request_context("/", json={"graph": empty_mapping(), "features": empty_mapping()}),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(DraftWorkflowNotSync):
method(api, user, pipeline)
def test_sync_invalid_text_plain(self, app: Flask) -> None:
api = DraftRagPipelineApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
with (
app.test_request_context("/", data="bad-json", headers={"Content-Type": "text/plain"}),
):
response, status = method(api, user, pipeline)
assert status == 400
def test_restore_published_workflow_to_draft_success(self, app: Flask) -> None:
api = RagPipelineDraftWorkflowRestoreApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account(id="account-1")
workflow = make_workflow(
graph=json.dumps({"nodes": [{"id": "restored"}], "edges": []}),
created_at=datetime(2024, 1, 1),
)
service = MagicMock()
service.restore_published_workflow_to_draft.return_value = workflow
with (
app.test_request_context("/", method="POST"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline, "published-workflow")
assert result["result"] == "success"
assert result["hash"] == workflow.unique_hash
def test_restore_published_workflow_to_draft_not_found(self, app: Flask) -> None:
api = RagPipelineDraftWorkflowRestoreApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account(id="account-1")
service = MagicMock()
service.restore_published_workflow_to_draft.side_effect = WorkflowNotFoundError("Workflow not found")
with (
app.test_request_context("/", method="POST"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(NotFound):
method(api, user, pipeline, "published-workflow")
def test_restore_published_workflow_to_draft_returns_400_for_draft_source(self, app: Flask) -> None:
api = RagPipelineDraftWorkflowRestoreApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account(id="account-1")
service = MagicMock()
service.restore_published_workflow_to_draft.side_effect = IsDraftWorkflowError(
"source workflow must be published"
)
with (
app.test_request_context("/", method="POST"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(HTTPException) as exc:
method(api, user, pipeline, "draft-workflow")
assert exc.value.code == 400
assert exc.value.description == "source workflow must be published"
class TestDraftRunNodes:
def test_iteration_node_success(self, app: Flask) -> None:
api = RagPipelineDraftRunIterationNodeApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
with (
app.test_request_context("/", json={"inputs": empty_mapping()}),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_iteration",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response",
return_value={"ok": True},
),
):
result = method(api, user, pipeline, "node")
assert result == {"ok": True}
def test_iteration_node_conversation_not_exists(self, app: Flask) -> None:
api = RagPipelineDraftRunIterationNodeApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
with (
app.test_request_context("/", json={"inputs": empty_mapping()}),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_iteration",
side_effect=services.errors.conversation.ConversationNotExistsError(),
),
):
with pytest.raises(NotFound):
method(api, user, pipeline, "node")
def test_loop_node_success(self, app: Flask) -> None:
api = RagPipelineDraftRunLoopNodeApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
with (
app.test_request_context("/", json={"inputs": empty_mapping()}),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_loop",
return_value=MagicMock(),
),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response",
return_value={"ok": True},
),
):
assert method(api, user, pipeline, "node") == {"ok": True}
class TestDraftNodeRun:
def test_execution_not_found(self, app: Flask) -> None:
api = RagPipelineDraftNodeRunApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
service = MagicMock()
service.run_draft_workflow_node.return_value = None
with (
app.test_request_context("/", json={"inputs": empty_mapping()}),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(ValueError):
method(api, user, pipeline, "node")
class TestPublishedPipelineApis:
def test_publish_success(self, app: Flask) -> None:
api = PublishedRagPipelineApi()
method = unwrap(api.post)
tenant_id = str(uuid4())
pipeline = Pipeline(
tenant_id=tenant_id,
name="test-pipeline",
description="test",
created_by=str(uuid4()),
)
user = make_account(id="u1")
workflow = make_workflow(id=str(uuid4()), created_at=naive_utc_now())
service = MagicMock()
service.publish_workflow.return_value = workflow
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline)
assert result["result"] == "success"
assert "created_at" in result
class TestMiscApis:
def test_task_stop(self, app: Flask) -> None:
api = RagPipelineTaskStopApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account(id="u1")
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.AppQueueManager.set_stop_flag"
) as stop_mock,
):
result = method(api, user, pipeline, "task-1")
stop_mock.assert_called_once()
assert result["result"] == "success"
def test_recommended_plugins(self, app: Flask) -> None:
api = RagPipelineRecommendedPluginApi()
method = unwrap(api.get)
service = MagicMock()
recommended_plugins = {
"installed_recommended_plugins": [{"id": "p1"}],
"uninstalled_recommended_plugins": [{"id": "p2"}],
}
service.get_recommended_plugins.return_value = recommended_plugins
user = make_account()
tenant_id = "tenant-1"
with (
app.test_request_context("/?type=all"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, tenant_id, user)
assert result == recommended_plugins
service.get_recommended_plugins.assert_called_once_with("all", user, tenant_id)
class TestDefaultBlockConfigApi:
def test_get_block_config_success(self, app: Flask) -> None:
api = DefaultRagPipelineBlockConfigApi()
method = unwrap(api.get)
pipeline = make_pipeline()
service = MagicMock()
service.get_default_block_config.return_value = {"k": "v"}
with (
app.test_request_context("/?q={}"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, pipeline, "llm")
assert result == {"k": "v"}
def test_get_block_config_invalid_json(self, app: Flask) -> None:
api = DefaultRagPipelineBlockConfigApi()
method = unwrap(api.get)
pipeline = make_pipeline()
with app.test_request_context("/?q=bad-json"):
with pytest.raises(ValueError):
method(api, pipeline, "llm")
class TestPublishedAllRagPipelineApi:
def test_get_published_workflows_success(self, app: Flask) -> None:
api = PublishedAllRagPipelineApi()
method = unwrap(api.get)
pipeline = make_pipeline()
user = make_account(id="u1")
service = MagicMock()
service.get_all_published_workflow.return_value = ([make_workflow(id="w1")], False)
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline)
assert result["items"][0]["id"] == "w1"
assert result["items"][0]["graph"] == {"nodes": [], "edges": []}
assert result["has_more"] is False
def test_get_published_workflows_forbidden(self, app: Flask) -> None:
api = PublishedAllRagPipelineApi()
method = unwrap(api.get)
pipeline = make_pipeline()
user = make_account(id="u1")
with (
app.test_request_context("/?user_id=u2"),
):
with pytest.raises(Forbidden):
method(api, user, pipeline)
class TestRagPipelineByIdApi:
def test_patch_success(self, app: Flask) -> None:
api = RagPipelineByIdApi()
method = unwrap(api.patch)
pipeline = make_pipeline(tenant_id="t1")
user = make_account(id="u1")
workflow = make_workflow(id="w1", marked_name="test")
service = MagicMock()
service.update_workflow.return_value = workflow
payload = {"marked_name": "test"}
with (
app.test_request_context("/", json=payload),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline, "w1")
assert result["id"] == "w1"
assert result["marked_name"] == "test"
assert result["hash"] == workflow.unique_hash
def test_patch_no_fields(self, app: Flask) -> None:
api = RagPipelineByIdApi()
method = unwrap(api.patch)
pipeline = make_pipeline()
user = make_account()
with app.test_request_context("/", json={}):
result, status = method(api, user, pipeline, "w1")
assert status == 400
def test_delete_success(self, app: Flask) -> None:
api = RagPipelineByIdApi()
method = unwrap(api.delete)
pipeline = make_pipeline(tenant_id="t1", workflow_id="active-workflow")
workflow_service = MagicMock()
with (
app.test_request_context("/", method="DELETE"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.WorkflowService",
return_value=workflow_service,
),
):
result = method(api, pipeline, "old-workflow")
workflow_service.delete_workflow.assert_called_once()
assert result == (None, 204)
def test_delete_active_workflow_rejected(self, app: Flask) -> None:
api = RagPipelineByIdApi()
method = unwrap(api.delete)
pipeline = make_pipeline(tenant_id="t1", workflow_id="active-workflow")
with app.test_request_context("/", method="DELETE"):
with pytest.raises(BadRequest, match="currently in use by pipeline"):
method(api, pipeline, "active-workflow")
class TestRagPipelineWorkflowLastRunApi:
def test_last_run_success(self, app: Flask) -> None:
api = RagPipelineWorkflowLastRunApi()
method = unwrap(api.get)
pipeline = make_pipeline()
workflow = make_workflow()
node_exec = make_node_execution()
service = MagicMock()
service.get_draft_workflow.return_value = workflow
service.get_node_last_run.return_value = node_exec
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, pipeline, "node1")
assert result["id"] == "node-exec-1"
assert result["inputs"] == {"query": "hello"}
assert result["outputs"] == {"answer": "world"}
def test_last_run_not_found(self, app: Flask) -> None:
api = RagPipelineWorkflowLastRunApi()
method = unwrap(api.get)
pipeline = make_pipeline()
service = MagicMock()
service.get_draft_workflow.return_value = None
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
with pytest.raises(NotFound):
method(api, pipeline, "node1")
class TestRagPipelineWorkflowRunNodeExecutionListApi:
def test_get_node_executions_passes_current_user(self, app: Flask) -> None:
api = RagPipelineWorkflowRunNodeExecutionListApi()
method = unwrap(api.get)
user = make_account()
pipeline = make_pipeline()
run_id = uuid4()
node_exec = make_node_execution(workflow_run_id=str(run_id))
service = MagicMock()
service.get_rag_pipeline_workflow_run_node_executions.return_value = [node_exec]
with (
app.test_request_context("/"),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline, run_id)
service.get_rag_pipeline_workflow_run_node_executions.assert_called_once_with(
pipeline=pipeline,
run_id=str(run_id),
user=user,
)
assert result["data"][0]["id"] == "node-exec-1"
assert result["data"][0]["inputs"] == {"query": "hello"}
assert result["data"][0]["outputs"] == {"answer": "world"}
class TestRagPipelineDatasourceVariableApi:
def test_set_datasource_variables_success(self, app: Flask) -> None:
api = RagPipelineDatasourceVariableApi()
method = unwrap(api.post)
pipeline = make_pipeline()
user = make_account()
payload = {
"datasource_type": "db",
"datasource_info": empty_mapping(),
"start_node_id": "n1",
"start_node_title": "Node",
}
service = MagicMock()
service.set_datasource_variables.return_value = make_node_execution(node_id="n1")
with (
app.test_request_context("/", json=payload),
patch(
"controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService",
return_value=service,
),
):
result = method(api, user, pipeline)
assert result["node_id"] == "n1"
assert result["process_data"] == {}