diff --git a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py index 129220cbc9c..c5072ec2e70 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py @@ -14,6 +14,8 @@ from unittest.mock import ANY, MagicMock, Mock import pytest from flask import Flask +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker import services.app_generate_service as ags_module from controllers.service_api.app.workflow_events import WorkflowEventsApi @@ -31,7 +33,7 @@ from core.app.entities.task_entities import ( from core.app.layers.pause_state_persist_layer import WorkflowResumptionContext, _WorkflowGenerateEntityWrapper from core.workflow.human_input_policy import FormDisposition, HumanInputSurface from core.workflow.nodes.human_input.entities import ParagraphInputConfig, UserActionConfig -from core.workflow.nodes.human_input.enums import FormInputType +from core.workflow.nodes.human_input.enums import FormInputType, HumanInputFormKind, HumanInputFormStatus from core.workflow.nodes.human_input.pause_reason import DifyHITLEventType, HumanInputRequired from core.workflow.system_variables import build_system_variables from graphon.entities import WorkflowStartReason @@ -39,6 +41,7 @@ from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus from graphon.runtime import GraphRuntimeState, VariablePool from models.account import Account from models.enums import CreatorUserRole +from models.human_input import HumanInputForm from models.model import AppMode from models.workflow import WorkflowRun from repositories.api_workflow_node_execution_repository import WorkflowNodeExecutionSnapshot @@ -66,7 +69,7 @@ class _DummyRateLimit: return generator -def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run): +def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run, sqlite_engine: Engine): workflow_events_module = sys.modules["controllers.service_api.app.workflow_events"] repo = SimpleNamespace(get_workflow_run_by_id_and_tenant_id=lambda **_kwargs: workflow_run) monkeypatch.setattr( @@ -74,10 +77,33 @@ def _mock_repo_for_run(monkeypatch: pytest.MonkeyPatch, workflow_run): "create_api_workflow_run_repository", lambda *_args, **_kwargs: repo, ) - monkeypatch.setattr(workflow_events_module, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(workflow_events_module, "db", SimpleNamespace(engine=sqlite_engine)) return workflow_events_module +def _persist_human_input_form( + sqlite_session: Session, + *, + expiration_time: datetime, +) -> HumanInputForm: + form = HumanInputForm( + id="form-1", + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + conversation_id=None, + form_kind=HumanInputFormKind.RUNTIME, + node_id="node-1", + form_definition=json.dumps({"display_in_ui": True}), + rendered_content="Rendered", + status=HumanInputFormStatus.WAITING, + expiration_time=expiration_time, + ) + sqlite_session.add(form) + sqlite_session.commit() + return form + + def _build_service_api_pause_converter() -> WorkflowResponseConverter: application_generate_entity = SimpleNamespace( inputs={}, @@ -257,7 +283,10 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext: class TestHitlServiceApi: # Service API event-stream continuation def test_workflow_events_continue_on_pause_keeps_stream_open( - self, app: Flask, monkeypatch: pytest.MonkeyPatch + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, ) -> None: workflow_run = SimpleNamespace( id="run-1", @@ -266,7 +295,11 @@ class TestHitlServiceApi: created_by="end-user-1", finished_at=None, ) - workflow_events_module = _mock_repo_for_run(monkeypatch, workflow_run=workflow_run) + workflow_events_module = _mock_repo_for_run( + monkeypatch, + workflow_run=workflow_run, + sqlite_engine=sqlite_engine, + ) msg_generator = Mock() msg_generator.retrieve_events.return_value = ["raw-event"] workflow_generator = Mock() @@ -291,7 +324,10 @@ class TestHitlServiceApi: workflow_generator.convert_to_event_stream.assert_called_once_with(["raw-event"]) def test_workflow_events_snapshot_continue_on_pause_keeps_pause_open( - self, app: Flask, monkeypatch: pytest.MonkeyPatch + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, ) -> None: workflow_run = SimpleNamespace( id="run-1", @@ -300,7 +336,11 @@ class TestHitlServiceApi: created_by="end-user-1", finished_at=None, ) - workflow_events_module = _mock_repo_for_run(monkeypatch, workflow_run=workflow_run) + workflow_events_module = _mock_repo_for_run( + monkeypatch, + workflow_run=workflow_run, + sqlite_engine=sqlite_engine, + ) msg_generator = Mock() workflow_generator = Mock() workflow_generator.convert_to_event_stream.return_value = iter(["data: snapshot\n\n"]) @@ -331,16 +371,24 @@ class TestHitlServiceApi: human_input_surface=HumanInputSurface.SERVICE_API, close_on_pause=False, ) + snapshot_session_maker = snapshot_builder.call_args.kwargs["session_maker"] + assert isinstance(snapshot_session_maker, sessionmaker) + assert snapshot_session_maker.kw["bind"] is sqlite_engine workflow_generator.convert_to_event_stream.assert_called_once_with(["snapshot-events"]) - def test_advanced_chat_blocking_injects_pause_state_config(self, monkeypatch: pytest.MonkeyPatch) -> None: + def test_advanced_chat_blocking_injects_pause_state_config( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + ) -> None: monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", False) monkeypatch.setattr(ags_module, "RateLimit", _DummyRateLimit) workflow = MagicMock() workflow.created_by = "owner-id" monkeypatch.setattr(AppGenerateService, "_get_workflow", lambda *args, **kwargs: workflow) - monkeypatch.setattr(ags_module.session_factory, "get_session_maker", lambda: "session-maker") + sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(ags_module.session_factory, "get_session_maker", lambda: sqlite_session_maker) generator_instance = MagicMock() generator_instance.generate.return_value = {"result": "advanced-blocking"} @@ -358,20 +406,21 @@ class TestHitlServiceApi: user = MagicMock() user.id = "user-id" - result = AppGenerateService.generate( - session=Mock(), - app_model=app_model, - user=user, - args={"workflow_id": None, "query": "hi", "inputs": {}}, - invoke_from=InvokeFrom.SERVICE_API, - streaming=False, - ) + with sqlite_session_maker() as session: + result = AppGenerateService.generate( + session=session, + app_model=app_model, + user=user, + args={"workflow_id": None, "query": "hi", "inputs": {}}, + invoke_from=InvokeFrom.SERVICE_API, + streaming=False, + ) assert result == {"result": "advanced-blocking"} call_kwargs = generator_instance.generate.call_args.kwargs assert call_kwargs["streaming"] is False assert call_kwargs["pause_state_config"] is not None - assert call_kwargs["pause_state_config"].session_factory == "session-maker" + assert call_kwargs["pause_state_config"].session_factory is sqlite_session_maker assert call_kwargs["pause_state_config"].state_owner_user_id == "owner-id" # Blocking payload contract @@ -569,7 +618,13 @@ class TestHitlServiceApi: assert response.data.paused_nodes == ["node-1"] assert response.data.reasons == [{"TYPE": "human_input_required", "form_id": "form-1", "expiration_time": 1}] - def test_service_api_pause_event_serializes_hitl_reason(self, monkeypatch: pytest.MonkeyPatch) -> None: + @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) + def test_service_api_pause_event_serializes_hitl_reason( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + ) -> None: converter = _build_service_api_pause_converter() converter.workflow_start_to_stream_response( task_id="task", @@ -578,20 +633,10 @@ class TestHitlServiceApi: reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1, tzinfo=UTC) + expiration_time = datetime(2024, 1, 1) + _persist_human_input_form(sqlite_session, expiration_time=expiration_time) - class _FakeSession: - def execute(self, _stmt): - return [("form-1", expiration_time, '{"display_in_ui": true}')] - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(workflow_response_converter, "Session", lambda **_: _FakeSession()) - monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=object())) + monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_engine)) monkeypatch.setattr( workflow_response_converter, "load_form_dispositions_by_form_id", @@ -651,10 +696,18 @@ class TestHitlServiceApi: assert hi_resp.data.expiration_time == int(expiration_time.timestamp()) # Snapshot payload contract - def test_snapshot_events_include_pause_payload_contract(self, monkeypatch: pytest.MonkeyPatch) -> None: + @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) + def test_snapshot_events_include_pause_payload_contract( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + ) -> None: 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) + _persist_human_input_form(sqlite_session, expiration_time=expiration_time) monkeypatch.setattr( "services.workflow_event_snapshot_service.load_form_dispositions_by_form_id", lambda form_ids, session=None, surface=None: { @@ -662,22 +715,7 @@ class TestHitlServiceApi: }, ) - class _SessionContext: - def __init__(self, session): - self._session = session - - def __enter__(self): - return self._session - - def __exit__(self, exc_type, exc, tb): - return False - - def session_maker() -> _SessionContext: - return _SessionContext( - SimpleNamespace( - execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')], - ) - ) + sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) pause_entity = _FakePauseEntity( pause_id="pause-1", @@ -701,7 +739,7 @@ class TestHitlServiceApi: message_context=None, pause_entity=pause_entity, resumption_context=resumption_context, - session_maker=session_maker, + session_maker=sqlite_session_maker, ) assert [event["event"] for event in events] == [ @@ -713,13 +751,13 @@ class TestHitlServiceApi: ] assert events[2]["data"]["status"] == WorkflowNodeExecutionStatus.PAUSED.value assert events[3]["data"]["form_token"] == "wtok" - assert events[3]["data"]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp()) + assert events[3]["data"]["expiration_time"] == int(expiration_time.timestamp()) pause_data = events[-1]["data"] assert pause_data["paused_nodes"] == ["node-1"] assert pause_data["outputs"] == {"result": "value"} assert pause_data["reasons"][0]["TYPE"] == "human_input_required" 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"] == int(expiration_time.timestamp()) assert pause_data["status"] == WorkflowExecutionStatus.PAUSED.value assert pause_data["created_at"] == int(workflow_run.created_at.timestamp()) assert pause_data["elapsed_time"] == workflow_run.elapsed_time