fix(api): scope HITL snapshot message lookup to utilize db index (#39351)

This commit is contained in:
QuantumGhost 2026-07-21 17:47:17 +08:00 committed by GitHub
parent 62bbc0dbeb
commit af465fb527
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 346 additions and 33 deletions

View File

@ -13,6 +13,7 @@ from sqlalchemy import desc, select
from sqlalchemy.orm import Session, sessionmaker
from core.app.apps.message_generator import MessageGenerator
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity
from core.app.entities.task_entities import (
HumanInputRequiredResponse,
MessageReplaceStreamResponse,
@ -84,9 +85,6 @@ def build_workflow_event_stream(
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
node_execution_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository(session_maker)
message_context = (
_get_message_context(session_maker, workflow_run.id) if app_mode == AppMode.ADVANCED_CHAT else None
)
pause_entity: WorkflowPauseEntity | None = None
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
@ -97,6 +95,38 @@ def build_workflow_event_stream(
pause_entity = None
resumption_context = _load_resumption_context(pause_entity)
message_context: MessageContext | None = None
if app_mode == AppMode.ADVANCED_CHAT:
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
if resumption_context is None:
raise AssertionError(
"WorkflowResumptionContext is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}"
)
generate_entity = resumption_context.get_generate_entity()
if not isinstance(generate_entity, AdvancedChatAppGenerateEntity):
raise AssertionError(
"AdvancedChatAppGenerateEntity is required for advanced-chat snapshot replay, "
f"workflow_run_id={workflow_run.id}, generate_entity_type={type(generate_entity).__name__}"
)
if not generate_entity.conversation_id:
raise AssertionError(
f"conversation_id is required for advanced-chat snapshot replay, workflow_run_id={workflow_run.id}"
)
message_context = _get_message_context_by_conversation(
session_maker,
conversation_id=generate_entity.conversation_id,
workflow_run_id=workflow_run.id,
)
else:
# Compatibility fallback for non-suspended snapshot requests. This app-scoped lookup is not optimal;
# a dedicated index or stronger lookup key would be preferable.
message_context = _get_message_context_by_app(
session_maker,
app_id=app_id,
workflow_run_id=workflow_run.id,
)
node_snapshots = node_execution_repo.get_execution_snapshots_by_workflow_run(
tenant_id=tenant_id,
app_id=app_id,
@ -175,19 +205,68 @@ def build_workflow_event_stream(
return _generate()
def _get_message_context(session_maker: sessionmaker[Session], workflow_run_id: str) -> MessageContext | None:
def _get_message_context_by_conversation(
session_maker: sessionmaker[Session],
*,
conversation_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a paused or suspended Advanced Chat snapshot message by conversation and workflow run.
Use this exact lookup after recovering ``conversation_id`` from persisted resumption context. Its predicates match
``message_workflow_run_id_idx``.
"""
with session_maker() as session:
stmt = select(Message).where(Message.workflow_run_id == workflow_run_id).order_by(desc(Message.created_at))
stmt = (
select(Message)
.where(
Message.conversation_id == conversation_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
)
message = session.scalar(stmt)
if message is None:
return None
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
return _to_message_context(message)
def _get_message_context_by_app(
session_maker: sessionmaker[Session],
*,
app_id: str,
workflow_run_id: str,
) -> MessageContext | None:
"""Look up a non-suspended or running Advanced Chat reconnect snapshot by app and workflow run.
This compatibility path applies only when no resumption context is expected. The app-scoped query is not optimal;
a dedicated index or stronger lookup key would be preferable.
"""
with session_maker() as session:
stmt = (
select(Message)
.where(
Message.app_id == app_id,
Message.workflow_run_id == workflow_run_id,
)
.order_by(desc(Message.created_at))
.limit(1)
)
message = session.scalar(stmt)
if message is None:
return None
return _to_message_context(message)
def _to_message_context(message: Message) -> MessageContext:
created_at = int(message.created_at.timestamp()) if message.created_at else 0
return MessageContext(
conversation_id=message.conversation_id,
message_id=message.id,
created_at=created_at,
answer=message.answer,
)
def _load_resumption_context(pause_entity: WorkflowPauseEntity | None) -> WorkflowResumptionContext | None:

View File

@ -13,9 +13,13 @@ import pytest
from sqlalchemy.orm import Session, sessionmaker
from core.app.app_config.entities import WorkflowUIBasedAppConfig
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.task_entities import StreamEvent
from core.app.layers.pause_state_persist_layer import WorkflowResumptionContext, _WorkflowGenerateEntityWrapper
from core.app.layers.pause_state_persist_layer import (
WorkflowResumptionContext,
_AdvancedChatAppGenerateEntityWrapper,
_WorkflowGenerateEntityWrapper,
)
from core.workflow.human_input_policy import FormDisposition, HumanInputSurface
from core.workflow.nodes.human_input.entities import SelectInputConfig, StringListSource
from core.workflow.nodes.human_input.enums import ValueSourceType
@ -251,6 +255,34 @@ def _build_resumption_context_additional(task_id: str) -> WorkflowResumptionCont
)
def _build_advanced_chat_resumption_context(conversation_id: str | None) -> WorkflowResumptionContext:
app_config = WorkflowUIBasedAppConfig(
tenant_id="tenant-1",
app_id="app-1",
app_mode=AppMode.ADVANCED_CHAT,
workflow_id="workflow-1",
)
generate_entity = AdvancedChatAppGenerateEntity(
task_id="task-ctx",
app_config=app_config,
inputs={},
files=[],
user_id="user-1",
stream=True,
invoke_from=InvokeFrom.EXPLORE,
call_depth=0,
conversation_id=conversation_id,
workflow_run_id="run-1",
query="hello",
)
runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0)
wrapper = _AdvancedChatAppGenerateEntityWrapper(entity=generate_entity)
return WorkflowResumptionContext(
generate_entity=wrapper,
serialized_graph_runtime_state=runtime_state.dumps(),
)
class _SessionContext:
def __init__(self, session: Any) -> None:
self._session = session
@ -327,19 +359,69 @@ class _PauseEntity(WorkflowPauseEntity):
return []
def test_get_message_context_should_return_none_when_no_message() -> None:
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)
# Act
result = service_module._get_message_context(cast(sessionmaker[Session], session_maker), "run-1")
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
# Assert
assert result is None
def test_get_message_context_should_default_created_at_to_zero_when_message_has_no_timestamp() -> None:
def test_get_message_context_by_conversation_should_scope_and_bound_message_lookup() -> None:
# Arrange
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
# Act
service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
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")
def test_get_message_context_by_app_should_scope_and_bound_compatibility_lookup() -> None:
# Arrange
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
# Act
service_module._get_message_context_by_app(
cast(sessionmaker[Session], session_maker),
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")
def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp() -> None:
# Arrange
message = SimpleNamespace(
id="msg-1",
@ -351,7 +433,11 @@ def test_get_message_context_should_default_created_at_to_zero_when_message_has_
session_maker = _SessionMaker(session)
# Act
result = service_module._get_message_context(cast(sessionmaker[Session], session_maker), "run-1")
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
# Assert
assert result is not None
@ -559,9 +645,14 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
topic = _Topic(_StaticSubscription())
workflow_run_repo = SimpleNamespace(get_workflow_pause=MagicMock())
pause_entity = _PauseEntity(state=b"state")
resumption_context = _build_advanced_chat_resumption_context(conversation_id="conv-1")
call_order: list[str] = []
workflow_run_repo = SimpleNamespace(
get_workflow_pause=MagicMock(side_effect=lambda _run_id: call_order.append("pause") or pause_entity)
)
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
factory = SimpleNamespace(
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
@ -569,12 +660,19 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
)
monkeypatch.setattr(service_module, "DifyAPIRepositoryFactory", factory)
monkeypatch.setattr(service_module.MessageGenerator, "get_response_topic", MagicMock(return_value=topic))
message_context_lookup = MagicMock(side_effect=lambda *_args, **_kwargs: call_order.append("message") or None)
app_lookup = MagicMock(return_value=None)
monkeypatch.setattr(
service_module,
"_get_message_context",
MagicMock(return_value=MessageContext("conv-1", "msg-1", 1700000000)),
"_get_message_context_by_conversation",
message_context_lookup,
)
monkeypatch.setattr(service_module, "_get_message_context_by_app", app_lookup)
monkeypatch.setattr(
service_module,
"_load_resumption_context",
MagicMock(side_effect=lambda _pause_entity: call_order.append("state") or resumption_context),
)
monkeypatch.setattr(service_module, "_load_resumption_context", MagicMock(return_value=None))
buffer_state = BufferState(
queue=queue.Queue(),
stop_event=Event(),
@ -589,6 +687,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
"_build_snapshot_events",
MagicMock(return_value=[{"event": StreamEvent.WORKFLOW_FINISHED, "task_id": "task-1"}]),
)
session_maker = MagicMock()
# Act
events = list(
@ -597,7 +696,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=MagicMock(),
session_maker=session_maker,
)
)
@ -609,6 +708,119 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
node_repo.get_execution_snapshots_by_workflow_run.assert_called_once()
called_kwargs = node_repo.get_execution_snapshots_by_workflow_run.call_args.kwargs
assert called_kwargs["workflow_run_id"] == "run-1"
assert call_order == ["pause", "state", "message"]
message_context_lookup.assert_called_once_with(
session_maker,
conversation_id="conv-1",
workflow_run_id="run-1",
)
app_lookup.assert_not_called()
@pytest.mark.parametrize(
("resumption_context", "expected_error"),
[
pytest.param(None, "WorkflowResumptionContext.*workflow_run_id=run-1", id="missing-state"),
pytest.param(
_build_resumption_context_additional(task_id="task-ctx"),
"AdvancedChatAppGenerateEntity.*workflow_run_id=run-1",
id="wrong-entity-type",
),
pytest.param(
_build_advanced_chat_resumption_context(conversation_id=None),
"conversation_id.*workflow_run_id=run-1",
id="missing-conversation-id",
),
pytest.param(
_build_advanced_chat_resumption_context(conversation_id=""),
"conversation_id.*workflow_run_id=run-1",
id="empty-conversation-id",
),
],
)
def test_build_advanced_chat_snapshot_requires_conversation_context(
monkeypatch: pytest.MonkeyPatch,
resumption_context: WorkflowResumptionContext | None,
expected_error: str,
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
pause_entity = _PauseEntity(state=b"state")
workflow_run_repo = SimpleNamespace(get_workflow_pause=MagicMock(return_value=pause_entity))
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
factory = SimpleNamespace(
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
create_api_workflow_node_execution_repository=MagicMock(return_value=node_repo),
)
monkeypatch.setattr(service_module, "DifyAPIRepositoryFactory", factory)
monkeypatch.setattr(service_module.MessageGenerator, "get_response_topic", MagicMock())
monkeypatch.setattr(service_module, "_load_resumption_context", MagicMock(return_value=resumption_context))
conversation_lookup = MagicMock(return_value=None)
app_lookup = MagicMock(return_value=None)
monkeypatch.setattr(
service_module,
"_get_message_context_by_conversation",
conversation_lookup,
)
monkeypatch.setattr(service_module, "_get_message_context_by_app", app_lookup)
# Act / Assert
with pytest.raises(AssertionError, match=expected_error):
build_workflow_event_stream(
app_mode=AppMode.ADVANCED_CHAT,
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=MagicMock(),
)
conversation_lookup.assert_not_called()
app_lookup.assert_not_called()
def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Arrange
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
workflow_run_repo = SimpleNamespace(get_workflow_pause=MagicMock())
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
factory = SimpleNamespace(
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
create_api_workflow_node_execution_repository=MagicMock(return_value=node_repo),
)
monkeypatch.setattr(service_module, "DifyAPIRepositoryFactory", factory)
monkeypatch.setattr(service_module.MessageGenerator, "get_response_topic", MagicMock())
load_resumption_context = MagicMock(return_value=None)
monkeypatch.setattr(service_module, "_load_resumption_context", load_resumption_context)
conversation_lookup = MagicMock(return_value=None)
app_lookup = MagicMock(return_value=MessageContext("conv-1", "msg-1", 1700000000))
monkeypatch.setattr(
service_module,
"_get_message_context_by_conversation",
conversation_lookup,
)
monkeypatch.setattr(service_module, "_get_message_context_by_app", app_lookup)
session_maker = MagicMock()
# Act
event_stream = build_workflow_event_stream(
app_mode=AppMode.ADVANCED_CHAT,
workflow_run=workflow_run,
tenant_id="tenant-1",
app_id="app-1",
session_maker=session_maker,
)
# Assert
assert event_stream is not None
workflow_run_repo.get_workflow_pause.assert_not_called()
load_resumption_context.assert_called_once_with(None)
conversation_lookup.assert_not_called()
app_lookup.assert_called_once_with(
session_maker,
app_id="app-1",
workflow_run_id="run-1",
)
def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_idle_timeout(

View File

@ -13,9 +13,12 @@ import pytest
from sqlalchemy.orm import Session, sessionmaker
from core.app.app_config.entities import WorkflowUIBasedAppConfig
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.task_entities import StreamEvent
from core.app.layers.pause_state_persist_layer import WorkflowResumptionContext, _WorkflowGenerateEntityWrapper
from core.app.layers.pause_state_persist_layer import (
WorkflowResumptionContext,
_WorkflowGenerateEntityWrapper,
)
from graphon.enums import WorkflowExecutionStatus
from graphon.runtime import GraphRuntimeState, VariablePool
from models.enums import CreatorUserRole
@ -147,15 +150,21 @@ class _PauseEntity(WorkflowPauseEntity):
class TestWorkflowEventSnapshotHelpers:
def test_get_message_context_should_return_none_when_no_message(self) -> None:
def test_get_message_context_by_conversation_should_return_none_when_no_message(self) -> None:
session = SimpleNamespace(scalar=MagicMock(return_value=None))
session_maker = _SessionMaker(session)
result = service_module._get_message_context(cast(sessionmaker[Session], session_maker), "run-1")
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
assert result is None
def test_get_message_context_should_default_created_at_to_zero_when_message_has_no_timestamp(self) -> None:
def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp(
self,
) -> None:
message = SimpleNamespace(
id="msg-1",
conversation_id="conv-1",
@ -165,7 +174,11 @@ class TestWorkflowEventSnapshotHelpers:
session = SimpleNamespace(scalar=MagicMock(return_value=message))
session_maker = _SessionMaker(session)
result = service_module._get_message_context(cast(sessionmaker[Session], session_maker), "run-1")
result = service_module._get_message_context_by_conversation(
cast(sessionmaker[Session], session_maker),
conversation_id="conv-1",
workflow_run_id="run-1",
)
assert result is not None
assert result.created_at == 0
@ -324,9 +337,10 @@ class TestBuildWorkflowEventStream:
self,
monkeypatch: pytest.MonkeyPatch,
) -> None:
workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.RUNNING)
workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.PAUSED)
topic = _Topic(_StaticSubscription())
workflow_run_repo = SimpleNamespace(get_workflow_pause=MagicMock())
pause_entity = _PauseEntity(state=b"state")
workflow_run_repo = SimpleNamespace(get_workflow_pause=MagicMock(return_value=pause_entity))
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
factory = SimpleNamespace(
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
@ -336,10 +350,18 @@ class TestBuildWorkflowEventStream:
monkeypatch.setattr(service_module.MessageGenerator, "get_response_topic", MagicMock(return_value=topic))
monkeypatch.setattr(
service_module,
"_get_message_context",
"_get_message_context_by_conversation",
MagicMock(return_value=MessageContext("conv-1", "msg-1", 1700000000)),
)
monkeypatch.setattr(service_module, "_load_resumption_context", MagicMock(return_value=None))
generate_entity = AdvancedChatAppGenerateEntity.model_construct(conversation_id="conv-1")
resumption_context = SimpleNamespace(
get_generate_entity=MagicMock(return_value=generate_entity),
)
monkeypatch.setattr(
service_module,
"_load_resumption_context",
MagicMock(return_value=resumption_context),
)
buffer_state = BufferState(
queue=queue.Queue(),
stop_event=Event(),