diff --git a/app/backend/main.py b/app/backend/main.py index 0d31060d..ca110bbb 100644 --- a/app/backend/main.py +++ b/app/backend/main.py @@ -4,6 +4,7 @@ import os from contextlib import asynccontextmanager from pathlib import Path +from ag_ui.core import EventType, StateSnapshotEvent from dotenv import find_dotenv, load_dotenv from pydantic_ai import Agent from pydantic_ai.ui import SSE_CONTENT_TYPE @@ -102,11 +103,33 @@ async def stream_chat(request: Request) -> Response: adapter = AGUIAdapter(agent=agent, run_input=run_input, accept=accept) + incoming_state = run_input.state if isinstance(run_input.state, dict) else {} + async def event_stream(): async with run_agui_stream( - adapter, toolset=toolset, deps=SkillDeps() + adapter, toolset=toolset, deps=SkillDeps(state=incoming_state) ) as stream: - async for chunk in adapter.encode_stream(stream): + # Emit a STATE_SNAPSHOT after RUN_STARTED so the client holds every + # namespace object before any STATE_DELTA patches into it. Without it, + # the first `add /rag//...` fails against a missing parent. + if incoming_state: + toolset.restore_state_snapshot(incoming_state) + snapshot = StateSnapshotEvent( + type=EventType.STATE_SNAPSHOT, + snapshot=toolset.build_state_snapshot(), + ) + + async def with_state_snapshot(): + emitted = False + async for event in stream: + yield event + if not emitted and getattr(event, "type", None) == ( + EventType.RUN_STARTED + ): + yield snapshot + emitted = True + + async for chunk in adapter.encode_stream(with_state_snapshot()): yield chunk return StreamingResponse( @@ -216,7 +239,7 @@ async def visualize_chunk(request: Request) -> JSONResponse: @asynccontextmanager -async def lifespan(app: Starlette): +async def lifespan(_app: Starlette): """Shut down the cached HaikuRAG client cleanly on app exit. Awaits any in-flight background vacuum tasks and closes the LanceDB diff --git a/app/frontend/components/Chat.tsx b/app/frontend/components/Chat.tsx index a2f2f6b5..b1d14fed 100644 --- a/app/frontend/components/Chat.tsx +++ b/app/frontend/components/Chat.tsx @@ -438,13 +438,12 @@ function ChatContentInner({ useEffect(() => { if (agent.messages.length > 0) return; const session = getSession(sessionId); - if (!session) return; - if (session.ragState) { - agent.setState({ - [AGUI_STATE_KEY]: normalizeRAGState(session.ragState), - }); - } - if (session.messages.length > 0) { + // Seed the namespaced AG-UI state so the backend's first STATE_DELTA + // (e.g. add /rag/searches/...) has a namespace object to patch into. + agent.setState({ + [AGUI_STATE_KEY]: normalizeRAGState(session?.ragState), + }); + if (session && session.messages.length > 0) { // biome-ignore lint/suspicious/noExplicitAny: AG-UI Message type is a broad union agent.setMessages(session.messages as any[]); }