diff --git a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py index 4260f0064c6..80bc90b72b9 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py @@ -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)