test: use sqlite3 session in test_workflow_run_service (#38701)

This commit is contained in:
Asuka Minato 2026-07-27 17:19:32 +09:00 committed by GitHub
parent a752f43b8e
commit 460efbf285
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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)