mirror of
https://github.com/langgenius/dify.git
synced 2026-07-28 23:59:34 +08:00
test: use sqlite3 session in test_workflow_run_service (#38701)
This commit is contained in:
parent
a752f43b8e
commit
460efbf285
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user