test: use sqlite3 session in test_hitl_service_api (#38685)

This commit is contained in:
Asuka Minato 2026-07-23 12:33:01 +09:00 committed by GitHub
parent 78023bf50d
commit 82c8741c73
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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