mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
fix(api): scope HITL snapshot message lookup to utilize db index (#39351)
This commit is contained in:
parent
62bbc0dbeb
commit
af465fb527
@ -13,6 +13,7 @@ from sqlalchemy import desc, select
|
|||||||
from sqlalchemy.orm import Session, sessionmaker
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
from core.app.apps.message_generator import MessageGenerator
|
from core.app.apps.message_generator import MessageGenerator
|
||||||
|
from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity
|
||||||
from core.app.entities.task_entities import (
|
from core.app.entities.task_entities import (
|
||||||
HumanInputRequiredResponse,
|
HumanInputRequiredResponse,
|
||||||
MessageReplaceStreamResponse,
|
MessageReplaceStreamResponse,
|
||||||
@ -84,9 +85,6 @@ def build_workflow_event_stream(
|
|||||||
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
|
topic = MessageGenerator.get_response_topic(app_mode, workflow_run.id)
|
||||||
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
workflow_run_repo = DifyAPIRepositoryFactory.create_api_workflow_run_repository(session_maker)
|
||||||
node_execution_repo = DifyAPIRepositoryFactory.create_api_workflow_node_execution_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
|
pause_entity: WorkflowPauseEntity | None = None
|
||||||
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
|
if workflow_run.status == WorkflowExecutionStatus.PAUSED:
|
||||||
@ -97,6 +95,38 @@ def build_workflow_event_stream(
|
|||||||
pause_entity = None
|
pause_entity = None
|
||||||
|
|
||||||
resumption_context = _load_resumption_context(pause_entity)
|
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(
|
node_snapshots = node_execution_repo.get_execution_snapshots_by_workflow_run(
|
||||||
tenant_id=tenant_id,
|
tenant_id=tenant_id,
|
||||||
app_id=app_id,
|
app_id=app_id,
|
||||||
@ -175,19 +205,68 @@ def build_workflow_event_stream(
|
|||||||
return _generate()
|
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:
|
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)
|
message = session.scalar(stmt)
|
||||||
if message is None:
|
if message is None:
|
||||||
return None
|
return None
|
||||||
created_at = int(message.created_at.timestamp()) if message.created_at else 0
|
return _to_message_context(message)
|
||||||
return MessageContext(
|
|
||||||
conversation_id=message.conversation_id,
|
|
||||||
message_id=message.id,
|
def _get_message_context_by_app(
|
||||||
created_at=created_at,
|
session_maker: sessionmaker[Session],
|
||||||
answer=message.answer,
|
*,
|
||||||
|
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:
|
def _load_resumption_context(pause_entity: WorkflowPauseEntity | None) -> WorkflowResumptionContext | None:
|
||||||
|
|||||||
@ -13,9 +13,13 @@ import pytest
|
|||||||
from sqlalchemy.orm import Session, sessionmaker
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
from core.app.app_config.entities import WorkflowUIBasedAppConfig
|
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.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.human_input_policy import FormDisposition, HumanInputSurface
|
||||||
from core.workflow.nodes.human_input.entities import SelectInputConfig, StringListSource
|
from core.workflow.nodes.human_input.entities import SelectInputConfig, StringListSource
|
||||||
from core.workflow.nodes.human_input.enums import ValueSourceType
|
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:
|
class _SessionContext:
|
||||||
def __init__(self, session: Any) -> None:
|
def __init__(self, session: Any) -> None:
|
||||||
self._session = session
|
self._session = session
|
||||||
@ -327,19 +359,69 @@ class _PauseEntity(WorkflowPauseEntity):
|
|||||||
return []
|
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
|
# Arrange
|
||||||
session = SimpleNamespace(scalar=MagicMock(return_value=None))
|
session = SimpleNamespace(scalar=MagicMock(return_value=None))
|
||||||
session_maker = _SessionMaker(session)
|
session_maker = _SessionMaker(session)
|
||||||
|
|
||||||
# Act
|
# 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
|
||||||
assert result is None
|
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
|
# Arrange
|
||||||
message = SimpleNamespace(
|
message = SimpleNamespace(
|
||||||
id="msg-1",
|
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)
|
session_maker = _SessionMaker(session)
|
||||||
|
|
||||||
# Act
|
# 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
|
||||||
assert result is not None
|
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,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
# Arrange
|
# Arrange
|
||||||
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING)
|
workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED)
|
||||||
topic = _Topic(_StaticSubscription())
|
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=[]))
|
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
|
||||||
factory = SimpleNamespace(
|
factory = SimpleNamespace(
|
||||||
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
|
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, "DifyAPIRepositoryFactory", factory)
|
||||||
monkeypatch.setattr(service_module.MessageGenerator, "get_response_topic", MagicMock(return_value=topic))
|
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(
|
monkeypatch.setattr(
|
||||||
service_module,
|
service_module,
|
||||||
"_get_message_context",
|
"_get_message_context_by_conversation",
|
||||||
MagicMock(return_value=MessageContext("conv-1", "msg-1", 1700000000)),
|
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(
|
buffer_state = BufferState(
|
||||||
queue=queue.Queue(),
|
queue=queue.Queue(),
|
||||||
stop_event=Event(),
|
stop_event=Event(),
|
||||||
@ -589,6 +687,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
|
|||||||
"_build_snapshot_events",
|
"_build_snapshot_events",
|
||||||
MagicMock(return_value=[{"event": StreamEvent.WORKFLOW_FINISHED, "task_id": "task-1"}]),
|
MagicMock(return_value=[{"event": StreamEvent.WORKFLOW_FINISHED, "task_id": "task-1"}]),
|
||||||
)
|
)
|
||||||
|
session_maker = MagicMock()
|
||||||
|
|
||||||
# Act
|
# Act
|
||||||
events = list(
|
events = list(
|
||||||
@ -597,7 +696,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even
|
|||||||
workflow_run=workflow_run,
|
workflow_run=workflow_run,
|
||||||
tenant_id="tenant-1",
|
tenant_id="tenant-1",
|
||||||
app_id="app-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()
|
node_repo.get_execution_snapshots_by_workflow_run.assert_called_once()
|
||||||
called_kwargs = node_repo.get_execution_snapshots_by_workflow_run.call_args.kwargs
|
called_kwargs = node_repo.get_execution_snapshots_by_workflow_run.call_args.kwargs
|
||||||
assert called_kwargs["workflow_run_id"] == "run-1"
|
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(
|
def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_idle_timeout(
|
||||||
|
|||||||
@ -13,9 +13,12 @@ import pytest
|
|||||||
from sqlalchemy.orm import Session, sessionmaker
|
from sqlalchemy.orm import Session, sessionmaker
|
||||||
|
|
||||||
from core.app.app_config.entities import WorkflowUIBasedAppConfig
|
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.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.enums import WorkflowExecutionStatus
|
||||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||||
from models.enums import CreatorUserRole
|
from models.enums import CreatorUserRole
|
||||||
@ -147,15 +150,21 @@ class _PauseEntity(WorkflowPauseEntity):
|
|||||||
|
|
||||||
|
|
||||||
class TestWorkflowEventSnapshotHelpers:
|
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 = SimpleNamespace(scalar=MagicMock(return_value=None))
|
||||||
session_maker = _SessionMaker(session)
|
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
|
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(
|
message = SimpleNamespace(
|
||||||
id="msg-1",
|
id="msg-1",
|
||||||
conversation_id="conv-1",
|
conversation_id="conv-1",
|
||||||
@ -165,7 +174,11 @@ class TestWorkflowEventSnapshotHelpers:
|
|||||||
session = SimpleNamespace(scalar=MagicMock(return_value=message))
|
session = SimpleNamespace(scalar=MagicMock(return_value=message))
|
||||||
session_maker = _SessionMaker(session)
|
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 is not None
|
||||||
assert result.created_at == 0
|
assert result.created_at == 0
|
||||||
@ -324,9 +337,10 @@ class TestBuildWorkflowEventStream:
|
|||||||
self,
|
self,
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.RUNNING)
|
workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.PAUSED)
|
||||||
topic = _Topic(_StaticSubscription())
|
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=[]))
|
node_repo = SimpleNamespace(get_execution_snapshots_by_workflow_run=MagicMock(return_value=[]))
|
||||||
factory = SimpleNamespace(
|
factory = SimpleNamespace(
|
||||||
create_api_workflow_run_repository=MagicMock(return_value=workflow_run_repo),
|
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.MessageGenerator, "get_response_topic", MagicMock(return_value=topic))
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
service_module,
|
service_module,
|
||||||
"_get_message_context",
|
"_get_message_context_by_conversation",
|
||||||
MagicMock(return_value=MessageContext("conv-1", "msg-1", 1700000000)),
|
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(
|
buffer_state = BufferState(
|
||||||
queue=queue.Queue(),
|
queue=queue.Queue(),
|
||||||
stop_event=Event(),
|
stop_event=Event(),
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user