diff --git a/backend/apps/agents/agent_manager.py b/backend/apps/agents/agent_manager.py index a6d816b5..be5b4f15 100644 --- a/backend/apps/agents/agent_manager.py +++ b/backend/apps/agents/agent_manager.py @@ -58,6 +58,7 @@ from backend.apps.agents.manager.streaming import stop_hook as stop_hook_mod from backend.apps.agents.manager.streaming import stream_event from backend.apps.agents.manager.streaming import assistant_message from backend.apps.agents.manager.streaming import result_message +from backend.apps.agents.manager.streaming.LivePartial import LivePartial from backend.apps.agents.manager.streaming.upsert_message import upsert_message from backend.apps.agents.manager.prompt.system_prompt import compose_turn_system_prompt from backend.apps.agents.tools.web import should_register_web_mcp @@ -99,7 +100,7 @@ class AgentManager(SessionLifecycleMixin, MessagingMixin): # Live mirror of the in-flight streamed assistant text per session, so a # stop can persist the partial reply instantly instead of waiting out the # multi-second SDK teardown the cancel handler sits behind. - self._live_partial: dict[str, dict] = {} + self._live_partial: Dict[str, LivePartial] = {} async def _build_mcp_servers( self, @@ -1909,8 +1910,8 @@ class AgentManager(SessionLifecycleMixin, MessagingMixin): live = self._live_partial.pop(session.id, None) if not live: return False - text = live.get("text") or "" - msg_id = live.get("msg_id") + text = live.text or "" + msg_id = live.msg_id if not msg_id or not text.strip(): return False if any(getattr(m, "id", None) == msg_id for m in session.messages): @@ -1919,7 +1920,7 @@ class AgentManager(SessionLifecycleMixin, MessagingMixin): id=msg_id, role="assistant", content=text, - branch_id=live.get("branch_id") or session.active_branch_id, + branch_id=live.branch_id or session.active_branch_id, ) upsert_message(session, partial) try: diff --git a/backend/apps/agents/manager/streaming/LivePartial.py b/backend/apps/agents/manager/streaming/LivePartial.py new file mode 100644 index 00000000..94b5e042 --- /dev/null +++ b/backend/apps/agents/manager/streaming/LivePartial.py @@ -0,0 +1,15 @@ +"""The in-flight streamed assistant text for one session, mirrored off the stream so a stop can +commit the partial reply instantly instead of waiting out the SDK teardown. A fixed-shape +record, so it's a model, not a dict.""" + +from typing import Optional + +from pydantic import BaseModel, ConfigDict + + +class LivePartial(BaseModel): + model_config = ConfigDict(validate_assignment=True) + + msg_id: Optional[str] = None + text: str = "" + branch_id: Optional[str] = None diff --git a/backend/apps/agents/manager/streaming/assistant_message.py b/backend/apps/agents/manager/streaming/assistant_message.py index 0373eb4f..c82bcd81 100644 --- a/backend/apps/agents/manager/streaming/assistant_message.py +++ b/backend/apps/agents/manager/streaming/assistant_message.py @@ -5,7 +5,7 @@ Lifted out of the agent loop; mutates the passed TurnState / ThinkingState by re through the manager's live-partial mirror + session registry, exactly as it did inline.""" import asyncio -from typing import Optional +from typing import Dict, Optional from uuid import uuid4 from typeguard import typechecked @@ -14,6 +14,7 @@ from backend.apps.agents.core.models import AgentSession, Message from backend.apps.agents.core.ws_manager import ws_manager from backend.apps.agents.manager.streaming.state import ThinkingState, TurnState from backend.apps.agents.manager.streaming.upsert_message import upsert_message +from backend.apps.agents.manager.streaming.LivePartial import LivePartial from backend.apps.agents.manager.streaming import thinking as thinking_mod try: @@ -30,8 +31,8 @@ async def handle_assistant_message( session_id: str, turn: TurnState, thinking: ThinkingState, - live_partial: dict, - sessions: dict, + live_partial: Dict[str, LivePartial], + sessions: Dict[str, AgentSession], ) -> None: content_parts = [] new_thinking_parts = [] diff --git a/backend/apps/agents/manager/streaming/result_message.py b/backend/apps/agents/manager/streaming/result_message.py index daee18c5..a428fbfd 100644 --- a/backend/apps/agents/manager/streaming/result_message.py +++ b/backend/apps/agents/manager/streaming/result_message.py @@ -6,7 +6,7 @@ inline. resolved_model / api_type / global_settings are the loop's per-run confi import asyncio import logging -from typing import Optional +from typing import Dict, Optional from typeguard import typechecked @@ -30,7 +30,7 @@ async def handle_result_message( session_id: str, turn: TurnState, thinking: ThinkingState, - sessions: dict, + sessions: Dict[str, AgentSession], resolved_model: object, api_type: Optional[str], global_settings: object, diff --git a/backend/apps/agents/manager/streaming/stream_event.py b/backend/apps/agents/manager/streaming/stream_event.py index 144b56f9..d300c0a6 100644 --- a/backend/apps/agents/manager/streaming/stream_event.py +++ b/backend/apps/agents/manager/streaming/stream_event.py @@ -5,6 +5,7 @@ writes the manager's live-partial mirror, exactly as it did inline.""" import time from datetime import datetime +from typing import Dict from uuid import uuid4 from typeguard import typechecked @@ -12,6 +13,7 @@ from typeguard import typechecked from backend.apps.agents.core.models import AgentSession from backend.apps.agents.core.ws_manager import ws_manager from backend.apps.agents.manager.streaming.state import ThinkingState, TurnState +from backend.apps.agents.manager.streaming.LivePartial import LivePartial try: from claude_agent_sdk.types import StreamEvent @@ -26,7 +28,7 @@ async def handle_stream_event( session_id: str, turn: TurnState, thinking: ThinkingState, - live_partial: dict, + live_partial: Dict[str, LivePartial], ) -> None: event = message.event event_type = event.get("type") @@ -111,11 +113,11 @@ async def handle_stream_event( text_chunk = delta.get("text", "") turn.assistant_text_chars += len(text_chunk) turn.stream_text_accum += text_chunk - live_partial[session_id] = { - "msg_id": turn.stream_text_msg_id, - "text": turn.stream_text_accum, - "branch_id": session.active_branch_id, - } + live_partial[session_id] = LivePartial( + msg_id=turn.stream_text_msg_id, + text=turn.stream_text_accum, + branch_id=session.active_branch_id, + ) await ws_manager.send_to_session(session_id, "agent:stream_delta", { "session_id": session_id, "message_id": msg_id, diff --git a/backend/tests/test_stream_event.py b/backend/tests/test_stream_event.py index 93737c4b..7926be1f 100644 --- a/backend/tests/test_stream_event.py +++ b/backend/tests/test_stream_event.py @@ -44,7 +44,7 @@ async def test_text_delta_accumulates_and_mirrors_live_partial(): session, session.id, turn, thinking, lp) assert turn.stream_text_accum == "Hello" assert turn.assistant_text_chars == 5 - assert lp[session.id]["text"] == "Hello" # the live-partial mirror the manager reads on resume + assert lp[session.id].text == "Hello" # the live-partial mirror the manager reads on resume @pytest.mark.asyncio