test: use sqlite3 session in test_workflow (#38686)

This commit is contained in:
Asuka Minato
2026-07-22 13:51:33 +09:00
committed by GitHub
parent 0aa04f610e
commit a4c7261bf9
@@ -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()