From c31a649c8054f448f8dcdab2b2c81a9ce8cbe534 Mon Sep 17 00:00:00 2001 From: haikdc Date: Thu, 2 Apr 2026 19:43:41 -0700 Subject: [PATCH] [Haik]: message persistance done --- backend/core/Agent/Agent.py | 24 +++-- .../helpers/handle_assistant_message.py | 46 +++++----- .../Agent/run_agent_loop/run_agent_loop.py | 87 +++++++++++++------ backend/core/events/events.py | 3 +- 4 files changed, 108 insertions(+), 52 deletions(-) diff --git a/backend/core/Agent/Agent.py b/backend/core/Agent/Agent.py index f2c2ce56..1bd6a4a6 100644 --- a/backend/core/Agent/Agent.py +++ b/backend/core/Agent/Agent.py @@ -1,7 +1,8 @@ +import asyncio +import logging +import os from copy import deepcopy from uuid import uuid4 -import asyncio -import os from claude_agent_sdk import ClaudeAgentOptions from pydantic import BaseModel, Field, InstanceOf @@ -19,6 +20,9 @@ from backend.core.events.events import ( os.environ.setdefault("CLAUDE_CODE_STREAM_CLOSE_TIMEOUT", "3600000") +logger = logging.getLogger(__name__) + + class Agent(BaseModel): model: str mode: str @@ -58,11 +62,19 @@ class Agent(BaseModel): if self.on_event: await self.on_event(event) + @typechecked + async def _handle_event(self, event: AnyEvent) -> None: + """Internal event handler that updates Agent state and forwards to on_event.""" + if isinstance(event, AgentStatusEvent): + self.status = event.status # type: ignore[assignment] + if self.on_event: + await self.on_event(event) + @typechecked async def send_message(self, msg: Message) -> None: async with self.lock: if self.task is not None and not self.task.done(): - print("[Agent.send_message] Agent is already running") + logger.warning("[Agent.send_message] Agent %s is already running", self.session_id) return await self.emit(AgentMessageEvent( @@ -76,10 +88,12 @@ class Agent(BaseModel): self.messages.append(msg) self.task = asyncio.create_task(run_agent_loop( - msg=msg, + prompt_msg=msg.to_prompt(), + messages=self.messages, options=self.config, + session_id=self.session_id, branch_id=self.branch_id, - emit=self.on_event, + emit=self._handle_event, )) @typechecked diff --git a/backend/core/Agent/run_agent_loop/helpers/handle_assistant_message.py b/backend/core/Agent/run_agent_loop/helpers/handle_assistant_message.py index 1a73c271..403314e0 100644 --- a/backend/core/Agent/run_agent_loop/helpers/handle_assistant_message.py +++ b/backend/core/Agent/run_agent_loop/helpers/handle_assistant_message.py @@ -1,27 +1,33 @@ -from backend.core.shared_structs.agent.Message.Message import Message +from backend.core.shared_structs.agent.Message.Message import ( + AnyMessage, AssistantMessage as AssistantMsg, ToolCallMessage, +) from backend.core.shared_structs.agent.Message.agent_outputs import ToolCallContent from backend.core.events.events import EventCallback, AgentMessageEvent -from claude_agent_sdk.types import TextBlock, ToolUseBlock, AssistantMessage +from claude_agent_sdk.types import TextBlock, ToolUseBlock, AssistantMessage as SDKAssistantMessage from typeguard import typechecked -from typing import List, Optional +from typing import List, Optional, Tuple from uuid import uuid4 +HandleAssistantResult = Tuple[Optional[str], List[str], dict, List[AnyMessage]] + @typechecked async def handle_assistant_message( session_id: str, branch_id: str, - message: AssistantMessage, - stream_text_msg_id: str, + message: SDKAssistantMessage, + stream_text_msg_id: Optional[str], stream_tool_ids: List[str], emit: Optional[EventCallback] = None, -): +) -> HandleAssistantResult: """Handle an assistant message from the Claude Agent SDK. - - NOTE: Does not save the message to the Agent's MessageLog yet. - The caller (run_agent_loop / Agent) is responsible for persistence. + + Returns (reset stream_text_msg_id, reset stream_tool_ids, reset block_map, + list of Message objects created during this turn). """ content_parts: List[str] = [] tool_uses: List[ToolCallContent] = [] + created: List[AnyMessage] = [] + for block in message.content: if isinstance(block, TextBlock): content_parts.append(block.text) @@ -29,22 +35,20 @@ async def handle_assistant_message( tool_uses.append(ToolCallContent(id=block.id, tool=block.name, input=block.input)) if content_parts: - asst_msg: Message = Message( + asst_msg = AssistantMsg( id=stream_text_msg_id or uuid4().hex, - role="assistant", content="\n".join(content_parts), + content="\n".join(content_parts), branch_id=branch_id, ) + created.append(asst_msg) if emit: - await emit(AgentMessageEvent( - session_id=session_id, - message=asst_msg, - )) + await emit(AgentMessageEvent(session_id=session_id, message=asst_msg)) - for i, tu in enumerate[ToolCallContent](tool_uses): + for i, tu in enumerate(tool_uses): mid: str = stream_tool_ids[i] if i < len(stream_tool_ids) else uuid4().hex - tool_msg: Message = Message(id=mid, role="tool_call", content=tu, branch_id=branch_id) + tool_msg = ToolCallMessage(id=mid, content=tu, branch_id=branch_id) + created.append(tool_msg) if emit: - await emit(AgentMessageEvent( - session_id=session_id, - message=tool_msg, - )) \ No newline at end of file + await emit(AgentMessageEvent(session_id=session_id, message=tool_msg)) + + return None, [], {}, created \ No newline at end of file diff --git a/backend/core/Agent/run_agent_loop/run_agent_loop.py b/backend/core/Agent/run_agent_loop/run_agent_loop.py index 86c8745c..11ff8bb8 100644 --- a/backend/core/Agent/run_agent_loop/run_agent_loop.py +++ b/backend/core/Agent/run_agent_loop/run_agent_loop.py @@ -1,3 +1,6 @@ +import asyncio +import logging + from typeguard import typechecked from claude_agent_sdk import ( @@ -6,20 +9,32 @@ from claude_agent_sdk import ( from claude_agent_sdk.types import StreamEvent from backend.core.Agent.run_agent_loop.helpers.handle_stream_event import handle_stream_event from backend.core.Agent.run_agent_loop.helpers.handle_assistant_message import handle_assistant_message -from backend.core.shared_structs.agent.Message.Message import Message, PromptMsgDict -from backend.core.events.events import EventCallback +from backend.core.shared_structs.agent.Message.Message import ( + SystemMessage, PromptMsgDict, +) +from backend.core.shared_structs.agent.MessageLog import MessageLog +from backend.core.events.events import ( + EventCallback, AgentStatusEvent, AgentMessageEvent, +) from typing import List, Dict, Optional +logger = logging.getLogger(__name__) + + @typechecked async def run_agent_loop( - msg: Message, - options: ClaudeAgentOptions | None = None, - branch_id: str | None = None, + prompt_msg: PromptMsgDict, + messages: MessageLog, + options: ClaudeAgentOptions, + session_id: str, + branch_id: str = "main", emit: Optional[EventCallback] = None, ): - """Run the Claude Agent SDK query loop for a session.""" + """Run the Claude Agent SDK query loop for a session. - prompt_msg: PromptMsgDict = msg.to_prompt() + Streams events to the frontend via `emit` and persists every + assistant / tool_call message into `messages`. + """ async def prompt_stream(): yield prompt_msg @@ -28,26 +43,48 @@ async def run_agent_loop( stream_tool_msg_ids_ordered: List[str] = [] stream_block_index_map: Dict[int, str] = {} - async for message in query(prompt=prompt_stream(), options=options): + try: + async for message in query(prompt=prompt_stream(), options=options): - if isinstance(message, StreamEvent): - stream_text_msg_id = await handle_stream_event( - session_id=options.session_id, - event=message.event, - stream_text_msg_id=stream_text_msg_id, - stream_tool_ids=stream_tool_msg_ids_ordered, - block_map=stream_block_index_map, - emit=emit, - ) - - elif isinstance(message, AssistantMessage): - stream_text_msg_id, stream_tool_msg_ids_ordered, stream_block_index_map = ( - await handle_assistant_message( - session_id=options.session_id, - branch_id=branch_id, - message=message, + if isinstance(message, StreamEvent): + stream_text_msg_id = await handle_stream_event( + session_id=session_id, + event=message.event, stream_text_msg_id=stream_text_msg_id, stream_tool_ids=stream_tool_msg_ids_ordered, + block_map=stream_block_index_map, emit=emit, ) - ) \ No newline at end of file + + elif isinstance(message, AssistantMessage): + stream_text_msg_id, stream_tool_msg_ids_ordered, stream_block_index_map, created = ( + await handle_assistant_message( + session_id=session_id, + branch_id=branch_id, + message=message, + stream_text_msg_id=stream_text_msg_id, + stream_tool_ids=stream_tool_msg_ids_ordered, + emit=emit, + ) + ) + for msg in created: + messages.append(msg) + + if emit: + await emit(AgentStatusEvent(session_id=session_id, status="completed")) + + except asyncio.CancelledError: + if emit: + await emit(AgentStatusEvent(session_id=session_id, status="stopped")) + raise + + except Exception as e: + logger.exception("Agent %s error: %s", session_id, e) + error_msg = SystemMessage( + content=f"Error: {e}", + branch_id=branch_id, + ) + messages.append(error_msg) + if emit: + await emit(AgentMessageEvent(session_id=session_id, message=error_msg)) + await emit(AgentStatusEvent(session_id=session_id, status="error")) \ No newline at end of file diff --git a/backend/core/events/events.py b/backend/core/events/events.py index ab975811..371e5ac0 100644 --- a/backend/core/events/events.py +++ b/backend/core/events/events.py @@ -66,7 +66,8 @@ AnyEvent = Annotated[ Union[ AgentStatusEvent, AgentMessageEvent, StreamStartEvent, StreamDeltaEvent, StreamEndEvent, - BranchSwitchedEvent, AgentClosedEvent, BrowserCardAddedEvent, + BranchSwitchedEvent, AgentClosedEvent, + BrowserCardAddedEvent, ], Field(discriminator="event"), ]