mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-11 12:17:45 +02:00
[Haik]: message persistance done
This commit is contained in:
@@ -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"))
|
||||
@@ -66,7 +66,8 @@ AnyEvent = Annotated[
|
||||
Union[
|
||||
AgentStatusEvent, AgentMessageEvent,
|
||||
StreamStartEvent, StreamDeltaEvent, StreamEndEvent,
|
||||
BranchSwitchedEvent, AgentClosedEvent, BrowserCardAddedEvent,
|
||||
BranchSwitchedEvent, AgentClosedEvent,
|
||||
BrowserCardAddedEvent,
|
||||
],
|
||||
Field(discriminator="event"),
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user