[eric] agents: model the live-partial stream mirror (LivePartial) + type the streaming dict params

This commit is contained in:
ciregenz
2026-06-23 04:29:39 -07:00
parent 5629f70c11
commit f04c050b16
6 changed files with 35 additions and 16 deletions
+5 -4
View File
@@ -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:
@@ -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
@@ -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 = []
@@ -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,
@@ -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,
+1 -1
View File
@@ -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