[Haik]: message persistance done

This commit is contained in:
haikdc
2026-04-02 19:43:41 -07:00
parent 48a4974fcf
commit c31a649c80
4 changed files with 108 additions and 52 deletions
+19 -5
View File
@@ -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
@@ -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,
))
await emit(AgentMessageEvent(session_id=session_id, message=tool_msg))
return None, [], {}, created
@@ -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,
)
)
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"))
+2 -1
View File
@@ -66,7 +66,8 @@ AnyEvent = Annotated[
Union[
AgentStatusEvent, AgentMessageEvent,
StreamStartEvent, StreamDeltaEvent, StreamEndEvent,
BranchSwitchedEvent, AgentClosedEvent, BrowserCardAddedEvent,
BranchSwitchedEvent, AgentClosedEvent,
BrowserCardAddedEvent,
],
Field(discriminator="event"),
]