diff --git a/backend/app/gateway/services.py b/backend/app/gateway/services.py index b03b8bd79..c444719e1 100644 --- a/backend/app/gateway/services.py +++ b/backend/app/gateway/services.py @@ -18,6 +18,7 @@ from typing import Any from fastapi import HTTPException, Request from langchain_core.messages import BaseMessage from langchain_core.messages.utils import convert_to_messages +from langgraph.types import Command from app.gateway.deps import get_checkpointer, get_run_context, get_run_manager, get_stream_bridge from app.gateway.internal_auth import INTERNAL_SYSTEM_ROLE, get_trusted_internal_owner_user_id @@ -457,7 +458,11 @@ async def start_run( logger.warning("Failed to upsert thread_meta for %s (non-fatal)", sanitize_log_param(thread_id)) agent_factory = resolve_agent_factory(body.assistant_id) - graph_input = normalize_input(body.input) + command = getattr(body, "command", None) + if command and command.get("resume") is not None: + graph_input = Command(resume=command["resume"]) + else: + graph_input = normalize_input(body.input) config = build_run_config(thread_id, body.config, body.metadata, assistant_id=body.assistant_id) await apply_checkpoint_to_run_config(config, body=body, thread_id=thread_id, request=request) diff --git a/backend/tests/test_gateway_services.py b/backend/tests/test_gateway_services.py index f47fee92c..5f3662752 100644 --- a/backend/tests/test_gateway_services.py +++ b/backend/tests/test_gateway_services.py @@ -612,6 +612,89 @@ def test_inject_authenticated_user_context_skips_internal_role(): assert config["context"]["user_id"] == "channel-user-7" +async def _capture_start_run_graph_input(body): + from types import SimpleNamespace + from unittest.mock import patch + + from langgraph.checkpoint.memory import InMemorySaver + from langgraph.store.memory import InMemoryStore + + from app.gateway.services import start_run + from deerflow.persistence.thread_meta.memory import MemoryThreadMetaStore + from deerflow.runtime import RunManager + from deerflow.runtime.runs.store.memory import MemoryRunStore + + run_manager = RunManager(store=MemoryRunStore()) + state = SimpleNamespace( + stream_bridge=SimpleNamespace(), + run_manager=run_manager, + checkpointer=InMemorySaver(), + store=InMemoryStore(), + run_event_store=SimpleNamespace(), + run_events_config=None, + thread_store=MemoryThreadMetaStore(InMemoryStore()), + ) + request = SimpleNamespace( + headers={}, + state=SimpleNamespace(), + app=SimpleNamespace(state=state), + ) + captured: dict[str, object] = {} + + async def fake_run_agent(*args, **kwargs): + captured["graph_input"] = kwargs["graph_input"] + + with ( + patch("app.gateway.services.resolve_agent_factory", return_value=object()), + patch("app.gateway.services.run_agent", side_effect=fake_run_agent), + ): + record = await start_run(body, "thread-command-test", request) + await record.task + + return captured["graph_input"] + + +def test_start_run_translates_resume_command_to_langgraph_command(_stub_app_config): + import asyncio + + from langgraph.types import Command + + from app.gateway.routers.thread_runs import RunCreateRequest + + graph_input = asyncio.run( + _capture_start_run_graph_input( + RunCreateRequest( + input=None, + command={"resume": {"answer": "approved"}}, + ) + ) + ) + + assert isinstance(graph_input, Command) + assert graph_input.resume == {"answer": "approved"} + + +def test_start_run_uses_normalized_input_without_command(_stub_app_config): + import asyncio + + from langchain_core.messages import HumanMessage + + from app.gateway.routers.thread_runs import RunCreateRequest + + graph_input = asyncio.run( + _capture_start_run_graph_input( + RunCreateRequest( + input={"messages": [{"role": "human", "content": "hi"}]}, + command=None, + ) + ) + ) + + assert isinstance(graph_input, dict) + assert isinstance(graph_input["messages"][0], HumanMessage) + assert graph_input["messages"][0].content == "hi" + + def test_start_run_uses_internal_owner_header_for_persistence(_stub_app_config): import asyncio from types import SimpleNamespace