mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-13 21:27:41 +02:00
[arnav] Re-adds backend/apps/service/models.py
Also removes tests that covered code no longer present in the codebase: - test_agent_loop.py, test_tools_unit.py — deleted (targets removed) - test_mcp_servers_unit.py — deleted (targets removed) - test_providers_anthropic_extra.py, test_providers_openai_compat.py, test_providers_registry.py — deleted (providers layer refactored away) - test_api_outputs.py — drops vibe-code tests (endpoint removed) - test_browser_agent_integration.py — drops run_browser_agents fanout tests (function removed) - test_mcp_preflight.py — drops _call_classifier tests (internal removed)
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
"""(Reserved for future use; intentionally empty.)
|
||||
|
||||
The service-sync layer ships opaque payload dicts through `submit()` —
|
||||
no Pydantic shape exposed in the public repo.
|
||||
"""
|
||||
@@ -1,712 +0,0 @@
|
||||
"""Unit tests for `backend.apps.agents.agent_loop.AgentLoop`.
|
||||
|
||||
The agent loop is the provider-agnostic streaming + tool-use + HITL
|
||||
core: it drives `BaseProvider.stream_message`, accumulates content
|
||||
blocks, executes tools (with HITL gating), emits the WebSocket events
|
||||
the frontend consumes, and persists the final messages.
|
||||
|
||||
The thinking-block path is already covered by `test_phase1_stress.py`.
|
||||
This file focuses on the rest of the surface — the `run()` control
|
||||
flow, `_execute_tools` (HITL deny / updated_input / executor errors /
|
||||
truncation / multi-tool / non-tool skip), `_stream_and_collect`
|
||||
JSON handling and stop_reason routing, `_emit_collected_messages`
|
||||
formatting, and token-usage accumulation.
|
||||
|
||||
Pure unit tests: no FastAPI client, no network. We script provider
|
||||
output via `_StubProvider` and capture WS emissions via `_WSRecorder`.
|
||||
The conftest bootstrap (run automatically by virtue of living in
|
||||
`backend/tests/`) redirects `OPENSWARM_DATA_DIR` and mocks PostHog so
|
||||
nothing here ever touches the real disk or external services.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents.agent_loop import AgentLoop
|
||||
from backend.apps.agents.providers.base import (
|
||||
ContentBlock,
|
||||
ModelResponse,
|
||||
ProviderMessage,
|
||||
StreamEvent,
|
||||
ToolCall,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _StubProvider:
|
||||
"""Scripted `BaseProvider` stand-in for the agent loop.
|
||||
|
||||
Each call to `stream_message` consumes one entry from `turns` and
|
||||
yields its `StreamEvent`s in order. The messages list passed in is
|
||||
snapshotted into `self.calls[i]` so tests can assert on the
|
||||
conversation-history shape the agent loop sent on each turn.
|
||||
|
||||
Duck-typed (not a `BaseProvider` subclass) — matches the
|
||||
convention in `test_phase1_stress.py` and avoids having to stub
|
||||
out `create_message` / `get_model_id` that the loop never calls.
|
||||
"""
|
||||
|
||||
def __init__(self, turns: list[list[StreamEvent]]):
|
||||
self._turns = list(turns)
|
||||
self.calls: list[list[ProviderMessage]] = []
|
||||
|
||||
async def stream_message(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
system: str | None,
|
||||
messages: list[ProviderMessage],
|
||||
tools: list,
|
||||
):
|
||||
# Snapshot at call time — the loop mutates `self.messages`
|
||||
# between turns, so a reference would get clobbered.
|
||||
self.calls.append(list(messages))
|
||||
if not self._turns:
|
||||
return
|
||||
for ev in self._turns.pop(0):
|
||||
yield ev
|
||||
|
||||
def format_user_message(self, content: Any) -> ProviderMessage:
|
||||
return ProviderMessage(role="user", content=content)
|
||||
|
||||
def format_assistant_message(self, response: ModelResponse) -> ProviderMessage:
|
||||
# Shape doesn't matter for these tests — the loop just appends
|
||||
# the message; we never feed it back through a real provider.
|
||||
return ProviderMessage(role="assistant", content=response.content)
|
||||
|
||||
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}
|
||||
|
||||
|
||||
class _WSRecorder:
|
||||
"""Async-callable that captures every (event, payload) pair the
|
||||
loop emits. Tests assert on `events` directly or via `of_type`."""
|
||||
|
||||
def __init__(self):
|
||||
self.events: list[tuple[str, dict]] = []
|
||||
|
||||
async def __call__(self, event_type: str, data: dict) -> None:
|
||||
self.events.append((event_type, data))
|
||||
|
||||
def of_type(self, event_type: str) -> list[dict]:
|
||||
return [d for (e, d) in self.events if e == event_type]
|
||||
|
||||
|
||||
# Block factories. A "block" is the start/delta/stop triple for one
|
||||
# content_block in the model's output; a "turn" is one full
|
||||
# stream_message call (one or more blocks + optional usage +
|
||||
# message_stop).
|
||||
|
||||
|
||||
def _text_block(text: str, *, index: int = 0) -> list[StreamEvent]:
|
||||
return [
|
||||
StreamEvent(type="content_block_start", index=index, block_type="text"),
|
||||
StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=index,
|
||||
delta_type="text_delta",
|
||||
text=text,
|
||||
),
|
||||
StreamEvent(type="content_block_stop", index=index),
|
||||
]
|
||||
|
||||
|
||||
def _tool_block(
|
||||
name: str,
|
||||
tool_id: str,
|
||||
json_input: str,
|
||||
*,
|
||||
index: int = 0,
|
||||
) -> list[StreamEvent]:
|
||||
return [
|
||||
StreamEvent(
|
||||
type="content_block_start",
|
||||
index=index,
|
||||
block_type="tool_use",
|
||||
tool_name=name,
|
||||
tool_id=tool_id,
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta",
|
||||
index=index,
|
||||
delta_type="input_json_delta",
|
||||
text=json_input,
|
||||
),
|
||||
StreamEvent(type="content_block_stop", index=index),
|
||||
]
|
||||
|
||||
|
||||
def _turn(
|
||||
*blocks: list[StreamEvent],
|
||||
usage: dict[str, int] | None = None,
|
||||
) -> list[StreamEvent]:
|
||||
"""Compose 1+ blocks into a complete turn (with usage + message_stop)."""
|
||||
events: list[StreamEvent] = []
|
||||
for block in blocks:
|
||||
events.extend(block)
|
||||
if usage:
|
||||
events.append(StreamEvent(type="usage", usage=usage))
|
||||
events.append(StreamEvent(type="message_stop"))
|
||||
return events
|
||||
|
||||
|
||||
def _text_turn(text: str, *, usage: dict[str, int] | None = None) -> list[StreamEvent]:
|
||||
return _turn(_text_block(text), usage=usage)
|
||||
|
||||
|
||||
def _tool_turn(
|
||||
name: str,
|
||||
tool_id: str,
|
||||
json_input: str,
|
||||
*,
|
||||
usage: dict[str, int] | None = None,
|
||||
) -> list[StreamEvent]:
|
||||
return _turn(_tool_block(name, tool_id, json_input), usage=usage)
|
||||
|
||||
|
||||
def _make_loop(
|
||||
*,
|
||||
provider: _StubProvider | None = None,
|
||||
executor=None,
|
||||
hitl=None,
|
||||
max_turns: int | None = None,
|
||||
system_prompt: str | None = None,
|
||||
) -> tuple[AgentLoop, _WSRecorder]:
|
||||
"""Build an `AgentLoop` plus its WS recorder. Sensible defaults
|
||||
for tests that don't care about a particular callback."""
|
||||
ws = _WSRecorder()
|
||||
loop = AgentLoop(
|
||||
session_id="test-session",
|
||||
provider=provider if provider is not None else _StubProvider([]),
|
||||
model="sonnet",
|
||||
system_prompt=system_prompt,
|
||||
tools=[],
|
||||
tool_executor=executor
|
||||
if executor is not None
|
||||
else AsyncMock(return_value=[{"type": "text", "text": "ok"}]),
|
||||
hitl_handler=hitl
|
||||
if hitl is not None
|
||||
else AsyncMock(return_value=(True, None)),
|
||||
ws_emitter=ws,
|
||||
max_turns=max_turns,
|
||||
)
|
||||
return loop, ws
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 1 — run() control flow
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_end_turn_exits_after_one_iteration():
|
||||
"""No tool_use → loop calls the provider once and stops with
|
||||
[user, assistant] in history."""
|
||||
provider = _StubProvider([_text_turn("hello!")])
|
||||
loop, _ = _make_loop(provider=provider)
|
||||
|
||||
await loop.run("hi")
|
||||
|
||||
assert len(provider.calls) == 1
|
||||
assert [m.role for m in loop.messages] == ["user", "assistant"]
|
||||
|
||||
|
||||
async def test_tool_use_continues_loop_then_terminates_on_end_turn():
|
||||
"""tool_use → execute → second model call → end_turn. Verifies the
|
||||
full conversation grows to [user, assistant, tool_result, assistant]
|
||||
AND that turn 2 sees the prior tool_result."""
|
||||
provider = _StubProvider([
|
||||
_tool_turn("Read", "t1", '{"path": "/x"}'),
|
||||
_text_turn("done"),
|
||||
])
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": "file contents"}])
|
||||
loop, _ = _make_loop(provider=provider, executor=executor)
|
||||
|
||||
await loop.run("hi")
|
||||
|
||||
assert len(provider.calls) == 2
|
||||
assert [m.role for m in loop.messages] == [
|
||||
"user", "assistant", "tool_result", "assistant",
|
||||
]
|
||||
# The second model call must see the tool_result we appended.
|
||||
roles_seen_on_turn_2 = [m.role for m in provider.calls[1]]
|
||||
assert roles_seen_on_turn_2 == ["user", "assistant", "tool_result"]
|
||||
executor.assert_awaited_once()
|
||||
|
||||
|
||||
async def test_max_turns_halts_loop_before_second_model_call():
|
||||
"""`max_turns=1` allows exactly one model call + tool execution,
|
||||
then the top-of-loop guard breaks before turn 2. The unused
|
||||
second turn stays scripted but un-consumed."""
|
||||
provider = _StubProvider([
|
||||
_tool_turn("Read", "t1", "{}"),
|
||||
_tool_turn("Read", "t2", "{}"), # would be consumed if guard failed
|
||||
])
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": "ok"}])
|
||||
loop, _ = _make_loop(provider=provider, executor=executor, max_turns=1)
|
||||
|
||||
await loop.run("hi")
|
||||
|
||||
assert len(provider.calls) == 1, "max_turns=1 must cap provider calls"
|
||||
assert executor.await_count == 1, "tool exec runs in turn 1, not gated by max_turns"
|
||||
|
||||
|
||||
async def test_no_tool_results_breaks_loop(monkeypatch):
|
||||
"""Defensive guard: if `_execute_tools` returns `[]` for any
|
||||
reason, the loop must exit without appending an empty
|
||||
`tool_result` and without re-calling the provider."""
|
||||
provider = _StubProvider([
|
||||
_tool_turn("Read", "t1", "{}"),
|
||||
_text_turn("never reached"),
|
||||
])
|
||||
loop, _ = _make_loop(provider=provider)
|
||||
|
||||
async def _empty_results(self, response):
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(AgentLoop, "_execute_tools", _empty_results)
|
||||
|
||||
await loop.run("hi")
|
||||
|
||||
assert len(provider.calls) == 1
|
||||
assert [m.role for m in loop.messages] == ["user", "assistant"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 2 — _execute_tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _tool_use_response(*calls: tuple[str, str, dict]) -> ModelResponse:
|
||||
"""Build a ModelResponse(stop_reason='tool_use') from (id, name, input) tuples."""
|
||||
return ModelResponse(
|
||||
content=[
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id=tid, name=name, input=inp),
|
||||
)
|
||||
for (tid, name, inp) in calls
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
)
|
||||
|
||||
|
||||
async def test_hitl_denial_skips_executor_and_returns_denial_text():
|
||||
executor = AsyncMock()
|
||||
hitl = AsyncMock(return_value=(False, None))
|
||||
loop, ws = _make_loop(executor=executor, hitl=hitl)
|
||||
|
||||
response = _tool_use_response(("t1", "Read", {"path": "/x"}))
|
||||
results = await loop._execute_tools(response)
|
||||
|
||||
executor.assert_not_called()
|
||||
assert len(results) == 1, "denied tools still produce a tool_result for the model"
|
||||
|
||||
tool_result_msgs = [
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "tool_result"
|
||||
]
|
||||
assert len(tool_result_msgs) == 1
|
||||
assert tool_result_msgs[0]["content"]["text"] == "Tool use was denied by the user."
|
||||
assert tool_result_msgs[0]["content"]["tool_name"] == "Read"
|
||||
|
||||
|
||||
async def test_hitl_updated_input_passed_to_executor():
|
||||
"""When HITL approves with an `updated_input`, the executor must
|
||||
see that dict — not the model's original input."""
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": "ok"}])
|
||||
hitl = AsyncMock(return_value=(True, {"path": "/y"}))
|
||||
loop, _ = _make_loop(executor=executor, hitl=hitl)
|
||||
|
||||
response = _tool_use_response(("t1", "Read", {"path": "/x"}))
|
||||
await loop._execute_tools(response)
|
||||
|
||||
executor.assert_awaited_once_with("Read", {"path": "/y"})
|
||||
|
||||
|
||||
async def test_executor_exception_is_caught_and_surfaced_as_error_text():
|
||||
async def boom(name, inp):
|
||||
raise RuntimeError("disk on fire")
|
||||
|
||||
loop, ws = _make_loop(executor=boom)
|
||||
response = _tool_use_response(("t1", "Read", {}))
|
||||
|
||||
results = await loop._execute_tools(response)
|
||||
|
||||
# Loop survives, returns a formatted error result for the provider.
|
||||
assert len(results) == 1
|
||||
tool_results = [
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "tool_result"
|
||||
]
|
||||
assert tool_results[0]["content"]["text"] == "Error executing Read: disk on fire"
|
||||
|
||||
|
||||
async def test_tool_result_text_truncated_to_15000_chars_in_emitted_message():
|
||||
"""The model gets the full tool output; the WS message bubble
|
||||
that the UI renders is sliced to 15K to avoid jank."""
|
||||
huge = "x" * 20_000
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": huge}])
|
||||
loop, ws = _make_loop(executor=executor)
|
||||
|
||||
response = _tool_use_response(("t1", "Read", {}))
|
||||
results = await loop._execute_tools(response)
|
||||
|
||||
# WS-emitted snippet capped at 15K.
|
||||
tr_msg = next(
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "tool_result"
|
||||
)
|
||||
assert len(tr_msg["content"]["text"]) == 15_000
|
||||
|
||||
# Provider-bound result is the raw, untruncated content.
|
||||
assert results[0]["content"][0]["text"] == huge
|
||||
|
||||
|
||||
async def test_multiple_tool_calls_in_one_response_all_execute():
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": "ok"}])
|
||||
loop, ws = _make_loop(executor=executor)
|
||||
|
||||
response = _tool_use_response(
|
||||
("t1", "Read", {"path": "/a"}),
|
||||
("t2", "Edit", {"path": "/b"}),
|
||||
)
|
||||
results = await loop._execute_tools(response)
|
||||
|
||||
assert executor.await_count == 2
|
||||
assert [r["tool_use_id"] for r in results] == ["t1", "t2"]
|
||||
tool_result_msgs = [
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "tool_result"
|
||||
]
|
||||
assert [m["content"]["tool_name"] for m in tool_result_msgs] == ["Read", "Edit"]
|
||||
|
||||
|
||||
async def test_non_tool_use_blocks_are_skipped_in_executor():
|
||||
"""A response with text + tool_use should only execute the tool block."""
|
||||
executor = AsyncMock(return_value=[{"type": "text", "text": "ok"}])
|
||||
loop, _ = _make_loop(executor=executor)
|
||||
|
||||
response = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(type="text", text="thinking aloud"),
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t1", name="Read", input={}),
|
||||
),
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
)
|
||||
|
||||
results = await loop._execute_tools(response)
|
||||
assert len(results) == 1
|
||||
executor.assert_awaited_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 3 — _stream_and_collect
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_invalid_tool_input_json_falls_back_to_empty_dict():
|
||||
"""Malformed JSON in input_json_delta must not crash; the
|
||||
resulting ToolCall.input is `{}`."""
|
||||
provider = _StubProvider([_tool_turn("Read", "t1", "{not valid json")])
|
||||
loop, _ = _make_loop(provider=provider)
|
||||
|
||||
response = await loop._stream_and_collect()
|
||||
|
||||
tool_blocks = [b for b in response.content if b.type == "tool_use"]
|
||||
assert len(tool_blocks) == 1
|
||||
assert tool_blocks[0].tool_call is not None
|
||||
assert tool_blocks[0].tool_call.input == {}
|
||||
|
||||
|
||||
async def test_tool_use_input_assembled_from_multiple_json_deltas():
|
||||
"""Real Anthropic streams ship tool input in multiple
|
||||
input_json_delta chunks — they must concatenate into one JSON
|
||||
parse."""
|
||||
provider = _StubProvider([[
|
||||
StreamEvent(
|
||||
type="content_block_start", index=0, block_type="tool_use",
|
||||
tool_name="Read", tool_id="t1",
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="input_json_delta", text='{"pa',
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="input_json_delta", text='th": "/x", "n": 7}',
|
||||
),
|
||||
StreamEvent(type="content_block_stop", index=0),
|
||||
StreamEvent(type="message_stop"),
|
||||
]])
|
||||
loop, _ = _make_loop(provider=provider)
|
||||
|
||||
response = await loop._stream_and_collect()
|
||||
|
||||
assert response.content[0].tool_call.input == {"path": "/x", "n": 7}
|
||||
|
||||
|
||||
async def test_stop_reason_routes_on_presence_of_tool_use_block():
|
||||
"""`response.stop_reason` is `"tool_use"` iff any collected
|
||||
block is tool_use, else `"end_turn"`."""
|
||||
text_provider = _StubProvider([_text_turn("hi")])
|
||||
loop1, _ = _make_loop(provider=text_provider)
|
||||
resp_text = await loop1._stream_and_collect()
|
||||
assert resp_text.stop_reason == "end_turn"
|
||||
|
||||
tool_provider = _StubProvider([_tool_turn("Read", "t1", "{}")])
|
||||
loop2, _ = _make_loop(provider=tool_provider)
|
||||
resp_tool = await loop2._stream_and_collect()
|
||||
assert resp_tool.stop_reason == "tool_use"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 3b — _stream_and_collect WS emissions
|
||||
# ---------------------------------------------------------------------------
|
||||
# Thinking-block WS emissions are covered in `test_phase1_stress.py`. The
|
||||
# tests below cover the text and tool_use streaming paths plus the
|
||||
# routing rule that determines WHERE `agent:stream_end` fires for each
|
||||
# block type — text waits for `message_stop`, tool_use/thinking close at
|
||||
# `content_block_stop`.
|
||||
|
||||
|
||||
async def test_text_block_streams_deltas_and_ends_at_message_stop():
|
||||
"""Text block emits one stream_start (role=assistant), one
|
||||
stream_delta per text_delta event (carrying the same message_id),
|
||||
and one stream_end deferred to `message_stop` — NOT to
|
||||
`content_block_stop`."""
|
||||
provider = _StubProvider([[
|
||||
StreamEvent(type="content_block_start", index=0, block_type="text"),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="text_delta", text="hel",
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="text_delta", text="lo",
|
||||
),
|
||||
StreamEvent(type="content_block_stop", index=0),
|
||||
StreamEvent(type="message_stop"),
|
||||
]])
|
||||
loop, ws = _make_loop(provider=provider)
|
||||
|
||||
await loop._stream_and_collect()
|
||||
|
||||
starts = ws.of_type("agent:stream_start")
|
||||
assert len(starts) == 1
|
||||
assert starts[0]["role"] == "assistant"
|
||||
text_msg_id = starts[0]["message_id"]
|
||||
|
||||
deltas = ws.of_type("agent:stream_delta")
|
||||
assert [d["delta"] for d in deltas] == ["hel", "lo"]
|
||||
assert all(d["message_id"] == text_msg_id for d in deltas)
|
||||
|
||||
ends = ws.of_type("agent:stream_end")
|
||||
assert len(ends) == 1
|
||||
assert ends[0] == {"message_id": text_msg_id}
|
||||
|
||||
# Stream-end placement: text's stream_end must come AFTER the
|
||||
# content_block_stop has already been processed — i.e. it's tied
|
||||
# to message_stop. Concretely, no agent:stream_delta or new
|
||||
# agent:stream_start can follow it for this same message_id.
|
||||
types_in_order = [e for (e, _) in ws.events]
|
||||
end_idx = types_in_order.index("agent:stream_end")
|
||||
assert "agent:stream_delta" not in types_in_order[end_idx + 1:]
|
||||
|
||||
|
||||
async def test_tool_use_block_streams_deltas_and_ends_at_block_stop():
|
||||
"""Tool_use block emits stream_start (role=tool_call, tool_name),
|
||||
one stream_delta per input_json_delta, and stream_end at
|
||||
`content_block_stop` so the UI can finalize the tool-call bubble
|
||||
before any subsequent text streams in."""
|
||||
provider = _StubProvider([[
|
||||
StreamEvent(
|
||||
type="content_block_start", index=0, block_type="tool_use",
|
||||
tool_name="Read", tool_id="t1",
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="input_json_delta", text='{"path"',
|
||||
),
|
||||
StreamEvent(
|
||||
type="content_block_delta", index=0,
|
||||
delta_type="input_json_delta", text=': "/x"}',
|
||||
),
|
||||
StreamEvent(type="content_block_stop", index=0),
|
||||
StreamEvent(type="message_stop"),
|
||||
]])
|
||||
loop, ws = _make_loop(provider=provider)
|
||||
|
||||
await loop._stream_and_collect()
|
||||
|
||||
starts = ws.of_type("agent:stream_start")
|
||||
assert len(starts) == 1
|
||||
assert starts[0]["role"] == "tool_call"
|
||||
assert starts[0]["tool_name"] == "Read"
|
||||
tool_msg_id = starts[0]["message_id"]
|
||||
|
||||
deltas = ws.of_type("agent:stream_delta")
|
||||
assert [d["delta"] for d in deltas] == ['{"path"', ': "/x"}']
|
||||
assert all(d["message_id"] == tool_msg_id for d in deltas)
|
||||
|
||||
ends = ws.of_type("agent:stream_end")
|
||||
assert len(ends) == 1
|
||||
assert ends[0]["message_id"] == tool_msg_id
|
||||
|
||||
# tool_use stream_end fires at content_block_stop — i.e. BEFORE
|
||||
# the message_stop housekeeping. Verify there's no subsequent
|
||||
# delta/end for this id and that an agent:message (the persisted
|
||||
# tool_call bubble) follows in `_emit_collected_messages`.
|
||||
types_in_order = [e for (e, _) in ws.events]
|
||||
stop_idx = types_in_order.index("agent:stream_end")
|
||||
assert "agent:stream_delta" not in types_in_order[stop_idx + 1:]
|
||||
assert "agent:message" in types_in_order[stop_idx + 1:], (
|
||||
"the persisted tool_call message must be emitted after stream_end"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 4 — _emit_collected_messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_text_blocks_joined_with_single_newline_and_msg_id_preserved():
|
||||
"""Multiple text blocks → one assistant Message whose `id` is the
|
||||
streamed msg_id (so the client can dedupe its optimistic bubble)."""
|
||||
loop, ws = _make_loop()
|
||||
content = [
|
||||
ContentBlock(type="text", text="line A"),
|
||||
ContentBlock(type="text", text="line B"),
|
||||
ContentBlock(type="text", text=""), # empty blocks dropped
|
||||
]
|
||||
|
||||
await loop._emit_collected_messages(
|
||||
content, text_msg_id="text-id-123", tool_msg_ids={},
|
||||
)
|
||||
|
||||
assistants = [
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "assistant"
|
||||
]
|
||||
assert len(assistants) == 1
|
||||
assert assistants[0]["content"] == "line A\nline B"
|
||||
assert assistants[0]["id"] == "text-id-123"
|
||||
|
||||
|
||||
async def test_thinking_blocks_joined_with_double_newline_and_metadata():
|
||||
"""Multiple thinking blocks → one persisted thinking Message
|
||||
joined by `\\n\\n`. `elapsed_ms` and `tokens` are derived from
|
||||
the server-stamped accumulators."""
|
||||
loop, ws = _make_loop()
|
||||
content = [
|
||||
ContentBlock(type="thinking", text="step one"),
|
||||
ContentBlock(type="thinking", text="step two"),
|
||||
ContentBlock(type="text", text="answer"),
|
||||
]
|
||||
|
||||
await loop._emit_collected_messages(
|
||||
content,
|
||||
text_msg_id="t1",
|
||||
tool_msg_ids={},
|
||||
thinking_elapsed_ms=1234,
|
||||
thinking_total_chars=20,
|
||||
)
|
||||
|
||||
msgs = [d["message"] for d in ws.of_type("agent:message")]
|
||||
thinking = [m for m in msgs if m["role"] == "thinking"]
|
||||
assert len(thinking) == 1
|
||||
assert thinking[0]["content"] == "step one\n\nstep two"
|
||||
assert thinking[0]["elapsed_ms"] == 1234
|
||||
assert thinking[0]["tokens"] == max(1, round(20 / 3.6))
|
||||
|
||||
|
||||
async def test_thinking_message_omits_tokens_and_elapsed_when_zero():
|
||||
"""`thinking_elapsed_ms=0` → `elapsed_ms=None`, and
|
||||
`thinking_total_chars=0` → `tokens=None`. Matches the
|
||||
`... or None` / `if thinking_total_chars` guards."""
|
||||
loop, ws = _make_loop()
|
||||
content = [ContentBlock(type="thinking", text="thoughts")]
|
||||
|
||||
await loop._emit_collected_messages(
|
||||
content,
|
||||
text_msg_id=None,
|
||||
tool_msg_ids={},
|
||||
thinking_elapsed_ms=0,
|
||||
thinking_total_chars=0,
|
||||
)
|
||||
|
||||
thinking = next(
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "thinking"
|
||||
)
|
||||
assert thinking["elapsed_ms"] is None
|
||||
assert thinking["tokens"] is None
|
||||
|
||||
|
||||
async def test_tool_call_messages_use_stream_msg_ids_in_index_order():
|
||||
"""`tool_msg_ids` is keyed by stream block index. Emission must
|
||||
sort by index so the first emitted tool_call gets the id from
|
||||
index 0, even if the dict was inserted out of order."""
|
||||
loop, ws = _make_loop()
|
||||
content = [
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t1", name="Read", input={"x": 1}),
|
||||
),
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t2", name="Edit", input={"y": 2}),
|
||||
),
|
||||
]
|
||||
|
||||
# Insert higher index first to guard against accidental
|
||||
# insertion-order semantics in the future.
|
||||
await loop._emit_collected_messages(
|
||||
content,
|
||||
text_msg_id=None,
|
||||
tool_msg_ids={1: "ID-B", 0: "ID-A"},
|
||||
)
|
||||
|
||||
tool_calls = [
|
||||
d["message"] for d in ws.of_type("agent:message")
|
||||
if d["message"]["role"] == "tool_call"
|
||||
]
|
||||
assert [m["id"] for m in tool_calls] == ["ID-A", "ID-B"]
|
||||
assert tool_calls[0]["content"] == {"id": "t1", "tool": "Read", "input": {"x": 1}}
|
||||
assert tool_calls[1]["content"] == {"id": "t2", "tool": "Edit", "input": {"y": 2}}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Group 5 — usage tracking
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_token_usage_accumulates_across_turns():
|
||||
"""Per-turn `usage` events must sum into `total_input_tokens` /
|
||||
`total_output_tokens` over the whole `run()`."""
|
||||
provider = _StubProvider([
|
||||
_tool_turn(
|
||||
"Read", "t1", "{}",
|
||||
usage={"input_tokens": 100, "output_tokens": 50},
|
||||
),
|
||||
_text_turn(
|
||||
"done",
|
||||
usage={"input_tokens": 30, "output_tokens": 20},
|
||||
),
|
||||
])
|
||||
loop, _ = _make_loop(provider=provider)
|
||||
|
||||
await loop.run("hi")
|
||||
|
||||
assert loop.total_input_tokens == 130
|
||||
assert loop.total_output_tokens == 70
|
||||
@@ -559,64 +559,6 @@ def test_execute_no_backend_code_returns_none_result(client):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_vibe_code_happy_path(client, monkeypatch):
|
||||
response_json = json.dumps({
|
||||
"frontend_code": "<html>new</html>",
|
||||
"backend_code": "result = {'k': 1}",
|
||||
"input_schema": {"type": "object"},
|
||||
"name": "Generated",
|
||||
"description": "by AI",
|
||||
"message": "All set.",
|
||||
})
|
||||
_patch_anthropic(monkeypatch, response_json)
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/vibe-code",
|
||||
json={"prompt": "make me a thing", "current_frontend_code": "<old/>"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["frontend_code"] == "<html>new</html>"
|
||||
assert body["backend_code"] == "result = {'k': 1}"
|
||||
assert body["name"] == "Generated"
|
||||
assert body["message"] == "All set."
|
||||
|
||||
|
||||
def test_vibe_code_strips_markdown_fences(client, monkeypatch):
|
||||
fenced = "```json\n" + json.dumps({
|
||||
"frontend_code": "<html>fenced</html>",
|
||||
"message": "ok",
|
||||
}) + "\n```"
|
||||
_patch_anthropic(monkeypatch, fenced)
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/vibe-code",
|
||||
json={"prompt": "hi"},
|
||||
)
|
||||
body = resp.json()
|
||||
assert body["frontend_code"] == "<html>fenced</html>"
|
||||
|
||||
|
||||
def test_vibe_code_json_decode_error_keeps_user_code(client, monkeypatch):
|
||||
_patch_anthropic(monkeypatch, "not json at all")
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/vibe-code",
|
||||
json={
|
||||
"prompt": "x",
|
||||
"current_frontend_code": "<keep/>",
|
||||
"current_backend_code": "result={}",
|
||||
},
|
||||
)
|
||||
body = resp.json()
|
||||
assert "couldn't parse" in body["message"]
|
||||
assert body["frontend_code"] == "<keep/>"
|
||||
assert body["backend_code"] == "result={}"
|
||||
|
||||
|
||||
def test_vibe_code_resolver_raises_value_error_returns_graceful(client, monkeypatch):
|
||||
"""resolve_aux_model raises ValueError when no model is connected.
|
||||
The route catches and returns a graceful error response without
|
||||
@@ -651,87 +593,11 @@ def test_vibe_code_anthropic_import_error(client, monkeypatch):
|
||||
assert body["frontend_code"] == "<keep/>"
|
||||
|
||||
|
||||
def test_vibe_code_anthropic_call_raises_returns_graceful(client, monkeypatch):
|
||||
"""Generic exception from messages.create is caught and surfaced as
|
||||
`message: "Error: ..."` while preserving the user's existing code."""
|
||||
from backend.apps.outputs import outputs as outputs_mod
|
||||
|
||||
failing_client = MagicMock()
|
||||
failing_client.messages.create = AsyncMock(side_effect=RuntimeError("api down"))
|
||||
monkeypatch.setattr(outputs_mod, "_get_anthropic_client", lambda: failing_client)
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/vibe-code",
|
||||
json={"prompt": "x", "current_frontend_code": "<keep/>"},
|
||||
)
|
||||
body = resp.json()
|
||||
assert "api down" in body["message"]
|
||||
assert body["frontend_code"] == "<keep/>"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /auto-run (Anthropic mocked)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_auto_run_happy_path_with_backend(client, monkeypatch):
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "string"}},
|
||||
"required": ["q"],
|
||||
}
|
||||
_patch_anthropic(monkeypatch, json.dumps({"q": "hello"}))
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={
|
||||
"prompt": "give me data",
|
||||
"input_schema": schema,
|
||||
"backend_code": "print('side'); result['echoed'] = input_data['q']",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["input_data"] == {"q": "hello"}
|
||||
assert body["backend_result"] == {"echoed": "hello"}
|
||||
assert "side" in (body["stdout"] or "")
|
||||
assert body["error"] is None
|
||||
|
||||
|
||||
def test_auto_run_schema_validation_failure(client, monkeypatch):
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "integer"}},
|
||||
"required": ["q"],
|
||||
}
|
||||
_patch_anthropic(monkeypatch, json.dumps({"q": "string-not-int"}))
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={"prompt": "x", "input_schema": schema},
|
||||
)
|
||||
body = resp.json()
|
||||
assert body["input_data"] == {"q": "string-not-int"}
|
||||
assert body["backend_result"] is None
|
||||
assert "Schema validation failed" in body["error"]
|
||||
|
||||
|
||||
def test_auto_run_json_decode_error(client, monkeypatch):
|
||||
_patch_anthropic(monkeypatch, "not json at all")
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={"prompt": "x", "input_schema": {"type": "object"}},
|
||||
)
|
||||
body = resp.json()
|
||||
assert body["error"] == "Failed to parse generated data as JSON"
|
||||
assert body["input_data"] is None
|
||||
|
||||
|
||||
def test_auto_run_resolver_value_error(client, monkeypatch):
|
||||
from backend.apps.agents.providers import registry
|
||||
|
||||
@@ -751,44 +617,6 @@ def test_auto_run_resolver_value_error(client, monkeypatch):
|
||||
assert body["input_data"] is None
|
||||
|
||||
|
||||
def test_auto_run_known_model_uses_resolve_for_sdk(client, monkeypatch):
|
||||
"""When body.model is a known builtin, the route calls
|
||||
`resolve_model_id_for_sdk` instead of `resolve_aux_model`. Patch
|
||||
both: only the sdk resolver should be hit."""
|
||||
from backend.apps.agents.providers import registry
|
||||
|
||||
monkeypatch.setattr(registry, "_find_builtin_model", lambda _name: {"value": "sonnet"})
|
||||
|
||||
resolved_calls = []
|
||||
|
||||
def _fake_resolve(short_name, _settings):
|
||||
resolved_calls.append(short_name)
|
||||
return "claude-sonnet-real"
|
||||
|
||||
monkeypatch.setattr(registry, "resolve_model_id_for_sdk", _fake_resolve)
|
||||
|
||||
aux_calls = []
|
||||
|
||||
async def _aux(_settings, preferred_tier="haiku"):
|
||||
aux_calls.append(preferred_tier)
|
||||
return ("should-not-be-used", None)
|
||||
|
||||
monkeypatch.setattr(registry, "resolve_aux_model", _aux)
|
||||
|
||||
mock_client = _patch_anthropic(monkeypatch, json.dumps({"x": 1}))
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={"prompt": "x", "input_schema": {"type": "object"}, "model": "sonnet"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resolved_calls == ["sonnet"]
|
||||
assert aux_calls == []
|
||||
# And the model id flowed into the Anthropic client call.
|
||||
call_kwargs = mock_client.messages.create.await_args.kwargs
|
||||
assert call_kwargs["model"] == "claude-sonnet-real"
|
||||
|
||||
|
||||
def test_auto_run_anthropic_import_error(client, monkeypatch):
|
||||
monkeypatch.setitem(sys.modules, "anthropic", None)
|
||||
resp = client.post(
|
||||
@@ -800,61 +628,6 @@ def test_auto_run_anthropic_import_error(client, monkeypatch):
|
||||
assert body["input_data"] is None
|
||||
|
||||
|
||||
def test_auto_run_strips_markdown_fences(client, monkeypatch):
|
||||
"""Markdown-fenced JSON from the model is stripped before parsing
|
||||
(covers the `if raw.startswith('```')` branch in auto_run_output)."""
|
||||
schema = {"type": "object", "properties": {"q": {"type": "string"}}}
|
||||
fenced = "```json\n" + json.dumps({"q": "ok"}) + "\n```"
|
||||
_patch_anthropic(monkeypatch, fenced)
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={"prompt": "x", "input_schema": schema},
|
||||
)
|
||||
body = resp.json()
|
||||
assert body["input_data"] == {"q": "ok"}
|
||||
assert body["error"] is None
|
||||
|
||||
|
||||
def test_auto_run_backend_code_raise_populates_error(client, monkeypatch):
|
||||
"""If `execute_backend_code` raises, the route catches the exception
|
||||
and stuffs it into `error` while keeping the validated input_data."""
|
||||
schema = {"type": "object", "properties": {"q": {"type": "string"}}}
|
||||
_patch_anthropic(monkeypatch, json.dumps({"q": "ok"}))
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={
|
||||
"prompt": "x",
|
||||
"input_schema": schema,
|
||||
"backend_code": "raise RuntimeError('exec boom')",
|
||||
},
|
||||
)
|
||||
body = resp.json()
|
||||
assert body["input_data"] == {"q": "ok"}
|
||||
assert body["backend_result"] is None
|
||||
assert "exec boom" in (body["error"] or "")
|
||||
|
||||
|
||||
def test_auto_run_anthropic_call_raises(client, monkeypatch):
|
||||
"""Generic exception from messages.create lands in `error`."""
|
||||
from backend.apps.outputs import outputs as outputs_mod
|
||||
|
||||
failing_client = MagicMock()
|
||||
failing_client.messages.create = AsyncMock(side_effect=RuntimeError("api down"))
|
||||
monkeypatch.setattr(outputs_mod, "_get_anthropic_client", lambda: failing_client)
|
||||
_patch_aux_model(monkeypatch)
|
||||
|
||||
resp = client.post(
|
||||
"/api/outputs/auto-run",
|
||||
json={"prompt": "x", "input_schema": {"type": "object"}},
|
||||
)
|
||||
body = resp.json()
|
||||
assert "api down" in body["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# /auto-run-agent (stub_agent_loop + AgentConfig spy)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -46,7 +46,6 @@ from backend.apps.agents.browser_agent import (
|
||||
_create_browser_card,
|
||||
_validate_message_pairing,
|
||||
run_browser_agent,
|
||||
run_browser_agents,
|
||||
)
|
||||
from backend.apps.agents.models import AgentSession
|
||||
|
||||
@@ -831,90 +830,3 @@ async def test_create_browser_card_defaults_url_when_blank(monkeypatch):
|
||||
assert card.tabs[0].url == "https://www.google.com"
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# run_browser_agents — parallel fanout
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
async def test_run_browser_agents_fanout_auto_creates_card_and_returns_results(monkeypatch):
|
||||
"""Two tasks: one with browser_id, one without. The latter should
|
||||
auto-create a card. Both results returned in input order."""
|
||||
from backend.apps.dashboards.dashboards import _save
|
||||
from backend.apps.dashboards.models import Dashboard
|
||||
|
||||
_save(Dashboard(id="d-fanout", name="t"))
|
||||
|
||||
# Don't actually wait 2s for "card settle".
|
||||
monkeypatch.setattr(ba.asyncio, "sleep", AsyncMock())
|
||||
monkeypatch.setattr(ba.ws_manager, "broadcast_global", AsyncMock())
|
||||
# Patch analytics so we don't depend on its plumbing here.
|
||||
monkeypatch.setattr(
|
||||
"backend.apps.analytics.collector.record",
|
||||
lambda *a, **kw: None,
|
||||
)
|
||||
|
||||
call_log: list[dict] = []
|
||||
|
||||
async def _fake_run_one(**kwargs):
|
||||
call_log.append(kwargs)
|
||||
return {
|
||||
"session_id": f"s-{kwargs['browser_id']}",
|
||||
"browser_id": kwargs["browser_id"],
|
||||
"summary": f"summary for {kwargs['browser_id']}",
|
||||
"action_log": [],
|
||||
"final_screenshot": None,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(ba, "run_browser_agent", _fake_run_one)
|
||||
|
||||
results = await run_browser_agents(
|
||||
tasks=[
|
||||
{"browser_id": "b-existing", "task": "do A"},
|
||||
{"task": "do B", "url": "https://example.com"},
|
||||
],
|
||||
model="sonnet",
|
||||
dashboard_id="d-fanout",
|
||||
)
|
||||
|
||||
assert len(results) == 2
|
||||
assert results[0]["summary"] == "summary for b-existing"
|
||||
# Auto-created browser_id matches the second call's argument.
|
||||
auto_browser_id = call_log[1]["browser_id"]
|
||||
assert auto_browser_id.startswith("browser-")
|
||||
assert results[1]["summary"] == f"summary for {auto_browser_id}"
|
||||
|
||||
|
||||
async def test_run_browser_agents_exception_per_task_surfaces_in_results(monkeypatch):
|
||||
"""If run_browser_agent raises for one task, the other still
|
||||
completes and the failed slot becomes {summary: 'Error: ...'}."""
|
||||
monkeypatch.setattr(ba.asyncio, "sleep", AsyncMock())
|
||||
monkeypatch.setattr(ba.ws_manager, "broadcast_global", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
"backend.apps.analytics.collector.record",
|
||||
lambda *a, **kw: None,
|
||||
)
|
||||
|
||||
async def _maybe_explode(**kwargs):
|
||||
if kwargs["browser_id"] == "boom":
|
||||
raise RuntimeError("kaboom")
|
||||
return {
|
||||
"session_id": "s-ok",
|
||||
"browser_id": kwargs["browser_id"],
|
||||
"summary": "ok-summary",
|
||||
"action_log": [],
|
||||
"final_screenshot": None,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(ba, "run_browser_agent", _maybe_explode)
|
||||
|
||||
results = await run_browser_agents(
|
||||
tasks=[
|
||||
{"browser_id": "ok-1", "task": "fine"},
|
||||
{"browser_id": "boom", "task": "explodes"},
|
||||
],
|
||||
model="sonnet",
|
||||
)
|
||||
assert results[0]["summary"] == "ok-summary"
|
||||
assert results[1]["summary"].startswith("Error: ")
|
||||
assert results[1]["action_log"] == []
|
||||
assert results[1]["final_screenshot"] is None
|
||||
|
||||
@@ -19,7 +19,6 @@ Coverage targets:
|
||||
- hallucinated id outside CURATED_SHORTLIST dropped
|
||||
- timeout → default (fail-open)
|
||||
- generic exception → default (fail-open)
|
||||
- `_call_classifier` JSON cleanup with code-fence wrapping
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -35,7 +34,6 @@ from backend.apps.agents import mcp_preflight as pf
|
||||
from backend.apps.agents.mcp_preflight import (
|
||||
CURATED_SHORTLIST,
|
||||
_build_available_shortlist,
|
||||
_call_classifier,
|
||||
_decorate,
|
||||
_is_obviously_local,
|
||||
run_preflight,
|
||||
@@ -300,92 +298,3 @@ async def test_run_preflight_no_provider_classifier_value_error_returns_default(
|
||||
assert out == {"is_vague": False, "suggestions": []}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _call_classifier — JSON cleanup paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_classifier_setup(text: str):
|
||||
"""Build the patches needed to drive _call_classifier with a fake
|
||||
Anthropic client returning `text` as the assistant content."""
|
||||
from backend.apps.agents.mcp_preflight import resolve_aux_model as _real
|
||||
|
||||
fake_resp = SimpleNamespace(
|
||||
content=[SimpleNamespace(text=text)],
|
||||
)
|
||||
fake_client = MagicMock()
|
||||
fake_client.messages = MagicMock()
|
||||
fake_client.messages.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
return (
|
||||
patch.object(pf, "resolve_aux_model",
|
||||
AsyncMock(return_value=("claude-haiku-4-5", None))),
|
||||
patch.object(pf, "get_anthropic_client", return_value=fake_client),
|
||||
)
|
||||
|
||||
|
||||
async def test_call_classifier_strips_markdown_code_fences():
|
||||
"""Some models wrap JSON in ```json fences — preflight must strip
|
||||
them before parsing."""
|
||||
fenced = '```json\n{"is_vague": true, "suggestions": []}\n```'
|
||||
aux_p, client_p = _make_classifier_setup(fenced)
|
||||
with aux_p, client_p:
|
||||
data = await _call_classifier(AppSettings(), "anything", [])
|
||||
assert data == {"is_vague": True, "suggestions": []}
|
||||
|
||||
|
||||
async def test_call_classifier_strips_plain_code_fences():
|
||||
"""``` (without `json` tag) also stripped."""
|
||||
fenced = '```\n{"is_vague": false, "suggestions": []}\n```'
|
||||
aux_p, client_p = _make_classifier_setup(fenced)
|
||||
with aux_p, client_p:
|
||||
data = await _call_classifier(AppSettings(), "anything", [])
|
||||
assert data["is_vague"] is False
|
||||
|
||||
|
||||
async def test_call_classifier_normalizes_non_list_suggestions():
|
||||
"""If the model returns suggestions as a non-list (e.g. None or
|
||||
dict), normalize to []."""
|
||||
text = '{"is_vague": true, "suggestions": null}'
|
||||
aux_p, client_p = _make_classifier_setup(text)
|
||||
with aux_p, client_p:
|
||||
data = await _call_classifier(AppSettings(), "anything", [])
|
||||
assert data["suggestions"] == []
|
||||
|
||||
|
||||
async def test_call_classifier_raises_on_non_object_root():
|
||||
text = '"not an object"'
|
||||
aux_p, client_p = _make_classifier_setup(text)
|
||||
with aux_p, client_p, pytest.raises(ValueError):
|
||||
await _call_classifier(AppSettings(), "anything", [])
|
||||
|
||||
|
||||
async def test_call_classifier_handles_string_content_response():
|
||||
"""Some translators return content as a single string instead of a
|
||||
list of blocks. Adapter must coerce gracefully."""
|
||||
fake_resp = SimpleNamespace(content='{"is_vague": false, "suggestions": []}')
|
||||
fake_client = MagicMock()
|
||||
fake_client.messages = MagicMock()
|
||||
fake_client.messages.create = AsyncMock(return_value=fake_resp)
|
||||
with patch.object(pf, "resolve_aux_model",
|
||||
AsyncMock(return_value=("claude-haiku-4-5", None))), \
|
||||
patch.object(pf, "get_anthropic_client", return_value=fake_client):
|
||||
data = await _call_classifier(AppSettings(), "anything", [])
|
||||
assert data == {"is_vague": False, "suggestions": []}
|
||||
|
||||
|
||||
async def test_call_classifier_passes_aux_model_into_request():
|
||||
"""Verify the resolved aux model id reaches the upstream call."""
|
||||
fake_resp = SimpleNamespace(content=[SimpleNamespace(text='{"is_vague": false}')])
|
||||
fake_client = MagicMock()
|
||||
fake_client.messages = MagicMock()
|
||||
fake_client.messages.create = AsyncMock(return_value=fake_resp)
|
||||
with patch.object(pf, "resolve_aux_model",
|
||||
AsyncMock(return_value=("cc/claude-haiku-4-5-20251001", None))), \
|
||||
patch.object(pf, "get_anthropic_client", return_value=fake_client):
|
||||
await _call_classifier(AppSettings(), "anything", [])
|
||||
|
||||
_, kwargs = fake_client.messages.create.call_args
|
||||
assert kwargs["model"] == "cc/claude-haiku-4-5-20251001"
|
||||
assert kwargs["max_tokens"] == 300
|
||||
assert "is_vague" in kwargs["system"]
|
||||
|
||||
@@ -1,623 +0,0 @@
|
||||
"""Direct handler tests for the stdio MCP meta-servers.
|
||||
|
||||
The CLI-side MCP servers are launched as standalone Python subprocesses
|
||||
by the SDK. Subprocess startup is the SDK's job; here we just exercise
|
||||
the per-tool handler functions in-process. Each server's `call_backend`
|
||||
helper goes through `urllib.request.urlopen`, which we mock with a
|
||||
thin shim returning canned JSON.
|
||||
|
||||
Coverage targets (all currently 0%):
|
||||
- `outputs_meta_server`: TOOLS shape, OutputList success + empty +
|
||||
error, OutputSearch missing query + matches, OutputActivate
|
||||
unknown / already_active / activated paths, format_outputs
|
||||
- `mcp_meta_server`: TOOLS shape, MCPList success + empty + error,
|
||||
MCPSearch missing query + matches, MCPActivate unknown /
|
||||
already_active / activated paths
|
||||
- `web_mcp_server`: WebSearch + WebFetch happy paths and error
|
||||
branches, schema validation
|
||||
- `invoke_agent_mcp_server`: TOOLS shape, missing args, success
|
||||
payload formatting (cost line + source_name)
|
||||
- `browser_mcp_server`: action_map dispatch, missing browser_id,
|
||||
screenshot too large, text fallback
|
||||
- `browser_agent_mcp_server`: format_result + format_batch_results,
|
||||
CreateBrowserAgent / BrowserAgent / BrowserAgents validation
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents import (
|
||||
browser_agent_mcp_server as ba_srv,
|
||||
browser_mcp_server as br_srv,
|
||||
invoke_agent_mcp_server as inv_srv,
|
||||
mcp_meta_server as mcp_srv,
|
||||
outputs_meta_server as out_srv,
|
||||
web_mcp_server as web_srv,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _mock_backend(module, payload: dict | list):
|
||||
"""Patch the module's `urllib.request.urlopen` to return `payload`
|
||||
as JSON. Works for every mcp meta-server because they all use the
|
||||
same stdlib request/json round-trip."""
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.read.return_value = json.dumps(payload).encode()
|
||||
fake_resp.__enter__ = MagicMock(return_value=fake_resp)
|
||||
fake_resp.__exit__ = MagicMock(return_value=False)
|
||||
with patch.object(module.urllib.request, "urlopen", return_value=fake_resp):
|
||||
yield
|
||||
|
||||
|
||||
def _text(blocks: list[dict]) -> str:
|
||||
return "".join(b.get("text", "") for b in blocks if b.get("type") == "text")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# outputs_meta_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_outputs_meta_tools_shape():
|
||||
"""Every TOOLS entry must have name + description + inputSchema."""
|
||||
names = {t["name"] for t in out_srv.TOOLS}
|
||||
assert names == {"OutputList", "OutputSearch", "OutputActivate"}
|
||||
for t in out_srv.TOOLS:
|
||||
assert "description" in t
|
||||
assert "inputSchema" in t
|
||||
schema = t["inputSchema"]
|
||||
assert schema["type"] == "object"
|
||||
|
||||
|
||||
def test_outputs_format_outputs_renders_status_and_use_count():
|
||||
out = out_srv.format_outputs(
|
||||
[{
|
||||
"id": "abc",
|
||||
"name": "My View",
|
||||
"description": "Does X",
|
||||
"status": "active",
|
||||
"use_count": 7,
|
||||
}],
|
||||
heading="Active:",
|
||||
)
|
||||
assert out.startswith("Active:")
|
||||
assert "`abc`" in out and "**My View**" in out
|
||||
assert "[active]" in out
|
||||
assert "(used 7×)" in out
|
||||
assert "Does X" in out
|
||||
|
||||
|
||||
def test_outputs_format_outputs_empty_returns_empty_string():
|
||||
assert out_srv.format_outputs([]) == ""
|
||||
|
||||
|
||||
def test_outputs_handle_list_empty():
|
||||
with _mock_backend(out_srv, {"active": [], "available": []}):
|
||||
out = out_srv.handle_tool_call("OutputList", {})
|
||||
assert "No Outputs / Views are defined" in _text(out["content"])
|
||||
assert "isError" not in out
|
||||
|
||||
|
||||
def test_outputs_handle_list_with_data():
|
||||
with _mock_backend(out_srv, {
|
||||
"active": [{"id": "a1", "name": "A", "description": "x", "status": "active"}],
|
||||
"available": [{"id": "b2", "name": "B", "description": "y", "status": "available"}],
|
||||
}):
|
||||
out = out_srv.handle_tool_call("OutputList", {})
|
||||
text = _text(out["content"])
|
||||
assert "Active" in text and "Available" in text
|
||||
assert "`a1`" in text and "`b2`" in text
|
||||
|
||||
|
||||
def test_outputs_handle_list_backend_error():
|
||||
with _mock_backend(out_srv, {"error": "backend down"}):
|
||||
out = out_srv.handle_tool_call("OutputList", {})
|
||||
assert out.get("isError") is True
|
||||
assert "backend down" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_search_missing_query():
|
||||
out = out_srv.handle_tool_call("OutputSearch", {})
|
||||
assert out.get("isError") is True
|
||||
assert "query is required" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_search_no_matches():
|
||||
with _mock_backend(out_srv, {"matches": []}):
|
||||
out = out_srv.handle_tool_call("OutputSearch", {"query": "anything"})
|
||||
assert "No Outputs matched" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_search_with_matches_includes_next_step():
|
||||
with _mock_backend(out_srv, {"matches": [
|
||||
{"id": "v1", "name": "View1", "description": "x", "status": "available"},
|
||||
]}):
|
||||
out = out_srv.handle_tool_call("OutputSearch", {"query": "view"})
|
||||
text = _text(out["content"])
|
||||
assert "`v1`" in text
|
||||
assert "OutputActivate" in text
|
||||
|
||||
|
||||
def test_outputs_handle_activate_missing_id():
|
||||
out = out_srv.handle_tool_call("OutputActivate", {})
|
||||
assert out.get("isError") is True
|
||||
assert "output_id is required" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_activate_unknown():
|
||||
with _mock_backend(out_srv, {
|
||||
"status": "unknown_output",
|
||||
"available": ["v1", "v2"],
|
||||
}):
|
||||
out = out_srv.handle_tool_call("OutputActivate", {"output_id": "phantom"})
|
||||
assert out.get("isError") is True
|
||||
text = _text(out["content"])
|
||||
assert "Unknown Output id" in text
|
||||
assert "`v1`" in text and "`v2`" in text
|
||||
|
||||
|
||||
def test_outputs_handle_activate_already_active():
|
||||
with _mock_backend(out_srv, {"status": "already_active"}):
|
||||
out = out_srv.handle_tool_call("OutputActivate", {"output_id": "v1"})
|
||||
assert "isError" not in out
|
||||
assert "already active" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_activate_activated():
|
||||
with _mock_backend(out_srv, {"status": "activated"}):
|
||||
out = out_srv.handle_tool_call("OutputActivate", {"output_id": "v1"})
|
||||
assert "isError" not in out
|
||||
assert "Activated Output `v1`" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_activate_unexpected_status():
|
||||
with _mock_backend(out_srv, {"status": "wat"}):
|
||||
out = out_srv.handle_tool_call("OutputActivate", {"output_id": "v1"})
|
||||
assert out.get("isError") is True
|
||||
assert "Unexpected response" in _text(out["content"])
|
||||
|
||||
|
||||
def test_outputs_handle_unknown_tool():
|
||||
out = out_srv.handle_tool_call("NotARealTool", {})
|
||||
assert out.get("isError") is True
|
||||
assert "Unknown tool" in _text(out["content"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# mcp_meta_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_mcp_meta_tools_shape():
|
||||
names = {t["name"] for t in mcp_srv.TOOLS}
|
||||
assert names == {"MCPList", "MCPSearch", "MCPActivate"}
|
||||
for t in mcp_srv.TOOLS:
|
||||
assert "inputSchema" in t
|
||||
|
||||
|
||||
def test_mcp_meta_format_servers_renders_status():
|
||||
out = mcp_srv.format_servers(
|
||||
[{"name": "slack", "description": "Slack tools", "status": "active"}],
|
||||
heading="Active:",
|
||||
)
|
||||
assert "Active:" in out
|
||||
assert "`slack`" in out and "[active]" in out
|
||||
|
||||
|
||||
def test_mcp_meta_handle_list_empty():
|
||||
with _mock_backend(mcp_srv, {"active": [], "available": []}):
|
||||
out = mcp_srv.handle_tool_call("MCPList", {})
|
||||
assert "No MCP servers are installed" in _text(out["content"])
|
||||
|
||||
|
||||
def test_mcp_meta_handle_list_with_data():
|
||||
with _mock_backend(mcp_srv, {
|
||||
"active": [{"name": "slack", "description": "x", "status": "active"}],
|
||||
"available": [{"name": "discord", "description": "y", "status": "available"}],
|
||||
}):
|
||||
out = mcp_srv.handle_tool_call("MCPList", {})
|
||||
text = _text(out["content"])
|
||||
assert "Active" in text and "Available" in text
|
||||
|
||||
|
||||
def test_mcp_meta_handle_list_backend_error():
|
||||
with _mock_backend(mcp_srv, {"error": "boom"}):
|
||||
out = mcp_srv.handle_tool_call("MCPList", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_mcp_meta_handle_search_missing_query():
|
||||
out = mcp_srv.handle_tool_call("MCPSearch", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_mcp_meta_handle_search_no_matches():
|
||||
with _mock_backend(mcp_srv, {"matches": []}):
|
||||
out = mcp_srv.handle_tool_call("MCPSearch", {"query": "x"})
|
||||
assert "No MCP servers matched" in _text(out["content"])
|
||||
|
||||
|
||||
def test_mcp_meta_handle_search_with_matches_includes_next_step():
|
||||
with _mock_backend(mcp_srv, {"matches": [
|
||||
{"name": "slack", "description": "S", "status": "available"},
|
||||
]}):
|
||||
out = mcp_srv.handle_tool_call("MCPSearch", {"query": "channel"})
|
||||
text = _text(out["content"])
|
||||
assert "MCPActivate" in text
|
||||
|
||||
|
||||
def test_mcp_meta_handle_activate_missing_name():
|
||||
out = mcp_srv.handle_tool_call("MCPActivate", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_mcp_meta_handle_activate_unknown_returns_valid_options():
|
||||
with _mock_backend(mcp_srv, {
|
||||
"status": "unknown_server",
|
||||
"available": ["slack", "notion"],
|
||||
}):
|
||||
out = mcp_srv.handle_tool_call("MCPActivate", {"server_name": "phantom"})
|
||||
text = _text(out["content"])
|
||||
assert out.get("isError") is True
|
||||
assert "`slack`" in text and "`notion`" in text
|
||||
|
||||
|
||||
def test_mcp_meta_handle_activate_already_active():
|
||||
with _mock_backend(mcp_srv, {"status": "already_active"}):
|
||||
out = mcp_srv.handle_tool_call("MCPActivate", {"server_name": "slack"})
|
||||
assert "isError" not in out
|
||||
assert "already active" in _text(out["content"])
|
||||
|
||||
|
||||
def test_mcp_meta_handle_activate_activated():
|
||||
with _mock_backend(mcp_srv, {"status": "activated"}):
|
||||
out = mcp_srv.handle_tool_call("MCPActivate", {"server_name": "slack"})
|
||||
assert "isError" not in out
|
||||
text = _text(out["content"])
|
||||
assert "mcp__slack__" in text # next-turn hint
|
||||
|
||||
|
||||
def test_mcp_meta_handle_unknown_tool():
|
||||
out = mcp_srv.handle_tool_call("NotReal", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# web_mcp_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_web_mcp_tools_shape():
|
||||
names = {t["name"] for t in web_srv.TOOLS}
|
||||
assert names == {"WebSearch", "WebFetch"}
|
||||
|
||||
|
||||
def test_web_mcp_websearch_missing_query():
|
||||
out = web_srv.handle_tool_call("WebSearch", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_web_mcp_websearch_returns_results():
|
||||
with _mock_backend(web_srv, {"results": "[1] Title\n https://example.com"}):
|
||||
out = web_srv.handle_tool_call("WebSearch", {"query": "openswarm"})
|
||||
text = _text(out["content"])
|
||||
assert "Title" in text
|
||||
|
||||
|
||||
def test_web_mcp_websearch_empty_results_falls_back_to_marker():
|
||||
with _mock_backend(web_srv, {"results": ""}):
|
||||
out = web_srv.handle_tool_call("WebSearch", {"query": "missing"})
|
||||
assert "No results for: missing" in _text(out["content"])
|
||||
|
||||
|
||||
def test_web_mcp_websearch_backend_error():
|
||||
with _mock_backend(web_srv, {"error": "ddg down"}):
|
||||
out = web_srv.handle_tool_call("WebSearch", {"query": "x"})
|
||||
assert out.get("isError") is True
|
||||
assert "Search failed" in _text(out["content"])
|
||||
|
||||
|
||||
def test_web_mcp_websearch_clamps_num_results():
|
||||
"""num_results > 10 is clamped down to 10."""
|
||||
captured: dict = {}
|
||||
|
||||
def _fake_post(url, body, timeout=45.0):
|
||||
captured.update(body)
|
||||
return {"results": "ok"}
|
||||
|
||||
with patch.object(web_srv, "_post", side_effect=_fake_post):
|
||||
web_srv.handle_tool_call("WebSearch", {"query": "x", "num_results": 50})
|
||||
assert captured["num_results"] == 10
|
||||
|
||||
|
||||
def test_web_mcp_webfetch_missing_url():
|
||||
out = web_srv.handle_tool_call("WebFetch", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_web_mcp_webfetch_invalid_scheme():
|
||||
out = web_srv.handle_tool_call("WebFetch", {"url": "ftp://example.com"})
|
||||
assert out.get("isError") is True
|
||||
assert "must start with http" in _text(out["content"])
|
||||
|
||||
|
||||
def test_web_mcp_webfetch_returns_content():
|
||||
with _mock_backend(web_srv, {"content": "Plain text content"}):
|
||||
out = web_srv.handle_tool_call(
|
||||
"WebFetch",
|
||||
{"url": "https://example.com", "prompt": "x"},
|
||||
)
|
||||
assert "Plain text content" in _text(out["content"])
|
||||
|
||||
|
||||
def test_web_mcp_webfetch_empty_content_falls_back_to_marker():
|
||||
with _mock_backend(web_srv, {"content": ""}):
|
||||
out = web_srv.handle_tool_call("WebFetch", {"url": "https://example.com"})
|
||||
assert "No content returned from" in _text(out["content"])
|
||||
|
||||
|
||||
def test_web_mcp_unknown_tool_returns_error():
|
||||
out = web_srv.handle_tool_call("Phantom", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# invoke_agent_mcp_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_invoke_agent_tools_shape():
|
||||
names = {t["name"] for t in inv_srv.TOOLS}
|
||||
assert names == {"InvokeAgent"}
|
||||
|
||||
|
||||
def test_invoke_agent_unknown_tool():
|
||||
out = inv_srv.handle_tool_call("NotReal", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_invoke_agent_missing_session_id():
|
||||
out = inv_srv.handle_tool_call("InvokeAgent", {"message": "hi"})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_invoke_agent_missing_message():
|
||||
out = inv_srv.handle_tool_call("InvokeAgent", {"session_id": "x"})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_invoke_agent_backend_error():
|
||||
with _mock_backend(inv_srv, {"error": "agent down"}):
|
||||
out = inv_srv.handle_tool_call("InvokeAgent", {
|
||||
"session_id": "x", "message": "hi",
|
||||
})
|
||||
assert out.get("isError") is True
|
||||
assert "agent down" in _text(out["content"])
|
||||
|
||||
|
||||
def test_invoke_agent_success_format_includes_cost_and_source_name():
|
||||
with _mock_backend(inv_srv, {
|
||||
"forked_session_id": "fork-1",
|
||||
"response": "Did the thing",
|
||||
"cost_usd": 0.01,
|
||||
"source_name": "Original Agent",
|
||||
}):
|
||||
out = inv_srv.handle_tool_call("InvokeAgent", {
|
||||
"session_id": "x", "message": "hi",
|
||||
})
|
||||
text = _text(out["content"])
|
||||
assert "Original Agent" in text
|
||||
assert "fork-1" in text
|
||||
assert "$0.0100" in text
|
||||
assert "Did the thing" in text
|
||||
|
||||
|
||||
def test_invoke_agent_zero_cost_omits_cost_line():
|
||||
with _mock_backend(inv_srv, {
|
||||
"forked_session_id": "fork-1",
|
||||
"response": "Result",
|
||||
"cost_usd": 0,
|
||||
}):
|
||||
out = inv_srv.handle_tool_call("InvokeAgent", {
|
||||
"session_id": "x", "message": "hi",
|
||||
})
|
||||
text = _text(out["content"])
|
||||
assert "Cost" not in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# browser_mcp_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_browser_mcp_handle_missing_browser_id():
|
||||
out = br_srv.handle_tool_call("BrowserScreenshot", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_mcp_unknown_tool():
|
||||
out = br_srv.handle_tool_call("Phantom", {"browser_id": "b1"})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_mcp_get_text_dispatches_action():
|
||||
captured: dict = {}
|
||||
|
||||
def _fake_call(action, browser_id, params=None, tab_id=""):
|
||||
captured["action"] = action
|
||||
captured["browser_id"] = browser_id
|
||||
captured["params"] = params
|
||||
return {"text": "page contents here"}
|
||||
|
||||
with patch.object(br_srv, "call_backend", side_effect=_fake_call):
|
||||
out = br_srv.handle_tool_call("BrowserGetText", {"browser_id": "b1"})
|
||||
|
||||
assert captured["action"] == "get_text"
|
||||
assert captured["browser_id"] == "b1"
|
||||
text = _text(out["content"])
|
||||
assert "page contents here" in text
|
||||
|
||||
|
||||
def test_browser_mcp_navigate_passes_url_in_params():
|
||||
captured: dict = {}
|
||||
|
||||
def _fake_call(action, browser_id, params=None, tab_id=""):
|
||||
captured["params"] = params
|
||||
return {"text": "ok"}
|
||||
|
||||
with patch.object(br_srv, "call_backend", side_effect=_fake_call):
|
||||
br_srv.handle_tool_call(
|
||||
"BrowserNavigate",
|
||||
{"browser_id": "b1", "url": "https://example.com"},
|
||||
)
|
||||
assert captured["params"] == {"url": "https://example.com"}
|
||||
|
||||
|
||||
def test_browser_mcp_screenshot_returns_image_block():
|
||||
with patch.object(br_srv, "call_backend", return_value={
|
||||
"image": "AA==",
|
||||
"url": "https://example.com",
|
||||
}):
|
||||
out = br_srv.handle_tool_call(
|
||||
"BrowserScreenshot",
|
||||
{"browser_id": "b1"},
|
||||
)
|
||||
types = [b["type"] for b in out["content"]]
|
||||
assert "image" in types
|
||||
assert "text" in types
|
||||
|
||||
|
||||
def test_browser_mcp_screenshot_too_large_returns_text_only():
|
||||
"""Massive base64 with PIL unavailable → return text-only fallback."""
|
||||
huge = "x" * (br_srv.MAX_IMAGE_B64_BYTES + 1)
|
||||
with patch.object(br_srv, "call_backend", return_value={"image": huge, "url": "x"}), \
|
||||
patch.object(br_srv, "compress_screenshot", return_value=None):
|
||||
out = br_srv.handle_tool_call("BrowserScreenshot", {"browser_id": "b1"})
|
||||
assert all(b["type"] == "text" for b in out["content"])
|
||||
assert "too large" in _text(out["content"])
|
||||
|
||||
|
||||
def test_browser_mcp_backend_error():
|
||||
with patch.object(br_srv, "call_backend", return_value={"error": "ws disconnected"}):
|
||||
out = br_srv.handle_tool_call("BrowserGetText", {"browser_id": "b1"})
|
||||
assert out.get("isError") is True
|
||||
assert "ws disconnected" in _text(out["content"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# browser_agent_mcp_server
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_browser_agent_tools_shape():
|
||||
names = {t["name"] for t in ba_srv.TOOLS}
|
||||
assert names == {"CreateBrowserAgent", "BrowserAgent", "BrowserAgents"}
|
||||
|
||||
|
||||
def test_browser_agent_format_result_text_only():
|
||||
out = ba_srv.format_result({
|
||||
"summary": "Did the thing",
|
||||
"session_id": "s1",
|
||||
"browser_id": "b1",
|
||||
"action_log": [
|
||||
{"tool": "BrowserNavigate", "input": {"url": "https://example.com"}, "elapsed_ms": 50},
|
||||
{"tool": "BrowserClick", "input": {"selector": "#go"}, "elapsed_ms": 10},
|
||||
],
|
||||
})
|
||||
text = _text(out["content"])
|
||||
assert "Browser Agent Result" in text
|
||||
assert "Did the thing" in text
|
||||
assert "BrowserNavigate" in text
|
||||
assert "BrowserClick" in text
|
||||
# No screenshot → no image content
|
||||
assert all(b["type"] == "text" for b in out["content"])
|
||||
|
||||
|
||||
def test_browser_agent_format_result_error():
|
||||
out = ba_srv.format_result({"error": "no browser"})
|
||||
assert out.get("isError") is True
|
||||
assert "no browser" in _text(out["content"])
|
||||
|
||||
|
||||
def test_browser_agent_format_batch_results_separates_with_divider():
|
||||
out = ba_srv.format_batch_results([
|
||||
{"summary": "A", "session_id": "s1", "browser_id": "b1", "action_log": []},
|
||||
{"summary": "B", "session_id": "s2", "browser_id": "b2", "action_log": []},
|
||||
])
|
||||
text = _text(out["content"])
|
||||
assert "A" in text and "B" in text
|
||||
assert "---" in text
|
||||
|
||||
|
||||
def test_browser_agent_format_batch_results_top_level_error():
|
||||
out = ba_srv.format_batch_results({"error": "all failed"})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_agent_create_calls_backend_and_formats_first_result():
|
||||
with patch.object(ba_srv, "call_backend", return_value={"results": [{
|
||||
"summary": "Done", "session_id": "s1", "browser_id": "b1", "action_log": [],
|
||||
}]}):
|
||||
out = ba_srv.handle_tool_call("CreateBrowserAgent", {"task": "fetch a page"})
|
||||
assert "Done" in _text(out["content"])
|
||||
|
||||
|
||||
def test_browser_agent_browser_agent_missing_browser_id():
|
||||
out = ba_srv.handle_tool_call("BrowserAgent", {"task": "x"})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_agent_browser_agents_empty_tasks_errors():
|
||||
out = ba_srv.handle_tool_call("BrowserAgents", {"tasks": []})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_agent_browser_agents_missing_browser_id_in_task():
|
||||
out = ba_srv.handle_tool_call("BrowserAgents", {"tasks": [
|
||||
{"task": "x"}, # no browser_id
|
||||
]})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_agent_browser_agents_success():
|
||||
with patch.object(ba_srv, "call_backend", return_value={"results": [
|
||||
{"summary": "A", "session_id": "s1", "browser_id": "b1", "action_log": []},
|
||||
]}):
|
||||
out = ba_srv.handle_tool_call("BrowserAgents", {"tasks": [
|
||||
{"browser_id": "b1", "task": "x"},
|
||||
]})
|
||||
assert "A" in _text(out["content"])
|
||||
|
||||
|
||||
def test_browser_agent_unknown_tool():
|
||||
out = ba_srv.handle_tool_call("Phantom", {})
|
||||
assert out.get("isError") is True
|
||||
|
||||
|
||||
def test_browser_agent_call_backend_http_error():
|
||||
"""call_backend's exception branch surfaces the error string."""
|
||||
import urllib.error
|
||||
err = urllib.error.HTTPError(
|
||||
url="x", code=500, msg="boom", hdrs=None, fp=io.BytesIO(b"server err"),
|
||||
)
|
||||
with patch.object(ba_srv.urllib.request, "urlopen", side_effect=err):
|
||||
out = ba_srv.call_backend([{"task": "x", "browser_id": "b1", "url": ""}])
|
||||
assert "error" in out
|
||||
assert "HTTP 500" in out["error"]
|
||||
|
||||
|
||||
def test_browser_agent_call_backend_generic_exception():
|
||||
with patch.object(ba_srv.urllib.request, "urlopen", side_effect=RuntimeError("dns")):
|
||||
out = ba_srv.call_backend([{"task": "x", "browser_id": "b1", "url": ""}])
|
||||
assert out == {"error": "dns"}
|
||||
@@ -1,530 +0,0 @@
|
||||
"""Tests for the uncovered branches of the Anthropic provider adapter.
|
||||
|
||||
`test_phase1_stress.py::test_anthropic_provider_forwards_thinking_blocks`
|
||||
already covers the streaming-thinking path. This file fills in the
|
||||
remaining branches: model id mapping, message-format helpers,
|
||||
non-streaming `create_message`, `_build_messages` (tool_result list
|
||||
vs. single-dict shape), and the `message_start` / `message_delta`
|
||||
usage-extraction code in `stream_message`.
|
||||
|
||||
The Anthropic SDK client is fully mocked — no network, no API key
|
||||
required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents.providers.anthropic import AnthropicProvider, MODEL_MAP
|
||||
from backend.apps.agents.providers.base import (
|
||||
ContentBlock,
|
||||
ModelResponse,
|
||||
ProviderMessage,
|
||||
ToolCall,
|
||||
ToolSchema,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lightweight fakes (mimics the SDK's duck-typed objects without pulling
|
||||
# in the real anthropic types — they're a heavy import path).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeAttr:
|
||||
"""Generic dot-attribute object for SDK-shaped responses."""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
def _make_provider() -> AnthropicProvider:
|
||||
"""Build a provider whose underlying SDK client is fully mocked."""
|
||||
p = AnthropicProvider(api_key="test-key")
|
||||
# Replace the AsyncAnthropic client wholesale; tests will set the
|
||||
# specific behaviour on `client.messages.create`.
|
||||
p.client = MagicMock()
|
||||
p.client.messages = MagicMock()
|
||||
p.client.messages.create = AsyncMock()
|
||||
return p
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_model_id
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("short,full", list(MODEL_MAP.items()))
|
||||
def test_get_model_id_short_name_resolves_to_full(short: str, full: str):
|
||||
p = _make_provider()
|
||||
assert p.get_model_id(short) == full
|
||||
|
||||
|
||||
def test_get_model_id_passthrough_for_unknown():
|
||||
"""Anything not in MODEL_MAP is returned verbatim."""
|
||||
p = _make_provider()
|
||||
assert p.get_model_id("claude-7-sonnet-20991231") == "claude-7-sonnet-20991231"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_user_message / format_assistant_message / format_tool_result
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_format_user_message_string_content():
|
||||
p = _make_provider()
|
||||
msg = p.format_user_message("hello")
|
||||
assert isinstance(msg, ProviderMessage)
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "hello"
|
||||
|
||||
|
||||
def test_format_user_message_multimodal_blocks():
|
||||
"""Image + text user message — content list should pass through unchanged."""
|
||||
p = _make_provider()
|
||||
blocks = [
|
||||
{"type": "text", "text": "look at this"},
|
||||
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}},
|
||||
]
|
||||
msg = p.format_user_message(blocks)
|
||||
assert msg.role == "user"
|
||||
assert msg.content is blocks
|
||||
|
||||
|
||||
def test_format_assistant_message_text_only():
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[ContentBlock(type="text", text="hello")],
|
||||
stop_reason="end_turn",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
assert msg.role == "assistant"
|
||||
assert msg.content == [{"type": "text", "text": "hello"}]
|
||||
|
||||
|
||||
def test_format_assistant_message_mixed_text_and_tool_use():
|
||||
"""ContentBlocks of type=text/tool_use must round-trip into the
|
||||
Anthropic message-content shape, preserving id+name+input."""
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(type="text", text="thinking…"),
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t1", name="Read", input={"path": "/tmp/x"}),
|
||||
),
|
||||
ContentBlock(type="text", text="done"),
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
assert msg.role == "assistant"
|
||||
assert msg.content == [
|
||||
{"type": "text", "text": "thinking…"},
|
||||
{"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "/tmp/x"}},
|
||||
{"type": "text", "text": "done"},
|
||||
]
|
||||
|
||||
|
||||
def test_format_assistant_message_skips_tool_use_without_call():
|
||||
"""A tool_use block with no ToolCall is dropped (defensive — should
|
||||
never happen in practice, but the conditional is in the source)."""
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(type="text", text="hi"),
|
||||
ContentBlock(type="tool_use", tool_call=None), # silently dropped
|
||||
],
|
||||
stop_reason="end_turn",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
assert msg.content == [{"type": "text", "text": "hi"}]
|
||||
|
||||
|
||||
def test_format_tool_result_shape():
|
||||
p = _make_provider()
|
||||
out = p.format_tool_result(
|
||||
"tool-id-7",
|
||||
[{"type": "text", "text": "result body"}],
|
||||
)
|
||||
assert out == {
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "tool-id-7",
|
||||
"content": [{"type": "text", "text": "result body"}],
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# clean_tool_schema
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_clean_tool_schema_returns_anthropic_format():
|
||||
p = _make_provider()
|
||||
schema = ToolSchema(
|
||||
name="Read",
|
||||
description="Read a file",
|
||||
input_schema={"type": "object", "properties": {"path": {"type": "string"}}},
|
||||
)
|
||||
out = p.clean_tool_schema(schema)
|
||||
assert out == {
|
||||
"name": "Read",
|
||||
"description": "Read a file",
|
||||
"input_schema": {"type": "object", "properties": {"path": {"type": "string"}}},
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _build_messages — every role + tool_result list/non-list shape
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_build_messages_passes_through_user_and_assistant():
|
||||
p = _make_provider()
|
||||
msgs = [
|
||||
ProviderMessage(role="user", content="hi"),
|
||||
ProviderMessage(role="assistant", content=[{"type": "text", "text": "hello"}]),
|
||||
]
|
||||
built = p._build_messages(msgs)
|
||||
assert built == [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "hello"}]},
|
||||
]
|
||||
|
||||
|
||||
def test_build_messages_tool_result_list_passes_unwrapped():
|
||||
"""Tool results delivered as a list of result blocks must be sent
|
||||
as-is under role=user."""
|
||||
p = _make_provider()
|
||||
blocks = [
|
||||
{"type": "tool_result", "tool_use_id": "t1", "content": [{"type": "text", "text": "ok"}]},
|
||||
{"type": "tool_result", "tool_use_id": "t2", "content": [{"type": "text", "text": "ok2"}]},
|
||||
]
|
||||
built = p._build_messages([ProviderMessage(role="tool_result", content=blocks)])
|
||||
assert built == [{"role": "user", "content": blocks}]
|
||||
|
||||
|
||||
def test_build_messages_tool_result_single_dict_gets_wrapped_in_list():
|
||||
"""A non-list tool_result content must be wrapped: the API expects
|
||||
`content` to always be a list at this level."""
|
||||
p = _make_provider()
|
||||
block = {"type": "tool_result", "tool_use_id": "t1", "content": "ok"}
|
||||
built = p._build_messages([ProviderMessage(role="tool_result", content=block)])
|
||||
assert built == [{"role": "user", "content": [block]}]
|
||||
|
||||
|
||||
def test_build_messages_drops_unknown_role():
|
||||
"""If somehow a ProviderMessage with role='something_weird' is
|
||||
passed in, the loop must skip it rather than crash."""
|
||||
p = _make_provider()
|
||||
built = p._build_messages([
|
||||
ProviderMessage(role="weird", content="x"),
|
||||
ProviderMessage(role="user", content="hi"),
|
||||
])
|
||||
assert built == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_message (non-streaming)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_create_message_returns_normalized_response_text_only():
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
content=[_FakeAttr(type="text", text="hello world")],
|
||||
stop_reason="end_turn",
|
||||
usage=_FakeAttr(input_tokens=42, output_tokens=7),
|
||||
)
|
||||
p.client.messages.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="sonnet",
|
||||
system="be useful",
|
||||
messages=[ProviderMessage(role="user", content="hi")],
|
||||
tools=[],
|
||||
)
|
||||
|
||||
# Kwargs were translated through get_model_id + clean_tool_schema
|
||||
args, kwargs = p.client.messages.create.call_args
|
||||
assert kwargs["model"] == MODEL_MAP["sonnet"]
|
||||
assert kwargs["max_tokens"] == 8192
|
||||
assert kwargs["system"] == "be useful"
|
||||
assert kwargs["messages"] == [{"role": "user", "content": "hi"}]
|
||||
assert "tools" not in kwargs # empty list = omit
|
||||
|
||||
assert isinstance(out, ModelResponse)
|
||||
assert out.stop_reason == "end_turn"
|
||||
assert out.usage == {"input_tokens": 42, "output_tokens": 7}
|
||||
assert len(out.content) == 1
|
||||
assert out.content[0].type == "text"
|
||||
assert out.content[0].text == "hello world"
|
||||
|
||||
|
||||
async def test_create_message_translates_tool_use_block():
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
content=[
|
||||
_FakeAttr(type="text", text="let me check"),
|
||||
_FakeAttr(
|
||||
type="tool_use",
|
||||
id="toolu_1",
|
||||
name="Read",
|
||||
input={"path": "/tmp/x"},
|
||||
),
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
usage=_FakeAttr(input_tokens=10, output_tokens=2),
|
||||
)
|
||||
p.client.messages.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="opus",
|
||||
system=None,
|
||||
messages=[ProviderMessage(role="user", content="x")],
|
||||
tools=[
|
||||
ToolSchema(name="Read", description="Read a file",
|
||||
input_schema={"type": "object"}),
|
||||
],
|
||||
)
|
||||
|
||||
_, kwargs = p.client.messages.create.call_args
|
||||
assert kwargs["model"] == MODEL_MAP["opus"]
|
||||
assert kwargs["tools"] == [{
|
||||
"name": "Read", "description": "Read a file",
|
||||
"input_schema": {"type": "object"},
|
||||
}]
|
||||
assert "system" not in kwargs # None is dropped
|
||||
|
||||
assert out.stop_reason == "tool_use"
|
||||
assert out.content[0].type == "text"
|
||||
assert out.content[1].type == "tool_use"
|
||||
assert out.content[1].tool_call is not None
|
||||
assert out.content[1].tool_call.id == "toolu_1"
|
||||
assert out.content[1].tool_call.name == "Read"
|
||||
assert out.content[1].tool_call.input == {"path": "/tmp/x"}
|
||||
|
||||
|
||||
async def test_create_message_max_tokens_passthrough():
|
||||
p = _make_provider()
|
||||
p.client.messages.create = AsyncMock(return_value=_FakeAttr(
|
||||
content=[_FakeAttr(type="text", text="ok")],
|
||||
stop_reason="end_turn",
|
||||
usage=_FakeAttr(input_tokens=1, output_tokens=1),
|
||||
))
|
||||
await p.create_message(
|
||||
model="sonnet", system=None,
|
||||
messages=[ProviderMessage(role="user", content="x")],
|
||||
tools=[], max_tokens=12_345,
|
||||
)
|
||||
_, kwargs = p.client.messages.create.call_args
|
||||
assert kwargs["max_tokens"] == 12_345
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# stream_message: message_start / message_delta usage extraction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_stream_message_extracts_usage_from_message_start():
|
||||
"""The first SSE event the SDK emits is `message_start` carrying
|
||||
initial input_tokens. Adapter must surface that as a `usage`
|
||||
StreamEvent before the message_stop sentinel."""
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
# message_start carries initial usage (input + cache + output start)
|
||||
yield _FakeAttr(
|
||||
type="message_start",
|
||||
message=_FakeAttr(
|
||||
usage=_FakeAttr(input_tokens=100, output_tokens=0),
|
||||
),
|
||||
)
|
||||
# text block
|
||||
yield _FakeAttr(type="content_block_start", index=0,
|
||||
content_block=_FakeAttr(type="text"))
|
||||
yield _FakeAttr(type="content_block_delta", index=0,
|
||||
delta=_FakeAttr(type="text_delta", text="hi"))
|
||||
yield _FakeAttr(type="content_block_stop", index=0)
|
||||
# message_delta carries final output_tokens
|
||||
yield _FakeAttr(
|
||||
type="message_delta",
|
||||
usage=_FakeAttr(output_tokens=25),
|
||||
)
|
||||
|
||||
p.client.messages.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(
|
||||
model="sonnet", system=None, messages=[], tools=[],
|
||||
):
|
||||
events.append(ev)
|
||||
|
||||
usage_events = [e for e in events if e.type == "usage"]
|
||||
assert len(usage_events) == 2
|
||||
assert usage_events[0].usage == {"input_tokens": 100}
|
||||
assert usage_events[1].usage == {"output_tokens": 25}
|
||||
|
||||
# Always closes with message_stop
|
||||
assert events[-1].type == "message_stop"
|
||||
|
||||
|
||||
async def test_stream_message_skips_message_start_without_usage():
|
||||
"""If `message_start.message.usage` is missing or zero, no `usage`
|
||||
event must fire (the source guards on truthy input/output tokens)."""
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
yield _FakeAttr(
|
||||
type="message_start",
|
||||
message=_FakeAttr(usage=_FakeAttr(input_tokens=0, output_tokens=0)),
|
||||
)
|
||||
yield _FakeAttr(
|
||||
type="message_delta",
|
||||
usage=None,
|
||||
)
|
||||
|
||||
p.client.messages.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(
|
||||
model="sonnet", system=None, messages=[], tools=[],
|
||||
):
|
||||
events.append(ev)
|
||||
|
||||
usage_events = [e for e in events if e.type == "usage"]
|
||||
assert usage_events == []
|
||||
|
||||
|
||||
async def test_stream_message_input_json_delta_streamed():
|
||||
"""The tool_use streaming path: input_json_delta chunks must be
|
||||
surfaced as content_block_delta with delta_type=input_json_delta."""
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
yield _FakeAttr(
|
||||
type="content_block_start", index=0,
|
||||
content_block=_FakeAttr(type="tool_use", name="Read", id="toolu_1"),
|
||||
)
|
||||
yield _FakeAttr(
|
||||
type="content_block_delta", index=0,
|
||||
delta=_FakeAttr(type="input_json_delta", partial_json='{"pa'),
|
||||
)
|
||||
yield _FakeAttr(
|
||||
type="content_block_delta", index=0,
|
||||
delta=_FakeAttr(type="input_json_delta", partial_json='th": "/x"}'),
|
||||
)
|
||||
yield _FakeAttr(type="content_block_stop", index=0)
|
||||
|
||||
p.client.messages.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(
|
||||
model="sonnet", system=None, messages=[], tools=[],
|
||||
):
|
||||
events.append(ev)
|
||||
|
||||
starts = [e for e in events if e.type == "content_block_start"]
|
||||
deltas = [e for e in events if e.type == "content_block_delta"
|
||||
and e.delta_type == "input_json_delta"]
|
||||
assert len(starts) == 1
|
||||
assert starts[0].block_type == "tool_use"
|
||||
assert starts[0].tool_name == "Read"
|
||||
assert starts[0].tool_id == "toolu_1"
|
||||
assert len(deltas) == 2
|
||||
assert "".join(d.text for d in deltas) == '{"path": "/x"}'
|
||||
|
||||
|
||||
async def test_stream_message_passes_system_and_tools_to_sdk():
|
||||
"""Smoke-test that system + tool schemas reach the SDK call."""
|
||||
p = _make_provider()
|
||||
|
||||
async def empty_stream():
|
||||
if False:
|
||||
yield None # never yields — empty generator
|
||||
return
|
||||
|
||||
p.client.messages.create = AsyncMock(return_value=empty_stream())
|
||||
|
||||
async for _ in p.stream_message(
|
||||
model="sonnet",
|
||||
system="be helpful",
|
||||
messages=[ProviderMessage(role="user", content="hi")],
|
||||
tools=[ToolSchema(name="Read", description="d", input_schema={"type": "object"})],
|
||||
max_tokens=2048,
|
||||
):
|
||||
pass
|
||||
|
||||
_, kwargs = p.client.messages.create.call_args
|
||||
assert kwargs["model"] == MODEL_MAP["sonnet"]
|
||||
assert kwargs["system"] == "be helpful"
|
||||
assert kwargs["max_tokens"] == 2048
|
||||
assert kwargs["stream"] is True
|
||||
assert kwargs["tools"] == [{"name": "Read", "description": "d", "input_schema": {"type": "object"}}]
|
||||
assert kwargs["messages"] == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# stream_and_collect default raises
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_stream_and_collect_raises_not_implemented():
|
||||
"""The helper isn't used directly by AgentLoop — provider hides it
|
||||
behind a NotImplementedError to prevent accidental adoption."""
|
||||
p = _make_provider()
|
||||
with pytest.raises(NotImplementedError):
|
||||
await p.stream_and_collect(
|
||||
model="sonnet", system=None, messages=[], tools=[],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constructor kwarg handling
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_constructor_prefers_auth_token_over_api_key():
|
||||
"""When both are passed, auth_token wins (the elif branch in the
|
||||
constructor); api_key is silently dropped."""
|
||||
import anthropic
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
class _Stub:
|
||||
def __init__(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
|
||||
real = anthropic.AsyncAnthropic
|
||||
anthropic.AsyncAnthropic = _Stub
|
||||
try:
|
||||
AnthropicProvider(api_key="key", auth_token="tok", base_url="http://x")
|
||||
finally:
|
||||
anthropic.AsyncAnthropic = real
|
||||
|
||||
assert captured.get("auth_token") == "tok"
|
||||
assert "api_key" not in captured
|
||||
assert captured.get("base_url") == "http://x"
|
||||
|
||||
|
||||
def test_constructor_no_creds_passes_no_kwargs():
|
||||
import anthropic
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
class _Stub:
|
||||
def __init__(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
|
||||
real = anthropic.AsyncAnthropic
|
||||
anthropic.AsyncAnthropic = _Stub
|
||||
try:
|
||||
AnthropicProvider()
|
||||
finally:
|
||||
anthropic.AsyncAnthropic = real
|
||||
|
||||
assert captured == {}
|
||||
@@ -1,676 +0,0 @@
|
||||
"""Tests for `backend.apps.agents.providers.openai_compat`.
|
||||
|
||||
The whole module currently sits at 0% coverage because no other test
|
||||
exercises an OpenAI-compatible provider. We mock the `AsyncOpenAI`
|
||||
client so all paths run with no network access:
|
||||
|
||||
- `format_user_message`: string + multimodal (text + image) blocks
|
||||
- `format_assistant_message`: text-only, mixed text+tool_use,
|
||||
tool_use only (content=None branch)
|
||||
- `format_tool_result`: text + image + raw json fallback
|
||||
- `_build_messages`: system prefix, assistant in OpenAI format,
|
||||
assistant in Anthropic-block format, tool_result list / single
|
||||
dict, user passthrough
|
||||
- `create_message`: text completion + tool_calls, finish_reason
|
||||
handling, usage extraction
|
||||
- `stream_message`: text-delta chunks, tool-call streaming with
|
||||
json delta accumulation, usage-only final chunk, finish_reason
|
||||
closing all open blocks
|
||||
- `clean_tool_schema`: OpenAI function-calling shape
|
||||
- `get_model_id`: passthrough (no short-name mapping)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents.providers.base import (
|
||||
ContentBlock,
|
||||
ModelResponse,
|
||||
ProviderMessage,
|
||||
ToolCall,
|
||||
ToolSchema,
|
||||
)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeAttr:
|
||||
"""Generic dot-attribute object for SDK-shaped responses."""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
def _make_provider() -> OpenAICompatProvider:
|
||||
p = OpenAICompatProvider(api_key="test-key", base_url="http://example.invalid")
|
||||
p.client = MagicMock()
|
||||
p.client.chat = MagicMock()
|
||||
p.client.chat.completions = MagicMock()
|
||||
p.client.chat.completions.create = AsyncMock()
|
||||
return p
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constructor + simple helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_model_id_is_passthrough():
|
||||
"""OpenAI-compatible doesn't do short-name mapping; user supplies
|
||||
the exact API model id."""
|
||||
p = _make_provider()
|
||||
assert p.get_model_id("gpt-5.4") == "gpt-5.4"
|
||||
assert p.get_model_id("anything-else") == "anything-else"
|
||||
|
||||
|
||||
def test_clean_tool_schema_returns_openai_function_format():
|
||||
p = _make_provider()
|
||||
schema = ToolSchema(
|
||||
name="Read",
|
||||
description="Read a file",
|
||||
input_schema={"type": "object", "properties": {"path": {"type": "string"}}},
|
||||
)
|
||||
out = p.clean_tool_schema(schema)
|
||||
assert out == {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "Read",
|
||||
"description": "Read a file",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"path": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_constructor_defaults_api_key_to_none_placeholder():
|
||||
"""Some endpoints don't need real keys; the adapter sends "none"
|
||||
rather than failing. Capture the kwargs to verify."""
|
||||
from openai import AsyncOpenAI as _RealOpenAI
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
class _Stub:
|
||||
def __init__(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
|
||||
import backend.apps.agents.providers.openai_compat as oc_mod
|
||||
real = oc_mod.AsyncOpenAI
|
||||
oc_mod.AsyncOpenAI = _Stub
|
||||
try:
|
||||
OpenAICompatProvider(api_key="", base_url=None)
|
||||
finally:
|
||||
oc_mod.AsyncOpenAI = real
|
||||
assert captured.get("api_key") == "none"
|
||||
assert "base_url" not in captured
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_user_message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_format_user_message_string():
|
||||
p = _make_provider()
|
||||
msg = p.format_user_message("hello")
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "hello"
|
||||
|
||||
|
||||
def test_format_user_message_multimodal_text_and_image():
|
||||
"""Text + Anthropic-style image blocks → OpenAI image_url with
|
||||
base64 data URL."""
|
||||
p = _make_provider()
|
||||
blocks = [
|
||||
{"type": "text", "text": "look:"},
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": "image/png", "data": "AA=="},
|
||||
},
|
||||
"trailing string", # str gets coerced to a text part
|
||||
]
|
||||
msg = p.format_user_message(blocks)
|
||||
assert msg.role == "user"
|
||||
assert msg.content == [
|
||||
{"type": "text", "text": "look:"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,AA=="},
|
||||
},
|
||||
{"type": "text", "text": "trailing string"},
|
||||
]
|
||||
|
||||
|
||||
def test_format_user_message_other_types_str_coerced():
|
||||
p = _make_provider()
|
||||
msg = p.format_user_message(42)
|
||||
assert msg.role == "user"
|
||||
assert msg.content == "42"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_assistant_message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_format_assistant_message_text_only():
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(type="text", text="line one"),
|
||||
ContentBlock(type="text", text="line two"),
|
||||
],
|
||||
stop_reason="end_turn",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
# Stored under content -> dict (already in OpenAI shape) so _build_messages
|
||||
# can pass it through unchanged.
|
||||
assert msg.role == "assistant"
|
||||
assert msg.content == {"role": "assistant", "content": "line one\nline two"}
|
||||
|
||||
|
||||
def test_format_assistant_message_tool_use_only_sets_content_to_none():
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t1", name="Read", input={"path": "/x"}),
|
||||
),
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
assert msg.content["content"] is None
|
||||
assert msg.content["tool_calls"] == [{
|
||||
"id": "t1",
|
||||
"type": "function",
|
||||
"function": {"name": "Read", "arguments": json.dumps({"path": "/x"})},
|
||||
}]
|
||||
|
||||
|
||||
def test_format_assistant_message_mixed_text_and_tool_use():
|
||||
p = _make_provider()
|
||||
resp = ModelResponse(
|
||||
content=[
|
||||
ContentBlock(type="text", text="thinking"),
|
||||
ContentBlock(
|
||||
type="tool_use",
|
||||
tool_call=ToolCall(id="t1", name="Bash", input={"cmd": "ls"}),
|
||||
),
|
||||
],
|
||||
stop_reason="tool_use",
|
||||
)
|
||||
msg = p.format_assistant_message(resp)
|
||||
assert msg.content["content"] == "thinking"
|
||||
assert len(msg.content["tool_calls"]) == 1
|
||||
tc = msg.content["tool_calls"][0]
|
||||
assert tc["function"]["name"] == "Bash"
|
||||
assert json.loads(tc["function"]["arguments"]) == {"cmd": "ls"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# format_tool_result
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_format_tool_result_collapses_text_blocks_to_single_string():
|
||||
p = _make_provider()
|
||||
out = p.format_tool_result("call_1", [
|
||||
{"type": "text", "text": "line a"},
|
||||
{"type": "text", "text": "line b"},
|
||||
])
|
||||
assert out == {
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": "line a\nline b",
|
||||
}
|
||||
|
||||
|
||||
def test_format_tool_result_image_blocks_become_placeholder():
|
||||
p = _make_provider()
|
||||
out = p.format_tool_result("call_2", [
|
||||
{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "AA=="}},
|
||||
])
|
||||
assert out["content"] == "[image]"
|
||||
|
||||
|
||||
def test_format_tool_result_unknown_block_falls_back_to_json():
|
||||
p = _make_provider()
|
||||
block = {"type": "custom", "x": 1}
|
||||
out = p.format_tool_result("call_3", [block])
|
||||
assert out["content"] == json.dumps(block)
|
||||
|
||||
|
||||
def test_format_tool_result_empty_returns_done_marker():
|
||||
"""Empty content list → "Done." so OpenAI doesn't reject the
|
||||
message for empty content."""
|
||||
p = _make_provider()
|
||||
out = p.format_tool_result("call_4", [])
|
||||
assert out["content"] == "Done."
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _build_messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_build_messages_includes_system_prefix():
|
||||
p = _make_provider()
|
||||
out = p._build_messages("you are helpful", [
|
||||
ProviderMessage(role="user", content="hi"),
|
||||
])
|
||||
assert out[0] == {"role": "system", "content": "you are helpful"}
|
||||
assert out[1] == {"role": "user", "content": "hi"}
|
||||
|
||||
|
||||
def test_build_messages_no_system_no_prefix():
|
||||
p = _make_provider()
|
||||
out = p._build_messages(None, [ProviderMessage(role="user", content="hi")])
|
||||
assert out == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
def test_build_messages_assistant_in_openai_shape_passes_through():
|
||||
"""`format_assistant_message` already produces OpenAI-shape dicts;
|
||||
`_build_messages` must pass them through unchanged."""
|
||||
p = _make_provider()
|
||||
asst_dict = {"role": "assistant", "content": "hello"}
|
||||
out = p._build_messages(None, [ProviderMessage(role="assistant", content=asst_dict)])
|
||||
assert out == [asst_dict]
|
||||
|
||||
|
||||
def test_build_messages_assistant_in_anthropic_block_format():
|
||||
"""Coming from a cross-provider session, the assistant content
|
||||
might still be in Anthropic block format. _build_messages must
|
||||
translate it."""
|
||||
p = _make_provider()
|
||||
blocks = [
|
||||
{"type": "text", "text": "thinking"},
|
||||
{"type": "tool_use", "id": "t1", "name": "Read", "input": {"path": "/x"}},
|
||||
]
|
||||
out = p._build_messages(None, [ProviderMessage(role="assistant", content=blocks)])
|
||||
assert out[0]["role"] == "assistant"
|
||||
assert out[0]["content"] == "thinking"
|
||||
assert out[0]["tool_calls"] == [{
|
||||
"id": "t1",
|
||||
"type": "function",
|
||||
"function": {"name": "Read", "arguments": json.dumps({"path": "/x"})},
|
||||
}]
|
||||
|
||||
|
||||
def test_build_messages_tool_result_list_each_appended():
|
||||
p = _make_provider()
|
||||
tool_results = [
|
||||
{"role": "tool", "tool_call_id": "t1", "content": "ok"},
|
||||
{"role": "tool", "tool_call_id": "t2", "content": "ok2"},
|
||||
]
|
||||
out = p._build_messages(None, [ProviderMessage(role="tool_result", content=tool_results)])
|
||||
assert out == tool_results
|
||||
|
||||
|
||||
def test_build_messages_tool_result_single_dict_appended():
|
||||
p = _make_provider()
|
||||
tr = {"role": "tool", "tool_call_id": "t1", "content": "ok"}
|
||||
out = p._build_messages(None, [ProviderMessage(role="tool_result", content=tr)])
|
||||
assert out == [tr]
|
||||
|
||||
|
||||
def test_build_messages_tool_result_without_tool_call_id_dropped():
|
||||
"""Defensive: a malformed tool_result without `tool_call_id` is
|
||||
silently dropped to avoid crashing the API call."""
|
||||
p = _make_provider()
|
||||
out = p._build_messages(None, [
|
||||
ProviderMessage(role="tool_result", content={"role": "tool", "content": "x"}),
|
||||
])
|
||||
assert out == []
|
||||
|
||||
|
||||
def test_build_messages_user_string_passthrough():
|
||||
p = _make_provider()
|
||||
out = p._build_messages(None, [ProviderMessage(role="user", content="hi")])
|
||||
assert out == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_message (non-streaming)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_create_message_text_only_response():
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
choices=[
|
||||
_FakeAttr(
|
||||
message=_FakeAttr(content="hello", tool_calls=None),
|
||||
finish_reason="stop",
|
||||
),
|
||||
],
|
||||
usage=_FakeAttr(prompt_tokens=12, completion_tokens=3),
|
||||
)
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="gpt-5.4",
|
||||
system="be useful",
|
||||
messages=[ProviderMessage(role="user", content="hi")],
|
||||
tools=[],
|
||||
)
|
||||
|
||||
_, kwargs = p.client.chat.completions.create.call_args
|
||||
assert kwargs["model"] == "gpt-5.4" # passthrough
|
||||
assert kwargs["max_tokens"] == 8192
|
||||
assert kwargs["messages"][0] == {"role": "system", "content": "be useful"}
|
||||
assert kwargs["messages"][1] == {"role": "user", "content": "hi"}
|
||||
assert "tools" not in kwargs
|
||||
|
||||
assert out.stop_reason == "end_turn"
|
||||
assert len(out.content) == 1
|
||||
assert out.content[0].type == "text"
|
||||
assert out.content[0].text == "hello"
|
||||
assert out.usage == {"input_tokens": 12, "output_tokens": 3}
|
||||
|
||||
|
||||
async def test_create_message_tool_calls_translation():
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
choices=[
|
||||
_FakeAttr(
|
||||
message=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[
|
||||
_FakeAttr(
|
||||
id="call_1",
|
||||
function=_FakeAttr(
|
||||
name="Read",
|
||||
arguments=json.dumps({"path": "/tmp/a"}),
|
||||
),
|
||||
),
|
||||
],
|
||||
),
|
||||
finish_reason="tool_calls",
|
||||
),
|
||||
],
|
||||
usage=_FakeAttr(prompt_tokens=5, completion_tokens=2),
|
||||
)
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="gpt-5.4",
|
||||
system=None,
|
||||
messages=[ProviderMessage(role="user", content="x")],
|
||||
tools=[ToolSchema(name="Read", description="d", input_schema={"type": "object"})],
|
||||
)
|
||||
|
||||
_, kwargs = p.client.chat.completions.create.call_args
|
||||
assert kwargs["tools"] == [{
|
||||
"type": "function",
|
||||
"function": {"name": "Read", "description": "d", "parameters": {"type": "object"}},
|
||||
}]
|
||||
|
||||
assert out.stop_reason == "tool_use"
|
||||
assert len(out.content) == 1
|
||||
assert out.content[0].type == "tool_use"
|
||||
assert out.content[0].tool_call.id == "call_1"
|
||||
assert out.content[0].tool_call.name == "Read"
|
||||
assert out.content[0].tool_call.input == {"path": "/tmp/a"}
|
||||
|
||||
|
||||
async def test_create_message_invalid_tool_args_json_falls_back_to_empty():
|
||||
"""Malformed JSON in `function.arguments` must NOT crash; the adapter
|
||||
swallows the JSONDecodeError and leaves input={}."""
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
choices=[
|
||||
_FakeAttr(
|
||||
message=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[
|
||||
_FakeAttr(
|
||||
id="call_1",
|
||||
function=_FakeAttr(name="Read", arguments="{not valid json"),
|
||||
),
|
||||
],
|
||||
),
|
||||
finish_reason="tool_calls",
|
||||
),
|
||||
],
|
||||
usage=None,
|
||||
)
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="gpt-5.4", system=None,
|
||||
messages=[ProviderMessage(role="user", content="x")], tools=[],
|
||||
)
|
||||
assert out.content[0].tool_call.input == {}
|
||||
assert out.usage == {}
|
||||
|
||||
|
||||
async def test_create_message_text_plus_tool_use_yields_tool_use_stop():
|
||||
"""Mixed content with tool_calls → stop_reason becomes tool_use even
|
||||
if finish_reason was 'stop' (defensive against models that report
|
||||
'stop' alongside tool_calls)."""
|
||||
p = _make_provider()
|
||||
fake_resp = _FakeAttr(
|
||||
choices=[
|
||||
_FakeAttr(
|
||||
message=_FakeAttr(
|
||||
content="thinking",
|
||||
tool_calls=[
|
||||
_FakeAttr(
|
||||
id="c1",
|
||||
function=_FakeAttr(name="Read", arguments="{}"),
|
||||
),
|
||||
],
|
||||
),
|
||||
finish_reason="stop", # not "tool_calls"
|
||||
),
|
||||
],
|
||||
usage=_FakeAttr(prompt_tokens=1, completion_tokens=1),
|
||||
)
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_resp)
|
||||
|
||||
out = await p.create_message(
|
||||
model="gpt-5.4", system=None,
|
||||
messages=[ProviderMessage(role="user", content="x")], tools=[],
|
||||
)
|
||||
assert out.stop_reason == "tool_use"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# stream_message
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_stream_message_text_only_chunks():
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(content="hel", tool_calls=None),
|
||||
finish_reason=None,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(content="lo", tool_calls=None),
|
||||
finish_reason=None,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(content=None, tool_calls=None),
|
||||
finish_reason="stop",
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
# Final usage-only chunk
|
||||
yield _FakeAttr(
|
||||
choices=[],
|
||||
usage=_FakeAttr(prompt_tokens=10, completion_tokens=2),
|
||||
)
|
||||
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(model="gpt-5.4", system=None, messages=[], tools=[]):
|
||||
events.append(ev)
|
||||
|
||||
starts = [e for e in events if e.type == "content_block_start"]
|
||||
deltas = [e for e in events if e.type == "content_block_delta"]
|
||||
stops = [e for e in events if e.type == "content_block_stop"]
|
||||
|
||||
assert len(starts) == 1
|
||||
assert starts[0].block_type == "text"
|
||||
assert "".join(d.text for d in deltas) == "hello"
|
||||
assert len(stops) == 1
|
||||
assert any(e.type == "message_stop" for e in events)
|
||||
usage = [e for e in events if e.type == "usage"]
|
||||
assert usage and usage[0].usage == {"input_tokens": 10, "output_tokens": 2}
|
||||
|
||||
|
||||
async def test_stream_message_tool_call_streamed():
|
||||
"""Tool-call streaming: name comes in chunk 1, arguments stream in
|
||||
multiple JSON deltas. Adapter accumulates and emits normalized
|
||||
StreamEvents."""
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
# chunk 1: tool_call begin (id + name)
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[_FakeAttr(
|
||||
index=0,
|
||||
id="call_1",
|
||||
function=_FakeAttr(name="Read", arguments=""),
|
||||
)],
|
||||
),
|
||||
finish_reason=None,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
# chunk 2: arguments part 1
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[_FakeAttr(
|
||||
index=0,
|
||||
id=None,
|
||||
function=_FakeAttr(name=None, arguments='{"pa'),
|
||||
)],
|
||||
),
|
||||
finish_reason=None,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
# chunk 3: arguments part 2 + finish
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[_FakeAttr(
|
||||
index=0,
|
||||
id=None,
|
||||
function=_FakeAttr(name=None, arguments='th": "/x"}'),
|
||||
)],
|
||||
),
|
||||
finish_reason="tool_calls",
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(model="gpt-5.4", system=None, messages=[], tools=[]):
|
||||
events.append(ev)
|
||||
|
||||
starts = [e for e in events if e.type == "content_block_start"]
|
||||
deltas = [e for e in events if e.type == "content_block_delta"]
|
||||
stops = [e for e in events if e.type == "content_block_stop"]
|
||||
|
||||
assert len(starts) == 1
|
||||
assert starts[0].block_type == "tool_use"
|
||||
assert starts[0].tool_id == "call_1"
|
||||
assert starts[0].tool_name == "Read"
|
||||
json_deltas = [d for d in deltas if d.delta_type == "input_json_delta"]
|
||||
assert "".join(d.text for d in json_deltas) == '{"path": "/x"}'
|
||||
assert len(stops) == 1
|
||||
assert any(e.type == "message_stop" for e in events)
|
||||
|
||||
|
||||
async def test_stream_message_text_then_tool_use_closes_text_block():
|
||||
"""When text was already streaming and a tool_call begins, the
|
||||
text block must be closed first so the frontend's UI logic sees a
|
||||
clean handoff."""
|
||||
p = _make_provider()
|
||||
|
||||
async def fake_stream():
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(content="thinking ", tool_calls=None),
|
||||
finish_reason=None,
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
yield _FakeAttr(
|
||||
choices=[_FakeAttr(
|
||||
delta=_FakeAttr(
|
||||
content=None,
|
||||
tool_calls=[_FakeAttr(
|
||||
index=0, id="call_1",
|
||||
function=_FakeAttr(name="Read", arguments="{}"),
|
||||
)],
|
||||
),
|
||||
finish_reason="tool_calls",
|
||||
)],
|
||||
usage=None,
|
||||
)
|
||||
|
||||
p.client.chat.completions.create = AsyncMock(return_value=fake_stream())
|
||||
|
||||
events = []
|
||||
async for ev in p.stream_message(model="gpt-5.4", system=None, messages=[], tools=[]):
|
||||
events.append(ev)
|
||||
|
||||
block_types = [e.block_type for e in events if e.type == "content_block_start"]
|
||||
stops = [e for e in events if e.type == "content_block_stop"]
|
||||
assert block_types == ["text", "tool_use"]
|
||||
# Text close fires when tool_call starts; tool_use close fires at finish.
|
||||
assert len(stops) == 2
|
||||
|
||||
|
||||
async def test_stream_message_includes_usage_options_kwargs():
|
||||
"""`stream_options.include_usage` MUST be passed so the final
|
||||
chunk carries token counts."""
|
||||
p = _make_provider()
|
||||
|
||||
async def empty_stream():
|
||||
if False:
|
||||
yield None
|
||||
return
|
||||
|
||||
p.client.chat.completions.create = AsyncMock(return_value=empty_stream())
|
||||
|
||||
async for _ in p.stream_message(model="gpt-5.4", system=None, messages=[], tools=[]):
|
||||
pass
|
||||
|
||||
_, kwargs = p.client.chat.completions.create.call_args
|
||||
assert kwargs["stream"] is True
|
||||
assert kwargs["stream_options"] == {"include_usage": True}
|
||||
@@ -1,616 +0,0 @@
|
||||
"""Tests for `backend.apps.agents.providers.registry`.
|
||||
|
||||
The registry is currently 16% covered. This file fills in:
|
||||
|
||||
- `_find_builtin_model` known + unknown
|
||||
- `get_api_type` over every value in BUILTIN_MODELS + unknown default
|
||||
- `resolve_model_id_for_sdk` across every routing branch:
|
||||
- openswarm-pro mode → bare model_id
|
||||
- direct anthropic_api_key → bare model_id
|
||||
- explicit `route="cc"` → router_model_id (subscription)
|
||||
- explicit `route="api"` → bare model_id
|
||||
- gemini-cli + google_api_key → `gemini/<suffix>`
|
||||
- gemini-cli + Antigravity active (mock httpx 200) → `ag/<mapped>`
|
||||
- gemini-cli fallthrough → `gc/<suffix>`
|
||||
- openai/codex/gemini fallthrough → router_model_id
|
||||
- unknown short_name → passthrough
|
||||
- `resolve_aux_model` every priority branch + ValueError fallthrough
|
||||
- `create_provider` every api_type + 9Router fallback + missing-key raise
|
||||
- `thinking_params_for(api, level)` full matrix
|
||||
- `get_available_models` `configured` flag correctness
|
||||
- `get_context_window` known/custom/default
|
||||
- `calculate_cost` known + unknown rates + case-insensitive provider
|
||||
|
||||
Live network calls (httpx, 9Router) are fully mocked.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents.providers import registry as reg
|
||||
from backend.apps.agents.providers.registry import (
|
||||
BUILTIN_MODELS,
|
||||
_find_builtin_model,
|
||||
_get_api_type,
|
||||
_has_credentials,
|
||||
calculate_cost,
|
||||
create_provider,
|
||||
get_api_type,
|
||||
get_available_models,
|
||||
get_context_window,
|
||||
resolve_aux_model,
|
||||
resolve_model_id_for_sdk,
|
||||
thinking_params_for,
|
||||
)
|
||||
from backend.apps.settings.models import AppSettings, CustomProvider
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _find_builtin_model + get_api_type
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_find_builtin_model_known():
|
||||
entry = _find_builtin_model("sonnet")
|
||||
assert entry is not None
|
||||
assert entry["api"] == "anthropic"
|
||||
assert entry["value"] == "sonnet"
|
||||
|
||||
|
||||
def test_find_builtin_model_unknown_returns_none():
|
||||
assert _find_builtin_model("not-a-real-model") is None
|
||||
|
||||
|
||||
def test_get_api_type_unknown_defaults_to_anthropic():
|
||||
assert get_api_type("not-a-real-model") == "anthropic"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"short,expected_api",
|
||||
[
|
||||
("sonnet", "anthropic"),
|
||||
("opus", "anthropic"),
|
||||
("haiku", "anthropic"),
|
||||
("sonnet-cc", "anthropic"),
|
||||
("sonnet-api", "anthropic"),
|
||||
("gpt-5.4", "codex"),
|
||||
("gpt-5.4-mini", "codex"),
|
||||
("gpt-5.4-api", "openai"),
|
||||
("gpt-5.3-codex-api", "openai"),
|
||||
("gemini-3-pro", "gemini-cli"),
|
||||
("gemini-2.5-flash", "gemini-cli"),
|
||||
("gemini-3-pro-api", "gemini"),
|
||||
("gemini-2.5-flash-api", "gemini"),
|
||||
],
|
||||
)
|
||||
def test_get_api_type_known_models(short: str, expected_api: str):
|
||||
assert get_api_type(short) == expected_api
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _get_api_type — provider-name dispatch
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name,expected",
|
||||
[
|
||||
("Anthropic", "anthropic"),
|
||||
("OpenAI", "codex"), # the OpenAI tier's first entry uses api=codex
|
||||
("Google", "gemini-cli"),
|
||||
("anthropic", "anthropic"), # case-insensitive
|
||||
("OPENAI", "openai"),
|
||||
("google", "gemini"), # lowercase 'google' -> gemini via _API_NAME_MAP
|
||||
("openrouter", "openrouter"),
|
||||
("UnknownProvider", "openrouter"), # fallthrough default
|
||||
],
|
||||
)
|
||||
def test_underscore_get_api_type_dispatch(name: str, expected: str):
|
||||
assert _get_api_type(name) == expected
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_model_id_for_sdk
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_unknown_passthrough():
|
||||
s = AppSettings()
|
||||
assert resolve_model_id_for_sdk("not-a-real-model", s) == "not-a-real-model"
|
||||
|
||||
|
||||
def test_resolve_route_cc_uses_router_id():
|
||||
s = AppSettings()
|
||||
assert resolve_model_id_for_sdk("sonnet-cc", s) == "cc/claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_resolve_route_api_uses_bare_model_id():
|
||||
s = AppSettings()
|
||||
assert resolve_model_id_for_sdk("sonnet-api", s) == "claude-sonnet-4-6"
|
||||
assert resolve_model_id_for_sdk("gpt-5.4-api", s) == "gpt-5.4"
|
||||
|
||||
|
||||
def test_resolve_anthropic_with_openswarm_pro_returns_bare():
|
||||
s = AppSettings(connection_mode="openswarm-pro")
|
||||
assert resolve_model_id_for_sdk("sonnet", s) == "claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_resolve_anthropic_with_api_key_returns_bare():
|
||||
s = AppSettings(anthropic_api_key="sk-test")
|
||||
assert resolve_model_id_for_sdk("sonnet", s) == "claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_resolve_anthropic_no_creds_returns_router_id():
|
||||
"""Without anthropic_api_key + own_key mode → 9Router cc/ prefix."""
|
||||
s = AppSettings()
|
||||
assert resolve_model_id_for_sdk("sonnet", s) == "cc/claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_resolve_gemini_cli_with_google_api_key_uses_gemini_prefix():
|
||||
"""google_api_key set → AI Studio direct path."""
|
||||
s = AppSettings(google_api_key="AIza-test")
|
||||
assert resolve_model_id_for_sdk("gemini-3-pro", s) == "gemini/gemini-3-pro-preview"
|
||||
assert resolve_model_id_for_sdk("gemini-2.5-pro", s) == "gemini/gemini-2.5-pro"
|
||||
|
||||
|
||||
def test_resolve_gemini_cli_antigravity_active_returns_ag_prefix():
|
||||
"""Without google_api_key but with Antigravity connected on 9Router,
|
||||
map to ag/<mapped_suffix>."""
|
||||
s = AppSettings()
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.status_code = 200
|
||||
fake_resp.json.return_value = {
|
||||
"connections": [
|
||||
{"provider": "antigravity", "isActive": True},
|
||||
],
|
||||
}
|
||||
with patch("httpx.get", return_value=fake_resp):
|
||||
assert resolve_model_id_for_sdk("gemini-3-pro", s) == "ag/gemini-3.1-pro-high"
|
||||
assert resolve_model_id_for_sdk("gemini-3-flash", s) == "ag/gemini-3-flash"
|
||||
|
||||
|
||||
def test_resolve_gemini_cli_antigravity_active_but_unmapped_falls_through():
|
||||
"""Even with Antigravity active, models not in the _ANTIGRAVITY_MAP
|
||||
(gemini-2.5-*) must fall through to gc/."""
|
||||
s = AppSettings()
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.status_code = 200
|
||||
fake_resp.json.return_value = {
|
||||
"connections": [{"provider": "antigravity", "isActive": True}],
|
||||
}
|
||||
with patch("httpx.get", return_value=fake_resp):
|
||||
assert resolve_model_id_for_sdk("gemini-2.5-pro", s) == "gc/gemini-2.5-pro"
|
||||
|
||||
|
||||
def test_resolve_gemini_cli_no_creds_returns_gc():
|
||||
"""Default fallthrough — no API key, no Antigravity."""
|
||||
s = AppSettings()
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.status_code = 200
|
||||
fake_resp.json.return_value = {"connections": []}
|
||||
with patch("httpx.get", return_value=fake_resp):
|
||||
assert resolve_model_id_for_sdk("gemini-3-pro", s) == "gc/gemini-3-pro-preview"
|
||||
|
||||
|
||||
def test_resolve_gemini_cli_httpx_exception_falls_through_to_gc():
|
||||
"""If 9Router probe raises, fail open → gc/ prefix."""
|
||||
s = AppSettings()
|
||||
with patch("httpx.get", side_effect=Exception("boom")):
|
||||
assert resolve_model_id_for_sdk("gemini-3-pro", s) == "gc/gemini-3-pro-preview"
|
||||
|
||||
|
||||
def test_resolve_codex_returns_router_id():
|
||||
s = AppSettings()
|
||||
assert resolve_model_id_for_sdk("gpt-5.4", s) == "cx/gpt-5.4"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_aux_model
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_resolve_aux_model_openswarm_pro_returns_proxy_url():
|
||||
s = AppSettings(connection_mode="openswarm-pro", openswarm_proxy_url="https://proxy.test")
|
||||
model, base_url = await resolve_aux_model(s)
|
||||
assert "haiku" in model
|
||||
assert base_url == "https://proxy.test"
|
||||
|
||||
|
||||
async def test_resolve_aux_model_openswarm_pro_default_url():
|
||||
"""If openswarm_proxy_url isn't set, defaults to api.openswarm.com."""
|
||||
s = AppSettings(connection_mode="openswarm-pro")
|
||||
_model, base_url = await resolve_aux_model(s)
|
||||
assert base_url == "https://api.openswarm.com"
|
||||
|
||||
|
||||
async def test_resolve_aux_model_anthropic_api_key_returns_no_base_url():
|
||||
s = AppSettings(anthropic_api_key="sk-test")
|
||||
model, base_url = await resolve_aux_model(s)
|
||||
assert "haiku" in model
|
||||
assert base_url is None
|
||||
|
||||
|
||||
async def test_resolve_aux_model_sonnet_tier():
|
||||
s = AppSettings(anthropic_api_key="sk-test")
|
||||
model, _ = await resolve_aux_model(s, preferred_tier="sonnet")
|
||||
assert "sonnet" in model
|
||||
|
||||
|
||||
async def test_resolve_aux_model_9router_claude_connection():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_9r_running", create=True), \
|
||||
patch("backend.apps.nine_router.is_running", return_value=True), \
|
||||
patch("backend.apps.nine_router.get_providers", new_callable=AsyncMock,
|
||||
return_value=[{"provider": "claude", "isActive": True}]):
|
||||
model, base_url = await resolve_aux_model(s)
|
||||
assert model.startswith("cc/")
|
||||
assert base_url == "http://localhost:20128"
|
||||
|
||||
|
||||
async def test_resolve_aux_model_9router_codex_connection():
|
||||
s = AppSettings()
|
||||
with patch("backend.apps.nine_router.is_running", return_value=True), \
|
||||
patch("backend.apps.nine_router.get_providers", new_callable=AsyncMock,
|
||||
return_value=[{"provider": "codex", "isActive": True}]):
|
||||
model, base_url = await resolve_aux_model(s)
|
||||
assert model == "cx/gpt-5.4-mini"
|
||||
assert base_url == "http://localhost:20128"
|
||||
|
||||
|
||||
async def test_resolve_aux_model_9router_gemini_connection():
|
||||
s = AppSettings()
|
||||
with patch("backend.apps.nine_router.is_running", return_value=True), \
|
||||
patch("backend.apps.nine_router.get_providers", new_callable=AsyncMock,
|
||||
return_value=[{"provider": "gemini-cli", "isActive": True}]):
|
||||
model, base_url = await resolve_aux_model(s)
|
||||
assert model == "gc/gemini-2.5-flash"
|
||||
assert base_url == "http://localhost:20128"
|
||||
|
||||
|
||||
async def test_resolve_aux_model_9router_no_connections_raises():
|
||||
s = AppSettings()
|
||||
with patch("backend.apps.nine_router.is_running", return_value=True), \
|
||||
patch("backend.apps.nine_router.get_providers", new_callable=AsyncMock,
|
||||
return_value=[]):
|
||||
with pytest.raises(ValueError, match="No AI provider connected"):
|
||||
await resolve_aux_model(s)
|
||||
|
||||
|
||||
async def test_resolve_aux_model_9router_not_running_raises():
|
||||
s = AppSettings()
|
||||
with patch("backend.apps.nine_router.is_running", return_value=False):
|
||||
with pytest.raises(ValueError, match="No AI provider configured"):
|
||||
await resolve_aux_model(s)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_provider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_create_provider_anthropic_with_api_key():
|
||||
s = AppSettings(anthropic_api_key="sk-test")
|
||||
p = create_provider("Anthropic", s)
|
||||
from backend.apps.agents.providers.anthropic import AnthropicProvider
|
||||
assert isinstance(p, AnthropicProvider)
|
||||
|
||||
|
||||
def test_create_provider_anthropic_with_openswarm_pro():
|
||||
s = AppSettings(
|
||||
connection_mode="openswarm-pro",
|
||||
openswarm_bearer_token="bearer-x",
|
||||
openswarm_proxy_url="https://proxy.test",
|
||||
)
|
||||
p = create_provider("Anthropic", s)
|
||||
from backend.apps.agents.providers.anthropic import AnthropicProvider
|
||||
assert isinstance(p, AnthropicProvider)
|
||||
|
||||
|
||||
def test_create_provider_anthropic_falls_back_to_9router():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=True):
|
||||
p = create_provider("Anthropic", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
# The override remaps short names → 9Router prefix
|
||||
assert p.get_model_id("sonnet") == "cc/claude-sonnet-4-6"
|
||||
assert p.get_model_id("custom-id") == "cc/custom-id"
|
||||
assert p.get_model_id("cc/already") == "cc/already"
|
||||
|
||||
|
||||
def test_create_provider_anthropic_no_creds_raises():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=False):
|
||||
with pytest.raises(ValueError, match="Anthropic API key not configured"):
|
||||
create_provider("Anthropic", s)
|
||||
|
||||
|
||||
def test_create_provider_openai_with_key():
|
||||
s = AppSettings(openai_api_key="sk-openai-test")
|
||||
p = create_provider("OPENAI", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_openai_no_key_falls_back_to_9router():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=True):
|
||||
p = create_provider("OPENAI", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_openai_no_creds_raises():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=False):
|
||||
with pytest.raises(ValueError, match="OpenAI API key not configured"):
|
||||
create_provider("OPENAI", s)
|
||||
|
||||
|
||||
def test_create_provider_gemini_branch_imports_gemini_module():
|
||||
"""The gemini branch imports `backend.apps.agents.providers.gemini`,
|
||||
which doesn't currently ship in this repo. The branch is therefore
|
||||
only reachable once that module exists; verify the failure mode is
|
||||
`ModuleNotFoundError` (not silent), so we'll notice if the module
|
||||
is added without updating tests."""
|
||||
s = AppSettings(google_api_key="AIza-test")
|
||||
with pytest.raises(ModuleNotFoundError):
|
||||
create_provider("google", s)
|
||||
|
||||
|
||||
def test_create_provider_openrouter_with_key():
|
||||
s = AppSettings(openrouter_api_key="sk-or-test")
|
||||
p = create_provider("openrouter", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_openrouter_no_key_falls_back_to_9router():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=True):
|
||||
p = create_provider("openrouter", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_openrouter_no_creds_raises():
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_is_9router_available", return_value=False):
|
||||
with pytest.raises(ValueError, match="OpenRouter API key not configured"):
|
||||
create_provider("openrouter", s)
|
||||
|
||||
|
||||
def test_create_provider_9router_short_circuit():
|
||||
"""provider_name='9Router' takes the explicit early-return path —
|
||||
no settings needed."""
|
||||
s = AppSettings()
|
||||
p = create_provider("9Router", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_custom_provider_via_provider_config():
|
||||
"""Inline custom provider definition via the kwarg shortcut.
|
||||
|
||||
Unknown provider names default to api_type='openrouter' — to reach
|
||||
the `provider_config` branch we patch the api-type lookup to a
|
||||
value not in the known-api if-chain."""
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_get_api_type", return_value="custom-other"):
|
||||
p = create_provider(
|
||||
"MyCustom", s,
|
||||
provider_config={"api_key": "k", "base_url": "http://example.invalid/v1"},
|
||||
)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_custom_provider_lookup_in_settings():
|
||||
"""Same shape as above, but the custom provider lives on settings
|
||||
and is resolved by name."""
|
||||
s = AppSettings(custom_providers=[
|
||||
CustomProvider(name="MyCustom", base_url="http://example.invalid/v1", api_key="k"),
|
||||
])
|
||||
with patch.object(reg, "_get_api_type", return_value="custom-other"):
|
||||
p = create_provider("MyCustom", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
def test_create_provider_unknown_custom_name_raises():
|
||||
"""If the provider isn't found among any branch, raise ValueError."""
|
||||
s = AppSettings()
|
||||
with patch.object(reg, "_get_api_type", return_value="custom-other"):
|
||||
with pytest.raises(ValueError, match="Unknown provider"):
|
||||
create_provider("NotInSettings", s)
|
||||
|
||||
|
||||
def test_create_provider_unknown_provider_raises():
|
||||
"""No matching api_type, no provider_config, no custom provider → ValueError.
|
||||
But unknown providers default to api_type='openrouter' so they hit
|
||||
the openrouter branch first; we set up to reach the final unknown."""
|
||||
s = AppSettings(openrouter_api_key="sk-test")
|
||||
# With openrouter key, "Unknown" provider returns an OpenAI-compat
|
||||
# adapter via the openrouter branch — not the unknown raise.
|
||||
p = create_provider("Unknown", s)
|
||||
from backend.apps.agents.providers.openai_compat import OpenAICompatProvider
|
||||
assert isinstance(p, OpenAICompatProvider)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# thinking_params_for
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api,level,expected",
|
||||
[
|
||||
# auto → adaptive on Claude, defaults elsewhere
|
||||
("anthropic", "auto", {"thinking": {"type": "adaptive"}}),
|
||||
("codex", "auto", None),
|
||||
("gemini-cli", "auto", None),
|
||||
# off → explicit disable per provider
|
||||
("anthropic", "off", {"thinking": {"type": "disabled"}}),
|
||||
("codex", "off", {"reasoning": {"effort": "none"}}),
|
||||
("gemini-cli", "off", {"thinkingConfig": {"thinkingLevel": "LOW"}}),
|
||||
# Explicit levels
|
||||
("anthropic", "low", {"thinking": {"type": "adaptive"}}),
|
||||
("anthropic", "medium", {"thinking": {"type": "adaptive"}}),
|
||||
("anthropic", "high", {"thinking": {"type": "adaptive"}}),
|
||||
("codex", "low", {"reasoning": {"effort": "low"}}),
|
||||
("codex", "medium", {"reasoning": {"effort": "medium"}}),
|
||||
("codex", "high", {"reasoning": {"effort": "high"}}),
|
||||
("gemini-cli", "low", {"thinkingConfig": {"thinkingLevel": "LOW"}}),
|
||||
("gemini-cli", "medium", {"thinkingConfig": {"thinkingLevel": "MEDIUM"}}),
|
||||
("gemini-cli", "high", {"thinkingConfig": {"thinkingLevel": "HIGH"}}),
|
||||
# Unknown api → None
|
||||
("openai", "high", None),
|
||||
],
|
||||
)
|
||||
def test_thinking_params_for(api: str, level: str, expected):
|
||||
assert thinking_params_for(api, level) == expected
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_available_models / configured flag / _has_credentials
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_available_models_configured_flag_anthropic():
|
||||
"""anthropic_api_key set → Anthropic models marked configured."""
|
||||
s = AppSettings(anthropic_api_key="sk-test")
|
||||
out = get_available_models(s)
|
||||
assert all(m["configured"] for m in out["Anthropic"])
|
||||
# Other providers without keys → not configured
|
||||
assert all(not m["configured"] for m in out["OpenAI"])
|
||||
|
||||
|
||||
def test_get_available_models_configured_flag_openswarm_pro():
|
||||
"""In openswarm-pro mode, having a bearer token configures Anthropic."""
|
||||
s = AppSettings(connection_mode="openswarm-pro", openswarm_bearer_token="bearer-x")
|
||||
out = get_available_models(s)
|
||||
assert all(m["configured"] for m in out["Anthropic"])
|
||||
|
||||
|
||||
def test_get_available_models_includes_custom_providers():
|
||||
s = AppSettings(custom_providers=[
|
||||
CustomProvider(
|
||||
name="Local",
|
||||
base_url="http://localhost:8080/v1",
|
||||
api_key="x",
|
||||
models=[{"value": "phi-mini", "label": "Phi Mini", "context_window": 32_000}],
|
||||
),
|
||||
])
|
||||
out = get_available_models(s)
|
||||
assert "Local" in out
|
||||
assert out["Local"][0]["value"] == "phi-mini"
|
||||
assert out["Local"][0]["context_window"] == 32_000
|
||||
assert out["Local"][0]["configured"] is True
|
||||
|
||||
|
||||
def test_has_credentials_unknown_provider_returns_false():
|
||||
"""Unknown providers default to openrouter api_type, which checks
|
||||
openrouter_api_key — without it, returns False."""
|
||||
s = AppSettings()
|
||||
assert _has_credentials("Unknown", s) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_context_window
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_get_context_window_known_anthropic():
|
||||
assert get_context_window("Anthropic", "sonnet") == 1_000_000
|
||||
|
||||
|
||||
def test_get_context_window_known_haiku():
|
||||
assert get_context_window("Anthropic", "haiku") == 200_000
|
||||
|
||||
|
||||
def test_get_context_window_unknown_returns_default():
|
||||
assert get_context_window("Anthropic", "not-real") == 128_000
|
||||
|
||||
|
||||
def test_get_context_window_custom_provider_lookup():
|
||||
s = AppSettings(custom_providers=[
|
||||
CustomProvider(
|
||||
name="L",
|
||||
base_url="x",
|
||||
models=[{"value": "m", "context_window": 64_000}],
|
||||
),
|
||||
])
|
||||
assert get_context_window("L", "m", s) == 64_000
|
||||
|
||||
|
||||
def test_get_context_window_custom_provider_via_id_field():
|
||||
"""Some configs use `id` instead of `value` — both must resolve."""
|
||||
s = AppSettings(custom_providers=[
|
||||
CustomProvider(
|
||||
name="L",
|
||||
base_url="x",
|
||||
models=[{"id": "m", "context_window": 32_000}],
|
||||
),
|
||||
])
|
||||
assert get_context_window("L", "m", s) == 32_000
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# calculate_cost
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_calculate_cost_known_rates():
|
||||
"""Anthropic Sonnet: $3/M input, $15/M output. 1M of each → $18."""
|
||||
out = calculate_cost("Anthropic", "sonnet", 1_000_000, 1_000_000)
|
||||
assert out == 18.0
|
||||
|
||||
|
||||
def test_calculate_cost_case_insensitive_provider():
|
||||
out = calculate_cost("anthropic", "sonnet", 1_000_000, 0)
|
||||
assert out == 3.0
|
||||
|
||||
|
||||
def test_calculate_cost_unknown_returns_zero():
|
||||
assert calculate_cost("NobodyKnows", "made-up", 100_000, 50_000) == 0.0
|
||||
|
||||
|
||||
def test_calculate_cost_zero_token_count():
|
||||
assert calculate_cost("Anthropic", "sonnet", 0, 0) == 0.0
|
||||
|
||||
|
||||
def test_calculate_cost_subscription_path_zero_cost():
|
||||
"""Codex / Gemini CLI subscriptions are zero-cost to the user;
|
||||
confirm calculate_cost surfaces 0 even on heavy usage."""
|
||||
assert calculate_cost("OpenAI", "gpt-5.4", 1_000_000, 1_000_000) == 0.0
|
||||
assert calculate_cost("Google", "gemini-2.5-pro", 1_000_000, 1_000_000) == 0.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _is_9router_available cache
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_is_9router_available_caches_for_30s():
|
||||
"""Two consecutive calls within the 30s window should hit the cache
|
||||
after the first httpx.get."""
|
||||
reg._9router_cache["available"] = None
|
||||
reg._9router_cache["checked_at"] = 0
|
||||
fake_resp = MagicMock(status_code=200)
|
||||
with patch("httpx.get", return_value=fake_resp) as mock_get:
|
||||
a = reg._is_9router_available()
|
||||
b = reg._is_9router_available()
|
||||
assert a is True and b is True
|
||||
assert mock_get.call_count == 1
|
||||
|
||||
|
||||
def test_is_9router_available_handles_exception():
|
||||
"""Network error → cached False so we don't retry on every call."""
|
||||
reg._9router_cache["available"] = None
|
||||
reg._9router_cache["checked_at"] = 0
|
||||
with patch("httpx.get", side_effect=Exception("boom")):
|
||||
assert reg._is_9router_available() is False
|
||||
@@ -1,696 +0,0 @@
|
||||
"""Unit tests for the builtin agent tools.
|
||||
|
||||
These power the native agent loop's tool execution path. Currently 0%
|
||||
covered because the live CLI uses its own tool implementations. These
|
||||
tests pin the contract so the native loop can rely on it:
|
||||
|
||||
- `tools/registry`: register/get/get_all/init_tools roster
|
||||
- `tools/filesystem`: Read (text + image + offset/limit + missing),
|
||||
Write (creates parent dirs), Edit (exact + multi + replace_all),
|
||||
Glob (sorted matches + cap), Grep (rg path + Python fallback)
|
||||
- `tools/system`: Bash (echo + nonzero + timeout), AskUserQuestion
|
||||
- `tools/web`: WebSearch (mocked DuckDuckGo HTML), WebFetch (mocked
|
||||
httpx + html stripping + prompt header)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from backend.apps.agents.tools import registry as registry_mod
|
||||
from backend.apps.agents.tools.base import BaseTool, ToolContext
|
||||
from backend.apps.agents.tools.filesystem import (
|
||||
EditTool,
|
||||
GlobTool,
|
||||
GrepTool,
|
||||
ReadTool,
|
||||
WriteTool,
|
||||
_resolve,
|
||||
)
|
||||
from backend.apps.agents.tools.registry import (
|
||||
get_all_tool_schemas,
|
||||
get_all_tools,
|
||||
get_tool,
|
||||
init_tools,
|
||||
register_tool,
|
||||
)
|
||||
from backend.apps.agents.tools.system import AskUserQuestionTool, BashTool
|
||||
from backend.apps.agents.tools.web import WebFetchTool, WebSearchTool
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ctx(cwd: str) -> ToolContext:
|
||||
return ToolContext(cwd=cwd, session_id="test-sess")
|
||||
|
||||
|
||||
def _text(blocks: list[dict]) -> str:
|
||||
"""Pull the text content out of a tool result block list."""
|
||||
return "".join(b.get("text", "") for b in blocks if b.get("type") == "text")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# tools/registry
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_registry_init_tools_registers_full_roster():
|
||||
"""init_tools is run at import time. After import, all builtin
|
||||
tool names must be present in the registry."""
|
||||
init_tools() # idempotent
|
||||
expected = {
|
||||
"Read", "Write", "Edit", "Glob", "Grep",
|
||||
"Bash", "AskUserQuestion",
|
||||
"WebSearch", "WebFetch",
|
||||
}
|
||||
actual = {t.name for t in get_all_tools()}
|
||||
assert expected.issubset(actual)
|
||||
|
||||
|
||||
def test_register_tool_inserts_by_name():
|
||||
class FakeTool(BaseTool):
|
||||
name = "Fake_X"
|
||||
description = "fake"
|
||||
|
||||
def get_schema(self) -> dict:
|
||||
return {"type": "object"}
|
||||
|
||||
async def execute(self, input_data, context):
|
||||
return [{"type": "text", "text": "ok"}]
|
||||
|
||||
register_tool(FakeTool())
|
||||
try:
|
||||
assert get_tool("Fake_X") is not None
|
||||
assert get_tool("Fake_X").description == "fake"
|
||||
finally:
|
||||
registry_mod._TOOLS.pop("Fake_X", None)
|
||||
|
||||
|
||||
def test_get_tool_unknown_returns_none():
|
||||
assert get_tool("definitely-not-a-tool") is None
|
||||
|
||||
|
||||
def test_get_all_tool_schemas_returns_provider_agnostic_shape():
|
||||
schemas = get_all_tool_schemas()
|
||||
assert all(hasattr(s, "name") and hasattr(s, "input_schema") for s in schemas)
|
||||
by_name = {s.name: s for s in schemas}
|
||||
# Read tool's schema must require file_path
|
||||
assert "Read" in by_name
|
||||
assert by_name["Read"].input_schema["required"] == ["file_path"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# filesystem._resolve
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_relative_path_uses_cwd(tmp_path):
|
||||
p = _resolve("foo.txt", str(tmp_path))
|
||||
assert p == (tmp_path / "foo.txt").resolve()
|
||||
|
||||
|
||||
def test_resolve_absolute_path_passthrough(tmp_path):
|
||||
abs_path = str(tmp_path / "abs.txt")
|
||||
p = _resolve(abs_path, "/elsewhere")
|
||||
assert p == (tmp_path / "abs.txt").resolve()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ReadTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_read_tool_text_file_returns_numbered_lines(tmp_path):
|
||||
f = tmp_path / "hello.txt"
|
||||
f.write_text("line one\nline two\nline three\n")
|
||||
out = await ReadTool().execute({"file_path": str(f)}, _ctx(str(tmp_path)))
|
||||
text = _text(out)
|
||||
assert " 1\tline one" in text
|
||||
assert " 2\tline two" in text
|
||||
assert " 3\tline three" in text
|
||||
|
||||
|
||||
async def test_read_tool_offset_and_limit(tmp_path):
|
||||
"""offset is 1-based line number; limit caps total lines returned."""
|
||||
f = tmp_path / "many.txt"
|
||||
f.write_text("\n".join(f"row {i}" for i in range(1, 21)) + "\n")
|
||||
out = await ReadTool().execute(
|
||||
{"file_path": str(f), "offset": 5, "limit": 3},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
lines = [l for l in text.splitlines() if l.strip()]
|
||||
assert len(lines) == 3
|
||||
assert " 5\trow 5" in lines[0]
|
||||
assert " 7\trow 7" in lines[2]
|
||||
|
||||
|
||||
async def test_read_tool_missing_file_returns_error(tmp_path):
|
||||
out = await ReadTool().execute(
|
||||
{"file_path": str(tmp_path / "nope.txt")},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "Error: file not found" in _text(out)
|
||||
|
||||
|
||||
async def test_read_tool_empty_file_returns_marker(tmp_path):
|
||||
f = tmp_path / "empty.txt"
|
||||
f.write_text("")
|
||||
out = await ReadTool().execute({"file_path": str(f)}, _ctx(str(tmp_path)))
|
||||
assert "file is empty or offset beyond" in _text(out)
|
||||
|
||||
|
||||
async def test_read_tool_directory_path_returns_error(tmp_path):
|
||||
out = await ReadTool().execute(
|
||||
{"file_path": str(tmp_path)},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "not a regular file" in _text(out)
|
||||
|
||||
|
||||
async def test_read_tool_image_returns_base64_block(tmp_path):
|
||||
"""A PNG-extension file → image content block with base64 data."""
|
||||
f = tmp_path / "icon.png"
|
||||
raw = b"\x89PNG\r\n\x1a\nfake-png-bytes"
|
||||
f.write_bytes(raw)
|
||||
|
||||
out = await ReadTool().execute({"file_path": str(f)}, _ctx(str(tmp_path)))
|
||||
assert len(out) == 1
|
||||
assert out[0]["type"] == "image"
|
||||
assert out[0]["source"]["media_type"] == "image/png"
|
||||
assert out[0]["source"]["data"] == base64.b64encode(raw).decode("ascii")
|
||||
|
||||
|
||||
async def test_read_tool_offset_beyond_eof_returns_marker(tmp_path):
|
||||
f = tmp_path / "short.txt"
|
||||
f.write_text("only one line\n")
|
||||
out = await ReadTool().execute(
|
||||
{"file_path": str(f), "offset": 100},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "file is empty or offset beyond" in _text(out)
|
||||
|
||||
|
||||
async def test_read_tool_zero_limit_falls_back_to_default(tmp_path):
|
||||
"""limit<=0 → fall back to default 2000."""
|
||||
f = tmp_path / "two.txt"
|
||||
f.write_text("a\nb\n")
|
||||
out = await ReadTool().execute(
|
||||
{"file_path": str(f), "limit": 0},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert " 1\ta" in text and " 2\tb" in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WriteTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_write_tool_creates_file_and_parent_dirs(tmp_path):
|
||||
target = tmp_path / "deep" / "nested" / "file.txt"
|
||||
out = await WriteTool().execute(
|
||||
{"file_path": str(target), "content": "hello"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "Successfully wrote 5 bytes" in _text(out)
|
||||
assert target.read_text() == "hello"
|
||||
|
||||
|
||||
async def test_write_tool_overwrites_existing_file(tmp_path):
|
||||
f = tmp_path / "x.txt"
|
||||
f.write_text("old")
|
||||
await WriteTool().execute(
|
||||
{"file_path": str(f), "content": "new"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert f.read_text() == "new"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EditTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_edit_tool_unique_match_replaces(tmp_path):
|
||||
f = tmp_path / "edit.txt"
|
||||
f.write_text("hello world")
|
||||
out = await EditTool().execute(
|
||||
{"file_path": str(f), "old_string": "world", "new_string": "there"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "1 replacement" in _text(out)
|
||||
assert f.read_text() == "hello there"
|
||||
|
||||
|
||||
async def test_edit_tool_missing_string_errors(tmp_path):
|
||||
f = tmp_path / "edit.txt"
|
||||
f.write_text("nothing")
|
||||
out = await EditTool().execute(
|
||||
{"file_path": str(f), "old_string": "missing", "new_string": "x"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "old_string not found" in _text(out)
|
||||
|
||||
|
||||
async def test_edit_tool_multiple_matches_without_replace_all_errors(tmp_path):
|
||||
f = tmp_path / "edit.txt"
|
||||
f.write_text("aaaabbbb aaaa")
|
||||
out = await EditTool().execute(
|
||||
{"file_path": str(f), "old_string": "aaaa", "new_string": "X"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "appears 2 times" in _text(out)
|
||||
# File contents unchanged
|
||||
assert f.read_text() == "aaaabbbb aaaa"
|
||||
|
||||
|
||||
async def test_edit_tool_replace_all_replaces_every_match(tmp_path):
|
||||
f = tmp_path / "edit.txt"
|
||||
f.write_text("aaaa-aaaa-aaaa")
|
||||
out = await EditTool().execute(
|
||||
{
|
||||
"file_path": str(f),
|
||||
"old_string": "aaaa",
|
||||
"new_string": "X",
|
||||
"replace_all": True,
|
||||
},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "3 replacements" in _text(out)
|
||||
assert f.read_text() == "X-X-X"
|
||||
|
||||
|
||||
async def test_edit_tool_missing_file(tmp_path):
|
||||
out = await EditTool().execute(
|
||||
{"file_path": str(tmp_path / "nope.txt"), "old_string": "x", "new_string": "y"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "Error: file not found" in _text(out)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GlobTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_glob_tool_matches_files_sorted_by_mtime(tmp_path):
|
||||
older = tmp_path / "older.py"
|
||||
older.write_text("a")
|
||||
newer = tmp_path / "newer.py"
|
||||
newer.write_text("b")
|
||||
# Force older to be older than newer
|
||||
os.utime(older, (1, 1))
|
||||
|
||||
out = await GlobTool().execute(
|
||||
{"pattern": "*.py"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
# Newer first
|
||||
newer_idx = text.find("newer.py")
|
||||
older_idx = text.find("older.py")
|
||||
assert newer_idx >= 0 and older_idx >= 0
|
||||
assert newer_idx < older_idx
|
||||
|
||||
|
||||
async def test_glob_tool_no_matches_returns_marker(tmp_path):
|
||||
out = await GlobTool().execute(
|
||||
{"pattern": "*.nonexistent"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "No files matched" in _text(out)
|
||||
|
||||
|
||||
async def test_glob_tool_explicit_path_overrides_cwd(tmp_path):
|
||||
other = tmp_path / "other-dir"
|
||||
other.mkdir()
|
||||
(other / "x.md").write_text("x")
|
||||
out = await GlobTool().execute(
|
||||
{"pattern": "*.md", "path": str(other)},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "x.md" in _text(out)
|
||||
|
||||
|
||||
async def test_glob_tool_invalid_path_returns_error(tmp_path):
|
||||
out = await GlobTool().execute(
|
||||
{"pattern": "*", "path": str(tmp_path / "nope")},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "directory not found" in _text(out)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GrepTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_grep_tool_files_with_matches(tmp_path):
|
||||
a = tmp_path / "a.txt"
|
||||
a.write_text("the answer is 42")
|
||||
b = tmp_path / "b.txt"
|
||||
b.write_text("nothing here")
|
||||
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "answer", "path": str(tmp_path)},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "a.txt" in text
|
||||
assert "b.txt" not in text
|
||||
|
||||
|
||||
async def test_grep_tool_content_mode_includes_line_numbers(tmp_path):
|
||||
a = tmp_path / "a.txt"
|
||||
a.write_text("first line\nthe answer is 42\nthird line\n")
|
||||
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "answer", "path": str(tmp_path), "output_mode": "content"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
# rg prints `path:lineno:content`; python fallback uses same shape
|
||||
assert "answer is 42" in text
|
||||
|
||||
|
||||
async def test_grep_tool_count_mode(tmp_path):
|
||||
a = tmp_path / "a.txt"
|
||||
a.write_text("answer\nanswer\nnope\nanswer\n")
|
||||
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "answer", "path": str(tmp_path), "output_mode": "count"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "3" in text
|
||||
|
||||
|
||||
async def test_grep_tool_python_fallback_invalid_regex(tmp_path):
|
||||
"""When ripgrep isn't available and the regex is invalid, the
|
||||
Python fallback returns a clean error block."""
|
||||
# Force the rg attempt to raise FileNotFoundError so we hit fallback.
|
||||
with patch("asyncio.create_subprocess_exec", side_effect=FileNotFoundError):
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "[unclosed", "path": str(tmp_path)},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "Invalid regex" in _text(out)
|
||||
|
||||
|
||||
async def test_grep_tool_python_fallback_no_matches(tmp_path):
|
||||
a = tmp_path / "x.txt"
|
||||
a.write_text("nothing relevant")
|
||||
with patch("asyncio.create_subprocess_exec", side_effect=FileNotFoundError):
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "definitely-not-found", "path": str(tmp_path)},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "No matches found" in _text(out)
|
||||
|
||||
|
||||
async def test_grep_tool_python_fallback_path_not_found(tmp_path):
|
||||
with patch("asyncio.create_subprocess_exec", side_effect=FileNotFoundError):
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "anything", "path": str(tmp_path / "missing")},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
assert "path not found" in _text(out)
|
||||
|
||||
|
||||
async def test_grep_tool_python_fallback_glob_filter(tmp_path):
|
||||
"""Glob pattern restricts the file set the fallback scans."""
|
||||
(tmp_path / "match.py").write_text("found here")
|
||||
(tmp_path / "ignored.txt").write_text("found here too")
|
||||
|
||||
with patch("asyncio.create_subprocess_exec", side_effect=FileNotFoundError):
|
||||
out = await GrepTool().execute(
|
||||
{"pattern": "found", "path": str(tmp_path), "glob": "*.py"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "match.py" in text
|
||||
assert "ignored.txt" not in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BashTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_bash_tool_echo_round_trip(tmp_path):
|
||||
out = await BashTool().execute(
|
||||
{"command": "echo hello"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "hello" in text
|
||||
|
||||
|
||||
async def test_bash_tool_nonzero_exit_includes_code(tmp_path):
|
||||
out = await BashTool().execute(
|
||||
{"command": "exit 7"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "Exit code: 7" in text
|
||||
|
||||
|
||||
async def test_bash_tool_runs_in_session_cwd(tmp_path):
|
||||
(tmp_path / "marker.txt").write_text("x")
|
||||
out = await BashTool().execute(
|
||||
{"command": "ls"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "marker.txt" in text
|
||||
|
||||
|
||||
async def test_bash_tool_timeout_kills_process(tmp_path):
|
||||
"""timeout in milliseconds; passing 50ms forces the timeout path."""
|
||||
out = await BashTool().execute(
|
||||
{"command": "sleep 5", "timeout": 50},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "timed out" in text.lower()
|
||||
|
||||
|
||||
async def test_bash_tool_empty_output_with_zero_exit_includes_marker(tmp_path):
|
||||
"""Silent commands (e.g. `true`) get a synthetic completion marker."""
|
||||
out = await BashTool().execute(
|
||||
{"command": "true"},
|
||||
_ctx(str(tmp_path)),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "exit code 0" in text
|
||||
|
||||
|
||||
def test_bash_tool_truncate_helper_caps_long_output():
|
||||
"""_truncate adds a marker when the body is >100KB."""
|
||||
long = "x" * (101 * 1024)
|
||||
truncated = BashTool._truncate(long)
|
||||
assert truncated.endswith("(output truncated)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AskUserQuestionTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_ask_user_question_returns_question_text():
|
||||
out = await AskUserQuestionTool().execute(
|
||||
{"question": "Which file?"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
assert _text(out) == "Which file?"
|
||||
|
||||
|
||||
def test_ask_user_question_schema_requires_question():
|
||||
schema = AskUserQuestionTool().get_schema()
|
||||
assert schema["required"] == ["question"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSearchTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ddg_html(num: int = 3) -> str:
|
||||
"""Minimal DuckDuckGo HTML result page."""
|
||||
blocks = []
|
||||
for i in range(num):
|
||||
blocks.append(
|
||||
f'<div class="result results_links results_links_deep">'
|
||||
f'<a class="result__a" href="//duckduckgo.com/l/?uddg=https%3A%2F%2Fexample.com%2F{i}">Title {i}</a>'
|
||||
f'<a class="result__snippet" href="https://example.com/{i}">Snippet text {i}</a>'
|
||||
f'</div>'
|
||||
)
|
||||
return "".join(blocks)
|
||||
|
||||
|
||||
async def test_web_search_tool_parses_ddg_results():
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.text = _ddg_html(num=2)
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
|
||||
fake_client = MagicMock()
|
||||
fake_client.post = AsyncMock(return_value=fake_resp)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebSearchTool().execute(
|
||||
{"query": "openswarm"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "[1] Title 0" in text
|
||||
assert "https://example.com/0" in text
|
||||
assert "Snippet text 0" in text
|
||||
|
||||
|
||||
async def test_web_search_tool_empty_results_returns_marker():
|
||||
fake_resp = MagicMock(text="<html></html>")
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
fake_client = MagicMock()
|
||||
fake_client.post = AsyncMock(return_value=fake_resp)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebSearchTool().execute(
|
||||
{"query": "no-such-thing"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
assert "No search results" in _text(out)
|
||||
|
||||
|
||||
async def test_web_search_tool_exception_returns_error():
|
||||
fake_client = MagicMock()
|
||||
fake_client.post = AsyncMock(side_effect=RuntimeError("boom"))
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebSearchTool().execute(
|
||||
{"query": "x"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
assert "Web search error" in _text(out)
|
||||
|
||||
|
||||
async def test_web_search_tool_num_results_caps_returned_entries():
|
||||
fake_resp = MagicMock(text=_ddg_html(num=10))
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
fake_client = MagicMock()
|
||||
fake_client.post = AsyncMock(return_value=fake_resp)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebSearchTool().execute(
|
||||
{"query": "x", "num_results": 2},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "[1]" in text
|
||||
assert "[2]" in text
|
||||
assert "[3]" not in text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebFetchTool
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_web_fetch_tool_strips_html_to_plain_text():
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.text = "<html><body><p>Hello <b>world</b></p><script>x()</script></body></html>"
|
||||
fake_resp.headers = {"content-type": "text/html"}
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
|
||||
fake_client = MagicMock()
|
||||
fake_client.get = AsyncMock(return_value=fake_resp)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebFetchTool().execute(
|
||||
{"url": "https://example.com"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "Contents of https://example.com" in text
|
||||
assert "Hello" in text
|
||||
assert "world" in text
|
||||
assert "<script>" not in text # script blocks removed
|
||||
assert "<p>" not in text
|
||||
|
||||
|
||||
async def test_web_fetch_tool_includes_prompt_in_header():
|
||||
fake_resp = MagicMock()
|
||||
fake_resp.text = "raw text content"
|
||||
fake_resp.headers = {"content-type": "text/plain"}
|
||||
fake_resp.raise_for_status = MagicMock()
|
||||
|
||||
fake_client = MagicMock()
|
||||
fake_client.get = AsyncMock(return_value=fake_resp)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebFetchTool().execute(
|
||||
{"url": "https://example.com/x.txt", "prompt": "find the answer"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "Looking for: find the answer" in text
|
||||
assert "raw text content" in text
|
||||
|
||||
|
||||
async def test_web_fetch_tool_http_error_returns_marker():
|
||||
import httpx as _httpx
|
||||
|
||||
err_resp = MagicMock()
|
||||
err_resp.status_code = 404
|
||||
err = _httpx.HTTPStatusError("404", request=MagicMock(), response=err_resp)
|
||||
|
||||
fake_client = MagicMock()
|
||||
fake_client.get = AsyncMock(side_effect=err)
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebFetchTool().execute(
|
||||
{"url": "https://example.com/missing"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "HTTP error 404" in text
|
||||
|
||||
|
||||
async def test_web_fetch_tool_generic_exception_returns_marker():
|
||||
fake_client = MagicMock()
|
||||
fake_client.get = AsyncMock(side_effect=RuntimeError("DNS exploded"))
|
||||
fake_client.__aenter__ = AsyncMock(return_value=fake_client)
|
||||
fake_client.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
with patch("backend.apps.agents.tools.web.httpx.AsyncClient", return_value=fake_client):
|
||||
out = await WebFetchTool().execute(
|
||||
{"url": "https://example.com/x"},
|
||||
_ctx("/tmp"),
|
||||
)
|
||||
text = _text(out)
|
||||
assert "Error fetching" in text
|
||||
assert "DNS exploded" in text
|
||||
Reference in New Issue
Block a user