mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-11 12:17:45 +02:00
[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:
@@ -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]]
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+11
-7
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user