From 460efbf285b558dd7a109705a70adb87a6afe070 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Mon, 27 Jul 2026 17:19:32 +0900 Subject: [PATCH] test: use sqlite3 session in test_workflow_run_service (#38701) --- .../services/test_workflow_run_service.py | 134 +++++++++++------- 1 file changed, 82 insertions(+), 52 deletions(-) diff --git a/api/tests/unit_tests/services/test_workflow_run_service.py b/api/tests/unit_tests/services/test_workflow_run_service.py index fcfa9992cd1..b5902354165 100644 --- a/api/tests/unit_tests/services/test_workflow_run_service.py +++ b/api/tests/unit_tests/services/test_workflow_run_service.py @@ -1,11 +1,16 @@ +"""Workflow-run service tests with real SQLite-bound session factories.""" + +from decimal import Decimal from types import SimpleNamespace from typing import Any, cast from unittest.mock import MagicMock import pytest -from sqlalchemy import Engine +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker -from models import Account, App, EndUser, WorkflowRunTriggeredFrom +from models import Account, App, EndUser, Message, WorkflowRunTriggeredFrom +from models.enums import ConversationFromSource from services import workflow_run_service as service_module from services.workflow_run_service import WorkflowRunService @@ -22,6 +27,11 @@ def repository_factory_mocks(monkeypatch: pytest.MonkeyPatch) -> tuple[MagicMock return node_repo, workflow_run_repo, factory +@pytest.fixture +def sqlalchemy_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) + + def _app_model(**kwargs: Any) -> App: return cast(App, SimpleNamespace(**kwargs)) @@ -34,13 +44,22 @@ def _end_user(**kwargs: Any) -> EndUser: return cast(EndUser, SimpleNamespace(**kwargs)) -def _fake_session_factory_returning_messages(messages: list[Any]) -> tuple[MagicMock, MagicMock]: - """Build a session factory whose session returns the given messages.""" - session = MagicMock() - session.scalars.return_value.all.return_value = messages - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session - return session_factory, session +def _message(*, message_id: str, workflow_run_id: str, conversation_id: str) -> Message: + message = Message( + app_id="app-1", + conversation_id=conversation_id, + query="query", + message={"role": "user", "content": "query"}, + answer="answer", + message_unit_price=Decimal("0.0001"), + answer_unit_price=Decimal("0.0001"), + currency="USD", + from_source=ConversationFromSource.API, + ) + message.id = message_id + message._inputs = {} + message.workflow_run_id = workflow_run_id + return message class TestWorkflowRunServiceInitialization: @@ -48,59 +67,51 @@ class TestWorkflowRunServiceInitialization: self, monkeypatch: pytest.MonkeyPatch, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlite_engine: Engine, ) -> None: - session_factory = MagicMock(name="session_factory") - sessionmaker_mock = MagicMock(return_value=session_factory) - monkeypatch.setattr(service_module, "sessionmaker", sessionmaker_mock) - monkeypatch.setattr(service_module, "db", SimpleNamespace(engine="db-engine")) + monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_engine)) service = WorkflowRunService() - sessionmaker_mock.assert_called_once_with(bind="db-engine", expire_on_commit=False) - assert service._session_factory is session_factory + assert isinstance(service._session_factory, sessionmaker) + assert service._session_factory.kw["bind"] is sqlite_engine + assert service._session_factory.kw["expire_on_commit"] is False def test___init___should_create_sessionmaker_when_engine_is_provided( self, - monkeypatch: pytest.MonkeyPatch, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlite_engine: Engine, ) -> None: - class FakeEngine: - pass + service = WorkflowRunService(session_factory=sqlite_engine) - session_factory = MagicMock(name="session_factory") - sessionmaker_mock = MagicMock(return_value=session_factory) - monkeypatch.setattr(service_module, "Engine", FakeEngine) - monkeypatch.setattr(service_module, "sessionmaker", sessionmaker_mock) - engine = cast(Engine, FakeEngine()) - - service = WorkflowRunService(session_factory=engine) - - sessionmaker_mock.assert_called_once_with(bind=engine, expire_on_commit=False) - assert service._session_factory is session_factory + assert isinstance(service._session_factory, sessionmaker) + assert service._session_factory.kw["bind"] is sqlite_engine + assert service._session_factory.kw["expire_on_commit"] is False def test___init___should_keep_provided_sessionmaker_and_create_repositories( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, workflow_run_repo, factory = repository_factory_mocks - session_factory = MagicMock(name="session_factory") - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) - assert service._session_factory is session_factory + assert service._session_factory is sqlalchemy_session_factory assert service._node_execution_service_repo is node_repo assert service._workflow_run_repo is workflow_run_repo - factory.create_api_workflow_node_execution_repository.assert_called_once_with(session_factory) - factory.create_api_workflow_run_repository.assert_called_once_with(session_factory) + factory.create_api_workflow_node_execution_repository.assert_called_once_with(sqlalchemy_session_factory) + factory.create_api_workflow_run_repository.assert_called_once_with(sqlalchemy_session_factory) class TestWorkflowRunServiceQueries: def test_get_paginate_workflow_runs_should_forward_filters_and_parse_limit( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = MagicMock(name="pagination") workflow_run_repo.get_paginated_workflow_runs.return_value = expected @@ -122,20 +133,24 @@ class TestWorkflowRunServiceQueries: status="succeeded", ) + @pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) def test_get_paginate_advanced_chat_workflow_runs_should_attach_message_fields_when_message_exists( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: - message = SimpleNamespace(id="msg-1", conversation_id="conv-1", workflow_run_id="run-1") - session_factory, session = _fake_session_factory_returning_messages([message]) - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") run_with_message = SimpleNamespace(id="run-1", status="running") run_without_message = SimpleNamespace(id="run-2", status="succeeded") pagination = SimpleNamespace(data=[run_with_message, run_without_message]) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) + sqlite_session.add(_message(message_id="msg-1", conversation_id="conv-1", workflow_run_id="run-1")) + sqlite_session.commit() + result = service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={"limit": "2"}) assert result is pagination @@ -145,39 +160,49 @@ class TestWorkflowRunServiceQueries: assert result.data[0].status == "running" assert not hasattr(result.data[1], "message_id") assert result.data[1].id == "run-2" - # Messages are batch-loaded in a single query, not one per run. - session_factory.assert_called_once_with() - session.scalars.assert_called_once() + @pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) def test_get_paginate_advanced_chat_workflow_runs_batch_loads_messages_without_n_plus_one( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: """Messages must load with a constant query count regardless of run count. Previously the deprecated WorkflowRun.message property issued one query per run (N+1); they are now batch-loaded in a single query. """ - session_factory, session = _fake_session_factory_returning_messages([]) - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") runs = [SimpleNamespace(id=f"run-{i}", status="succeeded") for i in range(5)] pagination = SimpleNamespace(data=runs) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) - service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={}) + message_query_count = 0 - # Exactly one message query for the whole page, independent of run count. - session_factory.assert_called_once_with() - assert session.scalars.call_count == 1 + def count_message_query(*_args: object) -> None: + nonlocal message_query_count + message_query_count += 1 + + engine = sqlite_session.get_bind() + event.listen(engine, "before_cursor_execute", count_message_query) + try: + service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={}) + finally: + event.remove(engine, "before_cursor_execute", count_message_query) + + assert all(not hasattr(run, "message_id") for run in runs) + assert message_query_count == 1 def test_get_workflow_run_should_delegate_to_repository_by_tenant_and_app( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = MagicMock(name="workflow_run") workflow_run_repo.get_workflow_run_by_id.return_value = expected @@ -194,9 +219,10 @@ class TestWorkflowRunServiceQueries: def test_get_workflow_runs_count_should_forward_optional_filters( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = {"total": 3, "succeeded": 2} workflow_run_repo.get_workflow_runs_count.return_value = expected @@ -221,8 +247,9 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=None)) app_model = _app_model(id="app-1") user = _account(current_tenant_id="tenant-1") @@ -235,9 +262,10 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, _, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) class FakeEndUser: @@ -267,9 +295,10 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, _, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) app_model = _app_model(id="app-1") user = _account(current_tenant_id="tenant-account") @@ -293,8 +322,9 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) app_model = _app_model(id="app-1") user = _account(current_tenant_id=None)