diff --git a/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py b/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py index ccdb658b491..150f05081ea 100644 --- a/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py +++ b/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py @@ -1,9 +1,13 @@ import logging +from collections.abc import Iterator +from dataclasses import dataclass from datetime import UTC, datetime, timedelta from types import SimpleNamespace from unittest.mock import Mock, patch import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker from core.app.layers.trigger_post_layer import TriggerPostLayer from core.workflow.system_variables import build_system_variables @@ -13,19 +17,63 @@ from graphon.graph_events import ( GraphRunSucceededEvent, ) from graphon.runtime import VariablePool -from models.enums import WorkflowTriggerStatus +from models.enums import AppTriggerType, CreatorUserRole, WorkflowTriggerStatus +from models.trigger import WorkflowTriggerLog + + +@dataclass(frozen=True) +class TriggerDatabase: + session: Session + statements: list[str] + + +@pytest.fixture(autouse=True) +def trigger_database(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> Iterator[TriggerDatabase]: + """Create the trigger-log table and bind layer-owned sessions to SQLite.""" + WorkflowTriggerLog.metadata.create_all(sqlite_engine, tables=[WorkflowTriggerLog.__table__]) + sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr("core.db.session_factory._session_maker", sqlite_session_maker) + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + with sqlite_session_maker() as session: + try: + yield TriggerDatabase(session=session, statements=statements) + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) + + +def _persist_trigger_log(database: TriggerDatabase, *, trigger_log_id: str = "log-1") -> WorkflowTriggerLog: + trigger_log = WorkflowTriggerLog( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id=None, + root_node_id=None, + trigger_metadata="{}", + trigger_type=AppTriggerType.TRIGGER_WEBHOOK, + trigger_data="{}", + inputs="{}", + outputs=None, + status=WorkflowTriggerStatus.RUNNING, + error=None, + queue_name="workflow", + celery_task_id=None, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + ) + trigger_log.id = trigger_log_id + database.session.add(trigger_log) + database.session.commit() + return trigger_log class TestTriggerPostLayer: - def test_on_event_updates_trigger_log(self): - trigger_log = SimpleNamespace( - status=None, - workflow_run_id=None, - outputs=None, - elapsed_time=None, - total_tokens=None, - finished_at=None, - ) + def test_on_event_updates_trigger_log(self, trigger_database: TriggerDatabase): + trigger_log = _persist_trigger_log(trigger_database) runtime_state = SimpleNamespace( outputs={"answer": "ok"}, variable_pool=VariablePool.from_bootstrap( @@ -35,19 +83,10 @@ class TestTriggerPostLayer: ) with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, patch("core.app.layers.trigger_post_layer.datetime") as mock_datetime, ): mock_datetime.now.return_value = datetime(2026, 2, 20, tzinfo=UTC) - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session - - repo = Mock() - repo.get_by_id.return_value = trigger_log - mock_repo_cls.return_value = repo - layer = TriggerPostLayer( cfs_plan_scheduler_entity=Mock(), start_time=datetime(2026, 2, 20, tzinfo=UTC) - timedelta(seconds=10), @@ -57,25 +96,18 @@ class TestTriggerPostLayer: layer.on_event(GraphRunSucceededEvent()) - assert trigger_log.status == WorkflowTriggerStatus.SUCCEEDED - assert trigger_log.workflow_run_id == "run-1" - assert trigger_log.outputs is not None - assert trigger_log.elapsed_time is not None - assert trigger_log.total_tokens == 12 - assert trigger_log.finished_at is not None - repo.update.assert_called_once_with(trigger_log) - session.commit.assert_called_once() + trigger_database.session.expire_all() + persisted_log = trigger_database.session.get(WorkflowTriggerLog, trigger_log.id) + assert persisted_log is not None + assert persisted_log.status == WorkflowTriggerStatus.SUCCEEDED + assert persisted_log.workflow_run_id == "run-1" + assert persisted_log.outputs == '{"answer":"ok"}' + assert persisted_log.elapsed_time == 10 + assert persisted_log.total_tokens == 12 + assert persisted_log.finished_at is not None - def test_on_event_updates_trigger_log_for_aborted_event(self): - trigger_log = SimpleNamespace( - status=None, - workflow_run_id=None, - outputs=None, - error=None, - elapsed_time=None, - total_tokens=None, - finished_at=None, - ) + def test_on_event_updates_trigger_log_for_aborted_event(self, trigger_database: TriggerDatabase): + trigger_log = _persist_trigger_log(trigger_database) runtime_state = SimpleNamespace( outputs={"partial": "ok"}, variable_pool=VariablePool.from_bootstrap( @@ -85,19 +117,10 @@ class TestTriggerPostLayer: ) with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, patch("core.app.layers.trigger_post_layer.datetime") as mock_datetime, ): mock_datetime.now.return_value = datetime(2026, 2, 20, tzinfo=UTC) - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session - - repo = Mock() - repo.get_by_id.return_value = trigger_log - mock_repo_cls.return_value = repo - layer = TriggerPostLayer( cfs_plan_scheduler_entity=Mock(), start_time=datetime(2026, 2, 20, tzinfo=UTC) - timedelta(seconds=10), @@ -107,17 +130,22 @@ class TestTriggerPostLayer: layer.on_event(GraphRunAbortedEvent(reason="timeout")) - assert trigger_log.status == WorkflowTriggerStatus.FAILED - assert trigger_log.workflow_run_id == "run-1" - assert trigger_log.outputs is not None - assert trigger_log.error == "timeout" - assert trigger_log.elapsed_time is not None - assert trigger_log.total_tokens == 7 - assert trigger_log.finished_at is not None - repo.update.assert_called_once_with(trigger_log) - session.commit.assert_called_once() + trigger_database.session.expire_all() + persisted_log = trigger_database.session.get(WorkflowTriggerLog, trigger_log.id) + assert persisted_log is not None + assert persisted_log.status == WorkflowTriggerStatus.FAILED + assert persisted_log.workflow_run_id == "run-1" + assert persisted_log.outputs == '{"partial":"ok"}' + assert persisted_log.error == "timeout" + assert persisted_log.elapsed_time == 10 + assert persisted_log.total_tokens == 7 + assert persisted_log.finished_at is not None - def test_on_event_handles_missing_trigger_log(self, caplog: pytest.LogCaptureFixture): + def test_on_event_handles_missing_trigger_log( + self, + caplog: pytest.LogCaptureFixture, + trigger_database: TriggerDatabase, + ): runtime_state = SimpleNamespace( outputs={}, variable_pool=VariablePool.from_bootstrap( @@ -126,31 +154,20 @@ class TestTriggerPostLayer: total_tokens=0, ) - with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, - ): - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session + layer = TriggerPostLayer( + cfs_plan_scheduler_entity=Mock(), + start_time=datetime(2026, 2, 20, tzinfo=UTC), + trigger_log_id="missing", + ) + layer.initialize(runtime_state, Mock()) - repo = Mock() - repo.get_by_id.return_value = None - mock_repo_cls.return_value = repo - - layer = TriggerPostLayer( - cfs_plan_scheduler_entity=Mock(), - start_time=datetime(2026, 2, 20, tzinfo=UTC), - trigger_log_id="missing", - ) - layer.initialize(runtime_state, Mock()) - - with caplog.at_level(logging.ERROR, logger="core.app.layers.trigger_post_layer"): - layer.on_event(GraphRunFailedEvent(error="boom")) + with caplog.at_level(logging.ERROR, logger="core.app.layers.trigger_post_layer"): + layer.on_event(GraphRunFailedEvent(error="boom")) assert any(record.levelno == logging.ERROR for record in caplog.records) - session.commit.assert_not_called() + assert trigger_database.session.get(WorkflowTriggerLog, "missing") is None - def test_on_event_ignores_non_status_events(self): + def test_on_event_ignores_non_status_events(self, trigger_database: TriggerDatabase): runtime_state = SimpleNamespace( outputs={}, variable_pool=VariablePool.from_bootstrap( @@ -159,14 +176,14 @@ class TestTriggerPostLayer: total_tokens=0, ) - with patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory: - layer = TriggerPostLayer( - cfs_plan_scheduler_entity=Mock(), - start_time=datetime(2026, 2, 20, tzinfo=UTC), - trigger_log_id="log-1", - ) - layer.initialize(runtime_state, Mock()) + layer = TriggerPostLayer( + cfs_plan_scheduler_entity=Mock(), + start_time=datetime(2026, 2, 20, tzinfo=UTC), + trigger_log_id="log-1", + ) + layer.initialize(runtime_state, Mock()) - layer.on_event(Mock()) + trigger_database.statements.clear() + layer.on_event(Mock()) - mock_session_factory.create_session.assert_not_called() + assert trigger_database.statements == []