test: use SQLite sessions in services workflow (#39087)

This commit is contained in:
Asuka Minato 2026-07-31 21:54:41 +09:00 committed by GitHub
parent f6bbb94112
commit 33518407c9
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -3,6 +3,7 @@ import queue
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import UTC, datetime
from decimal import Decimal
from itertools import cycle
from threading import Event
from types import SimpleNamespace
@ -10,6 +11,7 @@ from typing import Any, cast, override
from unittest.mock import MagicMock
import pytest
from sqlalchemy import event as orm_event
from sqlalchemy.orm import Session, sessionmaker
from core.app.app_config.entities import WorkflowUIBasedAppConfig
@ -26,9 +28,10 @@ from core.workflow.nodes.human_input.enums import ValueSourceType
from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus
from graphon.runtime import GraphRuntimeState, VariablePool
from models.enums import CreatorUserRole
from models.human_input import RecipientType
from models.model import AppMode
from libs.datetime_utils import to_utc_timestamp
from models.enums import ConversationFromSource, CreatorUserRole
from models.human_input import HumanInputForm, HumanInputFormRecipient, RecipientType
from models.model import AppMode, Message
from models.workflow import WorkflowRun
from repositories.api_workflow_node_execution_repository import WorkflowNodeExecutionSnapshot
from repositories.entities.workflow_pause import WorkflowPauseEntity
@ -283,23 +286,57 @@ def _build_advanced_chat_resumption_context(conversation_id: str | None) -> Work
)
class _SessionContext:
def __init__(self, session: Any) -> None:
self._session = session
def __enter__(self) -> Any:
return self._session
def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool:
return False
def _persist_message(
session_maker: sessionmaker[Session],
*,
message_id: str = "msg-1",
app_id: str = "app-1",
conversation_id: str = "conv-1",
workflow_run_id: str = "run-1",
answer: str = "answer",
created_at: datetime | None = None,
) -> None:
message = Message(
id=message_id,
app_id=app_id,
conversation_id=conversation_id,
query="question",
message={"role": "user", "content": "question"},
answer=answer,
message_unit_price=Decimal(0),
answer_unit_price=Decimal(0),
currency="USD",
from_source=ConversationFromSource.API,
workflow_run_id=workflow_run_id,
)
message.inputs = {}
if created_at is not None:
message.created_at = created_at
with session_maker.begin() as session:
session.add(message)
class _SessionMaker:
def __init__(self, session: Any) -> None:
self._session = session
def __call__(self) -> _SessionContext:
return _SessionContext(self._session)
def _persist_human_input_form(
session_maker: sessionmaker[Session],
*,
recipients: Sequence[HumanInputFormRecipient] = (),
) -> datetime:
expiration_time = datetime(2024, 1, 1)
form = HumanInputForm(
id="form-1",
tenant_id="tenant-1",
app_id="app-1",
workflow_run_id="run-1",
conversation_id=None,
node_id="node-1",
form_definition='{"display_in_ui": true}',
rendered_content="content",
expiration_time=expiration_time,
)
with session_maker.begin() as session:
session.add(form)
session.add_all(recipients)
return expiration_time
class _SubscriptionContext:
@ -359,14 +396,12 @@ class _PauseEntity(WorkflowPauseEntity):
return []
def test_get_message_context_by_conversation_should_return_none_when_no_message() -> None:
# Arrange
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
def test_get_message_context_by_conversation_should_return_none_when_no_message(
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Act
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
sqlite_session_factory,
conversation_id="conv-1",
workflow_run_id="run-1",
)
@ -375,69 +410,112 @@ def test_get_message_context_by_conversation_should_return_none_when_no_message(
assert result is None
def test_get_message_context_by_conversation_should_scope_and_bound_message_lookup() -> None:
def test_get_message_context_by_conversation_should_scope_and_bound_message_lookup(
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
_persist_message(
sqlite_session_factory,
message_id="older-match",
created_at=datetime(2024, 1, 1),
)
_persist_message(
sqlite_session_factory,
message_id="newest-match",
app_id="other-app",
answer="newest answer",
created_at=datetime(2024, 1, 2),
)
_persist_message(
sqlite_session_factory,
message_id="wrong-workflow",
workflow_run_id="run-2",
created_at=datetime(2024, 1, 3),
)
_persist_message(
sqlite_session_factory,
message_id="wrong-conversation",
conversation_id="conv-2",
created_at=datetime(2024, 1, 4),
)
# Act
service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
result = service_module._get_message_context_by_conversation(
sqlite_session_factory,
conversation_id="conv-1",
workflow_run_id="run-1",
)
# Assert
stmt = session.scalar.call_args.args[0]
compiled = " ".join(str(stmt.compile(compile_kwargs={"literal_binds": True})).split())
where_clause = compiled.split(" WHERE ", maxsplit=1)[1].split(" ORDER BY ", maxsplit=1)[0]
assert "messages.conversation_id = 'conv-1'" in compiled
assert "messages.workflow_run_id = 'run-1'" in compiled
assert "messages.app_id" not in where_clause
assert "ORDER BY messages.created_at DESC" in compiled
assert compiled.endswith("LIMIT 1")
assert result is not None
assert result.message_id == "newest-match"
assert result.conversation_id == "conv-1"
assert result.answer == "newest answer"
def test_get_message_context_by_app_should_scope_and_bound_compatibility_lookup() -> None:
def test_get_message_context_by_app_should_scope_and_bound_compatibility_lookup(
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
_persist_message(
sqlite_session_factory,
message_id="older-match",
created_at=datetime(2024, 1, 1),
)
_persist_message(
sqlite_session_factory,
message_id="newest-match",
conversation_id="conv-2",
answer="newest answer",
created_at=datetime(2024, 1, 2),
)
_persist_message(
sqlite_session_factory,
message_id="wrong-app",
app_id="app-2",
created_at=datetime(2024, 1, 3),
)
_persist_message(
sqlite_session_factory,
message_id="wrong-workflow",
workflow_run_id="run-2",
created_at=datetime(2024, 1, 4),
)
# Act
service_module._get_message_context_by_app(
cast(sessionmaker[Session], session_maker),
result = service_module._get_message_context_by_app(
sqlite_session_factory,
app_id="app-1",
workflow_run_id="run-1",
)
# Assert
stmt = session.scalar.call_args.args[0]
compiled = " ".join(str(stmt.compile(compile_kwargs={"literal_binds": True})).split())
where_clause = compiled.split(" WHERE ", maxsplit=1)[1].split(" ORDER BY ", maxsplit=1)[0]
assert "messages.app_id = 'app-1'" in where_clause
assert "messages.workflow_run_id = 'run-1'" in where_clause
assert "messages.conversation_id" not in where_clause
assert "ORDER BY messages.created_at DESC" in compiled
assert compiled.endswith("LIMIT 1")
assert result is not None
assert result.message_id == "newest-match"
assert result.conversation_id == "conv-2"
assert result.answer == "newest answer"
def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp() -> None:
def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp(
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
message = SimpleNamespace(
id="msg-1",
conversation_id="conv-1",
created_at=None,
answer="answer",
)
session = SimpleNamespace(scalar=MagicMock(return_value=message))
session_maker = _SessionMaker(session)
_persist_message(sqlite_session_factory)
def clear_created_at(message: Message, _context: Any) -> None:
# A load hook preserves coverage for legacy rows without replacing the real ORM query.
message.created_at = None
# Act
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
orm_event.listen(Message, "load", clear_created_at)
try:
result = service_module._get_message_context_by_conversation(
sqlite_session_factory,
conversation_id="conv-1",
workflow_run_id="run-1",
)
finally:
orm_event.remove(Message, "load", clear_created_at)
# Assert
assert result is not None
@ -825,6 +903,7 @@ def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback(
def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_idle_timeout(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
@ -866,7 +945,7 @@ def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_id
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=MagicMock(),
session_maker=sqlite_session_factory,
idle_timeout=20.0,
ping_interval=5.0,
)
@ -879,6 +958,7 @@ def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_id
def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
@ -911,7 +991,7 @@ def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty(
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=MagicMock(),
session_maker=sqlite_session_factory,
)
)
@ -922,6 +1002,7 @@ def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty(
def test_build_workflow_event_stream_should_continue_when_pause_loading_fails(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
@ -954,7 +1035,7 @@ def test_build_workflow_event_stream_should_continue_when_pause_loading_fails(
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=MagicMock(),
session_maker=sqlite_session_factory,
)
)
@ -972,7 +1053,10 @@ def test_is_terminal_event_respects_close_on_pause_flag() -> None:
assert _is_terminal_event(finish_event, close_on_pause=False) is True
def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.MonkeyPatch) -> None:
def test_build_snapshot_events_preserves_public_form_token(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED)
snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED)
resumption_context = _build_resumption_context("task-ctx")
@ -983,12 +1067,7 @@ def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.M
"form-1": FormDisposition(form_token="wtok", approval_channels=[])
},
)
session_maker = _SessionMaker(
SimpleNamespace(
# Persisted UTC datetimes can be loaded as naive values.
execute=lambda _stmt: [("form-1", datetime(2024, 1, 1), '{"display_in_ui": true}')],
)
)
expiration_time = _persist_human_input_form(sqlite_session_factory)
pause_entity = _FakePauseEntity(
pause_id="pause-1",
workflow_run_id="run-1",
@ -1011,34 +1090,30 @@ def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.M
message_context=None,
pause_entity=pause_entity,
resumption_context=resumption_context,
session_maker=cast(sessionmaker[Session], session_maker),
session_maker=sqlite_session_factory,
)
assert events[-2]["event"] == StreamEvent.HUMAN_INPUT_REQUIRED
assert events[-2]["data"]["form_token"] == "wtok"
assert events[-2]["data"]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp())
assert events[-2]["data"]["expiration_time"] == to_utc_timestamp(expiration_time)
pause_data = events[-1]["data"]
assert pause_data["reasons"][0]["form_token"] == "wtok"
assert pause_data["reasons"][0]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp())
assert pause_data["reasons"][0]["expiration_time"] == to_utc_timestamp(expiration_time)
def _build_recipient_snapshot_events(recipients: Sequence[Any]) -> list[Mapping[str, Any]]:
def _build_recipient_snapshot_events(
session_maker: sessionmaker[Session],
recipients: Sequence[HumanInputFormRecipient],
) -> list[Mapping[str, Any]]:
"""Drive the reconnect snapshot pause path for the OPENAPI surface.
Lets the real disposition loader run against a fake session whose ``scalars``
yields the given recipients, so the reconnect path derives the same token and
approval channels as the live path for the same recipient set.
Persisting the recipients lets the real disposition query derive the same token
and approval channels as the live path for the same recipient set.
"""
workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED)
snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED)
resumption_context = _build_resumption_context("task-ctx")
expiration_time = datetime(2024, 1, 1, tzinfo=UTC)
session_maker = _SessionMaker(
SimpleNamespace(
execute=lambda _stmt: [("form-1", expiration_time, '{"display_in_ui": true}')],
scalars=lambda _stmt: list(recipients),
)
)
expiration_time = _persist_human_input_form(session_maker, recipients=recipients)
pause_entity = _FakePauseEntity(
pause_id="pause-1",
workflow_run_id="run-1",
@ -1060,16 +1135,31 @@ def _build_recipient_snapshot_events(recipients: Sequence[Any]) -> list[Mapping[
message_context=None,
pause_entity=pause_entity,
resumption_context=resumption_context,
session_maker=cast(sessionmaker[Session], session_maker),
session_maker=session_maker,
human_input_surface=HumanInputSurface.OPENAPI,
)
def test_reconnect_pause_without_web_app_recipient_emits_approval_channels() -> None:
def test_reconnect_pause_without_web_app_recipient_emits_approval_channels(
sqlite_session_factory: sessionmaker[Session],
) -> None:
events = _build_recipient_snapshot_events(
sqlite_session_factory,
recipients=[
SimpleNamespace(form_id="form-1", recipient_type=RecipientType.EMAIL_MEMBER, access_token="email-token"),
SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"),
HumanInputFormRecipient(
form_id="form-1",
delivery_id="delivery-1",
recipient_type=RecipientType.EMAIL_MEMBER,
recipient_payload="{}",
access_token="email-token",
),
HumanInputFormRecipient(
form_id="form-1",
delivery_id="delivery-2",
recipient_type=RecipientType.BACKSTAGE,
recipient_payload="{}",
access_token="backstage-token",
),
],
)
@ -1083,15 +1173,26 @@ def test_reconnect_pause_without_web_app_recipient_emits_approval_channels() ->
assert pause_data["reasons"][0]["approval_channels"] == ["console", "email"]
def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels() -> None:
def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels(
sqlite_session_factory: sessionmaker[Session],
) -> None:
events = _build_recipient_snapshot_events(
sqlite_session_factory,
recipients=[
SimpleNamespace(
HumanInputFormRecipient(
form_id="form-1",
delivery_id="delivery-1",
recipient_type=RecipientType.STANDALONE_WEB_APP,
recipient_payload="{}",
access_token="web-app-token",
),
SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"),
HumanInputFormRecipient(
form_id="form-1",
delivery_id="delivery-2",
recipient_type=RecipientType.BACKSTAGE,
recipient_payload="{}",
access_token="backstage-token",
),
],
)
@ -1105,7 +1206,10 @@ def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels() -> Non
assert pause_data["reasons"][0]["approval_channels"] == ["console"]
def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch: pytest.MonkeyPatch) -> None:
def test_build_snapshot_events_resolves_pause_reason_select_options(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED)
snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED)
resumption_context = _build_resumption_context("task-ctx", select_options=["approve", "reject"])
@ -1116,11 +1220,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch:
"form-1": FormDisposition(form_token="wtok", approval_channels=[])
},
)
session_maker = _SessionMaker(
SimpleNamespace(
execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')],
)
)
_persist_human_input_form(sqlite_session_factory)
pause_entity = _FakePauseEntity(
pause_id="pause-1",
workflow_run_id="run-1",
@ -1152,7 +1252,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch:
message_context=None,
pause_entity=pause_entity,
resumption_context=resumption_context,
session_maker=cast(sessionmaker[Session], session_maker),
session_maker=sqlite_session_factory,
)
human_input_event = events[-2]
@ -1164,6 +1264,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch:
def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_context(
monkeypatch: pytest.MonkeyPatch,
sqlite_session_factory: sessionmaker[Session],
) -> None:
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
topic = _Topic(_StaticSubscription())
@ -1199,11 +1300,7 @@ def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_contex
},
)
session = SimpleNamespace(
scalar=MagicMock(return_value=None),
execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')],
)
session_maker = _SessionMaker(session)
expiration_time = _persist_human_input_form(sqlite_session_factory)
events = list(
build_workflow_event_stream(
@ -1211,11 +1308,11 @@ def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_contex
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=cast(sessionmaker[Session], session_maker),
session_maker=sqlite_session_factory,
)
)
pause_event = cast(Mapping[str, Any], events[-1])
assert pause_event["event"] == StreamEvent.WORKFLOW_PAUSED
assert pause_event["data"]["reasons"][0]["form_token"] == "wtok"
assert pause_event["data"]["reasons"][0]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp())
assert pause_event["data"]["reasons"][0]["expiration_time"] == to_utc_timestamp(expiration_time)