diff --git a/backend/apps/HaikFix/Agent/Agent.py b/backend/apps/HaikFix/Agent/Agent.py index 69b88042..557ed96e 100644 --- a/backend/apps/HaikFix/Agent/Agent.py +++ b/backend/apps/HaikFix/Agent/Agent.py @@ -1,5 +1,3 @@ -# TODO: NON HAIK DEPS: ws_manager - from copy import deepcopy from uuid import uuid4 import asyncio @@ -7,20 +5,20 @@ import os from claude_agent_sdk import ClaudeAgentOptions from pydantic import BaseModel, Field, InstanceOf -from typing import List, Literal, Optional +from typing import Awaitable, Callable, List, Literal, Optional from typeguard import typechecked -from backend.apps.agents.manager.ws_manager import ws_manager from backend.apps.HaikFix.Agent.run_agent_loop.run_agent_loop import run_agent_loop from backend.apps.HaikFix.Agent.shared_structs.Message.Message import Message from backend.apps.HaikFix.Agent.shared_structs.ApprovalRequest import ApprovalRequest from backend.apps.HaikFix.Agent.shared_structs.MessageLog import MessageLog +from backend.apps.HaikFix.Agent.shared_structs.events import ( + AnyEvent, AgentSnapshot, AgentStatusEvent, AgentMessageEvent, + EventCallback, +) os.environ.setdefault("CLAUDE_CODE_STREAM_CLOSE_TIMEOUT", "3600000") - -# NOTE and TODO: we shld remove the ws streaming from this class bc it conflicts with the browser agent - class Agent(BaseModel): model: str mode: str @@ -37,9 +35,29 @@ class Agent(BaseModel): sub_branches: List["Agent"] = Field(default_factory=list) parent_id: Optional[str] = None + on_event: Optional[EventCallback] = Field(default=None, exclude=True) + task: Optional[asyncio.Task] = None lock: InstanceOf[asyncio.Lock] = Field(default_factory=asyncio.Lock) + @typechecked + def snapshot(self) -> AgentSnapshot: + return AgentSnapshot( + session_id=self.session_id, + model=self.model, + mode=self.mode, + status=self.status, + branch_id=self.branch_id, + parent_id=self.parent_id, + messages=self.messages, + pending_approvals=self.pending_approvals, + ) + + @typechecked + async def emit(self, event: AnyEvent) -> None: + if self.on_event: + await self.on_event(event) + @typechecked async def send_message(self, msg: Message) -> None: async with self.lock: @@ -47,15 +65,21 @@ class Agent(BaseModel): print("[Agent.send_message] Agent is already running") return - await ws_manager.emit_message(self.session_id, msg) + await self.emit(AgentMessageEvent( + session_id=self.session_id, + message=msg, + )) self.status = "running" - await ws_manager.emit_status(self.session_id, "running") + await self.emit(AgentStatusEvent( + session_id=self.session_id, status="running", + )) self.messages.append(msg) self.task = asyncio.create_task(run_agent_loop( msg=msg, options=self.config, branch_id=self.branch_id, + emit=self.on_event, )) @typechecked @@ -70,13 +94,12 @@ class Agent(BaseModel): except asyncio.CancelledError: pass - for req in list[ApprovalRequest](self.pending_approvals): - ws_manager.resolve_approval(req.id, {"behavior": "deny", "message": "Agent stopped"}) self.pending_approvals = [] - self.status = "stopped" - await ws_manager.emit_status(self.session_id, "stopped") - + await self.emit(AgentStatusEvent( + session_id=self.session_id, status="stopped", + )) + @typechecked def branch(self, at_message_id: str) -> "Agent": branch_id = uuid4().hex @@ -89,6 +112,7 @@ class Agent(BaseModel): "messages": MessageLog(messages=branched_messages), "sub_agents": [], "pending_approvals": [], + "on_event": self.on_event, "task": None, "lock": asyncio.Lock(), }) diff --git a/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_assistant_message.py b/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_assistant_message.py index 29f161c1..fa4a7147 100644 --- a/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_assistant_message.py +++ b/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_assistant_message.py @@ -1,9 +1,9 @@ from backend.apps.HaikFix.Agent.shared_structs.Message.Message import Message from backend.apps.HaikFix.Agent.shared_structs.Message.agent_outputs import ToolCallContent +from backend.apps.HaikFix.Agent.shared_structs.events import EventCallback, AgentMessageEvent from claude_agent_sdk.types import TextBlock, ToolUseBlock, AssistantMessage -from backend.apps.agents.manager.ws_manager import ws_manager from typeguard import typechecked -from typing import List +from typing import List, Optional from uuid import uuid4 @typechecked @@ -13,10 +13,13 @@ async def handle_assistant_message( message: AssistantMessage, stream_text_msg_id: str, stream_tool_ids: List[str], + emit: Optional[EventCallback] = None, ): - """Handle an assistant message from the Claude Agent SDK.""" - # NOTE: Does not acctually save the message to the session. - # TODO: Save the message to the session -Haik + """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. + """ content_parts: List[str] = [] tool_uses: List[ToolCallContent] = [] for block in message.content: @@ -31,9 +34,17 @@ async def handle_assistant_message( role="assistant", content="\n".join(content_parts), branch_id=branch_id, ) - await ws_manager.emit_message(session_id, asst_msg) + if emit: + await emit(AgentMessageEvent( + session_id=session_id, + message=asst_msg, + )) for i, tu in enumerate[ToolCallContent](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) - await ws_manager.emit_message(session_id, tool_msg) \ No newline at end of file + if emit: + await emit(AgentMessageEvent( + session_id=session_id, + message=tool_msg, + )) \ No newline at end of file diff --git a/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_stream_event.py b/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_stream_event.py index 11e80d6c..9b7bc9f3 100644 --- a/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_stream_event.py +++ b/backend/apps/HaikFix/Agent/run_agent_loop/helpers/handle_stream_event.py @@ -1,9 +1,14 @@ -from backend.apps.agents.manager.ws_manager import ws_manager -from typing import Any, Dict, Optional +from typing import Any, Awaitable, Callable, Dict, Optional from typeguard import typechecked from uuid import uuid4 +from backend.apps.HaikFix.Agent.shared_structs.events import ( + AnyEvent, StreamStartEvent, StreamDeltaEvent, StreamEndEvent, +) +EventCallback = Callable[[AnyEvent], Awaitable[None]] + +# NOTE: if we wanna, we could abstract this into helper functions for each event type @typechecked async def handle_stream_event( session_id: str, @@ -11,6 +16,7 @@ async def handle_stream_event( stream_text_msg_id: Optional[str], stream_tool_ids: list[str], block_map: dict[int, str], + emit: Optional[EventCallback] = None, ) -> str | None: """Process a single StreamEvent and return the (possibly updated) text msg id.""" assert "type" in event, "Stream event missing 'type'" @@ -26,7 +32,10 @@ async def handle_stream_event( if block_type == "text": if stream_text_msg_id is None: stream_text_msg_id = uuid4().hex - await ws_manager.emit_stream_start(session_id, stream_text_msg_id, "assistant") + if emit: + await emit(StreamStartEvent( + session_id=session_id, message_id=stream_text_msg_id, role="assistant", + )) block_map[index] = stream_text_msg_id elif block_type == "tool_use": assert "name" in block, "tool_use content_block missing 'name'" @@ -34,7 +43,11 @@ async def handle_stream_event( tool_msg_id: str = uuid4().hex stream_tool_ids.append(tool_msg_id) block_map[index] = tool_msg_id - await ws_manager.emit_stream_start(session_id, tool_msg_id, "tool_call", tool_name=tool_name) + if emit: + await emit(StreamStartEvent( + session_id=session_id, message_id=tool_msg_id, + role="tool_call", tool_name=tool_name, + )) elif event_type == "content_block_delta": assert "index" in event, "content_block_delta missing 'index'" @@ -48,20 +61,32 @@ async def handle_stream_event( if delta_type == "text_delta": assert "text" in delta, "text_delta missing 'text'" text: str = delta["text"] - await ws_manager.emit_stream_delta(session_id, msg_id, text) + if emit: + await emit(StreamDeltaEvent( + session_id=session_id, message_id=msg_id, delta=text, + )) elif delta_type == "input_json_delta": assert "partial_json" in delta, "input_json_delta missing 'partial_json'" partial_json: str = delta["partial_json"] - await ws_manager.emit_stream_delta(session_id, msg_id, partial_json) + if emit: + await emit(StreamDeltaEvent( + session_id=session_id, message_id=msg_id, delta=partial_json, + )) elif event_type == "content_block_stop": assert "index" in event, "content_block_stop missing 'index'" msg_id: Optional[str] = block_map.get(event["index"]) if msg_id and msg_id != stream_text_msg_id: - await ws_manager.emit_stream_end(session_id, msg_id) + if emit: + await emit(StreamEndEvent( + session_id=session_id, message_id=msg_id, + )) elif event_type == "message_stop": if stream_text_msg_id: - await ws_manager.emit_stream_end(session_id, stream_text_msg_id) + if emit: + await emit(StreamEndEvent( + session_id=session_id, message_id=stream_text_msg_id, + )) return stream_text_msg_id \ No newline at end of file diff --git a/backend/apps/HaikFix/Agent/run_agent_loop/run_agent_loop.py b/backend/apps/HaikFix/Agent/run_agent_loop/run_agent_loop.py index 61a4be23..8eaaf5a0 100644 --- a/backend/apps/HaikFix/Agent/run_agent_loop/run_agent_loop.py +++ b/backend/apps/HaikFix/Agent/run_agent_loop/run_agent_loop.py @@ -1,4 +1,3 @@ -import logging from typeguard import typechecked from claude_agent_sdk import ( @@ -7,16 +6,16 @@ from claude_agent_sdk import ( from claude_agent_sdk.types import StreamEvent from backend.apps.HaikFix.Agent.run_agent_loop.helpers.handle_stream_event import handle_stream_event from backend.apps.HaikFix.Agent.run_agent_loop.helpers.handle_assistant_message import handle_assistant_message - -# from backend.apps.agents.manager.HaikFix.shared_structs.Message.sub_types import ImageChunk, TextChunk from backend.apps.HaikFix.Agent.shared_structs.Message.Message import Message, PromptMsgDict -from typing import List, Dict, Literal, Union, Optional +from backend.apps.HaikFix.Agent.shared_structs.events import EventCallback +from typing import List, Dict, Optional @typechecked async def run_agent_loop( msg: Message, options: ClaudeAgentOptions | None = None, branch_id: str | None = None, + emit: Optional[EventCallback] = None, ): """Run the Claude Agent SDK query loop for a session.""" @@ -28,8 +27,6 @@ async def run_agent_loop( stream_text_msg_id: Optional[str] = None stream_tool_msg_ids_ordered: List[str] = [] stream_block_index_map: Dict[int, str] = {} - _turn_number: int = 0 - _first_event: bool = True async for message in query(prompt=prompt_stream(), options=options): @@ -40,6 +37,7 @@ async def run_agent_loop( 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): @@ -50,6 +48,6 @@ async def run_agent_loop( message=message, stream_text_msg_id=stream_text_msg_id, stream_tool_ids=stream_tool_msg_ids_ordered, + emit=emit, ) - ) - _turn_number += 1 \ No newline at end of file + ) \ No newline at end of file diff --git a/backend/apps/HaikFix/Agent/shared_structs/events.py b/backend/apps/HaikFix/Agent/shared_structs/events.py new file mode 100644 index 00000000..308c1bfd --- /dev/null +++ b/backend/apps/HaikFix/Agent/shared_structs/events.py @@ -0,0 +1,86 @@ +from pydantic import BaseModel, Field +from typing import Annotated, List, Literal, Optional, Union, Callable, Awaitable + +from backend.apps.HaikFix.Agent.shared_structs.Message.Message import AnyMessage +from backend.apps.HaikFix.Agent.shared_structs.MessageLog import MessageLog +from backend.apps.HaikFix.Agent.shared_structs.ApprovalRequest import ApprovalRequest +from backend.apps.dashboards.models import BrowserCardPosition + + +class AgentSnapshot(BaseModel): + """Wire-format representation of an Agent — no runtime fields (task, lock, on_event).""" + session_id: str + model: str + mode: str + status: str + branch_id: str = "main" + parent_id: Optional[str] = None + messages: MessageLog = Field(default_factory=MessageLog) + pending_approvals: List[ApprovalRequest] = Field(default_factory=list) + sub_agents: list = Field(default_factory=list) + sub_branches: list = Field(default_factory=list) + + +class AgentStatusEvent(BaseModel): + event: Literal["agent:status"] = "agent:status" + session_id: str + status: str + session: Optional[AgentSnapshot] = None + + +class AgentMessageEvent(BaseModel): + event: Literal["agent:message"] = "agent:message" + session_id: str + message: AnyMessage + + +class StreamStartEvent(BaseModel): + event: Literal["agent:stream_start"] = "agent:stream_start" + session_id: str + message_id: str + role: str + tool_name: Optional[str] = None + + +class StreamDeltaEvent(BaseModel): + event: Literal["agent:stream_delta"] = "agent:stream_delta" + session_id: str + message_id: str + delta: str + + +class StreamEndEvent(BaseModel): + event: Literal["agent:stream_end"] = "agent:stream_end" + session_id: str + message_id: str + + +class BranchSwitchedEvent(BaseModel): + event: Literal["agent:branch_switched"] = "agent:branch_switched" + session_id: str + active_branch_id: str + + +class AgentClosedEvent(BaseModel): + event: Literal["agent:closed"] = "agent:closed" + session_id: str + status: str + closed_at: str + + +class BrowserCardAddedEvent(BaseModel): + event: Literal["dashboard:browser_card_added"] = "dashboard:browser_card_added" + dashboard_id: str + browser_card: BrowserCardPosition + + +AnyEvent = Annotated[ + Union[ + AgentStatusEvent, AgentMessageEvent, + StreamStartEvent, StreamDeltaEvent, StreamEndEvent, + BranchSwitchedEvent, AgentClosedEvent, BrowserCardAddedEvent, + ], + Field(discriminator="event"), +] + +EventCallback = Callable[[AnyEvent], Awaitable[None]] \ No newline at end of file diff --git a/backend/apps/HaikFix/agents.py b/backend/apps/HaikFix/agents.py index 2410795b..958fcfa5 100644 --- a/backend/apps/HaikFix/agents.py +++ b/backend/apps/HaikFix/agents.py @@ -2,10 +2,14 @@ Endpoints operate directly on a module-level sessions dict and the Agent class. No manager layer — Agent already encapsulates its own runtime state. + +ws_manager is used ONLY in this file — the Agent class and its internals +communicate via the on_event callback, never importing ws_manager directly. """ from contextlib import asynccontextmanager from datetime import datetime +from typeguard import typechecked from uuid import uuid4 import asyncio @@ -16,13 +20,25 @@ from typing import Optional, List from backend.config.Apps import SubApp from backend.apps.HaikFix.Agent.Agent import Agent from backend.apps.HaikFix.Agent.shared_structs.Message.Message import UserMessage +from backend.apps.HaikFix.Agent.shared_structs.events import ( + AnyEvent, AgentStatusEvent, AgentClosedEvent, BranchSwitchedEvent, + EventCallback, +) from backend.apps.HaikFix import session_store -from backend.apps.agents.manager.ws_manager import ws_manager # TODO: move to HaikFix +from backend.apps.agents.manager.ws_manager import ws_manager from claude_agent_sdk import ClaudeAgentOptions SESSIONS: dict[str, Agent] = {} + +@typechecked +def p_make_session_emitter(session_id: str) -> EventCallback: + """Create an event callback that routes typed events to ws_manager for a session.""" + async def emit(event: AnyEvent) -> None: + await ws_manager.send_to_session(session_id, event.event, event.model_dump(mode="json")) + return emit + def get_agent(session_id: str) -> Agent: agent: Optional[Agent] = SESSIONS.get(session_id) if not agent: @@ -45,6 +61,7 @@ async def agents_lifespan(): data.pop("lock", None) agent: Agent = Agent(**data) agent.status = "stopped" + agent.on_event = p_make_session_emitter(agent.session_id) SESSIONS[agent.session_id] = agent session_store.delete(sid) except Exception as e: @@ -90,7 +107,6 @@ class LaunchBody(BaseModel): @agents.router.post("/launch") async def launch(body: LaunchBody) -> dict: # TODO: build ClaudeAgentOptions from body once prompt/options builder exists - # For now this is a placeholder — the config must be constructed here agent: Agent = Agent( model=body.model, mode=body.mode, @@ -100,9 +116,13 @@ async def launch(body: LaunchBody) -> dict: max_turns=body.max_turns, ), ) + agent.on_event = p_make_session_emitter(agent.session_id) SESSIONS[agent.session_id] = agent - await ws_manager.emit_status(agent.session_id, "running", agent) - return {"session_id": agent.session_id, "session": agent.model_dump(mode="json")} + await agent._emit(AgentStatusEvent( + session_id=agent.session_id, status="running", + session=agent.snapshot(), + )) + return {"session_id": agent.session_id, "session": agent.snapshot().model_dump(mode="json")} class UpdateBody(BaseModel): @@ -116,7 +136,10 @@ async def update_session(session_id: str, body: UpdateBody) -> dict: agent.name = body.name if body.system_prompt is not None: agent.config.system_prompt = body.system_prompt - await ws_manager.emit_status(session_id, agent.status, agent) + await agent._emit(AgentStatusEvent( + session_id=session_id, status=agent.status, + session=agent.snapshot(), + )) return {"ok": True} @@ -216,7 +239,9 @@ class SwitchBranchBody(BaseModel): async def switch_branch(session_id: str, body: SwitchBranchBody) -> dict: agent: Agent = get_agent(session_id) agent.branch_id = body.branch_id - await ws_manager.emit_branch_switched(session_id, body.branch_id) + await agent.emit(BranchSwitchedEvent( + session_id=session_id, active_branch_id=body.branch_id, + )) return {"ok": True} @@ -230,11 +255,15 @@ async def close_session(session_id: str) -> dict: if not agent: raise HTTPException(status_code=404, detail="Session not found") await agent.stop_agent() + closed_at: str = datetime.now().isoformat() data: dict = agent.model_dump(mode="json") data["search_text"] = session_store.build_search_text(data) - data["closed_at"] = datetime.now().isoformat() + data["closed_at"] = closed_at session_store.save(session_id, data) - await ws_manager.emit_closed(session_id, agent) + await agent.emit(AgentClosedEvent( + session_id=session_id, status=agent.status, + closed_at=closed_at, + )) return {"ok": True} @@ -251,10 +280,14 @@ async def resume_session(session_id: str) -> dict: data.pop("closed_at", None) agent: Agent = Agent(**data) agent.status = "stopped" + agent.on_event = p_make_session_emitter(agent.session_id) SESSIONS[agent.session_id] = agent session_store.delete(session_id) - await ws_manager.emit_status(session_id, agent.status, agent) - return {"session": agent.model_dump(mode="json")} + await agent.emit(AgentStatusEvent( + session_id=session_id, status=agent.status, + session=agent.snapshot(), + )) + return {"session": agent.snapshot().model_dump(mode="json")} @agents.router.post("/SESSIONS/{session_id}/duplicate") @@ -266,9 +299,7 @@ async def duplicate_session(session_id: str, body: dict = {}) -> dict: raise HTTPException(status_code=404, detail="Session not found") data.pop("task", None) data.pop("lock", None) - new_source: Agent = Agent(**data) - source = new_source - assert source is not None, "Source agent not found" + source = Agent(**data) clone: Agent = source.model_copy(deep=True) clone.session_id = uuid4().hex @@ -277,9 +308,13 @@ async def duplicate_session(session_id: str, body: dict = {}) -> dict: clone.lock = asyncio.Lock() clone.pending_approvals = [] clone.sub_agents = [] + clone.on_event = p_make_session_emitter(clone.session_id) SESSIONS[clone.session_id] = clone - await ws_manager.emit_status(clone.session_id, clone.status, clone) - return {"session": clone.model_dump(mode="json")} + await clone.emit(AgentStatusEvent( + session_id=clone.session_id, status=clone.status, + session=clone.snapshot(), + )) + return {"session": clone.snapshot().model_dump(mode="json")} @agents.router.get("/history") diff --git a/backend/apps/HaikFix/session_store.py b/backend/apps/HaikFix/session_store.py index a506f507..56c868bb 100644 --- a/backend/apps/HaikFix/session_store.py +++ b/backend/apps/HaikFix/session_store.py @@ -5,7 +5,6 @@ The agents subapp calls these functions during close/resume/startup/shutdown. """ import json -import logging import os from typing import List, Tuple, Optional diff --git a/backend/apps/HaikFix/tools/make_builtin_toolkit/open_swarm_toolkits/browser_toolkit/make_browser_delegation_toolkit/handlers/utils/create_browser_card.py b/backend/apps/HaikFix/tools/make_builtin_toolkit/open_swarm_toolkits/browser_toolkit/make_browser_delegation_toolkit/handlers/utils/create_browser_card.py index c242536f..b5b0dc0b 100644 --- a/backend/apps/HaikFix/tools/make_builtin_toolkit/open_swarm_toolkits/browser_toolkit/make_browser_delegation_toolkit/handlers/utils/create_browser_card.py +++ b/backend/apps/HaikFix/tools/make_builtin_toolkit/open_swarm_toolkits/browser_toolkit/make_browser_delegation_toolkit/handlers/utils/create_browser_card.py @@ -2,14 +2,17 @@ from datetime import datetime from uuid import uuid4 from typeguard import typechecked +from typing import Optional -# NOTE: Legacy dependancy. TODO: fix this shit cuh -from backend.apps.agents.manager.ws_manager import ws_manager from backend.apps.dashboards.dashboards import _load, _save from backend.apps.dashboards.models import BrowserCardPosition, BrowserTab +from backend.apps.HaikFix.Agent.shared_structs.events import EventCallback, BrowserCardAddedEvent @typechecked -async def create_browser_card(dashboard_id: str) -> str: +async def create_browser_card( + dashboard_id: str, + emit: Optional[EventCallback] = None, +) -> str: dashboard = _load(dashboard_id) browser_id = f"browser-{uuid4().hex[:8]}" tab_id = f"tab-{uuid4().hex[:8]}" @@ -21,8 +24,9 @@ async def create_browser_card(dashboard_id: str) -> str: dashboard.layout.browser_cards[browser_id] = card dashboard.updated_at = datetime.now() _save(dashboard) - await ws_manager.broadcast_global("dashboard:browser_card_added", { - "dashboard_id": dashboard_id, - "browser_card": card.model_dump(mode="json"), - }) + if emit: + await emit(BrowserCardAddedEvent( + dashboard_id=dashboard_id, + browser_card=card, + )) return browser_id