mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 18:58:35 +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 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:
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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(),
|
||||
|
||||
Loading…
Reference in New Issue
Block a user