mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-30 13:34:50 +02:00
[eric] delete ~3.6k lines of dead code (old unused agent runtime, ghost fields, orphan scripts). all 660 tests still pass, no functional
change.
This commit is contained in:
@@ -1,440 +0,0 @@
|
||||
"""Owned agent loop — replaces claude_agent_sdk's query() function.
|
||||
|
||||
Generalizes the pattern from browser_agent.py (lines 243-334) into a
|
||||
provider-agnostic, streaming, HITL-aware tool-use loop.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Callable, Awaitable
|
||||
from uuid import uuid4
|
||||
|
||||
from backend.apps.agents.providers.base import (
|
||||
BaseProvider, ContentBlock, ModelResponse, ProviderMessage,
|
||||
StreamEvent, ToolCall, ToolSchema,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Type aliases for callbacks
|
||||
ToolExecutor = Callable[[str, dict], Awaitable[list[dict]]]
|
||||
# hitl_handler(tool_name, tool_input) -> (approved, updated_input_or_None)
|
||||
HITLHandler = Callable[[str, dict], Awaitable[tuple[bool, dict | None]]]
|
||||
# ws_emitter(event_type, data) -> None
|
||||
WSEmitter = Callable[[str, dict], Awaitable[None]]
|
||||
|
||||
|
||||
class AgentLoop:
|
||||
"""Provider-agnostic agent loop with streaming and HITL support.
|
||||
|
||||
The loop:
|
||||
1. Sends user message to the model
|
||||
2. Streams the response (emitting WebSocket events)
|
||||
3. If the model requests tool use:
|
||||
a. For each tool call: check HITL permission → execute → collect result
|
||||
b. Append tool results → go to step 2
|
||||
4. If the model stops (end_turn/max_tokens): done
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_id: str,
|
||||
provider: BaseProvider,
|
||||
model: str,
|
||||
system_prompt: str | None,
|
||||
tools: list[ToolSchema],
|
||||
tool_executor: ToolExecutor,
|
||||
hitl_handler: HITLHandler,
|
||||
ws_emitter: WSEmitter,
|
||||
max_turns: int | None = None,
|
||||
cwd: str | None = None,
|
||||
):
|
||||
self.session_id = session_id
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
self.system_prompt = system_prompt
|
||||
self.tools = tools
|
||||
self.tool_executor = tool_executor
|
||||
self.hitl_handler = hitl_handler
|
||||
self.ws_emitter = ws_emitter
|
||||
self.max_turns = max_turns
|
||||
self.cwd = cwd
|
||||
|
||||
# Conversation history in provider-agnostic format
|
||||
self.messages: list[ProviderMessage] = []
|
||||
|
||||
# Token tracking
|
||||
self.total_input_tokens = 0
|
||||
self.total_output_tokens = 0
|
||||
|
||||
async def run(self, user_content: Any) -> None:
|
||||
"""Run the agent loop for a single user turn."""
|
||||
# Append user message
|
||||
user_msg = self.provider.format_user_message(user_content)
|
||||
self.messages.append(user_msg)
|
||||
|
||||
turn = 0
|
||||
while True:
|
||||
if self.max_turns and turn >= self.max_turns:
|
||||
logger.info(f"Agent {self.session_id}: max turns ({self.max_turns}) reached")
|
||||
break
|
||||
turn += 1
|
||||
|
||||
# Stream the model response and collect it
|
||||
response = await self._stream_and_collect()
|
||||
|
||||
# Track usage
|
||||
self.total_input_tokens += response.usage.get("input_tokens", 0)
|
||||
self.total_output_tokens += response.usage.get("output_tokens", 0)
|
||||
|
||||
# Append assistant message to conversation history
|
||||
assistant_msg = self.provider.format_assistant_message(response)
|
||||
self.messages.append(assistant_msg)
|
||||
|
||||
# If no tool use, we're done
|
||||
if response.stop_reason != "tool_use":
|
||||
break
|
||||
|
||||
# Execute tools
|
||||
tool_results = await self._execute_tools(response)
|
||||
if not tool_results:
|
||||
break
|
||||
|
||||
# Append tool results
|
||||
self.messages.append(ProviderMessage(role="tool_result", content=tool_results))
|
||||
|
||||
async def _stream_and_collect(self) -> ModelResponse:
|
||||
"""Stream model output, emit WebSocket events, collect full response."""
|
||||
collected_content: list[ContentBlock] = []
|
||||
collected_usage: dict[str, int] = {}
|
||||
stop_reason = "end_turn"
|
||||
|
||||
# Track streaming state for WS emissions
|
||||
stream_text_msg_id: str | None = None
|
||||
stream_tool_msg_ids: dict[int, str] = {} # block index -> msg_id
|
||||
block_index_map: dict[int, str] = {} # block index -> msg_id
|
||||
|
||||
# Buffers for collecting content
|
||||
text_buffers: dict[int, str] = {}
|
||||
json_buffers: dict[int, str] = {}
|
||||
tool_names: dict[int, str] = {}
|
||||
tool_ids: dict[int, str] = {}
|
||||
block_types: dict[int, str] = {}
|
||||
# Wall-clock start time per content block (server-side stamps).
|
||||
# Used to compute elapsed_ms for thinking blocks so the persisted
|
||||
# ThinkingBubble can show the duration after streaming ends.
|
||||
block_start_ts: dict[int, float] = {}
|
||||
thinking_total_ms: int = 0
|
||||
thinking_total_chars: int = 0
|
||||
|
||||
async for event in self.provider.stream_message(
|
||||
model=self.model,
|
||||
system=self.system_prompt,
|
||||
messages=self.messages,
|
||||
tools=self.tools,
|
||||
):
|
||||
if event.type == "content_block_start":
|
||||
if event.block_type == "text":
|
||||
if stream_text_msg_id is None:
|
||||
stream_text_msg_id = uuid4().hex
|
||||
await self.ws_emitter("agent:stream_start", {
|
||||
"message_id": stream_text_msg_id,
|
||||
"role": "assistant",
|
||||
})
|
||||
block_index_map[event.index] = stream_text_msg_id
|
||||
block_types[event.index] = "text"
|
||||
text_buffers[event.index] = ""
|
||||
|
||||
elif event.block_type == "tool_use":
|
||||
tool_msg_id = uuid4().hex
|
||||
stream_tool_msg_ids[event.index] = tool_msg_id
|
||||
block_index_map[event.index] = tool_msg_id
|
||||
block_types[event.index] = "tool_use"
|
||||
tool_names[event.index] = event.tool_name
|
||||
tool_ids[event.index] = event.tool_id
|
||||
json_buffers[event.index] = ""
|
||||
|
||||
await self.ws_emitter("agent:stream_start", {
|
||||
"message_id": tool_msg_id,
|
||||
"role": "tool_call",
|
||||
"tool_name": event.tool_name,
|
||||
})
|
||||
|
||||
elif event.block_type == "thinking":
|
||||
# Extended-thinking content block. Emit a distinct
|
||||
# WS stream with role="thinking" so the frontend
|
||||
# renders the live ThinkingBubble pill (rising
|
||||
# token counter, auto-collapse on first text). Each
|
||||
# thinking block gets its own message id — multiple
|
||||
# interleaved thinking/text blocks remain
|
||||
# individually addressable.
|
||||
thinking_msg_id = uuid4().hex
|
||||
block_index_map[event.index] = thinking_msg_id
|
||||
block_types[event.index] = "thinking"
|
||||
text_buffers[event.index] = ""
|
||||
# Server-stamp the start so we can compute exact
|
||||
# elapsed_ms server-side at content_block_stop. Using
|
||||
# time.time() (not monotonic) is fine here — we only
|
||||
# subtract two values from the same clock.
|
||||
block_start_ts[event.index] = time.time()
|
||||
await self.ws_emitter("agent:stream_start", {
|
||||
"message_id": thinking_msg_id,
|
||||
"role": "thinking",
|
||||
})
|
||||
|
||||
elif event.type == "content_block_delta":
|
||||
msg_id = block_index_map.get(event.index)
|
||||
if not msg_id:
|
||||
continue
|
||||
|
||||
if event.delta_type == "text_delta":
|
||||
text_buffers.setdefault(event.index, "")
|
||||
text_buffers[event.index] += event.text
|
||||
await self.ws_emitter("agent:stream_delta", {
|
||||
"message_id": msg_id,
|
||||
"delta": event.text,
|
||||
})
|
||||
|
||||
elif event.delta_type == "input_json_delta":
|
||||
json_buffers.setdefault(event.index, "")
|
||||
json_buffers[event.index] += event.text
|
||||
await self.ws_emitter("agent:stream_delta", {
|
||||
"message_id": msg_id,
|
||||
"delta": event.text,
|
||||
})
|
||||
|
||||
elif event.delta_type == "thinking_delta":
|
||||
# Reuse the text buffer for thinking — same shape
|
||||
# (accumulated str), different sink.
|
||||
text_buffers.setdefault(event.index, "")
|
||||
text_buffers[event.index] += event.text
|
||||
await self.ws_emitter("agent:stream_delta", {
|
||||
"message_id": msg_id,
|
||||
"delta": event.text,
|
||||
})
|
||||
|
||||
elif event.type == "content_block_stop":
|
||||
msg_id = block_index_map.get(event.index)
|
||||
bt = block_types.get(event.index, "")
|
||||
|
||||
if bt == "text":
|
||||
collected_content.append(
|
||||
ContentBlock(type="text", text=text_buffers.get(event.index, ""))
|
||||
)
|
||||
elif bt == "tool_use":
|
||||
try:
|
||||
tool_input = json.loads(json_buffers.get(event.index, "{}"))
|
||||
except json.JSONDecodeError:
|
||||
tool_input = {}
|
||||
collected_content.append(ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(
|
||||
id=tool_ids.get(event.index, uuid4().hex),
|
||||
name=tool_names.get(event.index, ""),
|
||||
input=tool_input,
|
||||
),
|
||||
))
|
||||
elif bt == "thinking":
|
||||
thinking_text = text_buffers.get(event.index, "")
|
||||
collected_content.append(
|
||||
ContentBlock(type="thinking", text=thinking_text)
|
||||
)
|
||||
# Accumulate per-block duration + char count for the
|
||||
# eventual persisted Message. We sum across multiple
|
||||
# thinking blocks in the same turn so a complex
|
||||
# interleaved (think → tool → think → answer) turn
|
||||
# still reports total time spent reasoning.
|
||||
start_ts = block_start_ts.get(event.index)
|
||||
if start_ts is not None:
|
||||
thinking_total_ms += int((time.time() - start_ts) * 1000)
|
||||
thinking_total_chars += len(thinking_text)
|
||||
|
||||
# Send stream_end for tool + thinking blocks (text block
|
||||
# ends at message_stop). For thinking blocks we
|
||||
# intentionally DO NOT include per-block elapsed_ms /
|
||||
# tokens — those would freeze the pill early on a
|
||||
# multi-block turn (think → tool → think → answer)
|
||||
# showing only the first block's stats. Instead the pill
|
||||
# stays in "thinking…" until the per-turn aggregate
|
||||
# arrives via the agent:message event for this turn's
|
||||
# persisted Message(role="thinking"), which carries
|
||||
# thinking_total_ms and the heuristic-or-9Router-truthed
|
||||
# token count for the WHOLE turn.
|
||||
if msg_id and (bt == "tool_use" or bt == "thinking"):
|
||||
await self.ws_emitter("agent:stream_end", {"message_id": msg_id})
|
||||
|
||||
elif event.type == "usage":
|
||||
# Accumulate token usage from provider stream
|
||||
for k, v in event.usage.items():
|
||||
collected_usage[k] = collected_usage.get(k, 0) + v
|
||||
|
||||
elif event.type == "message_stop":
|
||||
# Check if any tool calls means stop_reason is tool_use
|
||||
has_tool_use = any(b.type == "tool_use" for b in collected_content)
|
||||
if has_tool_use:
|
||||
stop_reason = "tool_use"
|
||||
|
||||
# End text stream
|
||||
if stream_text_msg_id:
|
||||
await self.ws_emitter("agent:stream_end", {
|
||||
"message_id": stream_text_msg_id,
|
||||
})
|
||||
|
||||
# Build and emit the collected messages
|
||||
await self._emit_collected_messages(
|
||||
collected_content, stream_text_msg_id, stream_tool_msg_ids,
|
||||
thinking_elapsed_ms=thinking_total_ms,
|
||||
thinking_total_chars=thinking_total_chars,
|
||||
)
|
||||
|
||||
return ModelResponse(
|
||||
content=collected_content,
|
||||
stop_reason=stop_reason,
|
||||
usage=collected_usage,
|
||||
)
|
||||
|
||||
async def _emit_collected_messages(
|
||||
self,
|
||||
content: list[ContentBlock],
|
||||
text_msg_id: str | None,
|
||||
tool_msg_ids: dict[int, str],
|
||||
thinking_elapsed_ms: int = 0,
|
||||
thinking_total_chars: int = 0,
|
||||
) -> None:
|
||||
"""Emit finalized agent:message events for the collected response."""
|
||||
from backend.apps.agents.models import Message
|
||||
|
||||
# Emit thinking blocks (extended thinking). Persisted as their own
|
||||
# messages so a session reload still shows the reasoning trail.
|
||||
# Multiple thinking blocks per turn are concatenated into a single
|
||||
# persisted message — the streaming UI already showed each block
|
||||
# individually, this is just for the historical record.
|
||||
thinking_parts = [b.text for b in content if b.type == "thinking" and b.text]
|
||||
if thinking_parts:
|
||||
joined = "\n\n".join(thinking_parts)
|
||||
# Token count source preference:
|
||||
# 1. 9Router's stored reasoning_tokens / thoughtsTokenCount —
|
||||
# exact, available for OpenAI/Gemini/DeepSeek when routed
|
||||
# through 9Router.
|
||||
# 2. chars/3.6 heuristic — Anthropic's API doesn't expose a
|
||||
# per-block thinking-token count, so we fall back to the
|
||||
# same BPE-ish approximation we use for the live counter.
|
||||
tokens_value: int | None = None
|
||||
try:
|
||||
from backend.apps.nine_router import (
|
||||
get_latest_reasoning_tokens,
|
||||
is_running as _9r_running,
|
||||
)
|
||||
if _9r_running():
|
||||
rt = await get_latest_reasoning_tokens(model_hint=self.model)
|
||||
if rt and rt > 0:
|
||||
tokens_value = rt
|
||||
except Exception:
|
||||
# 9Router is best-effort; fall back to heuristic silently.
|
||||
pass
|
||||
if tokens_value is None and thinking_total_chars:
|
||||
tokens_value = max(1, round(thinking_total_chars / 3.6))
|
||||
|
||||
# Stamp duration + token estimate so the persisted bubble can
|
||||
# show "Thought for Ns · M tokens" on reload instead of the
|
||||
# generic "Thoughts" fallback. Use the server-side accumulated
|
||||
# times so multi-block turns aggregate correctly.
|
||||
msg = Message(
|
||||
role="thinking",
|
||||
content=joined,
|
||||
elapsed_ms=thinking_elapsed_ms or None,
|
||||
tokens=tokens_value,
|
||||
)
|
||||
await self.ws_emitter("agent:message", {
|
||||
"message": msg.model_dump(mode="json"),
|
||||
})
|
||||
|
||||
# Emit text message
|
||||
text_parts = [b.text for b in content if b.type == "text" and b.text]
|
||||
if text_parts:
|
||||
msg = Message(
|
||||
id=text_msg_id or uuid4().hex,
|
||||
role="assistant",
|
||||
content="\n".join(text_parts),
|
||||
)
|
||||
await self.ws_emitter("agent:message", {
|
||||
"message": msg.model_dump(mode="json"),
|
||||
})
|
||||
|
||||
# Emit tool call messages
|
||||
tool_blocks = [b for b in content if b.type == "tool_use" and b.tool_call]
|
||||
tool_id_list = sorted(tool_msg_ids.items(), key=lambda x: x[0])
|
||||
for i, block in enumerate(tool_blocks):
|
||||
tc = block.tool_call
|
||||
msg_id = tool_id_list[i][1] if i < len(tool_id_list) else uuid4().hex
|
||||
msg = Message(
|
||||
id=msg_id,
|
||||
role="tool_call",
|
||||
content={
|
||||
"id": tc.id,
|
||||
"tool": tc.name,
|
||||
"input": tc.input,
|
||||
},
|
||||
)
|
||||
await self.ws_emitter("agent:message", {
|
||||
"message": msg.model_dump(mode="json"),
|
||||
})
|
||||
|
||||
async def _execute_tools(self, response: ModelResponse) -> list[dict]:
|
||||
"""Execute all tool calls from a response, respecting HITL permissions.
|
||||
|
||||
Returns a list of tool result dicts formatted for the provider.
|
||||
"""
|
||||
from backend.apps.agents.models import Message
|
||||
|
||||
results = []
|
||||
for block in response.content:
|
||||
if block.type != "tool_use" or not block.tool_call:
|
||||
continue
|
||||
|
||||
tc = block.tool_call
|
||||
start_time = time.time()
|
||||
|
||||
# HITL permission check
|
||||
approved, updated_input = await self.hitl_handler(tc.name, tc.input)
|
||||
|
||||
if not approved:
|
||||
result_content = [{"type": "text", "text": "Tool use was denied by the user."}]
|
||||
else:
|
||||
tool_input = updated_input if updated_input else tc.input
|
||||
try:
|
||||
result_content = await self.tool_executor(tc.name, tool_input)
|
||||
except Exception as e:
|
||||
logger.warning(f"Tool execution error: {tc.name}: {e}")
|
||||
result_content = [{"type": "text", "text": f"Error executing {tc.name}: {e}"}]
|
||||
|
||||
elapsed_ms = int((time.time() - start_time) * 1000)
|
||||
|
||||
# Emit tool result to frontend
|
||||
result_text = ""
|
||||
for block_item in result_content:
|
||||
if isinstance(block_item, dict) and block_item.get("type") == "text":
|
||||
result_text = block_item.get("text", "")
|
||||
break
|
||||
|
||||
result_msg = Message(
|
||||
role="tool_result",
|
||||
content={
|
||||
"text": result_text[:15000] if result_text else "Done.",
|
||||
"tool_name": tc.name,
|
||||
"elapsed_ms": elapsed_ms,
|
||||
},
|
||||
)
|
||||
await self.ws_emitter("agent:message", {
|
||||
"message": result_msg.model_dump(mode="json"),
|
||||
})
|
||||
|
||||
# Format for provider
|
||||
results.append(
|
||||
self.provider.format_tool_result(tc.id, result_content)
|
||||
)
|
||||
|
||||
return results
|
||||
@@ -1598,57 +1598,20 @@ class AgentManager:
|
||||
"type": "stdio",
|
||||
}
|
||||
|
||||
# -----------------------------------------------------------------
|
||||
# openswarm-web MCP — DDG search + trafilatura fetch
|
||||
# -----------------------------------------------------------------
|
||||
# The CLI's built-in WebSearch/WebFetch wrap Anthropic's server-
|
||||
# side web_search_20250305. Verified against 9Router 0.3.60's
|
||||
# full chunk tree (grep returned zero hits for web_search,
|
||||
# googleSearch, grounding, retrieval — 9Router does NOT translate
|
||||
# WebSearch to any provider's native search tool). So for every
|
||||
# non-Anthropic primary, the CLI delegates WebSearch execution
|
||||
# back to Anthropic via ANTHROPIC_SMALL_FAST_MODEL. That path
|
||||
# needs a Claude credential; without one it fails with "no
|
||||
# credentials for provider: claude". When it succeeds it can
|
||||
# still break on Gemini 3 thinking-mode thought-signature
|
||||
# validation in subsequent turns.
|
||||
#
|
||||
# To sidestep all of that: register our own DDG-backed MCP for
|
||||
# every primary whose native Anthropic delegation is unreliable
|
||||
# or unreachable. Claude primaries (cc/ and openswarm-pro's
|
||||
# Anthropic adaptive path) keep the built-in Anthropic search
|
||||
# because it IS high-quality and works end-to-end for them.
|
||||
#
|
||||
# Free: DuckDuckGo HTML + trafilatura extraction run locally on
|
||||
# each user's machine. No API keys, no subscriptions, no rate
|
||||
# limits at per-user scale.
|
||||
# The CLI's built-in WebSearch/WebFetch wraps Anthropic's
|
||||
# web_search_20250305. For non-Claude primaries the CLI
|
||||
# delegates execution back to Anthropic via
|
||||
# ANTHROPIC_SMALL_FAST_MODEL — needs an Anthropic credential
|
||||
# or it 401s. We register our DDG-backed MCP only for users
|
||||
# with no Anthropic path; Anthropic's hosted search is
|
||||
# higher-quality so we prefer it whenever it's reachable.
|
||||
_m = _router_model_id if isinstance(_router_model_id, str) else ""
|
||||
# Decide whether to register our DDG/Gemini-grounded MCP.
|
||||
#
|
||||
# The CLI's built-in WebSearch/WebFetch wrap Anthropic's
|
||||
# server-side web_search_20250305 tool. For Claude primaries
|
||||
# it runs inline. For non-Claude primaries the CLI delegates
|
||||
# the search execution back to Anthropic via a small model
|
||||
# (ANTHROPIC_SMALL_FAST_MODEL → haiku). That delegation path
|
||||
# needs *some* Anthropic credential to reach Anthropic.
|
||||
#
|
||||
# So: if the user has ANY Anthropic path available (Claude
|
||||
# subscription via 9Router, openswarm-pro cloud proxy, or a
|
||||
# direct Anthropic API key), we prefer the built-in. It's
|
||||
# bundled into what they're already paying for and gives
|
||||
# real Anthropic-curated search results — strictly higher
|
||||
# quality than our DDG scrape. We only fall back to our MCP
|
||||
# for users with ZERO Anthropic access.
|
||||
_has_anthropic_path = (
|
||||
getattr(global_settings, "connection_mode", "own_key") == "openswarm-pro"
|
||||
or bool(getattr(global_settings, "anthropic_api_key", None))
|
||||
)
|
||||
# Check 9Router for any active connection that can serve
|
||||
# Anthropic-format requests. Both the subscription id
|
||||
# `claude` (OAuth'd Claude Code subscription) and the
|
||||
# direct-API id `anthropic` (apikey connection — which is
|
||||
# how we register OpenSwarm Pro as a Claude-compatible
|
||||
# route) satisfy this.
|
||||
# Both 9Router provider ids `claude` (subscription OAuth) and
|
||||
# `anthropic` (direct API / Pro proxy) satisfy this check.
|
||||
_9r_has_anthropic = False
|
||||
try:
|
||||
from backend.apps.nine_router import get_providers as _9r_providers
|
||||
@@ -1662,18 +1625,11 @@ class AgentManager:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# For Pro users WITHOUT a 9Router Claude/Anthropic connection
|
||||
# yet (sync not complete, or first run), the CLI's built-in
|
||||
# WebSearch delegation through 9Router would fail. Only
|
||||
# consider Anthropic reachable if 9Router can actually serve
|
||||
# the Anthropic-format request.
|
||||
# Deliberately exclude openswarm-pro from the "has anthropic
|
||||
# path" heuristic when the primary is non-Claude. Reason: if
|
||||
# a Pro user picks GPT or Gemini as their primary, we
|
||||
# shouldn't drag their WebSearch/subagent calls through our
|
||||
# Pro Anthropic pool — they're already paying for a
|
||||
# ChatGPT/Gemini subscription we can use for free. Pro still
|
||||
# kicks in when they switch the primary to a Claude model.
|
||||
# When the primary is non-Claude we deliberately don't count
|
||||
# OpenSwarm Pro as an Anthropic path — using the Pro pool for
|
||||
# WebSearch on a GPT/Gemini session would drain it for the
|
||||
# user's Claude turns. The user's GPT/Gemini subscription
|
||||
# serves their non-Claude turns at zero cost to us.
|
||||
_primary_is_claude = _m.startswith("cc/") or (
|
||||
isinstance(_router_model_id, str)
|
||||
and not _router_model_id.startswith(("cc/", "cx/", "gc/", "ag/", "gemini/"))
|
||||
@@ -2336,14 +2292,8 @@ class AgentManager:
|
||||
# "Thought signature is not valid" 400). None for providers
|
||||
# that don't use signatures.
|
||||
_turn_thought_signature: str | None = None
|
||||
# Per-turn delta baseline. session.tokens["input"]/["output"]
|
||||
# is the SDK's CUMULATIVE total across all turns (the SDK
|
||||
# reports running totals on each ResultMessage, not per-turn
|
||||
# deltas). To stamp the consolidated thinking pill with
|
||||
# *this turn's* tokens — not the cumulative session total —
|
||||
# we snapshot the cumulative values at turn start and
|
||||
# subtract them at emit time. Same for any subagent token
|
||||
# totals, which also accumulate across turns.
|
||||
# session.tokens accumulates SDK running totals across turns,
|
||||
# so subtract the turn-start baseline to get this turn's delta.
|
||||
_turn_baseline_session_in: int = 0
|
||||
_turn_baseline_session_out: int = 0
|
||||
_turn_baseline_children_in: int = 0
|
||||
@@ -2509,11 +2459,8 @@ class AgentManager:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Per-turn deltas. If baseline wasn't captured (rare
|
||||
# race: emit fired before any AssistantMessage on this
|
||||
# turn), fall back to cumulative values — better than
|
||||
# showing zero, and acceptable since this only happens
|
||||
# on degenerate empty turns.
|
||||
# Fall back to cumulative if the baseline wasn't captured
|
||||
# (degenerate empty turn — better than showing zero).
|
||||
if _turn_baseline_captured:
|
||||
_parent_in = max(0, _cum_in - _turn_baseline_session_in)
|
||||
_parent_out = max(0, _cum_out - _turn_baseline_session_out)
|
||||
@@ -2539,12 +2486,6 @@ class AgentManager:
|
||||
tokens=turn_tokens,
|
||||
input_tokens=_turn_total_tokens,
|
||||
tool_count=_turn_tool_count or None,
|
||||
# Persist Gemini thoughtSignature so we can re-attach
|
||||
# it on the next request — this is what stops Google
|
||||
# from rejecting multi-step turns with "Thought
|
||||
# signature is not valid" 400. None for any provider
|
||||
# that doesn't use signatures.
|
||||
thought_signature=_turn_thought_signature,
|
||||
)
|
||||
existing_idx = next(
|
||||
(i for i, m in enumerate(session.messages)
|
||||
@@ -2606,15 +2547,8 @@ class AgentManager:
|
||||
# + assistant text generation.
|
||||
if _turn_started_ts is None:
|
||||
_turn_started_ts = time.time()
|
||||
# Capture cumulative-token baselines at turn
|
||||
# start so the pill can stamp per-turn deltas
|
||||
# instead of session totals. Without this,
|
||||
# turn 2's pill would show turn-1 tokens +
|
||||
# turn-2 tokens combined, and turn 3's would
|
||||
# show turn-1 + turn-2 + turn-3 — making it
|
||||
# look like every turn is bigger than the
|
||||
# last and that work was being "added on top"
|
||||
# of the first pill.
|
||||
# Snapshot cumulative tokens at turn start;
|
||||
# subtracted at emit time for per-turn deltas.
|
||||
try:
|
||||
if isinstance(session.tokens, dict):
|
||||
_turn_baseline_session_in = int(session.tokens.get("input", 0) or 0)
|
||||
@@ -2887,17 +2821,11 @@ class AgentManager:
|
||||
|
||||
if content_parts:
|
||||
_asst_text = "\n".join(content_parts)
|
||||
# 9Router can deliver upstream auth failures
|
||||
# AS the assistant's reply ("Failed to
|
||||
# authenticate. API Error: 401 ... [codex/...]
|
||||
# Provided authentication token is expired").
|
||||
# When that happens, the SDK doesn't raise —
|
||||
# so our catch-all _is_auth_error path never
|
||||
# fires. Detect the pattern in the text
|
||||
# itself and substitute a friendly system
|
||||
# bubble + auth_error WS event so the user
|
||||
# gets an actionable message instead of a
|
||||
# raw 401 JSON dump in the chat.
|
||||
# 9Router sometimes returns upstream 401s as
|
||||
# the assistant reply (no SDK exception), so
|
||||
# the catch-all auth handler never fires.
|
||||
# Match the text pattern and surface a
|
||||
# friendly system bubble instead.
|
||||
_lower_text = _asst_text.lower()
|
||||
_looks_like_router_auth_error = (
|
||||
("failed to authenticate" in _lower_text and "401" in _lower_text)
|
||||
@@ -2906,8 +2834,6 @@ class AgentManager:
|
||||
or ("provided authentication token" in _lower_text and ("401" in _lower_text or "expired" in _lower_text))
|
||||
)
|
||||
if _looks_like_router_auth_error:
|
||||
# Build a friendly message keyed off the
|
||||
# provider name in the upstream error.
|
||||
if "codex/" in _lower_text or "[codex" in _lower_text:
|
||||
friendly = (
|
||||
"GPT subscription token expired. Open Settings → Models and click "
|
||||
|
||||
@@ -1,398 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Minimal stdio MCP server that exposes browser interaction tools.
|
||||
|
||||
Launched as a subprocess by the Claude Agent SDK. Proxies tool calls
|
||||
to the OpenSwarm backend via HTTP, which bridges them to the Electron
|
||||
frontend via WebSocket where the actual webview lives.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
import sys
|
||||
import os
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
from io import BytesIO
|
||||
|
||||
try:
|
||||
from PIL import Image
|
||||
HAS_PIL = True
|
||||
except ImportError:
|
||||
HAS_PIL = False
|
||||
|
||||
BACKEND_PORT = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
BACKEND_URL = f"http://127.0.0.1:{BACKEND_PORT}/api/browser/command"
|
||||
|
||||
TAB_ID_PROP = {
|
||||
"type": "string",
|
||||
"description": "Optional tab ID within the browser card. If omitted, targets the active tab.",
|
||||
}
|
||||
|
||||
TOOLS = [
|
||||
{
|
||||
"name": "BrowserScreenshot",
|
||||
"description": (
|
||||
"Capture a screenshot of the browser page. Returns the screenshot as a "
|
||||
"base64-encoded PNG image. Use this to see what is currently displayed."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID to capture. Use the ID from the selected browser card context.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
},
|
||||
"required": ["browser_id"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserGetText",
|
||||
"description": (
|
||||
"Get the visible text content of the browser page. Returns the page's "
|
||||
"innerText (up to 15000 characters)."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
},
|
||||
"required": ["browser_id"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserNavigate",
|
||||
"description": "Navigate the browser to a URL.",
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"url": {
|
||||
"type": "string",
|
||||
"description": "The URL to navigate to.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id", "url"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserClick",
|
||||
"description": (
|
||||
"Click an element in the browser page identified by a CSS selector."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"selector": {
|
||||
"type": "string",
|
||||
"description": "CSS selector of the element to click.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id", "selector"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserType",
|
||||
"description": (
|
||||
"Type text into an input element in the browser page. Clears the "
|
||||
"existing value first, then types the new text."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"selector": {
|
||||
"type": "string",
|
||||
"description": "CSS selector of the input element.",
|
||||
},
|
||||
"text": {
|
||||
"type": "string",
|
||||
"description": "The text to type.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id", "selector", "text"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserEvaluate",
|
||||
"description": (
|
||||
"Evaluate a JavaScript expression in the browser page and return the result. "
|
||||
"The expression is run via executeJavaScript on the webview."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"expression": {
|
||||
"type": "string",
|
||||
"description": "JavaScript expression to evaluate.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id", "expression"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserGetElements",
|
||||
"description": (
|
||||
"Get a list of interactive elements on the page with their CSS selectors. "
|
||||
"Returns clickable elements, inputs, links, and buttons with selector paths "
|
||||
"you can use with BrowserClick and BrowserType. Call this BEFORE attempting "
|
||||
"to click or type so you know which selectors are valid."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"selector": {
|
||||
"type": "string",
|
||||
"description": (
|
||||
"Optional CSS selector to scope the search "
|
||||
"(e.g. 'form', '#main'). Defaults to 'body'."
|
||||
),
|
||||
},
|
||||
},
|
||||
"required": ["browser_id"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserScroll",
|
||||
"description": (
|
||||
"Scroll the page up or down. Automatically finds the correct scrollable "
|
||||
"container (works on SPAs like Notion, Gmail, etc. that use nested scroll "
|
||||
"containers instead of window-level scrolling). Returns scroll position info "
|
||||
"including whether top/bottom has been reached."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"direction": {
|
||||
"type": "string",
|
||||
"enum": ["up", "down"],
|
||||
"description": "Scroll direction. Defaults to 'down'.",
|
||||
},
|
||||
"amount": {
|
||||
"type": "number",
|
||||
"description": "Pixels to scroll. Defaults to 500.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "BrowserWait",
|
||||
"description": (
|
||||
"Wait for a specified duration. Useful after navigation or actions that "
|
||||
"trigger page loads, animations, or async content rendering. "
|
||||
"Min 100ms, max 10000ms."
|
||||
),
|
||||
"inputSchema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"browser_id": {
|
||||
"type": "string",
|
||||
"description": "The browser card ID.",
|
||||
},
|
||||
"tab_id": TAB_ID_PROP,
|
||||
"milliseconds": {
|
||||
"type": "number",
|
||||
"description": "Duration to wait in milliseconds. Defaults to 1000.",
|
||||
},
|
||||
},
|
||||
"required": ["browser_id"],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def send_response(id_, result=None, error=None):
|
||||
msg = {"jsonrpc": "2.0", "id": id_}
|
||||
if error is not None:
|
||||
msg["error"] = error
|
||||
else:
|
||||
msg["result"] = result
|
||||
sys.stdout.write(json.dumps(msg) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
def send_notification(method, params=None):
|
||||
msg = {"jsonrpc": "2.0", "method": method}
|
||||
if params is not None:
|
||||
msg["params"] = params
|
||||
sys.stdout.write(json.dumps(msg) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
def call_backend(action: str, browser_id: str, params: dict | None = None, tab_id: str = "") -> dict:
|
||||
payload = json.dumps({
|
||||
"action": action,
|
||||
"browser_id": browser_id,
|
||||
"tab_id": tab_id,
|
||||
"params": params or {},
|
||||
}).encode()
|
||||
req = urllib.request.Request(
|
||||
BACKEND_URL,
|
||||
data=payload,
|
||||
headers={"Content-Type": "application/json"},
|
||||
method="POST",
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
||||
return json.loads(resp.read().decode())
|
||||
except urllib.error.HTTPError as e:
|
||||
body = e.read().decode() if e.fp else str(e)
|
||||
return {"error": f"HTTP {e.code}: {body}"}
|
||||
except Exception as e:
|
||||
return {"error": str(e)}
|
||||
|
||||
|
||||
MAX_IMAGE_B64_BYTES = 400_000
|
||||
|
||||
|
||||
def compress_screenshot(b64_png: str) -> tuple[str, str] | None:
|
||||
"""Resize and re-encode as JPEG to stay under the stdio buffer limit."""
|
||||
if not HAS_PIL:
|
||||
return None
|
||||
try:
|
||||
raw = base64.b64decode(b64_png)
|
||||
img = Image.open(BytesIO(raw))
|
||||
max_width = 1024
|
||||
if img.width > max_width:
|
||||
ratio = max_width / img.width
|
||||
img = img.resize((max_width, int(img.height * ratio)), Image.LANCZOS)
|
||||
buf = BytesIO()
|
||||
img.convert("RGB").save(buf, format="JPEG", quality=45)
|
||||
return base64.b64encode(buf.getvalue()).decode(), "image/jpeg"
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def handle_tool_call(tool_name: str, arguments: dict) -> dict:
|
||||
browser_id = arguments.get("browser_id", "")
|
||||
tab_id = arguments.get("tab_id", "")
|
||||
if not browser_id:
|
||||
return {"content": [{"type": "text", "text": "Error: browser_id is required"}], "isError": True}
|
||||
|
||||
action_map = {
|
||||
"BrowserScreenshot": "screenshot",
|
||||
"BrowserGetText": "get_text",
|
||||
"BrowserNavigate": "navigate",
|
||||
"BrowserClick": "click",
|
||||
"BrowserType": "type",
|
||||
"BrowserEvaluate": "evaluate",
|
||||
"BrowserGetElements": "get_elements",
|
||||
"BrowserScroll": "scroll",
|
||||
"BrowserWait": "wait",
|
||||
}
|
||||
action = action_map.get(tool_name)
|
||||
if not action:
|
||||
return {"content": [{"type": "text", "text": f"Unknown tool: {tool_name}"}], "isError": True}
|
||||
|
||||
params = {k: v for k, v in arguments.items() if k not in ("browser_id", "tab_id")}
|
||||
result = call_backend(action, browser_id, params, tab_id=tab_id)
|
||||
|
||||
if "error" in result:
|
||||
return {"content": [{"type": "text", "text": f"Error: {result['error']}"}], "isError": True}
|
||||
|
||||
if action == "screenshot" and result.get("image"):
|
||||
image_data = result["image"]
|
||||
mime_type = "image/png"
|
||||
|
||||
if len(image_data) > MAX_IMAGE_B64_BYTES:
|
||||
compressed = compress_screenshot(image_data)
|
||||
if compressed:
|
||||
image_data, mime_type = compressed
|
||||
|
||||
if len(image_data) > MAX_IMAGE_B64_BYTES:
|
||||
return {
|
||||
"content": [
|
||||
{"type": "text", "text": (
|
||||
f"Screenshot too large to return ({len(image_data)} bytes base64). "
|
||||
f"URL: {result.get('url', 'unknown')}. "
|
||||
"Use BrowserGetText to read the page content instead."
|
||||
)},
|
||||
],
|
||||
}
|
||||
|
||||
return {
|
||||
"content": [
|
||||
{"type": "image", "data": image_data, "mimeType": mime_type},
|
||||
{"type": "text", "text": f"Screenshot captured. URL: {result.get('url', 'unknown')}"},
|
||||
],
|
||||
}
|
||||
|
||||
text = result.get("text", result.get("data", json.dumps(result)))
|
||||
return {"content": [{"type": "text", "text": str(text)}]}
|
||||
|
||||
|
||||
def main():
|
||||
for line in sys.stdin:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
msg = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
method = msg.get("method")
|
||||
id_ = msg.get("id")
|
||||
params = msg.get("params", {})
|
||||
|
||||
if method == "initialize":
|
||||
send_response(id_, {
|
||||
"protocolVersion": "2024-11-05",
|
||||
"capabilities": {"tools": {}},
|
||||
"serverInfo": {
|
||||
"name": "openswarm-browser",
|
||||
"version": "1.0.0",
|
||||
},
|
||||
})
|
||||
elif method == "notifications/initialized":
|
||||
pass
|
||||
elif method == "tools/list":
|
||||
send_response(id_, {"tools": TOOLS})
|
||||
elif method == "tools/call":
|
||||
tool_name = params.get("name", "")
|
||||
arguments = params.get("arguments", {})
|
||||
result = handle_tool_call(tool_name, arguments)
|
||||
send_response(id_, result)
|
||||
elif method == "ping":
|
||||
send_response(id_, {})
|
||||
elif id_ is not None:
|
||||
send_response(id_, error={"code": -32601, "message": f"Method not found: {method}"})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,360 +0,0 @@
|
||||
"""Standalone MCP client manager for agent sessions.
|
||||
|
||||
Replaces claude_agent_sdk's internal MCP server management.
|
||||
One MCPClientManager instance per agent session — manages connections
|
||||
to stdio/http/sse MCP servers, discovers tools, and routes tool calls.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from contextlib import AsyncExitStack
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from backend.apps.agents.providers.base import ToolSchema
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MCPConnection:
|
||||
"""A live connection to an MCP server."""
|
||||
server_name: str
|
||||
session: Any # mcp.ClientSession
|
||||
tools: list[ToolSchema] = field(default_factory=list)
|
||||
|
||||
|
||||
class MCPClientManager:
|
||||
"""Manages connections to MCP servers for a single agent session."""
|
||||
|
||||
def __init__(self):
|
||||
self._connections: dict[str, MCPConnection] = {}
|
||||
self._exit_stack = AsyncExitStack()
|
||||
self._started = False
|
||||
|
||||
async def __aenter__(self):
|
||||
await self._exit_stack.__aenter__()
|
||||
self._started = True
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
await self.disconnect_all()
|
||||
try:
|
||||
await self._exit_stack.__aexit__(*exc)
|
||||
except (BaseExceptionGroup, ExceptionGroup, Exception) as e:
|
||||
# MCP subprocess cleanup errors are non-fatal
|
||||
logger.warning(f"MCP cleanup error (non-fatal): {e}")
|
||||
self._started = False
|
||||
|
||||
async def connect(self, server_name: str, config: dict, timeout: float = 30.0) -> list[ToolSchema]:
|
||||
"""Connect to an MCP server and return its available tools.
|
||||
|
||||
The tools are returned with names prefixed as mcp__<server_name>__<tool_name>.
|
||||
"""
|
||||
transport = config.get("type", "stdio")
|
||||
try:
|
||||
if transport == "stdio":
|
||||
coro = self._connect_stdio(server_name, config)
|
||||
elif transport == "sse":
|
||||
coro = self._connect_sse(server_name, config)
|
||||
elif transport == "http":
|
||||
coro = self._connect_http(server_name, config)
|
||||
else:
|
||||
logger.warning(f"Unsupported MCP transport: {transport} for {server_name}")
|
||||
return []
|
||||
|
||||
conn = await asyncio.wait_for(coro, timeout=timeout)
|
||||
self._connections[server_name] = conn
|
||||
logger.info(f"MCP connected: {server_name} ({len(conn.tools)} tools)")
|
||||
return conn.tools
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(f"MCP server {server_name} connection timed out after {timeout}s")
|
||||
return []
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to connect MCP server {server_name}: {e}")
|
||||
return []
|
||||
|
||||
async def _connect_stdio(self, server_name: str, config: dict) -> MCPConnection:
|
||||
"""Connect to a stdio MCP server (spawns a subprocess)."""
|
||||
from mcp import ClientSession
|
||||
from mcp.client.stdio import stdio_client, StdioServerParameters
|
||||
|
||||
command = config.get("command", "")
|
||||
args = config.get("args", [])
|
||||
env = config.get("env")
|
||||
|
||||
params = StdioServerParameters(
|
||||
command=command,
|
||||
args=args,
|
||||
env=env,
|
||||
)
|
||||
|
||||
transport = await self._exit_stack.enter_async_context(
|
||||
stdio_client(params)
|
||||
)
|
||||
read_stream, write_stream = transport
|
||||
session = await self._exit_stack.enter_async_context(
|
||||
ClientSession(read_stream, write_stream)
|
||||
)
|
||||
await session.initialize()
|
||||
|
||||
result = await session.list_tools()
|
||||
tools = [
|
||||
ToolSchema(
|
||||
name=f"mcp__{server_name}__{t.name}",
|
||||
description=t.description or "",
|
||||
input_schema=t.inputSchema if hasattr(t, "inputSchema") else (t.input_schema if hasattr(t, "input_schema") else {}),
|
||||
)
|
||||
for t in result.tools
|
||||
]
|
||||
|
||||
return MCPConnection(server_name=server_name, session=session, tools=tools)
|
||||
|
||||
async def _connect_sse(self, server_name: str, config: dict) -> MCPConnection:
|
||||
"""Connect to an SSE MCP server."""
|
||||
from mcp import ClientSession
|
||||
from mcp.client.sse import sse_client
|
||||
|
||||
url = config.get("url", "")
|
||||
headers = config.get("headers")
|
||||
|
||||
transport = await self._exit_stack.enter_async_context(
|
||||
sse_client(url=url, headers=headers, timeout=30, sse_read_timeout=300)
|
||||
)
|
||||
read_stream, write_stream = transport
|
||||
session = await self._exit_stack.enter_async_context(
|
||||
ClientSession(read_stream, write_stream)
|
||||
)
|
||||
await session.initialize()
|
||||
|
||||
result = await session.list_tools()
|
||||
tools = [
|
||||
ToolSchema(
|
||||
name=f"mcp__{server_name}__{t.name}",
|
||||
description=t.description or "",
|
||||
input_schema=t.inputSchema if hasattr(t, "inputSchema") else (t.input_schema if hasattr(t, "input_schema") else {}),
|
||||
)
|
||||
for t in result.tools
|
||||
]
|
||||
|
||||
return MCPConnection(server_name=server_name, session=session, tools=tools)
|
||||
|
||||
async def _connect_http(self, server_name: str, config: dict) -> MCPConnection:
|
||||
"""Connect to a Streamable HTTP MCP server.
|
||||
|
||||
Falls back to SSE if streamable HTTP fails.
|
||||
"""
|
||||
url = config.get("url", "")
|
||||
headers = config.get("headers")
|
||||
|
||||
# Try streamable HTTP first, fall back to SSE
|
||||
try:
|
||||
return await self._connect_http_streamable(server_name, url, headers)
|
||||
except Exception as e:
|
||||
logger.info(f"Streamable HTTP failed for {server_name}, trying SSE: {e}")
|
||||
return await self._connect_sse(server_name, config)
|
||||
|
||||
async def _connect_http_streamable(
|
||||
self, server_name: str, url: str, headers: dict | None,
|
||||
) -> MCPConnection:
|
||||
"""Connect via Streamable HTTP (JSON-RPC POST)."""
|
||||
import httpx
|
||||
from mcp import ClientSession
|
||||
|
||||
# Use httpx for streamable HTTP — keep client alive in the exit stack
|
||||
client = await self._exit_stack.enter_async_context(
|
||||
httpx.AsyncClient(timeout=30.0)
|
||||
)
|
||||
|
||||
h = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
**(headers or {}),
|
||||
}
|
||||
|
||||
# Initialize
|
||||
init_resp = await client.post(url, headers=h, json={
|
||||
"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": {},
|
||||
"clientInfo": {"name": "self-swarm", "version": "0.1.0"},
|
||||
},
|
||||
})
|
||||
if init_resp.status_code not in (200, 201):
|
||||
raise ConnectionError(f"MCP initialize failed: {init_resp.status_code}")
|
||||
|
||||
session_id = init_resp.headers.get("mcp-session-id", "")
|
||||
if session_id:
|
||||
h["mcp-session-id"] = session_id
|
||||
|
||||
# Notify initialized
|
||||
await client.post(url, headers=h, json={
|
||||
"jsonrpc": "2.0", "method": "notifications/initialized",
|
||||
})
|
||||
|
||||
# List tools
|
||||
list_resp = await client.post(url, headers=h, json={
|
||||
"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {},
|
||||
})
|
||||
if list_resp.status_code not in (200, 201):
|
||||
raise ConnectionError(f"MCP tools/list failed: {list_resp.status_code}")
|
||||
|
||||
ct = list_resp.headers.get("content-type", "")
|
||||
if "text/event-stream" in ct:
|
||||
data = self._parse_sse_json(list_resp.text)
|
||||
else:
|
||||
data = list_resp.json()
|
||||
|
||||
if not data:
|
||||
raise ConnectionError("Empty response from MCP server")
|
||||
|
||||
tools_list = data.get("result", {}).get("tools", [])
|
||||
tools = [
|
||||
ToolSchema(
|
||||
name=f"mcp__{server_name}__{t.get('name', '')}",
|
||||
description=t.get("description", ""),
|
||||
input_schema=t.get("inputSchema", t.get("input_schema", {})),
|
||||
)
|
||||
for t in tools_list
|
||||
]
|
||||
|
||||
# Store the HTTP client info for call_tool
|
||||
conn = MCPConnection(server_name=server_name, session=None, tools=tools)
|
||||
conn._http_client = client # type: ignore[attr-defined]
|
||||
conn._http_url = url # type: ignore[attr-defined]
|
||||
conn._http_headers = h # type: ignore[attr-defined]
|
||||
conn._next_id = 3 # type: ignore[attr-defined]
|
||||
return conn
|
||||
|
||||
@staticmethod
|
||||
def _parse_sse_json(text: str) -> dict | None:
|
||||
"""Extract JSON from an SSE response body."""
|
||||
for line in text.splitlines():
|
||||
stripped = line.strip()
|
||||
if stripped.startswith("data:"):
|
||||
payload = stripped[len("data:"):].strip()
|
||||
if payload:
|
||||
try:
|
||||
return json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
try:
|
||||
return json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
async def call_tool(
|
||||
self, server_name: str, tool_name: str, arguments: dict,
|
||||
) -> list[dict]:
|
||||
"""Call a tool on a specific MCP server.
|
||||
|
||||
Args:
|
||||
server_name: The MCP server name (e.g. "google-workspace")
|
||||
tool_name: The bare tool name (without mcp__prefix)
|
||||
arguments: Tool input arguments
|
||||
|
||||
Returns:
|
||||
List of content blocks: [{"type": "text", "text": "..."}]
|
||||
"""
|
||||
conn = self._connections.get(server_name)
|
||||
if not conn:
|
||||
return [{"type": "text", "text": f"MCP server {server_name} not connected"}]
|
||||
|
||||
try:
|
||||
if conn.session is not None:
|
||||
# stdio or SSE — use MCP ClientSession
|
||||
result = await conn.session.call_tool(tool_name, arguments)
|
||||
return self._format_mcp_result(result)
|
||||
elif hasattr(conn, "_http_client"):
|
||||
# Streamable HTTP — use JSON-RPC
|
||||
return await self._call_tool_http(conn, tool_name, arguments)
|
||||
else:
|
||||
return [{"type": "text", "text": f"No session for MCP server {server_name}"}]
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"MCP tool call failed: {server_name}/{tool_name}: {e}")
|
||||
return [{"type": "text", "text": f"Error calling {tool_name}: {e}"}]
|
||||
|
||||
async def _call_tool_http(
|
||||
self, conn: MCPConnection, tool_name: str, arguments: dict,
|
||||
) -> list[dict]:
|
||||
"""Call a tool via Streamable HTTP."""
|
||||
client = conn._http_client # type: ignore[attr-defined]
|
||||
url = conn._http_url # type: ignore[attr-defined]
|
||||
headers = conn._http_headers # type: ignore[attr-defined]
|
||||
req_id = conn._next_id # type: ignore[attr-defined]
|
||||
conn._next_id = req_id + 1 # type: ignore[attr-defined]
|
||||
|
||||
resp = await client.post(url, headers=headers, json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": req_id,
|
||||
"method": "tools/call",
|
||||
"params": {"name": tool_name, "arguments": arguments},
|
||||
}, timeout=300.0)
|
||||
|
||||
ct = resp.headers.get("content-type", "")
|
||||
if "text/event-stream" in ct:
|
||||
data = self._parse_sse_json(resp.text)
|
||||
else:
|
||||
data = resp.json()
|
||||
|
||||
if not data:
|
||||
return [{"type": "text", "text": "Empty response from MCP server"}]
|
||||
|
||||
if "error" in data:
|
||||
return [{"type": "text", "text": f"MCP error: {data['error']}"}]
|
||||
|
||||
result = data.get("result", {})
|
||||
content = result.get("content", [])
|
||||
return content if content else [{"type": "text", "text": json.dumps(result)}]
|
||||
|
||||
@staticmethod
|
||||
def _format_mcp_result(result: Any) -> list[dict]:
|
||||
"""Convert an MCP CallToolResult to content blocks."""
|
||||
if hasattr(result, "content"):
|
||||
blocks = []
|
||||
for item in result.content:
|
||||
if hasattr(item, "text"):
|
||||
blocks.append({"type": "text", "text": item.text})
|
||||
elif hasattr(item, "data"):
|
||||
blocks.append({
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": getattr(item, "mimeType", "image/png"),
|
||||
"data": item.data,
|
||||
},
|
||||
})
|
||||
else:
|
||||
blocks.append({"type": "text", "text": str(item)})
|
||||
return blocks if blocks else [{"type": "text", "text": "Done."}]
|
||||
|
||||
return [{"type": "text", "text": str(result)}]
|
||||
|
||||
def get_all_tool_schemas(self) -> list[ToolSchema]:
|
||||
"""Return tool schemas from all connected MCP servers."""
|
||||
schemas = []
|
||||
for conn in self._connections.values():
|
||||
schemas.extend(conn.tools)
|
||||
return schemas
|
||||
|
||||
def parse_mcp_tool_name(self, full_name: str) -> tuple[str, str] | None:
|
||||
"""Parse mcp__<server>__<tool> into (server_name, tool_name).
|
||||
|
||||
Returns None if the name doesn't match the MCP naming convention.
|
||||
"""
|
||||
import re
|
||||
m = re.match(r"mcp__([^_]+(?:-[^_]+)*)__(.+)", full_name)
|
||||
if m:
|
||||
return m.group(1), m.group(2)
|
||||
return None
|
||||
|
||||
async def disconnect_all(self):
|
||||
"""Disconnect all MCP servers. Called on session end."""
|
||||
self._connections.clear()
|
||||
# The AsyncExitStack handles actual cleanup of transports/sessions
|
||||
@@ -56,32 +56,10 @@ class Message(BaseModel):
|
||||
# number frozen on the persisted bubble matches what the user saw
|
||||
# rising during the stream. Pure display, not billing.
|
||||
tokens: Optional[int] = None
|
||||
# Richer thinking-pill label data. Populated only on Message(role=
|
||||
# "thinking") for live turns; legacy messages and non-thinking roles
|
||||
# leave them None. answer_tokens = total turn output - reasoning
|
||||
# tokens (the user-visible answer text + tool args). tool_count =
|
||||
# how many tools the model invoked on this turn. Drives the
|
||||
# "Thought for 18s · 430 reasoning · 2.4K answer · 3 tools" label.
|
||||
answer_tokens: Optional[int] = None
|
||||
# tool_count drives the "3 tools used" segment on the thinking pill.
|
||||
tool_count: Optional[int] = None
|
||||
# Combined input+output token total for the turn that produced this
|
||||
# thinking message — including all sub-work delegated to subagents
|
||||
# (browser, invoke-agent) and tool MCP servers that report their own
|
||||
# usage. Stored under `input_tokens` for back-compat with older
|
||||
# session JSONs even though the value is now the full
|
||||
# input+output+children sum. This is the "how big was this turn"
|
||||
# number that drives the pill's "M tokens" segment. None when no
|
||||
# usage data was captured.
|
||||
# combined input + output + children tokens for the turn (overloaded name).
|
||||
input_tokens: Optional[int] = None
|
||||
# Gemini 2.5/3.x emit a `thoughtSignature` (an opaque encrypted
|
||||
# blob) on each thinking block, and Google rejects subsequent
|
||||
# multi-step requests with a 400 if the signature isn't echoed
|
||||
# back in the conversation history. We capture it here on
|
||||
# role="thinking" messages so it survives serialization through
|
||||
# session.json AND can be re-attached to the assistant turn we
|
||||
# send back through the SDK on the next request. None for any
|
||||
# provider that doesn't use signatures (Anthropic, OpenAI).
|
||||
thought_signature: Optional[str] = None
|
||||
|
||||
class MessageBranch(BaseModel):
|
||||
id: str = Field(default_factory=lambda: uuid4().hex)
|
||||
|
||||
@@ -1,290 +0,0 @@
|
||||
"""Anthropic provider adapter using the native Anthropic SDK."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, AsyncIterator
|
||||
|
||||
import anthropic
|
||||
|
||||
from backend.apps.agents.providers.base import (
|
||||
BaseProvider, ContentBlock, ModelResponse, ProviderMessage,
|
||||
StreamEvent, ToolCall, ToolSchema,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MODEL_MAP = {
|
||||
"sonnet": "claude-sonnet-4-6",
|
||||
"opus": "claude-opus-4-6",
|
||||
"haiku": "claude-haiku-4-5",
|
||||
}
|
||||
|
||||
|
||||
class AnthropicProvider(BaseProvider):
|
||||
"""Provider adapter for Anthropic's Messages API."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
auth_token: str | None = None,
|
||||
base_url: str | None = None,
|
||||
):
|
||||
kwargs: dict[str, Any] = {}
|
||||
if auth_token:
|
||||
kwargs["auth_token"] = auth_token
|
||||
elif api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
if base_url:
|
||||
kwargs["base_url"] = base_url
|
||||
self.client = anthropic.AsyncAnthropic(**kwargs)
|
||||
|
||||
def get_model_id(self, short_name: str) -> str:
|
||||
return MODEL_MAP.get(short_name, short_name)
|
||||
|
||||
def clean_tool_schema(self, schema: ToolSchema) -> dict:
|
||||
return {
|
||||
"name": schema.name,
|
||||
"description": schema.description,
|
||||
"input_schema": schema.input_schema,
|
||||
}
|
||||
|
||||
def format_tool_result(self, tool_use_id: str, content: list[dict]) -> dict:
|
||||
return {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": tool_use_id,
|
||||
"content": content,
|
||||
}
|
||||
|
||||
def format_user_message(self, content: Any) -> ProviderMessage:
|
||||
return ProviderMessage(role="user", content=content)
|
||||
|
||||
def format_assistant_message(self, response: ModelResponse) -> ProviderMessage:
|
||||
blocks = []
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
blocks.append({"type": "text", "text": block.text})
|
||||
elif block.type == "tool_use" and block.tool_call:
|
||||
blocks.append({
|
||||
"type": "tool_use",
|
||||
"id": block.tool_call.id,
|
||||
"name": block.tool_call.name,
|
||||
"input": block.tool_call.input,
|
||||
})
|
||||
return ProviderMessage(role="assistant", content=blocks)
|
||||
|
||||
def _build_messages(self, messages: list[ProviderMessage]) -> list[dict]:
|
||||
"""Convert ProviderMessages to Anthropic API format."""
|
||||
result = []
|
||||
for msg in messages:
|
||||
if msg.role == "tool_result":
|
||||
# Tool results: content is a list of tool_result dicts
|
||||
if isinstance(msg.content, list):
|
||||
result.append({"role": "user", "content": msg.content})
|
||||
else:
|
||||
result.append({"role": "user", "content": [msg.content]})
|
||||
elif msg.role == "assistant":
|
||||
result.append({"role": "assistant", "content": msg.content})
|
||||
elif msg.role == "user":
|
||||
result.append({"role": "user", "content": msg.content})
|
||||
return result
|
||||
|
||||
async def create_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> ModelResponse:
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": self.get_model_id(model),
|
||||
"max_tokens": max_tokens,
|
||||
"messages": self._build_messages(messages),
|
||||
}
|
||||
if system:
|
||||
kwargs["system"] = system
|
||||
if tools:
|
||||
kwargs["tools"] = [self.clean_tool_schema(t) for t in tools]
|
||||
|
||||
resp = await self.client.messages.create(**kwargs)
|
||||
|
||||
content = []
|
||||
for block in resp.content:
|
||||
if block.type == "text":
|
||||
content.append(ContentBlock(type="text", text=block.text))
|
||||
elif block.type == "tool_use":
|
||||
content.append(ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(
|
||||
id=block.id,
|
||||
name=block.name,
|
||||
input=block.input,
|
||||
),
|
||||
))
|
||||
|
||||
return ModelResponse(
|
||||
content=content,
|
||||
stop_reason="tool_use" if resp.stop_reason == "tool_use" else "end_turn",
|
||||
usage={
|
||||
"input_tokens": resp.usage.input_tokens,
|
||||
"output_tokens": resp.usage.output_tokens,
|
||||
},
|
||||
)
|
||||
|
||||
async def stream_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> AsyncIterator[StreamEvent]:
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": self.get_model_id(model),
|
||||
"max_tokens": max_tokens,
|
||||
"messages": self._build_messages(messages),
|
||||
}
|
||||
if system:
|
||||
kwargs["system"] = system
|
||||
if tools:
|
||||
kwargs["tools"] = [self.clean_tool_schema(t) for t in tools]
|
||||
|
||||
# Use create() with stream=True for raw SSE events
|
||||
kwargs["stream"] = True
|
||||
raw_stream = await self.client.messages.create(**kwargs)
|
||||
|
||||
current_block_type: dict[int, str] = {}
|
||||
current_tool_name: dict[int, str] = {}
|
||||
current_tool_id: dict[int, str] = {}
|
||||
current_text: dict[int, str] = {}
|
||||
current_json: dict[int, str] = {}
|
||||
|
||||
async for event in raw_stream:
|
||||
event_type = getattr(event, "type", "")
|
||||
|
||||
if event_type == "content_block_start":
|
||||
index = event.index
|
||||
block = event.content_block
|
||||
block_type = block.type
|
||||
current_block_type[index] = block_type
|
||||
|
||||
if block_type == "text":
|
||||
current_text[index] = ""
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=index,
|
||||
block_type="text",
|
||||
)
|
||||
elif block_type == "tool_use":
|
||||
current_tool_name[index] = block.name
|
||||
current_tool_id[index] = block.id
|
||||
current_json[index] = ""
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=index,
|
||||
block_type="tool_use",
|
||||
tool_name=block.name,
|
||||
tool_id=block.id,
|
||||
)
|
||||
elif block_type == "thinking":
|
||||
# Extended-thinking content block. We track the
|
||||
# accumulated text in current_text just like a normal
|
||||
# text block, but tag it as "thinking" so the agent
|
||||
# loop emits a distinct WS event the frontend can
|
||||
# render in the ThinkingBubble pill.
|
||||
current_text[index] = ""
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=index,
|
||||
block_type="thinking",
|
||||
)
|
||||
|
||||
elif event_type == "content_block_delta":
|
||||
index = event.index
|
||||
delta = event.delta
|
||||
delta_type = delta.type
|
||||
|
||||
if delta_type == "text_delta":
|
||||
current_text.setdefault(index, "")
|
||||
current_text[index] += delta.text
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=index,
|
||||
delta_type="text_delta",
|
||||
text=delta.text,
|
||||
)
|
||||
elif delta_type == "input_json_delta":
|
||||
current_json.setdefault(index, "")
|
||||
current_json[index] += delta.partial_json
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=index,
|
||||
delta_type="input_json_delta",
|
||||
text=delta.partial_json,
|
||||
)
|
||||
elif delta_type == "thinking_delta":
|
||||
# Extended-thinking text streamed as it's produced.
|
||||
# Forward as a thinking_delta so the agent loop can
|
||||
# ship it to the frontend without conflating with
|
||||
# the assistant text stream.
|
||||
text_chunk = getattr(delta, "thinking", "") or ""
|
||||
current_text.setdefault(index, "")
|
||||
current_text[index] += text_chunk
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=index,
|
||||
delta_type="thinking_delta",
|
||||
text=text_chunk,
|
||||
)
|
||||
# Note: signature_delta (the cryptographic signature on
|
||||
# thinking blocks) is intentionally ignored — we don't
|
||||
# display it and it isn't needed for replay since we
|
||||
# never re-send thinking blocks to the model.
|
||||
|
||||
elif event_type == "content_block_stop":
|
||||
yield StreamEvent(type="content_block_stop", index=event.index)
|
||||
|
||||
elif event_type == "message_delta":
|
||||
# Extract output token usage from the final delta
|
||||
usage_data = {}
|
||||
delta_usage = getattr(event, "usage", None)
|
||||
if delta_usage:
|
||||
output_tokens = getattr(delta_usage, "output_tokens", 0)
|
||||
if output_tokens:
|
||||
usage_data["output_tokens"] = output_tokens
|
||||
if usage_data:
|
||||
yield StreamEvent(type="usage", usage=usage_data)
|
||||
|
||||
elif event_type == "message_start":
|
||||
# Extract input token usage from the message start
|
||||
msg = getattr(event, "message", None)
|
||||
if msg:
|
||||
msg_usage = getattr(msg, "usage", None)
|
||||
if msg_usage:
|
||||
usage_data = {}
|
||||
input_tokens = getattr(msg_usage, "input_tokens", 0)
|
||||
output_tokens = getattr(msg_usage, "output_tokens", 0)
|
||||
if input_tokens:
|
||||
usage_data["input_tokens"] = input_tokens
|
||||
if output_tokens:
|
||||
usage_data["output_tokens"] = output_tokens
|
||||
if usage_data:
|
||||
yield StreamEvent(type="usage", usage=usage_data)
|
||||
|
||||
yield StreamEvent(type="message_stop")
|
||||
|
||||
async def stream_and_collect(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> tuple[AsyncIterator[StreamEvent], ModelResponse]:
|
||||
"""Helper: stream events and also return the full collected response.
|
||||
|
||||
Not used directly — the AgentLoop handles collection.
|
||||
"""
|
||||
raise NotImplementedError("Use stream_message() directly; AgentLoop collects.")
|
||||
@@ -1,135 +0,0 @@
|
||||
"""Provider-agnostic base classes for multi-model support.
|
||||
|
||||
All provider adapters (Anthropic, OpenAI, Gemini, OpenAI-compatible)
|
||||
implement BaseProvider, translating their native APIs into these
|
||||
common data structures.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, AsyncIterator
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolSchema:
|
||||
"""Provider-agnostic tool definition."""
|
||||
name: str
|
||||
description: str
|
||||
input_schema: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolCall:
|
||||
"""A tool invocation requested by the model."""
|
||||
id: str
|
||||
name: str
|
||||
input: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContentBlock:
|
||||
"""A block of content from the model response."""
|
||||
type: str # "text" | "tool_use" | "thinking"
|
||||
text: str = ""
|
||||
tool_call: ToolCall | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelResponse:
|
||||
"""Complete (non-streaming) response from a provider."""
|
||||
content: list[ContentBlock]
|
||||
stop_reason: str # "end_turn" | "tool_use" | "max_tokens"
|
||||
usage: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamEvent:
|
||||
"""A single streaming event, normalized across providers.
|
||||
|
||||
The event types match what the frontend already expects via WebSocket:
|
||||
content_block_start, content_block_delta, content_block_stop, message_stop.
|
||||
"""
|
||||
type: str
|
||||
index: int = 0
|
||||
block_type: str = "" # "text" | "tool_use" | "thinking"
|
||||
delta_type: str = "" # "text_delta" | "input_json_delta" | "thinking_delta"
|
||||
text: str = ""
|
||||
tool_name: str = ""
|
||||
tool_id: str = ""
|
||||
usage: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProviderMessage:
|
||||
"""Provider-agnostic message for conversation history.
|
||||
|
||||
Each provider adapter converts these to/from its native format.
|
||||
"""
|
||||
role: str # "user" | "assistant" | "tool_result"
|
||||
content: Any # str, list[dict], or provider-specific content
|
||||
|
||||
|
||||
class BaseProvider(ABC):
|
||||
"""Abstract base for LLM provider adapters."""
|
||||
|
||||
@abstractmethod
|
||||
async def stream_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> AsyncIterator[StreamEvent]:
|
||||
"""Stream a model response, yielding normalized StreamEvents."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def create_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> ModelResponse:
|
||||
"""Non-streaming message creation."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def format_tool_result(
|
||||
self,
|
||||
tool_use_id: str,
|
||||
content: list[dict],
|
||||
) -> dict:
|
||||
"""Format a tool result in this provider's expected message format."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def format_user_message(self, content: Any) -> ProviderMessage:
|
||||
"""Wrap user content (str or multimodal blocks) into a ProviderMessage."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def format_assistant_message(self, response: ModelResponse) -> ProviderMessage:
|
||||
"""Convert a ModelResponse into a ProviderMessage for conversation history."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def get_model_id(self, short_name: str) -> str:
|
||||
"""Resolve a short model name to the full API model ID."""
|
||||
...
|
||||
|
||||
def clean_tool_schema(self, schema: ToolSchema) -> dict:
|
||||
"""Convert a ToolSchema to the provider's native tool format.
|
||||
|
||||
Default: Anthropic-style format. Override for providers that need
|
||||
different formats or schema cleaning (e.g. Gemini).
|
||||
"""
|
||||
return {
|
||||
"name": schema.name,
|
||||
"description": schema.description,
|
||||
"input_schema": schema.input_schema,
|
||||
}
|
||||
@@ -1,382 +0,0 @@
|
||||
"""OpenAI-compatible provider adapter.
|
||||
|
||||
Works with ANY endpoint that speaks the OpenAI Chat Completions API:
|
||||
OpenAI, OpenRouter, Together, Groq, Fireworks, Mistral, Ollama, vLLM, etc.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, AsyncIterator
|
||||
from uuid import uuid4
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from backend.apps.agents.providers.base import (
|
||||
BaseProvider, ContentBlock, ModelResponse, ProviderMessage,
|
||||
StreamEvent, ToolCall, ToolSchema,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OpenAICompatProvider(BaseProvider):
|
||||
"""Provider adapter for any OpenAI-compatible API endpoint."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str = "",
|
||||
base_url: str | None = None,
|
||||
):
|
||||
kwargs: dict[str, Any] = {}
|
||||
# Always set api_key — use "none" as placeholder if empty (some endpoints don't need real keys)
|
||||
kwargs["api_key"] = api_key if api_key else "none"
|
||||
if base_url:
|
||||
kwargs["base_url"] = base_url
|
||||
self.client = AsyncOpenAI(**kwargs)
|
||||
|
||||
def get_model_id(self, short_name: str) -> str:
|
||||
# Pass through — user selects exact model ID
|
||||
return short_name
|
||||
|
||||
def clean_tool_schema(self, schema: ToolSchema) -> dict:
|
||||
"""Convert to OpenAI function calling format."""
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": schema.name,
|
||||
"description": schema.description,
|
||||
"parameters": schema.input_schema,
|
||||
},
|
||||
}
|
||||
|
||||
def format_tool_result(self, tool_use_id: str, content: list[dict]) -> dict:
|
||||
"""Format tool result as OpenAI expects."""
|
||||
# OpenAI wants a single string for tool results
|
||||
text_parts = []
|
||||
for block in content:
|
||||
if block.get("type") == "text":
|
||||
text_parts.append(block.get("text", ""))
|
||||
elif block.get("type") == "image":
|
||||
text_parts.append("[image]")
|
||||
else:
|
||||
text_parts.append(json.dumps(block))
|
||||
return {
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_use_id,
|
||||
"content": "\n".join(text_parts) if text_parts else "Done.",
|
||||
}
|
||||
|
||||
def format_user_message(self, content: Any) -> ProviderMessage:
|
||||
"""Convert user content to OpenAI format."""
|
||||
if isinstance(content, str):
|
||||
return ProviderMessage(role="user", content=content)
|
||||
# Multimodal content (text + images)
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
for block in content:
|
||||
if isinstance(block, dict):
|
||||
if block.get("type") == "text":
|
||||
parts.append({"type": "text", "text": block["text"]})
|
||||
elif block.get("type") == "image":
|
||||
source = block.get("source", {})
|
||||
media_type = source.get("media_type", "image/png")
|
||||
data = source.get("data", "")
|
||||
parts.append({
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:{media_type};base64,{data}"},
|
||||
})
|
||||
elif isinstance(block, str):
|
||||
parts.append({"type": "text", "text": block})
|
||||
return ProviderMessage(role="user", content=parts)
|
||||
return ProviderMessage(role="user", content=str(content))
|
||||
|
||||
def format_assistant_message(self, response: ModelResponse) -> ProviderMessage:
|
||||
"""Convert ModelResponse to OpenAI assistant message format."""
|
||||
text_parts = []
|
||||
tool_calls = []
|
||||
for block in response.content:
|
||||
if block.type == "text":
|
||||
text_parts.append(block.text)
|
||||
elif block.type == "tool_use" and block.tool_call:
|
||||
tool_calls.append({
|
||||
"id": block.tool_call.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.tool_call.name,
|
||||
"arguments": json.dumps(block.tool_call.input),
|
||||
},
|
||||
})
|
||||
msg: dict[str, Any] = {"role": "assistant"}
|
||||
if text_parts:
|
||||
msg["content"] = "\n".join(text_parts)
|
||||
else:
|
||||
msg["content"] = None
|
||||
if tool_calls:
|
||||
msg["tool_calls"] = tool_calls
|
||||
return ProviderMessage(role="assistant", content=msg)
|
||||
|
||||
def _build_messages(
|
||||
self,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
) -> list[dict]:
|
||||
"""Convert ProviderMessages to OpenAI API format."""
|
||||
result = []
|
||||
if system:
|
||||
result.append({"role": "system", "content": system})
|
||||
|
||||
for msg in messages:
|
||||
if msg.role == "assistant":
|
||||
# Assistant messages are already in OpenAI format from format_assistant_message
|
||||
if isinstance(msg.content, dict) and "role" in msg.content:
|
||||
result.append(msg.content)
|
||||
else:
|
||||
# Raw content blocks from provider-agnostic format
|
||||
text_parts = []
|
||||
tool_calls = []
|
||||
if isinstance(msg.content, list):
|
||||
for block in msg.content:
|
||||
if isinstance(block, dict):
|
||||
if block.get("type") == "text":
|
||||
text_parts.append(block["text"])
|
||||
elif block.get("type") == "tool_use":
|
||||
tool_calls.append({
|
||||
"id": block.get("id", uuid4().hex),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.get("name", ""),
|
||||
"arguments": json.dumps(block.get("input", {})),
|
||||
},
|
||||
})
|
||||
api_msg: dict[str, Any] = {
|
||||
"role": "assistant",
|
||||
"content": "\n".join(text_parts) if text_parts else None,
|
||||
}
|
||||
if tool_calls:
|
||||
api_msg["tool_calls"] = tool_calls
|
||||
result.append(api_msg)
|
||||
|
||||
elif msg.role == "tool_result":
|
||||
# Tool results: content is a list of tool result dicts
|
||||
if isinstance(msg.content, list):
|
||||
for tr in msg.content:
|
||||
if isinstance(tr, dict) and "tool_call_id" in tr:
|
||||
result.append(tr)
|
||||
elif isinstance(msg.content, dict) and "tool_call_id" in msg.content:
|
||||
result.append(msg.content)
|
||||
|
||||
elif msg.role == "user":
|
||||
result.append({"role": "user", "content": msg.content})
|
||||
|
||||
return result
|
||||
|
||||
async def create_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> ModelResponse:
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": self.get_model_id(model),
|
||||
"max_tokens": max_tokens,
|
||||
"messages": self._build_messages(system, messages),
|
||||
}
|
||||
if tools:
|
||||
kwargs["tools"] = [self.clean_tool_schema(t) for t in tools]
|
||||
|
||||
resp = await self.client.chat.completions.create(**kwargs)
|
||||
choice = resp.choices[0]
|
||||
message = choice.message
|
||||
|
||||
content: list[ContentBlock] = []
|
||||
if message.content:
|
||||
content.append(ContentBlock(type="text", text=message.content))
|
||||
|
||||
if message.tool_calls:
|
||||
for tc in message.tool_calls:
|
||||
try:
|
||||
args = json.loads(tc.function.arguments)
|
||||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
content.append(ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(
|
||||
id=tc.id,
|
||||
name=tc.function.name,
|
||||
input=args,
|
||||
),
|
||||
))
|
||||
|
||||
stop = "end_turn"
|
||||
if choice.finish_reason == "tool_calls":
|
||||
stop = "tool_use"
|
||||
elif message.tool_calls:
|
||||
stop = "tool_use"
|
||||
|
||||
usage_dict = {}
|
||||
if resp.usage:
|
||||
usage_dict = {
|
||||
"input_tokens": resp.usage.prompt_tokens,
|
||||
"output_tokens": resp.usage.completion_tokens,
|
||||
}
|
||||
|
||||
return ModelResponse(content=content, stop_reason=stop, usage=usage_dict)
|
||||
|
||||
async def stream_message(
|
||||
self,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list[ToolSchema],
|
||||
max_tokens: int = 8192,
|
||||
) -> AsyncIterator[StreamEvent]:
|
||||
kwargs: dict[str, Any] = {
|
||||
"model": self.get_model_id(model),
|
||||
"max_tokens": max_tokens,
|
||||
"messages": self._build_messages(system, messages),
|
||||
"stream": True,
|
||||
"stream_options": {"include_usage": True},
|
||||
}
|
||||
if tools:
|
||||
kwargs["tools"] = [self.clean_tool_schema(t) for t in tools]
|
||||
|
||||
stream = await self.client.chat.completions.create(**kwargs)
|
||||
|
||||
# Track streaming state to emit normalized events
|
||||
text_started = False
|
||||
text_index = 0
|
||||
thinking_started = False
|
||||
thinking_index = 0
|
||||
tool_indices: dict[int, dict] = {} # openai tool_call index -> {name, id, json_buf}
|
||||
next_block_index = 0
|
||||
|
||||
async for chunk in stream:
|
||||
if not chunk.choices:
|
||||
# Usage-only chunk at the end
|
||||
if chunk.usage:
|
||||
yield StreamEvent(type="usage", usage={
|
||||
"input_tokens": chunk.usage.prompt_tokens or 0,
|
||||
"output_tokens": chunk.usage.completion_tokens or 0,
|
||||
})
|
||||
continue
|
||||
|
||||
delta = chunk.choices[0].delta
|
||||
finish_reason = chunk.choices[0].finish_reason
|
||||
|
||||
# Reasoning / thinking content. OpenAI o-series + GPT-5.x
|
||||
# (via Responses API → 9Router → Chat Completions shape),
|
||||
# DeepSeek-R1, and Gemini 2.5/3.x through 9Router all expose
|
||||
# their reasoning text on `delta.reasoning_content`. Forward
|
||||
# as a thinking content block so the frontend's existing
|
||||
# ThinkingBubble pill renders it just like Anthropic's
|
||||
# thinking_delta. The SDK's typed delta object doesn't
|
||||
# declare this field so we read it via getattr/dict access.
|
||||
reasoning_text: str | None = None
|
||||
try:
|
||||
reasoning_text = getattr(delta, "reasoning_content", None)
|
||||
if reasoning_text is None and isinstance(delta, dict):
|
||||
reasoning_text = delta.get("reasoning_content")
|
||||
except Exception:
|
||||
reasoning_text = None
|
||||
if reasoning_text:
|
||||
# Reasoning blocks always close before any visible
|
||||
# answer text starts; if we somehow got text first
|
||||
# (shouldn't happen with reasoning models), don't
|
||||
# interleave — just emit thinking after.
|
||||
if not thinking_started:
|
||||
thinking_started = True
|
||||
thinking_index = next_block_index
|
||||
next_block_index += 1
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=thinking_index,
|
||||
block_type="thinking",
|
||||
)
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=thinking_index,
|
||||
delta_type="thinking_delta",
|
||||
text=reasoning_text,
|
||||
)
|
||||
|
||||
# Text content
|
||||
if delta.content is not None:
|
||||
# Close any open thinking block before opening text — the
|
||||
# transition from "thinking" to "answer" is what triggers
|
||||
# the frontend pill to freeze. Mirrors Anthropic's
|
||||
# content_block_stop on thinking before the text block.
|
||||
if thinking_started:
|
||||
yield StreamEvent(type="content_block_stop", index=thinking_index)
|
||||
thinking_started = False
|
||||
if not text_started:
|
||||
text_started = True
|
||||
text_index = next_block_index
|
||||
next_block_index += 1
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=text_index,
|
||||
block_type="text",
|
||||
)
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=text_index,
|
||||
delta_type="text_delta",
|
||||
text=delta.content,
|
||||
)
|
||||
|
||||
# Tool calls
|
||||
if delta.tool_calls:
|
||||
for tc_delta in delta.tool_calls:
|
||||
tc_idx = tc_delta.index
|
||||
if tc_idx not in tool_indices:
|
||||
# New tool call starting — close any open
|
||||
# text or thinking block first.
|
||||
if text_started:
|
||||
yield StreamEvent(type="content_block_stop", index=text_index)
|
||||
text_started = False
|
||||
if thinking_started:
|
||||
yield StreamEvent(type="content_block_stop", index=thinking_index)
|
||||
thinking_started = False
|
||||
|
||||
block_idx = next_block_index
|
||||
next_block_index += 1
|
||||
tool_indices[tc_idx] = {
|
||||
"block_index": block_idx,
|
||||
"id": tc_delta.id or uuid4().hex,
|
||||
"name": tc_delta.function.name if tc_delta.function else "",
|
||||
"json_buf": "",
|
||||
}
|
||||
yield StreamEvent(
|
||||
type="content_block_start",
|
||||
index=block_idx,
|
||||
block_type="tool_use",
|
||||
tool_name=tool_indices[tc_idx]["name"],
|
||||
tool_id=tool_indices[tc_idx]["id"],
|
||||
)
|
||||
|
||||
info = tool_indices[tc_idx]
|
||||
if tc_delta.function and tc_delta.function.name:
|
||||
info["name"] = tc_delta.function.name
|
||||
if tc_delta.function and tc_delta.function.arguments:
|
||||
info["json_buf"] += tc_delta.function.arguments
|
||||
yield StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=info["block_index"],
|
||||
delta_type="input_json_delta",
|
||||
text=tc_delta.function.arguments,
|
||||
)
|
||||
|
||||
# Finish
|
||||
if finish_reason is not None:
|
||||
if thinking_started:
|
||||
yield StreamEvent(type="content_block_stop", index=thinking_index)
|
||||
thinking_started = False
|
||||
if text_started:
|
||||
yield StreamEvent(type="content_block_stop", index=text_index)
|
||||
for info in tool_indices.values():
|
||||
yield StreamEvent(type="content_block_stop", index=info["block_index"])
|
||||
yield StreamEvent(type="message_stop")
|
||||
@@ -1,14 +1,8 @@
|
||||
"""Provider registry and model catalog.
|
||||
|
||||
NOTE: `create_provider`, `BaseProvider`, `AnthropicProvider`, `OpenAICompatProvider`,
|
||||
and the native `AgentLoop` are currently unused. The live agent path is
|
||||
`claude_agent_sdk` via `agent_manager._run_agent_loop`. Kept as a foundation
|
||||
for a potential future native multi-provider loop.
|
||||
|
||||
Multi-model subscription support routes non-Anthropic models through 9Router's
|
||||
`/v1/messages` endpoint by passing prefixed model IDs (e.g. `cx/gpt-5.4`,
|
||||
`gc/gemini-3-pro-preview`). 9Router's translator converts the Anthropic-format
|
||||
request into the provider's native format transparently.
|
||||
Live agent path goes through claude_agent_sdk via agent_manager._run_agent_loop.
|
||||
Non-Anthropic models route through 9Router's /v1/messages endpoint with
|
||||
prefixed ids (cx/gpt-5.4, gc/gemini-3-pro-preview).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -16,8 +10,6 @@ from __future__ import annotations
|
||||
import logging
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
from backend.apps.agents.providers.base import BaseProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from backend.apps.settings.models import AppSettings
|
||||
|
||||
@@ -493,183 +485,6 @@ async def resolve_aux_model(
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Provider factory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def create_provider(
|
||||
provider_name: str,
|
||||
settings: AppSettings,
|
||||
provider_config: dict | None = None,
|
||||
) -> BaseProvider:
|
||||
"""Create a provider adapter.
|
||||
|
||||
Routes based on the 'api' field in BUILTIN_MODELS:
|
||||
- "anthropic" → native Anthropic SDK
|
||||
- "openai" → native OpenAI SDK (direct API)
|
||||
- "gemini" → native Google GenAI SDK
|
||||
- "openrouter" → OpenAI-compat via openrouter.ai (Meta, Mistral, DeepSeek, Qwen, xAI, etc.)
|
||||
Custom providers use OpenAI-compat with user's base_url.
|
||||
"""
|
||||
api_type = _get_api_type(provider_name)
|
||||
|
||||
# Check for 9Router first
|
||||
if provider_name in ("9Router", "9router"):
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
return OpenAICompatProvider(api_key="9router", base_url="http://localhost:20128/v1")
|
||||
|
||||
if api_type == "anthropic":
|
||||
from backend.apps.agents.providers.anthropic import AnthropicProvider
|
||||
if getattr(settings, "connection_mode", "own_key") == "openswarm-pro":
|
||||
return AnthropicProvider(
|
||||
auth_token=getattr(settings, "openswarm_bearer_token", None),
|
||||
base_url=getattr(settings, "openswarm_proxy_url", None) or "https://api.openswarm.com",
|
||||
)
|
||||
# Priority: API key → 9Router subscription
|
||||
if settings.anthropic_api_key:
|
||||
return AnthropicProvider(api_key=settings.anthropic_api_key)
|
||||
# No API key — try 9Router as fallback
|
||||
if _is_9router_available():
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
provider = OpenAICompatProvider(api_key="9router", base_url="http://localhost:20128/v1")
|
||||
# Override get_model_id to map our short names to 9Router's cc/ prefixed IDs
|
||||
_original_get_model = provider.get_model_id
|
||||
_9r_model_map = {
|
||||
"sonnet": "cc/claude-sonnet-4-6",
|
||||
"opus": "cc/claude-opus-4-6",
|
||||
"haiku": "cc/claude-haiku-4-5-20251001",
|
||||
}
|
||||
provider.get_model_id = lambda name: _9r_model_map.get(name, f"cc/{name}" if not name.startswith("cc/") else name)
|
||||
return provider
|
||||
raise ValueError("Anthropic API key not configured. Set it in Settings, or connect 9Router.")
|
||||
|
||||
if api_type == "openai":
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
if settings.openai_api_key:
|
||||
return OpenAICompatProvider(api_key=settings.openai_api_key, base_url="https://api.openai.com/v1")
|
||||
# No API key — try 9Router as fallback
|
||||
if _is_9router_available():
|
||||
return OpenAICompatProvider(api_key="9router", base_url="http://localhost:20128/v1")
|
||||
raise ValueError("OpenAI API key not configured. Set it in Settings, or connect 9Router.")
|
||||
|
||||
if api_type == "gemini":
|
||||
from backend.apps.agents.providers.gemini import GeminiProvider
|
||||
if settings.google_api_key:
|
||||
return GeminiProvider(api_key=settings.google_api_key)
|
||||
# No API key — try 9Router as fallback
|
||||
if _is_9router_available():
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
return OpenAICompatProvider(api_key="9router", base_url="http://localhost:20128/v1")
|
||||
raise ValueError("Google API key not configured. Set it in Settings, or connect 9Router.")
|
||||
|
||||
if api_type == "openrouter":
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
openrouter_key = getattr(settings, "openrouter_api_key", None)
|
||||
if openrouter_key:
|
||||
return OpenAICompatProvider(api_key=openrouter_key, base_url=OPENROUTER_BASE_URL)
|
||||
# No OpenRouter key — try 9Router as fallback
|
||||
if _is_9router_available():
|
||||
return OpenAICompatProvider(api_key="9router", base_url="http://localhost:20128/v1")
|
||||
raise ValueError(f"OpenRouter API key not configured for {provider_name}. Set it in Settings, or connect a subscription.")
|
||||
|
||||
# Custom provider — look up in settings.custom_providers
|
||||
if provider_config:
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
return OpenAICompatProvider(
|
||||
api_key=provider_config.get("api_key", ""),
|
||||
base_url=provider_config.get("base_url", ""),
|
||||
)
|
||||
|
||||
for cp in getattr(settings, "custom_providers", []):
|
||||
if cp.name == provider_name:
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
return OpenAICompatProvider(
|
||||
api_key=cp.api_key,
|
||||
base_url=cp.base_url,
|
||||
)
|
||||
|
||||
raise ValueError(f"Unknown provider: {provider_name}")
|
||||
|
||||
|
||||
def _get_api_type(provider_name: str) -> str:
|
||||
"""Get the API type for a provider from BUILTIN_MODELS.
|
||||
|
||||
Accepts both display names ('Anthropic') and lowercase API names ('anthropic').
|
||||
"""
|
||||
# Direct lookup first (display name like 'Anthropic', 'OpenAI', etc.)
|
||||
models = BUILTIN_MODELS.get(provider_name, [])
|
||||
if models:
|
||||
return models[0].get("api", "openrouter")
|
||||
|
||||
# Lowercase API name mapping
|
||||
_API_NAME_MAP = {
|
||||
"anthropic": "anthropic",
|
||||
"openai": "openai",
|
||||
"gemini": "gemini",
|
||||
"google": "gemini",
|
||||
"openrouter": "openrouter",
|
||||
}
|
||||
if provider_name.lower() in _API_NAME_MAP:
|
||||
return _API_NAME_MAP[provider_name.lower()]
|
||||
|
||||
# Case-insensitive lookup into BUILTIN_MODELS
|
||||
lower = provider_name.lower()
|
||||
for key, models in BUILTIN_MODELS.items():
|
||||
if key.lower() == lower:
|
||||
return models[0].get("api", "openrouter")
|
||||
|
||||
return "openrouter"
|
||||
|
||||
|
||||
def _has_credentials(provider_name: str, settings: AppSettings) -> bool:
|
||||
"""Check if a provider has credentials configured."""
|
||||
api_type = _get_api_type(provider_name)
|
||||
|
||||
if api_type == "anthropic":
|
||||
if getattr(settings, "connection_mode", "own_key") == "openswarm-pro":
|
||||
return bool(getattr(settings, "openswarm_bearer_token", None))
|
||||
return bool(settings.anthropic_api_key)
|
||||
if api_type == "openai":
|
||||
return bool(settings.openai_api_key)
|
||||
if api_type == "gemini":
|
||||
return bool(getattr(settings, "google_api_key", None))
|
||||
if api_type == "openrouter":
|
||||
return bool(getattr(settings, "openrouter_api_key", None))
|
||||
return False
|
||||
|
||||
|
||||
def get_available_models(settings: AppSettings) -> dict[str, list[dict]]:
|
||||
"""Return all models — always show everything, mark which have keys configured.
|
||||
|
||||
Like Cursor: show all models upfront, prompt for key when user tries to use one.
|
||||
Returns: {"provider_name": [{"value": ..., "label": ..., "context_window": ..., "configured": bool}, ...]}
|
||||
"""
|
||||
result: dict[str, list[dict]] = {}
|
||||
|
||||
# Built-in providers — always show all
|
||||
for provider_name, models in BUILTIN_MODELS.items():
|
||||
configured = _has_credentials(provider_name, settings)
|
||||
result[provider_name] = [
|
||||
{**m, "configured": configured}
|
||||
for m in models
|
||||
]
|
||||
|
||||
# Custom providers
|
||||
for cp in getattr(settings, "custom_providers", []):
|
||||
if cp.models:
|
||||
result[cp.name] = [
|
||||
{
|
||||
"value": m.get("value", m.get("id", "")),
|
||||
"label": m.get("label", m.get("value", m.get("id", ""))),
|
||||
"context_window": m.get("context_window", 128_000),
|
||||
"configured": True,
|
||||
}
|
||||
for m in cp.models
|
||||
]
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def get_context_window(provider: str, model: str, settings: AppSettings | None = None) -> int:
|
||||
"""Look up context window for any model."""
|
||||
# Check built-in models first
|
||||
|
||||
@@ -2,46 +2,22 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolContext:
|
||||
"""Runtime context passed to every tool execution."""
|
||||
cwd: str
|
||||
session_id: str
|
||||
|
||||
|
||||
class BaseTool(ABC):
|
||||
"""Abstract base for all builtin tools.
|
||||
|
||||
Subclasses must set ``name`` and ``description`` as class attributes and
|
||||
implement ``get_schema`` (JSON Schema for tool input) and ``execute``.
|
||||
"""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
|
||||
@abstractmethod
|
||||
def get_schema(self) -> dict:
|
||||
"""Return JSON Schema for this tool's input parameters."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
"""Execute the tool.
|
||||
|
||||
Returns a list of content blocks, e.g.
|
||||
``[{"type": "text", "text": "..."}]``.
|
||||
"""
|
||||
...
|
||||
|
||||
def to_tool_schema(self):
|
||||
"""Convert to the provider-agnostic ``ToolSchema`` used everywhere."""
|
||||
from backend.apps.agents.providers.base import ToolSchema
|
||||
|
||||
return ToolSchema(
|
||||
name=self.name,
|
||||
description=self.description,
|
||||
input_schema=self.get_schema(),
|
||||
)
|
||||
|
||||
@@ -1,476 +0,0 @@
|
||||
"""Filesystem tools: Read, Write, Edit, Glob, Grep."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from backend.apps.agents.tools.base import BaseTool, ToolContext
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".svg"}
|
||||
_MAX_OUTPUT_BYTES = 50 * 1024 # ~50 KB cap for grep output
|
||||
|
||||
|
||||
def _resolve(file_path: str, cwd: str) -> Path:
|
||||
"""Resolve *file_path* against *cwd* when it is relative."""
|
||||
p = Path(file_path)
|
||||
if not p.is_absolute():
|
||||
p = Path(cwd) / p
|
||||
return p.resolve()
|
||||
|
||||
|
||||
def _text_block(text: str) -> list[dict]:
|
||||
return [{"type": "text", "text": text}]
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
# ReadTool
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ReadTool(BaseTool):
|
||||
name = "Read"
|
||||
description = (
|
||||
"Read a file from the filesystem. Returns lines with line numbers "
|
||||
"(cat -n style). For image files returns base64 content. Supports "
|
||||
"offset and limit parameters for reading portions of large files."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"file_path": {
|
||||
"type": "string",
|
||||
"description": "Absolute or relative path to the file to read.",
|
||||
},
|
||||
"offset": {
|
||||
"type": "integer",
|
||||
"description": "1-based line number to start reading from.",
|
||||
},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Maximum number of lines to return (default 2000).",
|
||||
},
|
||||
},
|
||||
"required": ["file_path"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
file_path = _resolve(input_data["file_path"], context.cwd)
|
||||
|
||||
if not file_path.exists():
|
||||
return _text_block(f"Error: file not found: {file_path}")
|
||||
|
||||
if not file_path.is_file():
|
||||
return _text_block(f"Error: not a regular file: {file_path}")
|
||||
|
||||
# Binary / image files → base64
|
||||
ext = file_path.suffix.lower()
|
||||
if ext in _IMAGE_EXTENSIONS:
|
||||
try:
|
||||
raw = file_path.read_bytes()
|
||||
b64 = base64.b64encode(raw).decode("ascii")
|
||||
media = mimetypes.guess_type(str(file_path))[0] or "application/octet-stream"
|
||||
return [
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": media,
|
||||
"data": b64,
|
||||
},
|
||||
}
|
||||
]
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error reading image {file_path}: {exc}")
|
||||
|
||||
# Text files
|
||||
offset = max(input_data.get("offset", 1), 1)
|
||||
limit = input_data.get("limit", 2000)
|
||||
if limit <= 0:
|
||||
limit = 2000
|
||||
|
||||
try:
|
||||
with open(file_path, "r", errors="replace") as fh:
|
||||
lines: list[str] = []
|
||||
for lineno, line in enumerate(fh, start=1):
|
||||
if lineno < offset:
|
||||
continue
|
||||
if len(lines) >= limit:
|
||||
break
|
||||
# cat -n style: right-justified line number + tab + content
|
||||
lines.append(f"{lineno:>6}\t{line.rstrip()}")
|
||||
if not lines:
|
||||
return _text_block(f"(file is empty or offset beyond end of file: {file_path})")
|
||||
return _text_block("\n".join(lines))
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error reading {file_path}: {exc}")
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
# WriteTool
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class WriteTool(BaseTool):
|
||||
name = "Write"
|
||||
description = (
|
||||
"Write content to a file. Creates parent directories if they do not "
|
||||
"exist. Overwrites the file if it already exists."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"file_path": {
|
||||
"type": "string",
|
||||
"description": "Absolute or relative path to the file to write.",
|
||||
},
|
||||
"content": {
|
||||
"type": "string",
|
||||
"description": "The full content to write to the file.",
|
||||
},
|
||||
},
|
||||
"required": ["file_path", "content"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
file_path = _resolve(input_data["file_path"], context.cwd)
|
||||
content: str = input_data["content"]
|
||||
|
||||
try:
|
||||
file_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
file_path.write_text(content, encoding="utf-8")
|
||||
return _text_block(f"Successfully wrote {len(content)} bytes to {file_path}")
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error writing {file_path}: {exc}")
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
# EditTool
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class EditTool(BaseTool):
|
||||
name = "Edit"
|
||||
description = (
|
||||
"Perform exact string replacements in a file. By default the "
|
||||
"old_string must appear exactly once (not unique → error). Pass "
|
||||
"replace_all=true to replace every occurrence."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"file_path": {
|
||||
"type": "string",
|
||||
"description": "Absolute or relative path to the file to edit.",
|
||||
},
|
||||
"old_string": {
|
||||
"type": "string",
|
||||
"description": "The exact text to find in the file.",
|
||||
},
|
||||
"new_string": {
|
||||
"type": "string",
|
||||
"description": "The text to replace old_string with.",
|
||||
},
|
||||
"replace_all": {
|
||||
"type": "boolean",
|
||||
"description": "If true, replace all occurrences. Default false.",
|
||||
"default": False,
|
||||
},
|
||||
},
|
||||
"required": ["file_path", "old_string", "new_string"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
file_path = _resolve(input_data["file_path"], context.cwd)
|
||||
old_string: str = input_data["old_string"]
|
||||
new_string: str = input_data["new_string"]
|
||||
replace_all: bool = input_data.get("replace_all", False)
|
||||
|
||||
if not file_path.exists():
|
||||
return _text_block(f"Error: file not found: {file_path}")
|
||||
if not file_path.is_file():
|
||||
return _text_block(f"Error: not a regular file: {file_path}")
|
||||
|
||||
try:
|
||||
content = file_path.read_text(encoding="utf-8")
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error reading {file_path}: {exc}")
|
||||
|
||||
count = content.count(old_string)
|
||||
if count == 0:
|
||||
return _text_block(
|
||||
f"Error: old_string not found in {file_path}. "
|
||||
"Make sure the string matches exactly, including whitespace and indentation."
|
||||
)
|
||||
|
||||
if not replace_all and count > 1:
|
||||
return _text_block(
|
||||
f"Error: old_string appears {count} times in {file_path}. "
|
||||
"Provide more surrounding context to make the match unique, "
|
||||
"or set replace_all=true to replace every occurrence."
|
||||
)
|
||||
|
||||
if replace_all:
|
||||
new_content = content.replace(old_string, new_string)
|
||||
else:
|
||||
# Replace only the first (and only) occurrence
|
||||
new_content = content.replace(old_string, new_string, 1)
|
||||
|
||||
try:
|
||||
file_path.write_text(new_content, encoding="utf-8")
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error writing {file_path}: {exc}")
|
||||
|
||||
replacements = count if replace_all else 1
|
||||
return _text_block(
|
||||
f"Successfully edited {file_path} ({replacements} replacement{'s' if replacements != 1 else ''})."
|
||||
)
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
# GlobTool
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GlobTool(BaseTool):
|
||||
name = "Glob"
|
||||
description = (
|
||||
"Fast file pattern matching. Supports glob patterns like '**/*.py'. "
|
||||
"Returns matching file paths sorted by modification time (newest first)."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"pattern": {
|
||||
"type": "string",
|
||||
"description": "Glob pattern to match files (e.g. '**/*.py', 'src/**/*.ts').",
|
||||
},
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "Directory to search in. Defaults to the working directory.",
|
||||
},
|
||||
},
|
||||
"required": ["pattern"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
pattern: str = input_data["pattern"]
|
||||
base = Path(input_data.get("path") or context.cwd)
|
||||
|
||||
if not base.is_dir():
|
||||
return _text_block(f"Error: directory not found: {base}")
|
||||
|
||||
try:
|
||||
matches: list[Path] = []
|
||||
for p in base.glob(pattern):
|
||||
if p.is_file():
|
||||
matches.append(p)
|
||||
if len(matches) >= 500:
|
||||
break
|
||||
|
||||
# Sort by modification time, newest first
|
||||
matches.sort(key=lambda p: p.stat().st_mtime, reverse=True)
|
||||
|
||||
if not matches:
|
||||
return _text_block(f"No files matched pattern '{pattern}' in {base}")
|
||||
|
||||
result = "\n".join(str(p) for p in matches)
|
||||
return _text_block(result)
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error during glob '{pattern}' in {base}: {exc}")
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
# GrepTool
|
||||
# ───────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GrepTool(BaseTool):
|
||||
name = "Grep"
|
||||
description = (
|
||||
"Search file contents using regular expressions. Uses ripgrep (rg) "
|
||||
"when available, otherwise falls back to Python's re module. "
|
||||
"Supports output modes: files_with_matches, content, count."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"pattern": {
|
||||
"type": "string",
|
||||
"description": "Regular expression pattern to search for.",
|
||||
},
|
||||
"path": {
|
||||
"type": "string",
|
||||
"description": "File or directory to search in. Defaults to the working directory.",
|
||||
},
|
||||
"glob": {
|
||||
"type": "string",
|
||||
"description": "Glob pattern to filter files (e.g. '*.py', '*.{ts,tsx}').",
|
||||
},
|
||||
"output_mode": {
|
||||
"type": "string",
|
||||
"enum": ["files_with_matches", "content", "count"],
|
||||
"description": "Output mode. Default: files_with_matches.",
|
||||
"default": "files_with_matches",
|
||||
},
|
||||
},
|
||||
"required": ["pattern"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
pattern: str = input_data["pattern"]
|
||||
search_path: str = input_data.get("path") or context.cwd
|
||||
file_glob: str | None = input_data.get("glob")
|
||||
output_mode: str = input_data.get("output_mode", "files_with_matches")
|
||||
|
||||
# Try ripgrep first
|
||||
try:
|
||||
result = await self._run_rg(pattern, search_path, file_glob, output_mode)
|
||||
if result is not None:
|
||||
return result
|
||||
except FileNotFoundError:
|
||||
pass # rg not installed, fall through to Python fallback
|
||||
|
||||
# Python fallback
|
||||
return await self._python_grep(pattern, search_path, file_glob, output_mode)
|
||||
|
||||
async def _run_rg(
|
||||
self,
|
||||
pattern: str,
|
||||
search_path: str,
|
||||
file_glob: str | None,
|
||||
output_mode: str,
|
||||
) -> list[dict] | None:
|
||||
"""Run ripgrep and return results, or None if rg is not available."""
|
||||
cmd = ["rg", "--no-heading", "--color=never"]
|
||||
|
||||
if output_mode == "files_with_matches":
|
||||
cmd.append("--files-with-matches")
|
||||
elif output_mode == "count":
|
||||
cmd.append("--count")
|
||||
else:
|
||||
cmd.extend(["--line-number"])
|
||||
|
||||
if file_glob:
|
||||
cmd.extend(["--glob", file_glob])
|
||||
|
||||
cmd.append(pattern)
|
||||
cmd.append(search_path)
|
||||
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
*cmd,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=30)
|
||||
except FileNotFoundError:
|
||||
raise # re-raise so caller knows rg is missing
|
||||
except asyncio.TimeoutError:
|
||||
return _text_block("Error: grep timed out after 30 seconds.")
|
||||
except Exception as exc:
|
||||
return _text_block(f"Error running ripgrep: {exc}")
|
||||
|
||||
output = stdout.decode("utf-8", errors="replace")
|
||||
|
||||
if proc.returncode not in (0, 1):
|
||||
err = stderr.decode("utf-8", errors="replace").strip()
|
||||
if err:
|
||||
return _text_block(f"Grep error: {err}")
|
||||
|
||||
if not output.strip():
|
||||
return _text_block(f"No matches found for pattern '{pattern}'.")
|
||||
|
||||
# Truncate if too large
|
||||
if len(output) > _MAX_OUTPUT_BYTES:
|
||||
output = output[:_MAX_OUTPUT_BYTES] + "\n... (output truncated)"
|
||||
|
||||
return _text_block(output.rstrip())
|
||||
|
||||
async def _python_grep(
|
||||
self,
|
||||
pattern: str,
|
||||
search_path: str,
|
||||
file_glob: str | None,
|
||||
output_mode: str,
|
||||
) -> list[dict]:
|
||||
"""Pure-Python grep fallback using the re module."""
|
||||
try:
|
||||
regex = re.compile(pattern)
|
||||
except re.error as exc:
|
||||
return _text_block(f"Invalid regex pattern: {exc}")
|
||||
|
||||
base = Path(search_path)
|
||||
if base.is_file():
|
||||
files = [base]
|
||||
elif base.is_dir():
|
||||
glob_pat = file_glob or "**/*"
|
||||
files = [p for p in base.glob(glob_pat) if p.is_file()]
|
||||
else:
|
||||
return _text_block(f"Error: path not found: {search_path}")
|
||||
|
||||
lines_out: list[str] = []
|
||||
total_bytes = 0
|
||||
truncated = False
|
||||
|
||||
for fp in sorted(files):
|
||||
try:
|
||||
text = fp.read_text(encoding="utf-8", errors="replace")
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
file_matches: list[tuple[int, str]] = []
|
||||
for lineno, line in enumerate(text.splitlines(), start=1):
|
||||
if regex.search(line):
|
||||
file_matches.append((lineno, line))
|
||||
|
||||
if not file_matches:
|
||||
continue
|
||||
|
||||
if output_mode == "files_with_matches":
|
||||
entry = str(fp)
|
||||
elif output_mode == "count":
|
||||
entry = f"{fp}:{len(file_matches)}"
|
||||
else:
|
||||
parts = [f"{fp}:{ln}:{txt}" for ln, txt in file_matches]
|
||||
entry = "\n".join(parts)
|
||||
|
||||
total_bytes += len(entry)
|
||||
if total_bytes > _MAX_OUTPUT_BYTES:
|
||||
truncated = True
|
||||
break
|
||||
|
||||
lines_out.append(entry)
|
||||
|
||||
if not lines_out:
|
||||
return _text_block(f"No matches found for pattern '{pattern}'.")
|
||||
|
||||
result = "\n".join(lines_out)
|
||||
if truncated:
|
||||
result += "\n... (output truncated)"
|
||||
|
||||
return _text_block(result)
|
||||
@@ -1,61 +0,0 @@
|
||||
"""Central tool registry.
|
||||
|
||||
Importing this module automatically registers all builtin tools.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from backend.apps.agents.tools.base import BaseTool
|
||||
from backend.apps.agents.providers.base import ToolSchema
|
||||
|
||||
_TOOLS: dict[str, BaseTool] = {}
|
||||
|
||||
|
||||
def register_tool(tool: BaseTool) -> None:
|
||||
"""Register a tool instance by its name."""
|
||||
_TOOLS[tool.name] = tool
|
||||
|
||||
|
||||
def get_tool(name: str) -> BaseTool | None:
|
||||
"""Look up a registered tool by name. Returns None if not found."""
|
||||
return _TOOLS.get(name)
|
||||
|
||||
|
||||
def get_all_tools() -> list[BaseTool]:
|
||||
"""Return all registered tool instances."""
|
||||
return list(_TOOLS.values())
|
||||
|
||||
|
||||
def get_all_tool_schemas() -> list[ToolSchema]:
|
||||
"""Return provider-agnostic ToolSchema for every registered tool."""
|
||||
return [t.to_tool_schema() for t in _TOOLS.values()]
|
||||
|
||||
|
||||
def init_tools() -> None:
|
||||
"""Import and register all builtin tools."""
|
||||
from backend.apps.agents.tools.filesystem import (
|
||||
ReadTool,
|
||||
WriteTool,
|
||||
EditTool,
|
||||
GlobTool,
|
||||
GrepTool,
|
||||
)
|
||||
from backend.apps.agents.tools.system import BashTool, AskUserQuestionTool
|
||||
from backend.apps.agents.tools.web import WebSearchTool, WebFetchTool
|
||||
|
||||
for tool_cls in [
|
||||
ReadTool,
|
||||
WriteTool,
|
||||
EditTool,
|
||||
GlobTool,
|
||||
GrepTool,
|
||||
BashTool,
|
||||
AskUserQuestionTool,
|
||||
WebSearchTool,
|
||||
WebFetchTool,
|
||||
]:
|
||||
register_tool(tool_cls())
|
||||
|
||||
|
||||
# Auto-register on import
|
||||
init_tools()
|
||||
@@ -1,125 +0,0 @@
|
||||
"""System tools: Bash and AskUserQuestion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from backend.apps.agents.tools.base import BaseTool, ToolContext
|
||||
|
||||
_MAX_OUTPUT_BYTES = 100 * 1024 # ~100 KB cap
|
||||
|
||||
|
||||
class BashTool(BaseTool):
|
||||
name = "Bash"
|
||||
description = (
|
||||
"Execute a shell command and return its output. The command runs in "
|
||||
"the session's working directory. Supports an optional timeout "
|
||||
"(default 120 000 ms). Stdout and stderr are captured and returned."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"command": {
|
||||
"type": "string",
|
||||
"description": "The shell command to execute.",
|
||||
},
|
||||
"timeout": {
|
||||
"type": "integer",
|
||||
"description": "Timeout in milliseconds (default 120000, max 600000).",
|
||||
"default": 120000,
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "Optional human-readable description of what this command does.",
|
||||
},
|
||||
},
|
||||
"required": ["command"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
command: str = input_data["command"]
|
||||
timeout_ms: int = min(input_data.get("timeout", 120000), 600000)
|
||||
timeout_s: float = timeout_ms / 1000.0
|
||||
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_shell(
|
||||
command,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=context.cwd,
|
||||
)
|
||||
except Exception as exc:
|
||||
return [{"type": "text", "text": f"Error starting command: {exc}"}]
|
||||
|
||||
try:
|
||||
stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=timeout_s)
|
||||
except asyncio.TimeoutError:
|
||||
# Attempt to kill the process
|
||||
try:
|
||||
proc.kill()
|
||||
stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=5)
|
||||
except Exception:
|
||||
stdout, stderr = b"", b""
|
||||
|
||||
partial = self._decode(stdout, stderr)
|
||||
msg = (
|
||||
f"Command timed out after {timeout_ms}ms.\n"
|
||||
f"Partial output:\n{partial}"
|
||||
)
|
||||
return [{"type": "text", "text": self._truncate(msg)}]
|
||||
except Exception as exc:
|
||||
return [{"type": "text", "text": f"Error executing command: {exc}"}]
|
||||
|
||||
output = self._decode(stdout, stderr)
|
||||
|
||||
if proc.returncode != 0:
|
||||
output = f"Exit code: {proc.returncode}\n{output}"
|
||||
|
||||
if not output.strip():
|
||||
output = f"(command completed with exit code {proc.returncode})"
|
||||
|
||||
return [{"type": "text", "text": self._truncate(output)}]
|
||||
|
||||
@staticmethod
|
||||
def _decode(stdout: bytes, stderr: bytes) -> str:
|
||||
parts: list[str] = []
|
||||
if stdout:
|
||||
parts.append(stdout.decode("utf-8", errors="replace"))
|
||||
if stderr:
|
||||
parts.append(stderr.decode("utf-8", errors="replace"))
|
||||
return "\n".join(parts)
|
||||
|
||||
@staticmethod
|
||||
def _truncate(text: str) -> str:
|
||||
if len(text) > _MAX_OUTPUT_BYTES:
|
||||
return text[:_MAX_OUTPUT_BYTES] + "\n... (output truncated)"
|
||||
return text
|
||||
|
||||
|
||||
class AskUserQuestionTool(BaseTool):
|
||||
name = "AskUserQuestion"
|
||||
description = (
|
||||
"Ask the user a clarifying question. The actual blocking/HITL "
|
||||
"interaction is handled by the agent loop's hitl_handler; this tool "
|
||||
"simply surfaces the question text."
|
||||
)
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"question": {
|
||||
"type": "string",
|
||||
"description": "The question to ask the user.",
|
||||
},
|
||||
},
|
||||
"required": ["question"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
|
||||
async def execute(self, input_data: dict, context: ToolContext) -> list[dict]:
|
||||
question: str = input_data.get("question", "")
|
||||
return [{"type": "text", "text": question}]
|
||||
@@ -1,8 +0,0 @@
|
||||
"""Stdio MCP shim that forwards Discord tool calls to the OpenSwarm cloud.
|
||||
|
||||
Run as: python -m backend.apps.discord_mcp_shim
|
||||
"""
|
||||
from backend.apps.discord_mcp_shim.server import main
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -1,27 +1,8 @@
|
||||
"""Stress tests for the Phase 1 / 2 / 3 perceived-latency changes.
|
||||
|
||||
Hits everything we touched on the eric/v2 branch:
|
||||
"""Stress tests for live perceived-latency paths.
|
||||
|
||||
- Message.client_message_id round-trip (optimistic dedupe)
|
||||
- Mode migration: 'chat' -> 'ask' on session reconcile + lifespan
|
||||
deletion of stale built-in chat.json
|
||||
- ContentBlock + StreamEvent now accept type='thinking' /
|
||||
delta_type='thinking_delta' without breaking existing types
|
||||
- Anthropic provider forwards thinking content_block_start /
|
||||
content_block_delta with the right shape
|
||||
- Agent loop emits agent:stream_start{role:'thinking'},
|
||||
agent:stream_delta, agent:stream_end for thinking blocks AND
|
||||
persists a Message(role='thinking') after stream end
|
||||
- DashboardLayout serializes notes round-trip
|
||||
- exclude_dynamic_sections reaches the SDK kwargs (presence-only;
|
||||
we don't run the real CLI here)
|
||||
|
||||
Each test runs many randomized iterations to surface race conditions
|
||||
and bad assumptions. Stub the network and CLI throughout — these
|
||||
tests are pure logic, no real Anthropic calls.
|
||||
|
||||
Run:
|
||||
cd backend && .venv/bin/python -m pytest tests/test_phase1_stress.py -v
|
||||
- Mode migration: 'chat' -> 'ask' on reconcile + lifespan deletion
|
||||
- DashboardLayout notes round-trip
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -213,288 +194,6 @@ def test_reconcile_idempotent():
|
||||
assert mtime_after_first == mtime_after_second, "reconcile must be idempotent"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 3 — ContentBlock / StreamEvent thinking acceptance
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_content_block_thinking_type():
|
||||
from backend.apps.agents.providers.base import ContentBlock
|
||||
|
||||
cb = ContentBlock(type="thinking", text="some reasoning")
|
||||
assert cb.type == "thinking"
|
||||
assert cb.text == "some reasoning"
|
||||
assert cb.tool_call is None
|
||||
|
||||
|
||||
def test_stream_event_thinking_delta():
|
||||
from backend.apps.agents.providers.base import StreamEvent
|
||||
|
||||
e = StreamEvent(type="content_block_delta", delta_type="thinking_delta", text="hmm")
|
||||
assert e.delta_type == "thinking_delta"
|
||||
assert e.text == "hmm"
|
||||
|
||||
# Existing types still work — no regression
|
||||
e2 = StreamEvent(type="content_block_delta", delta_type="text_delta", text="hi")
|
||||
assert e2.delta_type == "text_delta"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 4 — Anthropic provider thinking forwarding
|
||||
#
|
||||
# We feed a fake raw_stream (mimicking the SDK's async generator) through
|
||||
# AnthropicProvider.stream_message and confirm the right StreamEvents come
|
||||
# out. No network.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeRawEvent:
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
class _FakeBlock:
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
class _FakeDelta:
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_provider_forwards_thinking_blocks():
|
||||
"""Mock the raw Anthropic stream with a thinking block + thinking_delta
|
||||
+ content_block_stop, and assert AnthropicProvider yields the
|
||||
normalized StreamEvents the agent_loop expects."""
|
||||
from backend.apps.agents.providers.anthropic import AnthropicProvider
|
||||
|
||||
raw_events = [
|
||||
# thinking block opens at index 0
|
||||
_FakeRawEvent(type="content_block_start", index=0,
|
||||
content_block=_FakeBlock(type="thinking")),
|
||||
_FakeRawEvent(type="content_block_delta", index=0,
|
||||
delta=_FakeDelta(type="thinking_delta", thinking="step 1, ")),
|
||||
_FakeRawEvent(type="content_block_delta", index=0,
|
||||
delta=_FakeDelta(type="thinking_delta", thinking="step 2.")),
|
||||
# signature_delta on thinking — must be ignored, not crash
|
||||
_FakeRawEvent(type="content_block_delta", index=0,
|
||||
delta=_FakeDelta(type="signature_delta", signature="abc==")),
|
||||
_FakeRawEvent(type="content_block_stop", index=0),
|
||||
# text block follows at index 1
|
||||
_FakeRawEvent(type="content_block_start", index=1,
|
||||
content_block=_FakeBlock(type="text")),
|
||||
_FakeRawEvent(type="content_block_delta", index=1,
|
||||
delta=_FakeDelta(type="text_delta", text="hi")),
|
||||
_FakeRawEvent(type="content_block_stop", index=1),
|
||||
]
|
||||
|
||||
async def fake_stream():
|
||||
for ev in raw_events:
|
||||
yield ev
|
||||
|
||||
# AnthropicProvider takes api_key/auth_token/base_url; we monkeypatch
|
||||
# its `client.messages.create` after construction so no real
|
||||
# SDK client is needed.
|
||||
provider = AnthropicProvider(api_key="test-key")
|
||||
provider.client.messages.create = AsyncMock(return_value=fake_stream())
|
||||
out_events = []
|
||||
async for ev in provider.stream_message(model="sonnet", system=None, messages=[], tools=[]):
|
||||
out_events.append(ev)
|
||||
|
||||
types = [(e.type, e.block_type, e.delta_type) for e in out_events]
|
||||
# Thinking block should produce: start, 2x delta, stop. signature_delta ignored.
|
||||
assert ("content_block_start", "thinking", "") in types
|
||||
assert types.count(("content_block_delta", "", "thinking_delta")) == 2
|
||||
assert ("content_block_start", "text", "") in types
|
||||
assert ("content_block_delta", "", "text_delta") in types
|
||||
|
||||
thinking_text = "".join(
|
||||
e.text for e in out_events
|
||||
if e.type == "content_block_delta" and e.delta_type == "thinking_delta"
|
||||
)
|
||||
assert thinking_text == "step 1, step 2."
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 5 — Agent loop end-to-end thinking → WS events + persisted message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_loop_emits_thinking_stream_and_persists_message():
|
||||
"""Drive the agent loop with a fake provider that yields thinking,
|
||||
text, and one tool_use. Verify it emits the right WS events AND
|
||||
persists a Message(role='thinking') via _emit_collected_messages."""
|
||||
from backend.apps.agents.providers.base import StreamEvent
|
||||
|
||||
captured_ws: list[tuple[str, dict]] = []
|
||||
|
||||
async def fake_emitter(event: str, payload: dict):
|
||||
captured_ws.append((event, payload))
|
||||
|
||||
# Build a fake provider yielding our normalized StreamEvents.
|
||||
class FakeProvider:
|
||||
async def stream_message(self, **kwargs):
|
||||
yield StreamEvent(type="content_block_start", index=0, block_type="thinking")
|
||||
yield StreamEvent(type="content_block_delta", index=0,
|
||||
delta_type="thinking_delta", text="reasoning… ")
|
||||
yield StreamEvent(type="content_block_delta", index=0,
|
||||
delta_type="thinking_delta", text="more.")
|
||||
yield StreamEvent(type="content_block_stop", index=0)
|
||||
yield StreamEvent(type="content_block_start", index=1, block_type="text")
|
||||
yield StreamEvent(type="content_block_delta", index=1,
|
||||
delta_type="text_delta", text="hello!")
|
||||
yield StreamEvent(type="content_block_stop", index=1)
|
||||
yield StreamEvent(type="message_stop")
|
||||
|
||||
from backend.apps.agents.agent_loop import AgentLoop
|
||||
|
||||
loop = AgentLoop(
|
||||
session_id="s1",
|
||||
provider=FakeProvider(),
|
||||
model="sonnet",
|
||||
system_prompt="x",
|
||||
tools=[],
|
||||
ws_emitter=fake_emitter,
|
||||
hitl_handler=AsyncMock(return_value=(True, None)),
|
||||
tool_executor=AsyncMock(return_value=[{"type": "text", "text": "ok"}]),
|
||||
)
|
||||
|
||||
response = await loop._stream_and_collect()
|
||||
|
||||
# Stream events: thinking start + 2 deltas + stream_end, then text start + delta + (text end at message_stop)
|
||||
events_by_type = {}
|
||||
for ev, payload in captured_ws:
|
||||
events_by_type.setdefault(ev, []).append(payload)
|
||||
|
||||
# Thinking should have its own stream_start with role='thinking'
|
||||
starts = events_by_type.get("agent:stream_start", [])
|
||||
thinking_starts = [s for s in starts if s.get("role") == "thinking"]
|
||||
assistant_starts = [s for s in starts if s.get("role") == "assistant"]
|
||||
assert len(thinking_starts) == 1, f"expected 1 thinking start, got {len(thinking_starts)}"
|
||||
assert len(assistant_starts) == 1, "expected 1 assistant text start"
|
||||
|
||||
# Two thinking deltas
|
||||
deltas = events_by_type.get("agent:stream_delta", [])
|
||||
thinking_msg_id = thinking_starts[0]["message_id"]
|
||||
thinking_deltas = [d for d in deltas if d.get("message_id") == thinking_msg_id]
|
||||
assert len(thinking_deltas) == 2
|
||||
assert "".join(d["delta"] for d in thinking_deltas) == "reasoning… more."
|
||||
|
||||
# Thinking stream_end fires (text doesn't get stream_end inside _stream_and_collect — closes at message_stop)
|
||||
ends = events_by_type.get("agent:stream_end", [])
|
||||
assert any(e["message_id"] == thinking_msg_id for e in ends), "thinking must emit stream_end"
|
||||
|
||||
# Now persist via _emit_collected_messages and verify a thinking
|
||||
# Message went out
|
||||
captured_ws.clear()
|
||||
await loop._emit_collected_messages(
|
||||
response.content,
|
||||
text_msg_id=assistant_starts[0]["message_id"],
|
||||
tool_msg_ids={},
|
||||
)
|
||||
persisted = [p for ev, p in captured_ws if ev == "agent:message"]
|
||||
roles = [p["message"]["role"] for p in persisted]
|
||||
assert "thinking" in roles, "thinking content must be persisted as a Message"
|
||||
assert "assistant" in roles
|
||||
thinking_msg = next(p for p in persisted if p["message"]["role"] == "thinking")
|
||||
assert thinking_msg["message"]["content"] == "reasoning… more."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_loop_handles_no_thinking_gracefully():
|
||||
"""Provider that emits zero thinking blocks must still work.
|
||||
Regression guard against the new branch breaking text-only paths."""
|
||||
from backend.apps.agents.providers.base import StreamEvent
|
||||
from backend.apps.agents.agent_loop import AgentLoop
|
||||
|
||||
captured_ws = []
|
||||
|
||||
async def fake_emitter(event, payload):
|
||||
captured_ws.append((event, payload))
|
||||
|
||||
class TextOnly:
|
||||
async def stream_message(self, **kwargs):
|
||||
yield StreamEvent(type="content_block_start", index=0, block_type="text")
|
||||
yield StreamEvent(type="content_block_delta", index=0,
|
||||
delta_type="text_delta", text="just text")
|
||||
yield StreamEvent(type="content_block_stop", index=0)
|
||||
yield StreamEvent(type="message_stop")
|
||||
|
||||
loop = AgentLoop(
|
||||
session_id="s2", provider=TextOnly(), model="sonnet", system_prompt=None,
|
||||
tools=[],
|
||||
ws_emitter=fake_emitter,
|
||||
hitl_handler=AsyncMock(return_value=(True, None)),
|
||||
tool_executor=AsyncMock(return_value=[]),
|
||||
)
|
||||
|
||||
resp = await loop._stream_and_collect()
|
||||
starts = [p for ev, p in captured_ws if ev == "agent:stream_start"]
|
||||
# Exactly one assistant start, zero thinking starts
|
||||
assert len([s for s in starts if s.get("role") == "thinking"]) == 0
|
||||
assert len([s for s in starts if s.get("role") == "assistant"]) == 1
|
||||
assert any(b.type == "text" for b in resp.content)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_loop_stress_many_thinking_blocks():
|
||||
"""Hammer the loop with a long sequence of interleaved thinking +
|
||||
text + tool blocks. Ensures the per-index buffers don't leak and
|
||||
every block gets the right WS events."""
|
||||
from backend.apps.agents.providers.base import StreamEvent
|
||||
from backend.apps.agents.agent_loop import AgentLoop
|
||||
|
||||
captured = []
|
||||
|
||||
async def fake_emitter(ev, p):
|
||||
captured.append((ev, p))
|
||||
|
||||
class Mix:
|
||||
async def stream_message(self, **kwargs):
|
||||
idx = 0
|
||||
for turn in range(40):
|
||||
yield StreamEvent(type="content_block_start", index=idx, block_type="thinking")
|
||||
for _ in range(random.randint(1, 5)):
|
||||
yield StreamEvent(type="content_block_delta", index=idx,
|
||||
delta_type="thinking_delta", text=f"t{idx} ")
|
||||
yield StreamEvent(type="content_block_stop", index=idx)
|
||||
idx += 1
|
||||
yield StreamEvent(type="content_block_start", index=idx, block_type="text")
|
||||
yield StreamEvent(type="content_block_delta", index=idx,
|
||||
delta_type="text_delta", text=f"text-{idx}")
|
||||
yield StreamEvent(type="content_block_stop", index=idx)
|
||||
idx += 1
|
||||
yield StreamEvent(type="message_stop")
|
||||
|
||||
loop = AgentLoop(
|
||||
session_id="s3", provider=Mix(), model="sonnet", system_prompt=None,
|
||||
tools=[],
|
||||
ws_emitter=fake_emitter,
|
||||
hitl_handler=AsyncMock(return_value=(True, None)),
|
||||
tool_executor=AsyncMock(return_value=[]),
|
||||
)
|
||||
resp = await loop._stream_and_collect()
|
||||
|
||||
starts = [p for ev, p in captured if ev == "agent:stream_start"]
|
||||
ends = [p for ev, p in captured if ev == "agent:stream_end"]
|
||||
|
||||
# 40 thinking + 1 assistant (text accumulates into one stream_text_msg_id)
|
||||
thinking_starts = [s for s in starts if s.get("role") == "thinking"]
|
||||
assistant_starts = [s for s in starts if s.get("role") == "assistant"]
|
||||
assert len(thinking_starts) == 40, f"got {len(thinking_starts)} thinking starts, want 40"
|
||||
assert len(assistant_starts) == 1, "all text blocks share one assistant stream id"
|
||||
|
||||
# Each thinking block must have its own stream_end
|
||||
thinking_ids = {s["message_id"] for s in thinking_starts}
|
||||
end_ids = {e["message_id"] for e in ends}
|
||||
assert thinking_ids.issubset(end_ids), "every thinking block needs a stream_end"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 6 — Notes layout serialization
|
||||
|
||||
@@ -32,10 +32,7 @@ const streamingCursorKeyframes = `
|
||||
}
|
||||
`;
|
||||
|
||||
// Claude.ai-style shimmer that sweeps left → right across text while the
|
||||
// model is actively thinking. Uses background-clip: text to mask a moving
|
||||
// linear gradient onto the text glyphs so the effect looks like a light
|
||||
// wave traveling through the letters.
|
||||
// shimmer-on-text effect for thinking. background-clip:text + sliding gradient.
|
||||
const thinkingShimmerKeyframes = `
|
||||
@keyframes thinking-shimmer {
|
||||
0% { background-position: 200% 0; }
|
||||
@@ -73,12 +70,9 @@ interface OpenSwarmErrorInfo {
|
||||
ctaAction?: 'upgrade' | 'retry' | 'settings' | 'waitlist';
|
||||
}
|
||||
|
||||
// Turn a raw Claude-CLI / cloud error string into a user-friendly card.
|
||||
// Returns null for things that aren't obviously our errors — those fall
|
||||
// through to normal markdown rendering.
|
||||
// raw error text into a friendly card. null = not ours, render as markdown.
|
||||
function parseOpenSwarmError(text: string): OpenSwarmErrorInfo | null {
|
||||
if (!text) return null;
|
||||
// Rate-limit cap from our cloud
|
||||
if (/rate_limit_error|reached your OpenSwarm.*plan limit|Usage cap exceeded/i.test(text)) {
|
||||
const reset = text.match(/Resets in ([\dhms\s]+)/)?.[1];
|
||||
return {
|
||||
@@ -91,14 +85,7 @@ function parseOpenSwarmError(text: string): OpenSwarmErrorInfo | null {
|
||||
ctaAction: 'upgrade',
|
||||
};
|
||||
}
|
||||
// Upstream capacity / 503 / transient. The backend already retries these
|
||||
// for ~5.5 minutes (5/15/45/90/180s) before bubbling up, so by the time a
|
||||
// user sees this the system has genuinely struggled — but it's almost
|
||||
// always recoverable on the next send, not a plan/billing issue. Show a
|
||||
// soft "connection hiccup" card instead of the waitlist/"servers maxed"
|
||||
// copy, which misleads Pro/Pro+/Ultra subscribers into thinking their
|
||||
// paid plan is out of capacity. The only real hard cap a user should see
|
||||
// is their own per-plan 5h limit (matched above as `kind: 'cap'`).
|
||||
// backend retried for ~5.5min before bubbling. show a soft hiccup, not a cap.
|
||||
if (/at capacity|Try again shortly|503|service unavailable/i.test(text)) {
|
||||
return {
|
||||
kind: 'network',
|
||||
@@ -106,11 +93,7 @@ function parseOpenSwarmError(text: string): OpenSwarmErrorInfo | null {
|
||||
detail: 'That request timed out after a few retries. Send the message again to continue.',
|
||||
};
|
||||
}
|
||||
// Too many MCP tool definitions for the chosen model's input window.
|
||||
// Classic case: user has 5+ apps connected (M365 alone has 141 actions),
|
||||
// chose Haiku (200K context), and even a one-line message can't fit
|
||||
// because the tool schemas alone push past the limit. Bigger models
|
||||
// (Sonnet/Opus, 1M) absorb it fine.
|
||||
// tool schemas overflowed the window. M365 alone is 141 actions.
|
||||
if (/Prompt is too long|prompt_too_long|input length and `max_tokens`|context length/i.test(text)) {
|
||||
return {
|
||||
kind: 'too_many_tools',
|
||||
@@ -125,7 +108,6 @@ function parseOpenSwarmError(text: string): OpenSwarmErrorInfo | null {
|
||||
ctaAction: 'settings',
|
||||
};
|
||||
}
|
||||
// Auth / subscription problems
|
||||
if (/No active subscription|Subscription canceled|Subscription past_due|Invalid.*token|Missing bearer token/i.test(text)) {
|
||||
return {
|
||||
kind: 'auth',
|
||||
@@ -135,15 +117,7 @@ function parseOpenSwarmError(text: string): OpenSwarmErrorInfo | null {
|
||||
ctaAction: 'settings',
|
||||
};
|
||||
}
|
||||
// Genuine, hard network failures only. The bare word `network` used to
|
||||
// match anything mentioning "network" (Python traces, MCP tool output,
|
||||
// ffmpeg lines, etc.), and `fetch failed` / `ETIMEDOUT` alone fire for
|
||||
// transient upstream blips the backend now silently retries — surfacing
|
||||
// a card for those just confuses the user. So: require the specific
|
||||
// errno codes at word boundaries, and only match `fetch failed` when
|
||||
// paired with a concrete cause so we don't swallow every Node-level
|
||||
// transient. The backend's capacity/transient retry layer handles the
|
||||
// rest without ever reaching this classifier.
|
||||
// strict matchers only. bare "network" used to false-match Python traces.
|
||||
if (/\b(?:ECONNREFUSED|ENETUNREACH|ENOTFOUND|EAI_AGAIN)\b|Could\s+not\s+reach\s+OpenSwarm|Unable\s+to\s+connect\s+to\s+OpenSwarm/i.test(text)) {
|
||||
return {
|
||||
kind: 'network',
|
||||
@@ -491,57 +465,33 @@ const MessageImageThumbnails: React.FC<{
|
||||
);
|
||||
};
|
||||
|
||||
// ── ThinkingBubble ──────────────────────────────────────────────────
|
||||
// Collapsible reasoning section styled after Claude.ai / ChatGPT /
|
||||
// Gemini. Defaults to expanded so thinking is always visible when
|
||||
// present. User can click the header to collapse. If we observed the
|
||||
// stream live we show "Thought for Ns"; otherwise (history replay) we
|
||||
// just show "Thoughts".
|
||||
// thinking pill. shows "Thought for Ns" if we caught it live, else just "Thoughts".
|
||||
const ThinkingBubble: React.FC<{
|
||||
content: string;
|
||||
isStreaming?: boolean;
|
||||
timestamp?: string;
|
||||
// Server-stamped duration / token count, populated on the persisted
|
||||
// Message at end-of-stream. When present, post-stream label uses these
|
||||
// exact values instead of the in-memory React-state estimates that
|
||||
// disappear when the streaming bubble unmounts.
|
||||
// server-stamped totals for the turn. survives unmount.
|
||||
persistedElapsedMs?: number;
|
||||
persistedTokens?: number;
|
||||
// Server-stamped input-side total for the turn (fresh + cache-creation
|
||||
// + cache-read). Used to render "M in" alongside the existing "K out"
|
||||
// segment so the pill honestly reflects the full turn cost, not just
|
||||
// output. Optional — turns with no SDK usage data (rare) skip it.
|
||||
persistedInputTokens?: number;
|
||||
// Tool invocation count for this turn — drives the "3 tools used"
|
||||
// segment of the post-stream label.
|
||||
persistedToolCount?: number;
|
||||
// Aux-LLM-generated dynamic label for the active turn ("Auditing the
|
||||
// pull request", "Drafting your email"). Replaces the static
|
||||
// "Thinking…" verb when present and the stream is still active.
|
||||
// aux-LLM label like "Auditing the pull request". null = use the heuristic.
|
||||
dynamicLabel?: string | null;
|
||||
}> = ({ content, isStreaming, persistedElapsedMs, persistedTokens, persistedInputTokens, persistedToolCount, dynamicLabel }) => {
|
||||
const c = useClaudeTokens();
|
||||
|
||||
// Live timer is only used as a fallback when we don't yet have
|
||||
// server-stamped persistedElapsedMs. The pill stays in "Thinking…"
|
||||
// for the entire duration of a multi-block turn (think → tool →
|
||||
// think → answer), and only swaps to "Thought for Ns · M tokens"
|
||||
// once persistedElapsedMs lands via the agent:message event for
|
||||
// role='thinking', which carries the per-turn aggregate (not the
|
||||
// per-block stats the live UI used to freeze on prematurely).
|
||||
// live timer is just the fallback. server-stamped values win.
|
||||
const [startedStreamingAt, setStartedStreamingAt] = useState<number | null>(
|
||||
isStreaming ? Date.now() : null
|
||||
);
|
||||
const [elapsed, setElapsed] = useState<number>(0);
|
||||
|
||||
// Record start time the first time we see streaming
|
||||
React.useEffect(() => {
|
||||
if (isStreaming && startedStreamingAt === null) {
|
||||
setStartedStreamingAt(Date.now());
|
||||
}
|
||||
}, [isStreaming, startedStreamingAt]);
|
||||
|
||||
// Tick the timer while streaming
|
||||
React.useEffect(() => {
|
||||
if (!isStreaming || startedStreamingAt === null) return;
|
||||
const iv = setInterval(() => {
|
||||
@@ -550,29 +500,15 @@ const ThinkingBubble: React.FC<{
|
||||
return () => clearInterval(iv);
|
||||
}, [isStreaming, startedStreamingAt]);
|
||||
|
||||
// Default behavior: expanded while streaming (so the user can watch
|
||||
// the model think live), collapsed after the turn ends (so the
|
||||
// transcript reads as answer-first, with reasoning available on click).
|
||||
// userOverride captures explicit clicks and pins the state — once the
|
||||
// user has chosen, we respect their pick across the streaming →
|
||||
// post-stream transition. This avoids the wall-of-text problem where
|
||||
// a 1.6K-token reasoning block stayed expanded after the turn finished.
|
||||
// expanded while streaming, collapsed after. userOverride pins explicit clicks.
|
||||
const [userOverride, setUserOverride] = useState<boolean | null>(null);
|
||||
const expanded = userOverride ?? !!isStreaming;
|
||||
const toggle = () => setUserOverride(!expanded);
|
||||
|
||||
const text = typeof content === 'string' ? content : JSON.stringify(content);
|
||||
// Live token estimate uses Anthropic's BPE-ish ratio for English prose
|
||||
// (~3.6 chars/token) instead of the cruder /4. Still an estimate — true
|
||||
// value lands via persistedTokens when the stream ends.
|
||||
// 3.6 chars/token for English. swap for persistedTokens once the stream ends.
|
||||
const liveTokenEstimate = isStreaming ? Math.max(0, Math.round(text.length / 3.6)) : 0;
|
||||
|
||||
// Post-stream label preference order:
|
||||
// 1. Server-stamped persisted values (per-turn aggregate, survives
|
||||
// reload — this is the truth source we actually want).
|
||||
// 2. Live React-state elapsed (only used if server values are
|
||||
// missing, e.g. legacy messages).
|
||||
// 3. Generic "Thoughts" fallback.
|
||||
const persistedSecs = persistedElapsedMs != null
|
||||
? Math.max(1, Math.round(persistedElapsedMs / 1000))
|
||||
: null;
|
||||
@@ -583,17 +519,10 @@ const ThinkingBubble: React.FC<{
|
||||
const finalTokens = persistedTokens
|
||||
?? (text && !isStreaming ? Math.max(1, Math.round(text.length / 3.6)) : null);
|
||||
|
||||
// Active-stream label preference:
|
||||
// 1. Aux-LLM dynamic label ("Auditing the pull request") when available.
|
||||
// 2. Heuristic "Thinking…" with token estimate as the fallback.
|
||||
// The dynamic label only replaces the verb part — token count chip
|
||||
// appends after, so users still see the live counter.
|
||||
const activeLabel = dynamicLabel
|
||||
? (liveTokenEstimate > 0 ? `${dynamicLabel}… · ~${liveTokenEstimate} tokens` : `${dynamicLabel}…`)
|
||||
: (liveTokenEstimate > 0 ? `Thinking… (~${liveTokenEstimate} tokens)` : 'Thinking…');
|
||||
|
||||
// Compact number formatter for the post-stream label — "2.4K" beats
|
||||
// "2400" once token counts get large.
|
||||
const fmtTokens = (n: number) => {
|
||||
if (n >= 1000) {
|
||||
const k = n / 1000;
|
||||
@@ -602,10 +531,7 @@ const ThinkingBubble: React.FC<{
|
||||
return String(n);
|
||||
};
|
||||
|
||||
// Duration formatter that rolls over at minute / hour boundaries so
|
||||
// "Thought for 251s" reads as "Thought for 4m 11s" — same shape the
|
||||
// header chip uses. Mirrors the AgentCard fmtSeconds helper but kept
|
||||
// local so the bubble stays self-contained.
|
||||
// 251s reads as "4m 11s". mirrors AgentCard's fmtSeconds.
|
||||
const fmtThoughtDuration = (sec: number) => {
|
||||
if (sec < 60) return `${sec}s`;
|
||||
const minutes = Math.floor(sec / 60);
|
||||
@@ -618,31 +544,16 @@ const ThinkingBubble: React.FC<{
|
||||
return remMin > 0 ? `${hours}h ${remMin}m` : `${hours}h`;
|
||||
};
|
||||
|
||||
// Post-stream label: "Thought for Ns · 32 tokens · 3 tools used".
|
||||
// The reasoning-token count is the honest signal of how much thinking
|
||||
// happened; tool count surfaces work done; duration surfaces wait time.
|
||||
// We deliberately omit a separate "answer tokens" number — earlier
|
||||
// experiments showed it confused users (it counted both visible reply
|
||||
// text AND tool-call JSON arguments, making tool-heavy turns
|
||||
// misleadingly look like long answers).
|
||||
// Backend stamps `input_tokens` as the all-in input+output+children
|
||||
// total (parent's primary call PLUS every subagent and tool MCP that
|
||||
// booked usage on this turn). Falls back to just-output (finalTokens)
|
||||
// for legacy thinking messages that predate the combined-total field.
|
||||
// input_tokens is the full turn cost (parent + subagents + tool MCPs).
|
||||
// legacy messages without it fall back to output-only.
|
||||
const combinedTotalTokens =
|
||||
persistedInputTokens != null && persistedInputTokens > 0
|
||||
? persistedInputTokens
|
||||
: finalTokens;
|
||||
// Input/output split shown in the breakdown tooltip on click. We
|
||||
// already have `finalTokens` (server-stamped output side) and
|
||||
// `combinedTotalTokens` (input + output + children sum). The
|
||||
// implied "input + children" portion is the difference. When the
|
||||
// backend hasn't separated them yet (legacy data), we still show
|
||||
// the total but skip the breakdown.
|
||||
// tooltip breakdown. legacy data without finalTokens shows total only.
|
||||
const tokenBreakdown = (() => {
|
||||
if (combinedTotalTokens == null || combinedTotalTokens <= 0) return null;
|
||||
if (finalTokens == null || finalTokens <= 0) {
|
||||
// Total-only case (rare). No split available.
|
||||
return { total: combinedTotalTokens, output: null as number | null, input: null as number | null };
|
||||
}
|
||||
const inputSide = Math.max(0, combinedTotalTokens - finalTokens);
|
||||
@@ -712,17 +623,9 @@ const ThinkingBubble: React.FC<{
|
||||
return segments;
|
||||
};
|
||||
|
||||
// Streaming gets a plain string label (the shimmer animation needs
|
||||
// the text to flow through a single gradient mask, which only works
|
||||
// on a flat string node). Post-stream uses the React-node renderer
|
||||
// so the tokens segment can be wrapped in a Tooltip with the
|
||||
// input/output breakdown.
|
||||
// shimmer needs a flat string. post-stream uses nodes for the tooltip.
|
||||
const label: React.ReactNode = isStreaming ? activeLabel : renderPostStreamLabel();
|
||||
|
||||
// Shimmer colors — use a bright mid-tone against the muted base to make
|
||||
// the sweep visible without being loud. The base color matches the
|
||||
// static "Thought for Ns" state so the only visible change is the moving
|
||||
// highlight band.
|
||||
const shimmerBase = c.text.tertiary;
|
||||
const shimmerHighlight = c.text.primary;
|
||||
|
||||
@@ -753,7 +656,6 @@ const ThinkingBubble: React.FC<{
|
||||
fontSize: '0.78rem',
|
||||
fontWeight: 500,
|
||||
...(isStreaming ? {
|
||||
// Moving gradient masked onto the text glyphs
|
||||
background: `linear-gradient(90deg, ${shimmerBase} 0%, ${shimmerBase} 40%, ${shimmerHighlight} 50%, ${shimmerBase} 60%, ${shimmerBase} 100%)`,
|
||||
backgroundSize: '200% 100%',
|
||||
WebkitBackgroundClip: 'text',
|
||||
@@ -809,14 +711,7 @@ const ThinkingBubble: React.FC<{
|
||||
);
|
||||
};
|
||||
|
||||
// Friendly explanation rendered in the expanded Thinking pill body when
|
||||
// the model thought but didn't return any reasoning text. Keeps the user
|
||||
// informed about *why* the panel is empty rather than leaving them
|
||||
// staring at a blank box. Three cases:
|
||||
// 1. Live-streaming, no text yet → "Reasoning..." with cursor
|
||||
// 2. Done, has reasoning tokens → explain that text isn't exposed by
|
||||
// the upstream provider but the model spent N tokens / Ms thinking
|
||||
// 3. Done, no signal at all → say so honestly
|
||||
// fallback when the model thought but the provider didn't expose the text.
|
||||
const ProviderReasoningExplanation: React.FC<{
|
||||
isStreaming: boolean;
|
||||
tokens: number | null;
|
||||
@@ -842,9 +737,6 @@ const ProviderReasoningExplanation: React.FC<{
|
||||
}
|
||||
return segs.join(' · ');
|
||||
})();
|
||||
// Stable per-mount variant pick. Each render of the same bubble keeps
|
||||
// its line; new bubbles get a fresh roll. Adds a touch of personality
|
||||
// without becoming repetitive across the transcript.
|
||||
const variants = [
|
||||
"It's still thinking — we just aren't allowed to peek behind the curtain.",
|
||||
"Wheels are turning, but this provider keeps its thoughts private.",
|
||||
@@ -869,8 +761,6 @@ interface Props {
|
||||
onSaveEdit?: (messageId: string, newContent: string) => void;
|
||||
onCancelEdit?: () => void;
|
||||
isStreaming?: boolean;
|
||||
// Session's current aux-LLM turn label, if any. Only meaningful when
|
||||
// this is the live-streaming thinking bubble; ignored otherwise.
|
||||
dynamicTurnLabel?: string | null;
|
||||
}
|
||||
|
||||
@@ -945,14 +835,10 @@ const MessageBubble: React.FC<Props> = React.memo(({ message, editing = false, o
|
||||
>{rawText}</ReactMarkdown>
|
||||
), [rawText]);
|
||||
|
||||
// Detect friendly OpenSwarm / upstream errors and render a card instead of
|
||||
// raw "API Error: ..." text. Checks both the wrapped format the Claude CLI
|
||||
// uses ("API Error: NNN …") and the raw JSON body.
|
||||
// upstream errors get a friendly card.
|
||||
const openswarmError = !isUser ? parseOpenSwarmError(rawText) : null;
|
||||
|
||||
// Fire subscription.rate_limit_hit exactly once per rate-limit error
|
||||
// card mount. Dependency on (message.id, kind) ensures we don't re-fire
|
||||
// on re-renders or content edits.
|
||||
// fire once per cap card. (message.id, kind) keeps it from re-firing on edits.
|
||||
React.useEffect(() => {
|
||||
if (openswarmError?.kind === 'cap') {
|
||||
trackEvent('subscription.rate_limit_hit', { message_id: message.id });
|
||||
@@ -981,8 +867,7 @@ const MessageBubble: React.FC<Props> = React.memo(({ message, editing = false, o
|
||||
? content.slice(0, 200)
|
||||
: JSON.stringify(content).slice(0, 200);
|
||||
|
||||
// Optimistic-bubble visuals: dim the bubble until the server echoes it
|
||||
// back (status: 'pending'), and tint it red on send failure.
|
||||
// pending = dim, failed = red tint.
|
||||
const optimisticStatus = (message as any).optimistic_status as 'pending' | 'failed' | undefined;
|
||||
const isPending = optimisticStatus === 'pending';
|
||||
const isFailed = optimisticStatus === 'failed';
|
||||
@@ -996,11 +881,7 @@ const MessageBubble: React.FC<Props> = React.memo(({ message, editing = false, o
|
||||
display: 'flex',
|
||||
justifyContent: isUser ? 'flex-end' : 'flex-start',
|
||||
my: 0.75,
|
||||
// Layout-style containment: any reflow inside this bubble (text
|
||||
// wrapping during streaming, tooltip popup, expand/collapse)
|
||||
// doesn't propagate to siblings. Without this, every delta in
|
||||
// a long assistant message reflowed the entire transcript.
|
||||
// Browser support is universal in modern Chromium/WebKit.
|
||||
// contain: reflow inside this bubble doesn't shake the transcript.
|
||||
contain: 'layout style',
|
||||
}}
|
||||
>
|
||||
@@ -1015,9 +896,6 @@ const MessageBubble: React.FC<Props> = React.memo(({ message, editing = false, o
|
||||
py: 1.25,
|
||||
boxShadow: isUser ? 'none' : c.shadow.sm,
|
||||
overflow: 'hidden',
|
||||
// Pending bubbles fade in at ~70% opacity until the server echo
|
||||
// resolves them; failed bubbles get a soft red tint so the user
|
||||
// can see the message didn't go through.
|
||||
opacity: isPending ? 0.7 : 1,
|
||||
transition: 'opacity 0.2s, border-color 0.2s',
|
||||
}}
|
||||
@@ -1188,12 +1066,9 @@ const MessageBubble: React.FC<Props> = React.memo(({ message, editing = false, o
|
||||
onClick={() => {
|
||||
const api = (window as any).openswarm;
|
||||
if (openswarmError.ctaAction === 'upgrade') {
|
||||
// Open the tier picker in a modal so the user can
|
||||
// choose Pro / Pro+ / Ultra + monthly/annual instead
|
||||
// of going directly to a hardcoded pro_plus checkout.
|
||||
// tier picker, not direct checkout.
|
||||
setPickerOpen(true);
|
||||
} else if (openswarmError.ctaAction === 'settings') {
|
||||
// Best-effort: dispatch a DOM event the Settings modal listens to
|
||||
window.dispatchEvent(new CustomEvent('openswarm:open-settings', { detail: { tab: 'models' } }));
|
||||
} else if (openswarmError.ctaAction === 'waitlist') {
|
||||
const url = 'https://discord.com/channels/1486442924391796896/1486442927554170892';
|
||||
|
||||
@@ -1427,12 +1427,8 @@ const ToolCallBubble: React.FC<ToolCallBubbleProps> = React.memo(
|
||||
const promptPrefix = getPromptPrefix(toolName);
|
||||
const shortAction = mcpInfo.isMcp ? getMcpShortAction(mcpInfo) : toolName;
|
||||
|
||||
// mcpCompact rows live INSIDE a ToolGroup whose header already shows
|
||||
// the brand + count. Render the friendly verb form here ("Sent message",
|
||||
// "Read 4 emails") so the row contributes a real noun instead of
|
||||
// repeating the brand or showing the raw action like "Send Slack Message".
|
||||
// Seed with call.id so each row picks a stable variant from the pool
|
||||
// (no flicker on re-render, but adjacent rows get different verbs).
|
||||
// mcpCompact rows live inside a ToolGroup whose header already shows
|
||||
// the brand + count, so the row uses the verb form, not the brand.
|
||||
const mcpVerbLabel = (() => {
|
||||
const lbl = getToolLabel(toolName, call.id);
|
||||
return result && !isDenied ? lbl.past : lbl.present;
|
||||
@@ -2013,18 +2009,7 @@ const ToolCallBubble: React.FC<ToolCallBubbleProps> = React.memo(
|
||||
}}
|
||||
>
|
||||
{(() => {
|
||||
// Verb-tense progression: "Reading" while pending, "Read" once
|
||||
// a tool_result has landed. Denied/streaming fall back to the
|
||||
// present participle since the action is in-flight.
|
||||
// Use the input-aware variant so MCPActivate shows the brand
|
||||
// ("Connecting to Gmail") and Bash derives a verb from the
|
||||
// command ("Deleted foo.ts" instead of "Ran command").
|
||||
// Seed with call.id so the verb pool picks a stable variant
|
||||
// per row (no flicker, variety across the transcript).
|
||||
// MCP tools (singleton rows that aren't grouped) ALSO go
|
||||
// through this path now so they get the friendly verb
|
||||
// pool ("Pulled up email") instead of "Gmail Get Message
|
||||
// Details" Title Case fallback.
|
||||
// call.id seeds the variant pool so re-renders are stable.
|
||||
const { present, past } = getToolLabelWithInput(toolName, input, call.id);
|
||||
return result && !isDenied ? past : present;
|
||||
})()}
|
||||
|
||||
@@ -99,11 +99,7 @@ const ToolGroupBubble: React.FC<Props> = React.memo(({ group, isSessionRunning =
|
||||
sx={{
|
||||
maxWidth: '85%',
|
||||
my: 0.5,
|
||||
// Layout containment: tool rows inserting inside this group
|
||||
// don't reflow the rest of the transcript. The header chip
|
||||
// count tabular-nums fix already handles the within-row
|
||||
// jitter; this stops the OUTER scroll container from
|
||||
// re-laying-out every other bubble when a new row appears.
|
||||
// contain: stops new tool rows from reflowing the whole transcript.
|
||||
contain: 'layout style',
|
||||
}}
|
||||
>
|
||||
@@ -200,10 +196,7 @@ const ToolGroupBubble: React.FC<Props> = React.memo(({ group, isSessionRunning =
|
||||
<Box
|
||||
sx={{
|
||||
borderTop: `0.5px solid ${c.border.medium}`,
|
||||
// Each tool row fades in over 140ms when inserted, instead
|
||||
// of jumping into place. Pure CSS — runs on the compositor
|
||||
// and pairs with the parent's contain:layout so the rest
|
||||
// of the transcript doesn't shift while the row settles.
|
||||
// 140ms fade so rows don't pop in.
|
||||
'& > *': {
|
||||
animation: 'toolRowFadeIn 140ms ease-out',
|
||||
},
|
||||
|
||||
@@ -1,25 +1,12 @@
|
||||
// Friendly verb-tense labels for tool calls. Replaces the raw tool name in
|
||||
// ToolCallBubble titles so the transcript reads as a narration of what the
|
||||
// agent is doing — "Reading foo.ts" while pending, "Read foo.ts" once done.
|
||||
//
|
||||
// Voice: a real person casually telling you what they did, in past tense.
|
||||
// Mixes Linear-style "cooking up your data" warmth with Fastmail-style
|
||||
// "Snoozed / Filed" physicality. Each tool has a small pool of variants;
|
||||
// a stable hash of the tool call's id picks one — same call always reads
|
||||
// the same way (no flicker on re-render), different calls get variety.
|
||||
//
|
||||
// **Confidence-scaled friendliness.** Read-side actions get the playful
|
||||
// variants. Destructive / irreversible operations (rm, git push, deletions)
|
||||
// stay flat and factual — quirky verbs on `rm -rf` would feel wrong.
|
||||
// Tool labels, with variant pools so the transcript reads like a person.
|
||||
// Destructive ops (rm, git push, delete) stay flat. quirky on rm felt off.
|
||||
|
||||
export interface ToolLabel {
|
||||
present: string;
|
||||
past: string;
|
||||
}
|
||||
|
||||
// Stable seeded pick. Same seed → same index across renders so the row
|
||||
// doesn't flicker between variants. Empty seed → always index 0 (the
|
||||
// "safe default" variant).
|
||||
// djb2. same seed always picks the same variant so rows don't flicker.
|
||||
function _stableIndex(seed: string | undefined, n: number): number {
|
||||
if (n <= 1 || !seed) return 0;
|
||||
let h = 5381;
|
||||
@@ -33,9 +20,7 @@ function _pick<T>(variants: T[], seed?: string): T {
|
||||
return variants[_stableIndex(seed, variants.length)];
|
||||
}
|
||||
|
||||
// Each entry is an array of variants. Index 0 is the safe-default
|
||||
// (used when no seed is supplied). Destructive / high-stakes tools
|
||||
// have a single-variant entry to keep them factual.
|
||||
// index 0 is the safe-default; single-entry pools = no seeded variation.
|
||||
const VARIANTS: Record<string, ToolLabel[]> = {
|
||||
read: [
|
||||
{ present: 'Reading', past: 'Read' },
|
||||
@@ -172,8 +157,7 @@ const VARIANTS: Record<string, ToolLabel[]> = {
|
||||
{ present: 'Browsing the toolbox', past: 'Browsed the toolbox' },
|
||||
{ present: 'Rummaging the toolbox', past: 'Rummaged the toolbox' },
|
||||
],
|
||||
// MCPActivate is brand-aware — see getToolLabelWithInput below. The
|
||||
// bare entry here is a fallback when input.server_name isn't available.
|
||||
// brand-aware version lives in getToolLabelWithInput; this is the fallback.
|
||||
mcpactivate: [
|
||||
{ present: 'Connecting', past: 'Connected' },
|
||||
{ present: 'Plugging in', past: 'Plugged in' },
|
||||
@@ -286,9 +270,7 @@ const VARIANTS: Record<string, ToolLabel[]> = {
|
||||
],
|
||||
};
|
||||
|
||||
// Brand names for MCP servers — what we want users to *see* instead of
|
||||
// the kebab-case server id. Keyed by the sanitized server name (matches
|
||||
// _sanitize_server_name in tools_lib).
|
||||
// keys match _sanitize_server_name in tools_lib.
|
||||
const MCP_SERVER_BRAND: Record<string, string> = {
|
||||
'google-workspace': 'Google Workspace',
|
||||
'microsoft-365': 'Microsoft 365',
|
||||
@@ -317,9 +299,7 @@ const MCP_SERVER_BRAND: Record<string, string> = {
|
||||
'openswarm-outputs-meta': 'views',
|
||||
};
|
||||
|
||||
// MCP sub-tool action verbs. Each verb class has variants — playful
|
||||
// where it's safe, factual for destructive (delete/remove). Ordered:
|
||||
// most-specific first.
|
||||
// most specific verb pattern wins, so order matters.
|
||||
interface McpVerbVariant { present: string; past: string; }
|
||||
const MCP_VERB_PATTERNS: Array<{ match: RegExp; variants: McpVerbVariant[] }> = [
|
||||
{ match: /^(send|new)_/, variants: [
|
||||
@@ -361,8 +341,7 @@ const MCP_VERB_PATTERNS: Array<{ match: RegExp; variants: McpVerbVariant[] }> =
|
||||
{ present: 'Refining', past: 'Refined' },
|
||||
{ present: 'Touching up', past: 'Touched up' },
|
||||
]},
|
||||
// Destructive — flat, no playful variants. Keep it factual so the
|
||||
// agent doesn't sound flippant about deletions.
|
||||
// delete = flat. don't be cute about deletions.
|
||||
{ match: /^(delete|remove|cancel|archive)_/, variants: [
|
||||
{ present: 'Deleting', past: 'Deleted' },
|
||||
]},
|
||||
@@ -427,18 +406,12 @@ const ACTION_OBJECTS: Array<{ match: RegExp; noun: string }> = [
|
||||
{ match: /(?:^|_)(?:task|todo)/, noun: 'task' },
|
||||
];
|
||||
|
||||
// Sentence case the input — first word capitalized, rest lowercased.
|
||||
// We deliberately avoid Title Case (capitalize-every-word) because it
|
||||
// reads marketing-y in agent narration; sentence case is Linear/Notion/
|
||||
// Stripe convention and feels more like a human wrote it.
|
||||
// sentence case (Linear/Notion vibe). title case felt too marketing-y.
|
||||
function _humanizeName(name: string): string {
|
||||
const spaced = name.replace(/[-_]+/g, ' ').toLowerCase();
|
||||
return spaced.charAt(0).toUpperCase() + spaced.slice(1);
|
||||
}
|
||||
|
||||
// Parse an mcp__<server>__<action> tool name into a friendly label.
|
||||
// Returns null when the input isn't an MCP-shaped tool name; callers
|
||||
// fall through to the builtin VARIANTS map.
|
||||
function _labelForMcpTool(toolName: string, seed?: string): ToolLabel | null {
|
||||
const parts = toolName.split('__');
|
||||
if (parts.length < 3 || parts[0] !== 'mcp') return null;
|
||||
@@ -446,22 +419,17 @@ function _labelForMcpTool(toolName: string, seed?: string): ToolLabel | null {
|
||||
const action = parts.slice(2).join('__').toLowerCase();
|
||||
const brand = MCP_SERVER_BRAND[server] || _humanizeName(server);
|
||||
|
||||
// Our internal openswarm-* meta-MCPs expose action names that already
|
||||
// map cleanly to VARIANTS keys (mcpsearch / mcpactivate / outputlist /
|
||||
// outputsearch / etc.). Without this short-circuit we hit the
|
||||
// verb-pattern fallback which renders "tools: Mcpsearch" — ugly.
|
||||
// our internal meta-MCPs go through VARIANTS so we don't render "tools: Mcpsearch".
|
||||
if (server.startsWith('openswarm-')) {
|
||||
const builtin = VARIANTS[action];
|
||||
if (builtin) return _pick(builtin, seed);
|
||||
}
|
||||
|
||||
// Find verb class (variant pool)
|
||||
let verbVariants: McpVerbVariant[] | null = null;
|
||||
for (const p of MCP_VERB_PATTERNS) {
|
||||
if (p.match.test(action)) { verbVariants = p.variants; break; }
|
||||
}
|
||||
|
||||
// Find object noun
|
||||
let noun = '';
|
||||
for (const a of ACTION_OBJECTS) {
|
||||
if (a.match.test(action)) { noun = a.noun; break; }
|
||||
@@ -475,7 +443,7 @@ function _labelForMcpTool(toolName: string, seed?: string): ToolLabel | null {
|
||||
return { present: `${verb.present} via ${brand}`, past: `${verb.past} via ${brand}` };
|
||||
}
|
||||
|
||||
// Verb didn't match — humanize the action name and pair with the brand.
|
||||
// no verb match. fall back to brand: action.
|
||||
const human = _humanizeName(action.replace(/^_+|_+$/g, ''));
|
||||
return { present: `${brand}: ${human}`, past: `${brand}: ${human}` };
|
||||
}
|
||||
@@ -487,20 +455,11 @@ export function getToolLabel(toolName: string, seed?: string): ToolLabel {
|
||||
const key = toolName.toLowerCase();
|
||||
const variants = VARIANTS[key];
|
||||
if (variants) return _pick(variants, seed);
|
||||
// Fallback: capitalize the raw name with neutral verbs that read OK
|
||||
// either way ("Running tool" / "Ran tool").
|
||||
const pretty = toolName.charAt(0).toUpperCase() + toolName.slice(1);
|
||||
return { present: `Running ${pretty}`, past: `Ran ${pretty}` };
|
||||
}
|
||||
|
||||
// Variant that consults the tool's input arguments so a few tools can
|
||||
// produce a more specific label. Falls back to getToolLabel for the
|
||||
// general case. Specifically:
|
||||
// - MCPActivate(server_name): "Connecting to Gmail" / "Connected to Gmail"
|
||||
// - Bash(command): derive verb from the leading binary
|
||||
//
|
||||
// `seed` should be the tool call's id so the variant pick is stable
|
||||
// across re-renders of the same row but varies between rows.
|
||||
// for tools where the input changes the label (MCPActivate, Bash).
|
||||
export function getToolLabelWithInput(toolName: string, input: any, seed?: string): ToolLabel {
|
||||
if (!toolName) return { present: 'Working', past: 'Done' };
|
||||
|
||||
@@ -531,8 +490,7 @@ export function getToolLabelWithInput(toolName: string, input: any, seed?: strin
|
||||
|
||||
// --- Bash verb extraction ---------------------------------------------------
|
||||
|
||||
// Single-variant entries are kept flat by design — destructive (rm,
|
||||
// chmod) or directional (cd) ops shouldn't get cute paraphrases.
|
||||
// rm and chmod don't get cute paraphrases for obvious reasons.
|
||||
const GIT_VERBS: Record<string, ToolLabel[]> = {
|
||||
commit: [
|
||||
{ present: 'Committing', past: 'Committed' },
|
||||
@@ -646,7 +604,6 @@ function _pkgVerb(sub: string, seed?: string): ToolLabel {
|
||||
type BinEntry = ToolLabel[] | ((sub: string, seed?: string) => ToolLabel | null);
|
||||
|
||||
const BIN_VERBS: Record<string, BinEntry> = {
|
||||
// Destructive — single variant, factual.
|
||||
rm: [{ present: 'Deleting', past: 'Deleted' }],
|
||||
rmdir: [{ present: 'Removing folder', past: 'Removed folder' }],
|
||||
chmod: [{ present: 'Changing permissions', past: 'Changed permissions' }],
|
||||
@@ -654,7 +611,6 @@ const BIN_VERBS: Record<string, BinEntry> = {
|
||||
killall: [{ present: 'Stopping process', past: 'Stopped process' }],
|
||||
kill: [{ present: 'Stopping process', past: 'Stopped process' }],
|
||||
|
||||
// Friendly — variant pool.
|
||||
mv: [
|
||||
{ present: 'Moving', past: 'Moved' },
|
||||
{ present: 'Shuffling', past: 'Shuffled' },
|
||||
@@ -864,7 +820,6 @@ const BIN_VERBS: Record<string, BinEntry> = {
|
||||
{ present: 'Listing processes', past: 'Listed processes' },
|
||||
{ present: 'Checking what is running', past: 'Checked what is running' },
|
||||
],
|
||||
// Subcommand-driven
|
||||
git: (sub: string, seed?: string) => {
|
||||
const v = GIT_VERBS[sub];
|
||||
if (v) return _pick<ToolLabel>(v, seed);
|
||||
|
||||
@@ -35,11 +35,7 @@ export interface AgentMessage {
|
||||
// messages so the pill can show "Thought for Ns · M in / K out"
|
||||
// — which is the only honest answer to "how big was this turn".
|
||||
input_tokens?: number;
|
||||
// Richer thinking-pill data: total post-thinking output tokens
|
||||
// (user-visible answer text + tool arguments) and tool invocation
|
||||
// count. Drives the "Thought for 18s · 430 reasoning · 2.4K answer
|
||||
// · 3 tools" label. Set on thinking messages only.
|
||||
answer_tokens?: number;
|
||||
// tool count drives the "3 tools used" segment on the thinking pill.
|
||||
tool_count?: number;
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -1,39 +0,0 @@
|
||||
#!/bin/bash
|
||||
# One-time Microsoft 365 authentication for OpenSwarm.
|
||||
# Run this once to cache your M365 token. After that, M365 works in OpenSwarm automatically.
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
PROJECT_ROOT="$(dirname "$SCRIPT_DIR")"
|
||||
|
||||
CACHE_DIR="$HOME/.openswarm"
|
||||
mkdir -p "$CACHE_DIR"
|
||||
|
||||
export MS365_MCP_TOKEN_CACHE_PATH="$CACHE_DIR/ms365-token-cache.json"
|
||||
export MS365_MCP_SELECTED_ACCOUNT_PATH="$CACHE_DIR/ms365-selected-account.json"
|
||||
|
||||
SERVER_SCRIPT="$PROJECT_ROOT/backend/npm-servers/softeria-ms-365-mcp-server/node_modules/@softeria/ms-365-mcp-server/dist/index.js"
|
||||
|
||||
if [ ! -f "$SERVER_SCRIPT" ]; then
|
||||
echo "M365 MCP server not found. Run 'cd backend/npm-servers/softeria-ms-365-mcp-server && npm install' first."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo ""
|
||||
echo " Microsoft 365 Login for OpenSwarm"
|
||||
echo " ─────────────────────────────────"
|
||||
echo " A browser window will open for you to sign in."
|
||||
echo " After login, the token is cached and M365 works in OpenSwarm automatically."
|
||||
echo ""
|
||||
|
||||
node "$SERVER_SCRIPT" --login
|
||||
|
||||
if [ -f "$MS365_MCP_TOKEN_CACHE_PATH" ]; then
|
||||
echo ""
|
||||
echo " ✓ Token cached at $MS365_MCP_TOKEN_CACHE_PATH"
|
||||
echo " ✓ M365 is ready to use in OpenSwarm!"
|
||||
echo ""
|
||||
else
|
||||
echo ""
|
||||
echo " ✗ Login may have failed — no token cache found."
|
||||
echo ""
|
||||
fi
|
||||
Reference in New Issue
Block a user