diff --git a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py index f381bd3fbc4..7975a935f93 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py @@ -16,20 +16,20 @@ Focus on: import json import sys import uuid -from dataclasses import dataclass, field from datetime import UTC, datetime from inspect import unwrap +from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask -from sqlalchemy.orm import sessionmaker +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, NotFound from controllers.service_api.app.error import NotWorkflowAppError, WorkflowVersionExecutionNotAllowedError from controllers.service_api.app.workflow import ( AppQueueManager, - DifyAPIRepositoryFactory, GraphEngineManager, WorkflowAppLogApi, WorkflowLogQuery, @@ -44,6 +44,7 @@ from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpErr from core.app.entities.app_invoke_entities import InvokeFrom from enums.cloud_plan import CloudPlan from graphon.enums import WorkflowExecutionStatus +from models import Account from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom from models.model import App, AppMode, EndUser from models.workflow import WorkflowAppLog, WorkflowAppLogCreatedFrom, WorkflowRun, WorkflowType @@ -51,58 +52,18 @@ from services.app_generate_service import AppGenerateService from services.billing_service import BillingService from services.errors.app import IsDraftWorkflowError, WorkflowNotFoundError from services.errors.llm import InvokeRateLimitError -from services.workflow_app_service import LogView, LogViewDetails, WorkflowAppService +from services.workflow_app_service import WorkflowAppService def _default_workflow_inputs() -> dict[str, object]: return {"input": "value"} -def _default_log_details() -> LogViewDetails: - return {"trigger_metadata": {"node": "answer", "latency": 1.25}} - - -class _DbSessionStub: - def get(self, *args: object, **kwargs: object) -> None: - return None - - -@dataclass -class _DbStub: - engine: object = field(default_factory=object) - session: _DbSessionStub = field(default_factory=_DbSessionStub) - - -@dataclass -class _WorkflowRunRepositoryStub: - run: WorkflowRun | None - - def get_workflow_run_by_id(self, *, tenant_id: str, app_id: str, run_id: str) -> WorkflowRun | None: - return self.run if tenant_id and app_id and run_id else None - - def get_workflow_run_by_id_without_tenant(self, *, run_id: str) -> WorkflowRun | None: - return self.run if run_id else None - - -class _BeginStub: - def __enter__(self) -> object: - return object() - - def __exit__(self, exc_type: object, exc: object, tb: object) -> bool: - return False - - -class _SessionMakerStub: - def __init__(self, *args: object, **kwargs: object) -> None: - pass - - def begin(self) -> _BeginStub: - return _BeginStub() - - def _make_workflow_run( run_id: str = "run-1", *, + tenant_id: str = "tenant-1", + app_id: str = "app-1", workflow_id: str = "wf-1", inputs: dict[str, object] | None = None, outputs: dict[str, object] | None = None, @@ -111,8 +72,8 @@ def _make_workflow_run( ) -> WorkflowRun: return WorkflowRun( id=run_id, - tenant_id="tenant-1", - app_id="app-1", + tenant_id=tenant_id, + app_id=app_id, workflow_id=workflow_id, type=WorkflowType.WORKFLOW, triggered_from=WorkflowRunTriggeredFrom.APP_RUN, @@ -133,12 +94,17 @@ def _make_workflow_run( ) -def _make_workflow_app_log() -> WorkflowAppLog: +def _make_workflow_app_log( + *, + tenant_id: str = "tenant-1", + app_id: str = "app-1", + workflow_run_id: str = "log-run-1", +) -> WorkflowAppLog: log = WorkflowAppLog( - tenant_id="tenant-1", - app_id="app-1", + tenant_id=tenant_id, + app_id=app_id, workflow_id="wf-1", - workflow_run_id="log-run-1", + workflow_run_id=workflow_run_id, created_from=WorkflowAppLogCreatedFrom.SERVICE_API, created_by_role=CreatorUserRole.ACCOUNT, created_by="account-1", @@ -148,16 +114,6 @@ def _make_workflow_app_log() -> WorkflowAppLog: return log -def _make_workflow_log_page() -> dict[str, object]: - return { - "page": 1, - "limit": 20, - "total": 1, - "has_more": False, - "data": [LogView(_make_workflow_app_log(), _default_log_details())], - } - - def _make_app_model( *, app_id: str = "app-1", @@ -177,6 +133,43 @@ def _make_end_user(user_id: str = "end-user-1") -> EndUser: return end_user +def _bind_sqlite_database( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, +) -> None: + """Bind controller- and model-owned database access to the test engine.""" + database = SimpleNamespace(engine=sqlite_engine, session=sqlite_session) + monkeypatch.setattr(sys.modules["controllers.service_api.app.workflow"], "db", database) + monkeypatch.setattr(sys.modules["models.workflow"], "db", database) + + +def _persist_workflow_log( + sqlite_session: Session, + *, + tenant_id: str, + app_id: str, +) -> None: + workflow_run_id = "log-run-1" + sqlite_session.add_all( + [ + _make_workflow_run( + run_id=workflow_run_id, + tenant_id=tenant_id, + app_id=app_id, + created_at=datetime(2026, 1, 1, 1, tzinfo=UTC), + finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC), + ), + _make_workflow_app_log( + tenant_id=tenant_id, + app_id=app_id, + workflow_run_id=workflow_run_id, + ), + ] + ) + sqlite_session.commit() + + def _expected_workflow_log_pagination_payload() -> dict[str, object]: return { "page": 1, @@ -195,16 +188,16 @@ def _expected_workflow_log_pagination_payload() -> dict[str, object]: "elapsed_time": 0.1, "total_tokens": 10, "total_steps": 1, - "created_at": 1767229200, - "finished_at": 1767229202, + "created_at": int(datetime(2026, 1, 1, 1).timestamp()), + "finished_at": int(datetime(2026, 1, 1, 1, 0, 2).timestamp()), "exceptions_count": 0, }, - "details": {"trigger_metadata": {"node": "answer", "latency": 1.25}}, + "details": None, "created_from": "service-api", "created_by_role": "account", "created_by_account": None, "created_by_end_user": None, - "created_at": 1767229203, + "created_at": int(datetime(2026, 1, 1, 1, 0, 3).timestamp()), } ], } @@ -364,15 +357,15 @@ class TestWorkflowAppService: assert hasattr(WorkflowAppService, "get_paginate_workflow_app_logs") assert callable(WorkflowAppService.get_paginate_workflow_app_logs) - @patch.object(WorkflowAppService, "get_paginate_workflow_app_logs") - def test_get_paginate_workflow_app_logs_returns_pagination(self, mock_get_logs): - """Test get_paginate_workflow_app_logs returns paginated result.""" - pagination = _make_workflow_log_page() - mock_get_logs.return_value = pagination - + @pytest.mark.parametrize("sqlite_session", [(WorkflowAppLog,)], indirect=True) + def test_get_paginate_workflow_app_logs_returns_pagination(self, sqlite_session: Session): + """Test pagination returns committed logs scoped to the requested app.""" + log = _make_workflow_app_log() + sqlite_session.add(log) + sqlite_session.commit() service = WorkflowAppService() result = service.get_paginate_workflow_app_logs( - session=Mock(), + session=sqlite_session, app_model=_make_app_model(), keyword=None, status=None, @@ -384,7 +377,11 @@ class TestWorkflowAppService: created_by_account=None, ) - assert result == pagination + assert result["page"] == 1 + assert result["limit"] == 20 + assert result["total"] == 1 + assert result["has_more"] is False + assert [item.id for item in result["data"]] == [log.id] class TestWorkflowExecutionStatus: @@ -409,8 +406,9 @@ class TestWorkflowExecutionStatus: class TestAppGenerateServiceWorkflow: """Test AppGenerateService workflow integration.""" + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) @patch.object(AppGenerateService, "generate") - def test_generate_accepts_workflow_args(self, mock_generate: MagicMock): + def test_generate_accepts_workflow_args(self, mock_generate: MagicMock, sqlite_session: Session): """Test generate accepts workflow-specific args.""" mock_generate.return_value = {"result": "success"} @@ -419,15 +417,17 @@ class TestAppGenerateServiceWorkflow: user=_make_end_user(), args={"inputs": {"key": "value"}, "workflow_id": "workflow_123"}, invoke_from=InvokeFrom.SERVICE_API, - session=MagicMock(), + session=sqlite_session, streaming=False, ) assert result == {"result": "success"} mock_generate.assert_called_once() + assert mock_generate.call_args.kwargs["session"] is sqlite_session + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) @patch.object(AppGenerateService, "generate") - def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock): + def test_generate_raises_workflow_not_found_error(self, mock_generate: MagicMock, sqlite_session: Session): """Test generate raises WorkflowNotFoundError.""" mock_generate.side_effect = WorkflowNotFoundError("Workflow not found") @@ -437,12 +437,13 @@ class TestAppGenerateServiceWorkflow: user=_make_end_user(), args={"workflow_id": "invalid_id"}, invoke_from=InvokeFrom.SERVICE_API, - session=MagicMock(), + session=sqlite_session, streaming=False, ) + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) @patch.object(AppGenerateService, "generate") - def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock): + def test_generate_raises_is_draft_workflow_error(self, mock_generate: MagicMock, sqlite_session: Session): """Test generate raises IsDraftWorkflowError.""" mock_generate.side_effect = IsDraftWorkflowError("Workflow is draft") @@ -452,12 +453,13 @@ class TestAppGenerateServiceWorkflow: user=_make_end_user(), args={"workflow_id": "draft_workflow"}, invoke_from=InvokeFrom.SERVICE_API, - session=MagicMock(), + session=sqlite_session, streaming=False, ) + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) @patch.object(AppGenerateService, "generate") - def test_generate_supports_streaming_mode(self, mock_generate: MagicMock): + def test_generate_supports_streaming_mode(self, mock_generate: MagicMock, sqlite_session: Session): """Test generate supports streaming response mode.""" mock_stream = Mock() mock_generate.return_value = mock_stream @@ -467,7 +469,7 @@ class TestAppGenerateServiceWorkflow: user=_make_end_user(), args={"inputs": {}, "response_mode": "streaming"}, invoke_from=InvokeFrom.SERVICE_API, - session=MagicMock(), + session=sqlite_session, streaming=True, ) @@ -499,19 +501,23 @@ class TestWorkflowRunRepository: assert hasattr(DifyAPIRepositoryFactory, "create_api_workflow_run_repository") - @patch("repositories.factory.DifyAPIRepositoryFactory.create_api_workflow_run_repository") - def test_workflow_run_repository_get_by_id(self, mock_factory): - """Test workflow run repository get_workflow_run_by_id method.""" + @pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True) + def test_workflow_run_repository_get_by_id(self, sqlite_engine: Engine, sqlite_session: Session): + """Test repository lookup against committed tenant-scoped state.""" run = _make_workflow_run(run_id=str(uuid.uuid4())) - mock_factory.return_value = _WorkflowRunRepositoryStub(run=run) - + sqlite_session.add(run) + sqlite_session.commit() from repositories.factory import DifyAPIRepositoryFactory - repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(sessionmaker()) + repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository( + sessionmaker(bind=sqlite_engine, expire_on_commit=False) + ) - result = repo.get_workflow_run_by_id(tenant_id="tenant_123", app_id="app_456", run_id="run_789") + result = repo.get_workflow_run_by_id(tenant_id="tenant-1", app_id="app-1", run_id=run.id) - assert result == run + assert result is not None + assert result.id == run.id + assert repo.get_workflow_run_by_id(tenant_id="other-tenant", app_id="app-1", run_id=run.id) is None class TestWorkflowRunDetailApi: @@ -524,16 +530,17 @@ class TestWorkflowRunDetailApi: with pytest.raises(NotWorkflowAppError): handler(api, app_model=app_model, workflow_run_id="run") - def test_success(self, monkeypatch: pytest.MonkeyPatch) -> None: - run = _make_workflow_run(run_id="run") - repo = _WorkflowRunRepositoryStub(run=run) - workflow_module = sys.modules["controllers.service_api.app.workflow"] - monkeypatch.setattr(workflow_module, "db", _DbStub()) - monkeypatch.setattr( - DifyAPIRepositoryFactory, - "create_api_workflow_run_repository", - lambda *_args, **_kwargs: repo, - ) + @pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True) + def test_success( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + ) -> None: + run = _make_workflow_run(run_id="run", tenant_id="t1", app_id="a1") + sqlite_session.add(run) + sqlite_session.commit() + _bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session) api = WorkflowRunDetailApi() handler = unwrap(api.get) @@ -546,7 +553,8 @@ class TestWorkflowRunDetailApi: class TestWorkflowRunApi: - def test_not_workflow_app(self, app: Flask) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_not_workflow_app(self, app: Flask, sqlite_session: Session) -> None: api = WorkflowRunApi() handler = unwrap(api.post) app_model = _make_app_model(mode=AppMode.CHAT) @@ -554,9 +562,10 @@ class TestWorkflowRunApi: with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}): with pytest.raises(NotWorkflowAppError): - handler(api, session=Mock(), app_model=app_model, end_user=end_user) + handler(api, session=sqlite_session, app_model=app_model, end_user=end_user) - def test_rate_limit(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_rate_limit(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: monkeypatch.setattr( AppGenerateService, "generate", @@ -570,7 +579,7 @@ class TestWorkflowRunApi: with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}): with pytest.raises(InvokeRateLimitHttpError): - handler(api, session=Mock(), app_model=app_model, end_user=end_user) + handler(api, session=sqlite_session, app_model=app_model, end_user=end_user) def test_sandbox_billing_does_not_gate_default_workflow_run( self, app: Flask, monkeypatch: pytest.MonkeyPatch @@ -680,7 +689,8 @@ class TestWorkflowRunByIdApi: else: billing_get_info.assert_not_called() - def test_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: workflow_module = sys.modules["controllers.service_api.app.workflow"] monkeypatch.setattr(workflow_module.dify_config, "BILLING_ENABLED", False) monkeypatch.setattr( @@ -696,9 +706,10 @@ class TestWorkflowRunByIdApi: with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}): with pytest.raises(NotFound): - handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1") + handler(api, session=sqlite_session, app_model=app_model, end_user=end_user, workflow_id="w1") - def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_draft_workflow(self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: workflow_module = sys.modules["controllers.service_api.app.workflow"] monkeypatch.setattr(workflow_module.dify_config, "BILLING_ENABLED", False) monkeypatch.setattr( @@ -714,7 +725,7 @@ class TestWorkflowRunByIdApi: with app.test_request_context("/workflows/1/run", method="POST", json={"inputs": {}}): with pytest.raises(BadRequest): - handler(api, session=Mock(), app_model=app_model, end_user=end_user, workflow_id="w1") + handler(api, session=sqlite_session, app_model=app_model, end_user=end_user, workflow_id="w1") class TestWorkflowTaskStopApi: @@ -748,28 +759,16 @@ class TestWorkflowTaskStopApi: class TestWorkflowAppLogApi: - def test_success(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - workflow_module = sys.modules["controllers.service_api.app.workflow"] - workflow_model_module = sys.modules["models.workflow"] - monkeypatch.setattr(workflow_module, "db", _DbStub()) - monkeypatch.setattr(workflow_model_module, "db", _DbStub()) - monkeypatch.setattr(workflow_module, "sessionmaker", _SessionMakerStub) - monkeypatch.setattr( - WorkflowAppService, - "get_paginate_workflow_app_logs", - lambda *_args, **_kwargs: _make_workflow_log_page(), - ) - monkeypatch.setattr( - DifyAPIRepositoryFactory, - "create_api_workflow_run_repository", - lambda *_args, **_kwargs: _WorkflowRunRepositoryStub( - run=_make_workflow_run( - run_id="log-run-1", - created_at=datetime(2026, 1, 1, 1, tzinfo=UTC), - finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC), - ) - ), - ) + @pytest.mark.parametrize("sqlite_session", [(WorkflowRun, WorkflowAppLog, Account)], indirect=True) + def test_success( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + ) -> None: + _persist_workflow_log(sqlite_session, tenant_id="tenant-1", app_id="a1") + _bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session) api = WorkflowAppLogApi() handler = unwrap(api.get) @@ -803,18 +802,24 @@ class TestWorkflowRunDetailApiGet: and we call the unwrapped method directly in tests. """ - @patch("controllers.service_api.app.workflow.DifyAPIRepositoryFactory") - @patch("controllers.service_api.app.workflow.db") + @pytest.mark.parametrize("sqlite_session", [(WorkflowRun,)], indirect=True) def test_get_workflow_run_success( self, - mock_db, - mock_repo_factory, app: Flask, workflow_app: App, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, ): """Test successful workflow run detail retrieval.""" - run = _make_workflow_run(run_id="run-1") - mock_repo_factory.create_api_workflow_run_repository.return_value = _WorkflowRunRepositoryStub(run=run) + run = _make_workflow_run( + run_id="run-1", + tenant_id=workflow_app.tenant_id, + app_id=workflow_app.id, + ) + sqlite_session.add(run) + sqlite_session.commit() + _bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session) from controllers.service_api.app.workflow import WorkflowRunDetailApi @@ -834,13 +839,12 @@ class TestWorkflowRunDetailApiGet: "error": None, "total_steps": 1, "total_tokens": 10, - "created_at": 1767225600, - "finished_at": 1767225600, + "created_at": int(datetime(2026, 1, 1).timestamp()), + "finished_at": int(datetime(2026, 1, 1).timestamp()), "elapsed_time": 0.1, } - @patch("controllers.service_api.app.workflow.db") - def test_get_workflow_run_wrong_app_mode(self, mock_db, app: Flask): + def test_get_workflow_run_wrong_app_mode(self, app: Flask): """Test NotWorkflowAppError when app mode is not workflow or advanced_chat.""" from controllers.service_api.app.workflow import WorkflowRunDetailApi @@ -902,46 +906,23 @@ class TestWorkflowAppLogApiGet: ``get`` is wrapped by ``@validate_app_token``. """ - @patch("controllers.service_api.app.workflow.WorkflowAppService") - @patch("controllers.service_api.app.workflow.db") + @pytest.mark.parametrize("sqlite_session", [(WorkflowRun, WorkflowAppLog, Account)], indirect=True) def test_get_workflow_logs_success( self, - mock_db, - mock_wf_svc_cls, app: Flask, workflow_app: App, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, ): """Test successful workflow log retrieval.""" - mock_svc_instance = Mock() - mock_svc_instance.get_paginate_workflow_app_logs.return_value = _make_workflow_log_page() - mock_wf_svc_cls.return_value = mock_svc_instance - mock_repo = _WorkflowRunRepositoryStub( - run=_make_workflow_run( - run_id="log-run-1", - created_at=datetime(2026, 1, 1, 1, tzinfo=UTC), - finished_at=datetime(2026, 1, 1, 1, 0, 2, tzinfo=UTC), - ) - ) - - # Mock sessionmaker(...).begin() context manager - mock_db.engine = object() - mock_db.session.get.return_value = None + _persist_workflow_log(sqlite_session, tenant_id=workflow_app.tenant_id, app_id=workflow_app.id) + _bind_sqlite_database(monkeypatch, sqlite_engine, sqlite_session) from controllers.service_api.app.workflow import WorkflowAppLogApi - with app.test_request_context( - "/workflows/logs?page=1&limit=20", - method="GET", - ): - with ( - patch("controllers.service_api.app.workflow.sessionmaker", _SessionMakerStub), - patch("models.workflow.db", _DbStub()), - patch( - "repositories.factory.DifyAPIRepositoryFactory.create_api_workflow_run_repository", - return_value=mock_repo, - ), - ): - api = WorkflowAppLogApi() - result = unwrap(api.get)(api, app_model=workflow_app) + with app.test_request_context("/workflows/logs?page=1&limit=20", method="GET"): + api = WorkflowAppLogApi() + result = unwrap(api.get)(api, app_model=workflow_app) assert result == _expected_workflow_log_pagination_payload()