[Haik]: ckpt, event callback done, now gonna refactor hella shit, then gonna work on the connection manager and future bridge

This commit is contained in:
haikdc
2026-04-02 17:47:00 -07:00
parent e209f7be39
commit d0f8cd94a6
8 changed files with 242 additions and 60 deletions
+38 -14
View File
@@ -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(),
})
@@ -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)
if emit:
await emit(AgentMessageEvent(
session_id=session_id,
message=tool_msg,
))
@@ -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
@@ -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
)
@@ -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]]
+50 -15
View File
@@ -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")
-1
View File
@@ -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
@@ -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