diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 47ccedf3..2976942c 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -42,5 +42,28 @@ jobs: -r backend/requirements.txt \ -r backend/requirements-dev.txt - - name: Run pytest - run: python -m pytest backend/tests/ -v + - name: Run pytest with coverage + run: | + python -m pytest backend/tests/ -v \ + --cov=backend \ + --cov-report=term-missing \ + --cov-report=html:backend/coverage_html \ + --cov-report=xml:backend/coverage.xml + + - name: Coverage summary + if: always() + run: | + echo "## Backend coverage" >> "$GITHUB_STEP_SUMMARY" + echo '```' >> "$GITHUB_STEP_SUMMARY" + python -m coverage report >> "$GITHUB_STEP_SUMMARY" || true + echo '```' >> "$GITHUB_STEP_SUMMARY" + + - name: Upload coverage artifacts + if: always() + uses: actions/upload-artifact@v4 + with: + name: backend-coverage + path: | + backend/coverage_html + backend/coverage.xml + if-no-files-found: ignore diff --git a/.gitignore b/.gitignore index bc7bfc3c..94e0f83a 100644 --- a/.gitignore +++ b/.gitignore @@ -33,3 +33,9 @@ openswarm-cloud .claude/ # Local-only operator helpers (never commit) scripts/set-fly-*.sh + +# Coverage reports (generated by scripts/test.sh and CI) +.coverage +.coverage.* +backend/coverage_html/ +backend/coverage.xml diff --git a/backend/requirements-dev.txt b/backend/requirements-dev.txt index 4f15a9fc..66d6ca12 100644 --- a/backend/requirements-dev.txt +++ b/backend/requirements-dev.txt @@ -9,5 +9,6 @@ pytest==8.3.4 pytest-asyncio==0.25.2 +pytest-cov==7.1.0 -e ./debugger \ No newline at end of file diff --git a/backend/tests/test_agent_manager_unit.py b/backend/tests/test_agent_manager_unit.py new file mode 100644 index 00000000..24121c07 --- /dev/null +++ b/backend/tests/test_agent_manager_unit.py @@ -0,0 +1,1392 @@ +"""Unit tests for `backend.apps.agents.agent_manager`. + +Covers the pure-logic helpers + the lifecycle methods on `AgentManager` +that don't require a live Claude Code subprocess. The streaming / +tool-execution path (`_run_agent_loop`) is deliberately out of scope — +it depends on the real CLI and lives in a future integration suite. + +Test groups: + - module-level helpers (`_save_session` / `_load_session_data` / + `_load_all_session_data`, error-classifier regex tables, + permission helpers, `_ensure_cwd_git_repo`) + - pure instance methods (`_resolve_mode`, `_compose_system_prompt`, + `_resolve_context_paths`, `_build_dir_tree`, `_resolve_forced_tools`, + `_resolve_attached_skills`, `_get_branch_messages`, + `_build_history_prefix`, `_approx_tokens`, + `_summarize_message_block`, `_truncate_large_tool_result`, + `_build_search_text`, `_maybe_compact`) + - lifecycle (launch/update/edit/switch_branch/duplicate/close/ + delete/resume/stop, history, browser-agent children, approval, + reconcile, persist+restore) + - cross-provider fork on `send_message` + +Tests rely on the conftest fixtures `tmp_data_dirs` (clears +`AgentManager.sessions/tasks` + per-feature data dirs) and where +needed, monkeypatches the SESSIONS_DIR module symbol so a single test +can write directly to a `tmp_path` without leaking into siblings. +""" + +from __future__ import annotations + +import asyncio +import json +import os +from datetime import datetime, timedelta +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from backend.apps.agents import agent_manager as am_mod +from backend.apps.agents.agent_manager import ( + AgentManager, + FULL_TOOLS, + _delete_session_file, + _ensure_cwd_git_repo, + _get_all_known_tool_names, + _get_denied_tool_names, + _is_auth_error, + _is_fully_denied, + _is_long_context_error, + _is_transient_capacity_error, + _load_all_session_data, + _load_session_data, + _save_session, + get_all_tool_names, +) +from backend.apps.agents.models import ( + AgentConfig, + AgentSession, + ApprovalRequest, + Message, + MessageBranch, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _seed_session( + manager: AgentManager, + *, + name: str = "Test", + model: str = "sonnet", + mode: str = "agent", + dashboard_id: str | None = None, + parent_session_id: str | None = None, + messages: list[Message] | None = None, +) -> AgentSession: + """Insert a synthetic session straight into manager.sessions. + + Bypasses launch_agent so we can test the post-launch methods + deterministically without exercising the lifespan / mode store. + """ + session = AgentSession( + name=name, + model=model, + mode=mode, + dashboard_id=dashboard_id, + parent_session_id=parent_session_id, + messages=messages or [], + ) + manager.sessions[session.id] = session + return session + + +# --------------------------------------------------------------------------- +# Module helpers: session file IO +# --------------------------------------------------------------------------- + + +def test_save_load_session_round_trip(tmp_data_dirs): + payload = { + "id": "abc", + "name": "round", + "model": "sonnet", + "mode": "agent", + "messages": [], + } + _save_session("abc", payload) + loaded = _load_session_data("abc") + assert loaded == payload + + +def test_load_session_missing_returns_none(tmp_data_dirs): + assert _load_session_data("nope-does-not-exist") is None + + +def test_delete_session_file_idempotent(tmp_data_dirs): + _save_session("d1", {"id": "d1"}) + _delete_session_file("d1") + _delete_session_file("d1") # second call must not raise + assert _load_session_data("d1") is None + + +def test_load_all_session_data_empty_dir(tmp_path, monkeypatch): + """When SESSIONS_DIR doesn't exist, must return [] rather than raise.""" + monkeypatch.setattr(am_mod, "SESSIONS_DIR", str(tmp_path / "nope")) + assert _load_all_session_data() == [] + + +def test_load_all_session_data_returns_id_and_payload(tmp_data_dirs): + _save_session("a1", {"id": "a1", "name": "A"}) + _save_session("b2", {"id": "b2", "name": "B"}) + pairs = _load_all_session_data() + by_id = {sid: data for sid, data in pairs} + assert by_id["a1"]["name"] == "A" + assert by_id["b2"]["name"] == "B" + + +# --------------------------------------------------------------------------- +# Error classifiers +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "msg", + [ + "HTTP 429 rate_limit_error", + "HTTP 503 service unavailable", + "HTTP 502 bad gateway", + "Anthropic is overloaded", + "model is at capacity", + "Try again shortly", + "internal server error", + "ECONNRESET on upstream", + "ETIMEDOUT", + "fetch failed", + "upstream connect error", + "No pool capacity available. Try again shortly.", + ], +) +def test_is_transient_capacity_error_matches_known_signals(msg: str): + """Every entry in `_TRANSIENT_CAPACITY_PATTERNS` (+ the explicit + no-pool-capacity check) must classify as transient.""" + assert _is_transient_capacity_error(RuntimeError(msg)) + + +@pytest.mark.parametrize( + "msg", + [ + "Usage cap exceeded for this billing period", + "You've reached your OpenSwarm Pro plan limit", + "no active subscription", + "subscription canceled", + "subscription past_due", + "Invalid token supplied", + "missing bearer token", + "401 Unauthorized", + "403 Forbidden", + "extra usage is required for long context", + ], +) +def test_is_transient_capacity_error_rejects_non_transient(msg: str): + """`_NON_TRANSIENT_PATTERNS` short-circuits even when the same + message also matches a transient pattern (via the explicit + `_NON_TRANSIENT_PATTERNS.search → return False` early-out).""" + assert not _is_transient_capacity_error(RuntimeError(msg)) + + +def test_is_transient_capacity_error_uses_extra_text(): + """The CLI's ProcessError stringifies generically; the real cause + only shows up via the `extra_text` channel (subprocess stderr). + Both must be classified.""" + exc = RuntimeError("Command failed with exit code 1") + assert _is_transient_capacity_error(exc, extra_text="rate_limit_error from upstream") + assert _is_transient_capacity_error(exc, extra_text="Overloaded; please retry") + + +def test_is_transient_capacity_error_empty_returns_false(): + assert not _is_transient_capacity_error(RuntimeError("")) + + +@pytest.mark.parametrize( + "msg,expected", + [ + ("extra usage is required for long context", True), + ("long context request requires premium tier", True), + ("long context not available on this plan", True), + ("regular 429 rate limit", False), + ("HTTP 503", False), + ], +) +def test_is_long_context_error(msg: str, expected: bool): + assert _is_long_context_error(RuntimeError(msg)) is expected + + +@pytest.mark.parametrize( + "msg,expected", + [ + ("HTTP 401 Unauthorized", True), + ("Invalid authentication credentials", True), + ("invalid api-key", True), + ("missing bearer token", True), + ("403 Forbidden", True), + ("no credentials for provider: claude", True), + ("provider not configured", True), + ("provider not connected", True), + ("HTTP 500 server error", False), + ("rate_limit_error", False), + ], +) +def test_is_auth_error(msg: str, expected: bool): + assert _is_auth_error(RuntimeError(msg)) is expected + + +def test_is_auth_error_empty_returns_false(): + assert not _is_auth_error(RuntimeError("")) + + +# --------------------------------------------------------------------------- +# Permission helpers + get_all_tool_names +# --------------------------------------------------------------------------- + + +def _make_tool(perms: dict) -> object: + """Synthetic ToolDefinition stand-in for permission helpers.""" + obj = MagicMock() + obj.tool_permissions = perms + return obj + + +def test_get_denied_tool_names_filters_underscored_keys(): + tool = _make_tool({ + "list_files": "deny", + "read_file": "always_allow", + "write_file": "deny", + "_tool_descriptions": {"list_files": "x", "read_file": "y"}, # must be skipped + }) + assert _get_denied_tool_names(tool) == {"list_files", "write_file"} + + +def test_get_all_known_tool_names_reads_descriptions_map(): + tool = _make_tool({ + "_tool_descriptions": {"a": "d", "b": "d", "c": "d"}, + }) + assert _get_all_known_tool_names(tool) == {"a", "b", "c"} + + +def test_is_fully_denied_true_when_every_known_subtool_denied(): + tool = _make_tool({ + "_tool_descriptions": {"a": "d", "b": "d"}, + "a": "deny", + "b": "deny", + }) + assert _is_fully_denied(tool) is True + + +def test_is_fully_denied_false_when_partial_or_unknown(): + partial = _make_tool({ + "_tool_descriptions": {"a": "d", "b": "d"}, + "a": "deny", # b not denied + }) + assert _is_fully_denied(partial) is False + + no_known = _make_tool({"_tool_descriptions": {}}) + assert _is_fully_denied(no_known) is False + + +def test_get_all_tool_names_returns_full_tools_when_no_perms_set(tmp_data_dirs): + """With an empty SETTINGS/TOOLS dir (tmp_data_dirs wipes both), + no builtin permissions are set → every entry in FULL_TOOLS is + surfaced. No MCP tools because the tool dir is empty too.""" + names = get_all_tool_names() + assert set(FULL_TOOLS).issubset(set(names)) + + +def test_get_all_tool_names_drops_explicitly_denied_builtins(tmp_data_dirs): + """Writes a builtin-permissions file to mark Bash as denied. The + file lives outside `tmp_data_dirs`'s wipe set, so we restore it + afterwards to avoid leaking state into sibling tests.""" + from backend.apps.tools_lib.tools_lib import save_builtin_permissions + from backend.config.paths import BUILTIN_PERMISSIONS_PATH + + try: + save_builtin_permissions({"Bash": "deny"}) + names = get_all_tool_names() + assert "Bash" not in names + assert "Read" in names # other tools survive + finally: + if os.path.exists(BUILTIN_PERMISSIONS_PATH): + os.remove(BUILTIN_PERMISSIONS_PATH) + + +# --------------------------------------------------------------------------- +# _ensure_cwd_git_repo +# --------------------------------------------------------------------------- + + +def test_ensure_cwd_git_repo_creates_repo_when_missing(tmp_path): + """Fresh tmp dir with no .git → function inits a repo + empty commit.""" + cwd = tmp_path / "fresh" + cwd.mkdir() + _ensure_cwd_git_repo(str(cwd), home=str(tmp_path)) + # Either .git lives here directly, or git decided we're already in + # a parent repo (the test runner's repo, for instance). Both are + # valid healthy outcomes. + if (cwd / ".git").exists(): + # Verify HEAD resolves — the function commits an empty seed. + import subprocess + head = subprocess.run( + ["git", "rev-parse", "--verify", "HEAD"], + cwd=str(cwd), + capture_output=True, + ) + assert head.returncode == 0 + + +def test_ensure_cwd_git_repo_skips_risky_roots(tmp_path): + """Calling on $HOME / / / parent-of-home must short-circuit and + leave the directory untouched.""" + home = str(tmp_path) + _ensure_cwd_git_repo(home, home=home) + assert not (tmp_path / ".git").exists() + + +def test_ensure_cwd_git_repo_silent_on_missing_dir(tmp_path): + """Nonexistent path → silent return, no exception.""" + _ensure_cwd_git_repo(str(tmp_path / "nope"), home=str(tmp_path)) + + +# --------------------------------------------------------------------------- +# Pure instance methods +# --------------------------------------------------------------------------- + + +def test_resolve_mode_unknown_returns_full_tools_and_no_prompt(tmp_data_dirs): + """No mode file on disk → fallback returns (get_all_tool_names(), + None, None).""" + mgr = AgentManager() + tools, prompt, folder = mgr._resolve_mode("definitely-not-a-mode") + assert prompt is None and folder is None + assert set(FULL_TOOLS).issubset(set(tools)) + + +def test_resolve_mode_known_returns_mode_definition(tmp_data_dirs): + """`ask` mode (built-in) has a fixed tool list + a system prompt.""" + from backend.apps.modes.modes import _save as save_mode + from backend.apps.modes.models import BUILTIN_MODES + + ask = next(m for m in BUILTIN_MODES if m.id == "ask") + save_mode(ask) + + mgr = AgentManager() + tools, prompt, folder = mgr._resolve_mode("ask") + assert "Read" in tools + assert "Edit" not in tools + assert prompt and "Ask mode" in prompt + assert folder is None + + +def test_resolve_mode_with_null_tools_returns_get_all_tool_names(tmp_data_dirs): + """`agent` mode ships with `tools=None`, which means 'all available'.""" + from backend.apps.modes.modes import _save as save_mode + from backend.apps.modes.models import BUILTIN_MODES + + agent_mode = next(m for m in BUILTIN_MODES if m.id == "agent") + save_mode(agent_mode) + + mgr = AgentManager() + tools, _prompt, _folder = mgr._resolve_mode("agent") + assert set(FULL_TOOLS).issubset(set(tools)) + + +def test_compose_system_prompt_joins_truthy_parts(): + mgr = AgentManager() + out = mgr._compose_system_prompt( + "default", + "mode", + "session", + connected_tools_ctx="tools", + outputs_ctx="outputs", + browser_ctx="browser", + mcp_registry_ctx="registry", + ) + # Order: default, mode, session, tools, registry, outputs, browser + assert out is not None + parts = out.split("\n\n") + assert parts == ["default", "mode", "session", "tools", "registry", "outputs", "browser"] + + +def test_compose_system_prompt_drops_falsy_parts(): + mgr = AgentManager() + out = mgr._compose_system_prompt(None, "", "session", connected_tools_ctx="ctx") + assert out == "session\n\nctx" + + +def test_compose_system_prompt_all_none_returns_none(): + assert AgentManager()._compose_system_prompt(None, None, None) is None + + +def test_resolve_context_paths_empty_returns_empty(): + assert AgentManager()._resolve_context_paths(None) == "" + assert AgentManager()._resolve_context_paths([]) == "" + + +def test_resolve_context_paths_file_round_trip(tmp_path): + f = tmp_path / "note.txt" + f.write_text("hello world") + out = AgentManager()._resolve_context_paths( + [{"path": str(f), "type": "file"}] + ) + assert "" in out and "" in out + assert "- Read" in out + assert "- Bash" in out + + +def test_resolve_attached_skills_empty_returns_empty(): + assert AgentManager()._resolve_attached_skills(None) == "" + assert AgentManager()._resolve_attached_skills([]) == "" + + +def test_resolve_attached_skills_emits_block_per_skill(): + out = AgentManager()._resolve_attached_skills([ + {"name": "Skill A", "content": "do X"}, + {"name": "Skill B", "content": "do Y"}, + {"name": "Empty", "content": ""}, # silently dropped + ]) + assert "[Using skill: Skill A]" in out + assert "do X" in out + assert "[Using skill: Skill B]" in out + assert "Empty" not in out + + +# --------------------------------------------------------------------------- +# _get_branch_messages +# --------------------------------------------------------------------------- + + +def test_get_branch_messages_main_only(): + """No branches beyond main → returns the main-branch messages.""" + s = AgentSession(name="x", model="sonnet") + s.messages = [ + Message(role="user", content="hi", branch_id="main"), + Message(role="assistant", content="hello", branch_id="main"), + ] + out = AgentManager._get_branch_messages(s) + assert [m.content for m in out] == ["hi", "hello"] + + +def test_get_branch_messages_walks_fork_lineage(): + """Branch B forks off main at the second user message; switching + active to B should yield main-up-to-fork + B's own messages.""" + s = AgentSession(name="x", model="sonnet") + u1 = Message(id="u1", role="user", content="first", branch_id="main") + a1 = Message(id="a1", role="assistant", content="reply1", branch_id="main") + u2 = Message(id="u2", role="user", content="second", branch_id="main") + a2 = Message(id="a2", role="assistant", content="reply2", branch_id="main") + # B forks at u2 + bu = Message(id="bu", role="user", content="fork-msg", branch_id="b1") + ba = Message(id="ba", role="assistant", content="fork-reply", branch_id="b1") + s.messages = [u1, a1, u2, a2, bu, ba] + s.branches["b1"] = MessageBranch(id="b1", parent_branch_id="main", fork_point_message_id="u2") + s.active_branch_id = "b1" + + out = AgentManager._get_branch_messages(s) + contents = [m.content for m in out] + # main slice (u1, a1) + b1 segment (bu, ba). u2 / a2 must be excluded. + assert "first" in contents and "reply1" in contents + assert "fork-msg" in contents and "fork-reply" in contents + assert "second" not in contents and "reply2" not in contents + + +# --------------------------------------------------------------------------- +# _build_history_prefix / _approx_tokens / _summarize_message_block +# --------------------------------------------------------------------------- + + +def test_build_history_prefix_skips_non_user_assistant_and_hidden(): + msgs = [ + Message(role="user", content="hi"), + Message(role="tool_call", content={"tool": "Read"}), # skipped + Message(role="assistant", content="hello"), + Message(role="user", content="hidden", hidden=True), # skipped + ] + out = AgentManager._build_history_prefix(msgs) + assert "" in out + assert "User: hi" in out + assert "Assistant: hello" in out + assert "Read" not in out + assert "hidden" not in out + + +def test_build_history_prefix_empty_returns_empty(): + assert AgentManager._build_history_prefix([]) == "" + + +@pytest.mark.parametrize("text,expected", [ + ("", 1), + ("a" * 4, 1), + ("a" * 16, 4), + ("a" * 100, 25), +]) +def test_approx_tokens_chars_over_four(text, expected): + assert AgentManager._approx_tokens(text) == expected + + +def test_approx_tokens_handles_none(): + assert AgentManager._approx_tokens(None) == 1 + + +def test_summarize_message_block_empty_returns_empty(): + assert AgentManager._summarize_message_block([]) == "" + + +def test_summarize_message_block_extracts_initial_task_and_counts(): + msgs = [ + Message(role="user", content="please do the thing"), + Message(role="tool_call", content={"tool": "Read", "input": {}}), + Message(role="tool_call", content={"tool": "Bash", "input": {}}), + Message(role="tool_call", content={"tool": "Read", "input": {}}), + Message(role="tool_result", content="ok"), + Message(role="assistant", content="done"), + ] + out = AgentManager._summarize_message_block(msgs) + assert "" in out + assert "please do the thing" in out + # Tool counts + assert "Read×2" in out + assert "Bash×1" in out + assert "Tool calls so far (3 total)" in out + assert "Tool results received: 1" in out + assert "Last assistant message:" in out + assert "done" in out + + +def test_summarize_message_block_assistant_list_content(): + """Assistant messages with list content (Anthropic block shape) + should still surface their text.""" + msgs = [ + Message(role="user", content="task"), + Message(role="assistant", content=[ + {"type": "text", "text": "the answer"}, + {"type": "tool_use", "id": "t1"}, + ]), + ] + out = AgentManager._summarize_message_block(msgs) + assert "the answer" in out + + +# --------------------------------------------------------------------------- +# _truncate_large_tool_result +# --------------------------------------------------------------------------- + + +def test_truncate_large_tool_result_under_threshold_unchanged(tmp_data_dirs): + content, blob = AgentManager._truncate_large_tool_result( + "small", "sess1", "msg1", max_bytes=100, + ) + assert content == "small" + assert blob is None + + +def test_truncate_large_tool_result_spills_over_threshold(tmp_data_dirs): + """Content over threshold → first 4K kept inline + saved to disk + under SESSIONS_DIR//blobs/.txt.""" + huge = "x" * 80_000 + replacement, blob_path = AgentManager._truncate_large_tool_result( + huge, "sessA", "msgA", max_bytes=50_000, + ) + assert blob_path is not None + assert os.path.exists(blob_path) + with open(blob_path) as fh: + assert fh.read() == huge + assert isinstance(replacement, str) + assert "[truncated" in replacement + # Inline keeps a 4K head. + assert replacement.startswith("x" * 4_000) + + +def test_truncate_large_tool_result_serializes_non_string(tmp_data_dirs): + big_dict = {"k": "v" * 40_000} + replacement, blob_path = AgentManager._truncate_large_tool_result( + big_dict, "sessB", "msgB", max_bytes=10_000, + ) + assert blob_path is not None + assert "[truncated" in replacement + + +# --------------------------------------------------------------------------- +# _build_search_text +# --------------------------------------------------------------------------- + + +def test_build_search_text_concatenates_user_assistant_only(): + s = AgentSession(name="My Session", model="sonnet") + s.messages = [ + Message(role="user", content="user msg"), + Message(role="assistant", content="asst msg"), + Message(role="tool_call", content={"tool": "Read"}), # skipped (dict content) + Message(role="tool_result", content="result"), # skipped + ] + out = AgentManager._build_search_text(s) + assert "My Session" in out + assert "user msg" in out + assert "asst msg" in out + assert "Read" not in out + assert "result" not in out + + +def test_build_search_text_truncates_to_max_len(): + s = AgentSession(name="x", model="sonnet") + s.messages = [Message(role="user", content="a" * 10_000)] + out = AgentManager._build_search_text(s, max_len=200) + assert len(out) == 200 + + +# --------------------------------------------------------------------------- +# _maybe_compact +# --------------------------------------------------------------------------- + + +def test_maybe_compact_below_threshold_returns_false(): + s = AgentSession(name="x", model="sonnet", compact_threshold_pct=0.65, context_window=200_000) + s.tokens["input"] = 1_000 # way under threshold + s.messages = [Message(role="user", content="hi") for _ in range(20)] + assert AgentManager()._maybe_compact(s) is False + + +def test_maybe_compact_force_with_few_messages_returns_false(): + """Even with force=True, a session with <4 messages can't compact.""" + s = AgentSession(name="x", model="sonnet") + s.messages = [Message(role="user", content="hi")] + assert AgentManager()._maybe_compact(s, force=True) is False + + +def test_maybe_compact_force_advances_compacted_through_id(): + s = AgentSession(name="x", model="sonnet") + msgs = [Message(role="user", content=f"m{i}") for i in range(20)] + s.messages = msgs + assert AgentManager()._maybe_compact(s, force=True) is True + # compacted_through_msg_id is the message at len-6 - 1 = 13. + assert s.compacted_through_msg_id == msgs[13].id + + +def test_maybe_compact_idempotent_when_already_compacted(): + s = AgentSession(name="x", model="sonnet") + msgs = [Message(role="user", content=f"m{i}") for i in range(20)] + s.messages = msgs + AgentManager()._maybe_compact(s, force=True) + # Second call with same state and not forced past the saved id → no-op + snapshot = s.compacted_through_msg_id + assert AgentManager()._maybe_compact(s) is False + assert s.compacted_through_msg_id == snapshot + + +# --------------------------------------------------------------------------- +# _build_prompt_content +# --------------------------------------------------------------------------- + + +def test_build_prompt_content_no_images_returns_string(): + out = AgentManager()._build_prompt_content("hello") + assert out == "hello" + + +def test_build_prompt_content_with_images_returns_blocks(): + out = AgentManager()._build_prompt_content( + "describe this", + images=[{"data": "base64bytes", "media_type": "image/png"}], + ) + assert isinstance(out, list) + assert out[0] == {"type": "text", "text": "describe this"} + assert out[1]["type"] == "image" + assert out[1]["source"]["data"] == "base64bytes" + + +def test_build_prompt_content_combines_context_and_forced(tmp_path): + f = tmp_path / "ctx.txt" + f.write_text("context content") + out = AgentManager()._build_prompt_content( + "real prompt", + context_paths=[{"path": str(f), "type": "file"}], + forced_tools=["Read"], + ) + assert "" in out + assert " to avoid writing into $HOME.""" + home = os.environ["HOME"] + monkeypatch.setattr("backend.apps.agents.agent_manager._ensure_cwd_git_repo", lambda *a, **kw: None) + + mgr = AgentManager() + cfg = AgentConfig(name="HomeFallback", model="sonnet", mode="agent") + session = await mgr.launch_agent(cfg) + + assert session.cwd != home + assert session.cwd.startswith(os.path.join(home, ".openswarm", "workspaces")) + + +# --------------------------------------------------------------------------- +# Lifecycle: update_session +# --------------------------------------------------------------------------- + + +async def test_update_session_allowlist_updates_only_known_fields(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr, name="Original", model="sonnet") + + await mgr.update_session( + s.id, + name="Renamed", + system_prompt="new sys", + thinking_level="high", + model="ignored-not-allowed", + cost_usd=999.0, # also ignored + ) + + assert s.name == "Renamed" + assert s.system_prompt == "new sys" + assert s.thinking_level == "high" + assert s.model == "sonnet" # not changed + assert s.cost_usd == 0.0 # not changed + + +async def test_update_session_rejects_invalid_thinking_level(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + s.thinking_level = "auto" + + await mgr.update_session(s.id, thinking_level="extreme") + assert s.thinking_level == "auto" + + +async def test_update_session_unknown_id_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().update_session("nope", name="x") + + +# --------------------------------------------------------------------------- +# Lifecycle: switch_branch +# --------------------------------------------------------------------------- + + +async def test_switch_branch_unknown_raises(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + with pytest.raises(ValueError): + await mgr.switch_branch(s.id, "no-such-branch") + + +async def test_switch_branch_to_existing_updates_active(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + s.branches["b1"] = MessageBranch(id="b1", parent_branch_id="main", fork_point_message_id="abc") + + await mgr.switch_branch(s.id, "b1") + assert s.active_branch_id == "b1" + + +async def test_switch_branch_unknown_session_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().switch_branch("nope", "main") + + +# --------------------------------------------------------------------------- +# Lifecycle: edit_message +# --------------------------------------------------------------------------- + + +async def test_edit_message_unknown_session_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().edit_message("nope", "m1", "x") + + +async def test_edit_message_non_user_role_raises(tmp_data_dirs, stub_agent_loop): + mgr = AgentManager() + s = _seed_session(mgr) + asst = Message(id="m-asst", role="assistant", content="reply", branch_id="main") + s.messages.append(asst) + with pytest.raises(ValueError): + await mgr.edit_message(s.id, "m-asst", "new content") + + +async def test_edit_message_creates_new_branch(tmp_data_dirs, stub_agent_loop): + mgr = AgentManager() + s = _seed_session(mgr) + user_msg = Message(id="u1", role="user", content="orig", branch_id="main") + s.messages.append(user_msg) + + await mgr.edit_message(s.id, "u1", "edited content") + + # New branch created and is active. main remains a key. + assert s.active_branch_id != "main" + new_branch = s.branches[s.active_branch_id] + assert new_branch.parent_branch_id == "main" + assert new_branch.fork_point_message_id == "u1" + # New user message appended on the new branch with the edited content. + edited = next(m for m in s.messages if m.branch_id == s.active_branch_id and m.role == "user") + assert edited.content == "edited content" + + +async def test_edit_message_on_branched_msg_uses_parent_fork_point(tmp_data_dirs, stub_agent_loop): + """Editing the FIRST user message of a forked branch should fold + the new branch back to the parent's fork_point_message_id, not + the message we're editing — otherwise re-edits chain forever.""" + mgr = AgentManager() + s = _seed_session(mgr) + s.branches["b1"] = MessageBranch( + id="b1", parent_branch_id="main", fork_point_message_id="orig-fork", + ) + s.active_branch_id = "b1" + branch_first = Message(id="bf", role="user", content="first on b1", branch_id="b1") + s.messages.append(branch_first) + + await mgr.edit_message(s.id, "bf", "new") + + new_branch = s.branches[s.active_branch_id] + assert new_branch.parent_branch_id == "main" # not "b1" + assert new_branch.fork_point_message_id == "orig-fork" # parent's, not "bf" + + +# --------------------------------------------------------------------------- +# Lifecycle: stop_agent / handle_approval +# --------------------------------------------------------------------------- + + +async def test_stop_agent_sets_status_stopped_and_drains_approvals(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + s.status = "running" + req = ApprovalRequest(session_id=s.id, tool_name="Bash", tool_input={"cmd": "x"}) + s.pending_approvals.append(req) + + await mgr.stop_agent(s.id) + + assert s.status == "stopped" + assert s.pending_approvals == [] + assert s.closed_at is not None + + +async def test_stop_agent_unknown_session_no_op(tmp_data_dirs): + """No raise, just a no-op.""" + await AgentManager().stop_agent("nope") + + +async def test_stop_agent_cancels_browser_children(tmp_data_dirs): + mgr = AgentManager() + parent = _seed_session(mgr, name="parent") + child = _seed_session(mgr, name="child", mode="browser-agent", parent_session_id=parent.id) + child.status = "running" + + await mgr.stop_agent(parent.id) + + assert child.status == "stopped" + assert parent.status == "stopped" + + +def test_handle_approval_resolves_pending_future(tmp_data_dirs): + from backend.apps.agents.ws_manager import ws_manager + + async def runner(): + loop = asyncio.get_event_loop() + fut = loop.create_future() + ws_manager.pending_futures["req-1"] = fut + + AgentManager().handle_approval("req-1", {"behavior": "allow"}) + result = await asyncio.wait_for(fut, timeout=1.0) + return result + + decision = asyncio.run(runner()) + assert decision == {"behavior": "allow"} + + +def test_handle_approval_unknown_request_id_no_op(): + """Resolving an unknown id must not raise.""" + AgentManager().handle_approval("nonexistent-id", {"behavior": "deny"}) + + +# --------------------------------------------------------------------------- +# Lifecycle: close / delete / resume +# --------------------------------------------------------------------------- + + +async def test_close_session_persists_and_evicts_from_memory(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr, name="ToClose") + s.messages = [ + Message(role="user", content="hello"), + Message(role="assistant", content="hi back"), + ] + + await mgr.close_session(s.id) + + assert s.id not in mgr.sessions + # Persisted to disk with search_text injected + data = _load_session_data(s.id) + assert data is not None + assert data["name"] == "ToClose" + assert "search_text" in data + assert "hello" in data["search_text"] + + +async def test_close_session_unknown_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().close_session("nope") + + +async def test_delete_session_removes_memory_and_disk(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + _save_session(s.id, {"id": s.id, "name": s.name, "model": s.model}) + + await mgr.delete_session(s.id) + + assert s.id not in mgr.sessions + assert _load_session_data(s.id) is None + + +async def test_delete_session_unknown_silent(tmp_data_dirs): + """Hard-delete is best-effort; deleting an unknown id is a no-op.""" + await AgentManager().delete_session("nope") # must not raise + + +async def test_resume_session_loads_from_disk(tmp_data_dirs): + """Round-trip: close → resume returns the same session, file is + deleted so it doesn't show up in /history any more.""" + mgr = AgentManager() + s = _seed_session(mgr, name="Resume me") + await mgr.close_session(s.id) + + restored = await mgr.resume_session(s.id) + assert restored.id == s.id + assert restored.name == "Resume me" + assert restored.closed_at is None + assert _load_session_data(s.id) is None # file gone + + +async def test_resume_session_already_in_memory_returns_existing(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + out = await mgr.resume_session(s.id) + assert out is s + + +async def test_resume_session_unknown_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().resume_session("nope") + + +# --------------------------------------------------------------------------- +# Lifecycle: duplicate_session +# --------------------------------------------------------------------------- + + +async def test_duplicate_session_clones_messages_and_appends_copy_suffix(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr, name="Original") + s.messages = [ + Message(id="m1", role="user", content="hi"), + Message(id="m2", role="assistant", content="hello"), + ] + + new = await mgr.duplicate_session(s.id) + + assert new.id != s.id + assert new.name == "Original (copy)" + assert new.needs_fork is True + # New ids on each cloned message + new_ids = {m.id for m in new.messages} + assert "m1" not in new_ids and "m2" not in new_ids + # Same content, mapped order + assert [m.content for m in new.messages] == ["hi", "hello"] + + +async def test_duplicate_session_up_to_message_truncates(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr) + s.messages = [ + Message(id="m1", role="user", content="A"), + Message(id="m2", role="assistant", content="B"), + Message(id="m3", role="user", content="C"), + ] + + new = await mgr.duplicate_session(s.id, up_to_message_id="m2") + + assert [m.content for m in new.messages] == ["A", "B"] + + +async def test_duplicate_session_unknown_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().duplicate_session("nope") + + +# --------------------------------------------------------------------------- +# Lifecycle: get_history / get_browser_agent_children +# --------------------------------------------------------------------------- + + +def test_get_history_paginates_and_filters_by_dashboard(tmp_data_dirs): + """Seed three closed sessions on disk + one on a different + dashboard. Filter by dashboard_id and assert only matches return.""" + for i, dash in enumerate(["A", "A", "B"]): + _save_session(f"sess-{i}", { + "id": f"sess-{i}", + "name": f"name-{i}", + "model": "sonnet", + "mode": "agent", + "status": "stopped", + "closed_at": f"2026-04-29T00:00:0{i}", + "dashboard_id": dash, + "search_text": f"some text-{i}", + "messages": [], + }) + + history_a = AgentManager().get_history(dashboard_id="A") + assert history_a["total"] == 2 + assert all(s["dashboard_id"] == "A" for s in history_a["sessions"]) + + +def test_get_history_search_matches_name_and_search_text(tmp_data_dirs): + _save_session("alpha", { + "id": "alpha", "name": "alpha-name", "model": "sonnet", + "mode": "agent", "status": "stopped", "closed_at": "2026-04-30T00:00:00", + "search_text": "lorem", "messages": [], + }) + _save_session("beta", { + "id": "beta", "name": "unrelated", "model": "sonnet", + "mode": "agent", "status": "stopped", "closed_at": "2026-04-30T00:00:01", + "search_text": "ALPHA hidden in body", "messages": [], + }) + + out = AgentManager().get_history(q="alpha") + ids = {s["id"] for s in out["sessions"]} + assert ids == {"alpha", "beta"} + + +def test_get_history_pagination_math(tmp_data_dirs): + for i in range(5): + _save_session(f"s{i}", { + "id": f"s{i}", "name": f"n{i}", "model": "sonnet", + "mode": "agent", "status": "stopped", + "closed_at": f"2026-04-30T00:00:0{i}", + "messages": [], + }) + + page = AgentManager().get_history(limit=2, offset=2) + assert page["total"] == 5 + assert len(page["sessions"]) == 2 + assert page["has_more"] is True + + last = AgentManager().get_history(limit=2, offset=4) + assert len(last["sessions"]) == 1 + assert last["has_more"] is False + + +def test_get_browser_agent_children_combines_memory_and_disk(tmp_data_dirs): + mgr = AgentManager() + parent = _seed_session(mgr, name="parent") + + # In-memory child + in_mem = _seed_session(mgr, name="in-mem", mode="browser-agent", parent_session_id=parent.id) + + # Disk-only child + _save_session("disk-child", { + "id": "disk-child", "name": "disk", "model": "sonnet", + "mode": "browser-agent", "parent_session_id": parent.id, + "status": "stopped", "messages": [], + }) + + out = mgr.get_browser_agent_children(parent.id) + ids = {c["id"] for c in out} + assert in_mem.id in ids + assert "disk-child" in ids + + +def test_get_browser_agent_children_dedupes_by_id(tmp_data_dirs): + """If the same child is in memory AND on disk, memory wins; the + disk row is dropped to avoid double-listing.""" + mgr = AgentManager() + parent = _seed_session(mgr, name="parent") + child = _seed_session(mgr, name="child", mode="browser-agent", parent_session_id=parent.id) + _save_session(child.id, { + "id": child.id, "name": "stale-disk-copy", + "model": "sonnet", "mode": "browser-agent", + "parent_session_id": parent.id, "messages": [], + }) + + out = mgr.get_browser_agent_children(parent.id) + matching = [c for c in out if c["id"] == child.id] + assert len(matching) == 1 + # Memory copy wins + assert matching[0]["name"] == "child" + + +def test_get_all_sessions_filters_by_dashboard(tmp_data_dirs): + mgr = AgentManager() + a = _seed_session(mgr, dashboard_id="A") + b = _seed_session(mgr, dashboard_id="B") + none = _seed_session(mgr, dashboard_id=None) + + all_sessions = mgr.get_all_sessions() + assert {s.id for s in all_sessions} == {a.id, b.id, none.id} + + just_a = mgr.get_all_sessions(dashboard_id="A") + assert {s.id for s in just_a} == {a.id} + + +def test_get_session_returns_none_for_unknown(tmp_data_dirs): + assert AgentManager().get_session("nope") is None + + +# --------------------------------------------------------------------------- +# Lifecycle: reconcile + persist + restore +# --------------------------------------------------------------------------- + + +async def test_reconcile_on_startup_flips_stale_running_to_stopped(tmp_data_dirs): + _save_session("s1", { + "id": "s1", "name": "x", "model": "sonnet", "mode": "agent", + "status": "running", "messages": [], + }) + + await AgentManager().reconcile_on_startup() + + data = _load_session_data("s1") + assert data["status"] == "stopped" + + +async def test_reconcile_on_startup_migrates_chat_to_ask(tmp_data_dirs): + _save_session("s2", { + "id": "s2", "name": "x", "model": "sonnet", "mode": "chat", + "status": "stopped", "messages": [], + }) + + await AgentManager().reconcile_on_startup() + + data = _load_session_data("s2") + assert data["mode"] == "ask" + + +async def test_reconcile_on_startup_idempotent(tmp_data_dirs): + """Two consecutive reconciles must not rewrite a stable file.""" + _save_session("s3", { + "id": "s3", "name": "x", "model": "sonnet", "mode": "ask", + "status": "stopped", "messages": [], + }) + mgr = AgentManager() + await mgr.reconcile_on_startup() + from backend.config.paths import SESSIONS_DIR + path = os.path.join(SESSIONS_DIR, "s3.json") + mtime1 = os.path.getmtime(path) + await mgr.reconcile_on_startup() + assert os.path.getmtime(path) == mtime1 + + +async def test_persist_all_sessions_writes_and_clears(tmp_data_dirs): + mgr = AgentManager() + s = _seed_session(mgr, name="persist me") + s.status = "running" + + await mgr.persist_all_sessions() + + assert mgr.sessions == {} + assert mgr.tasks == {} + data = _load_session_data(s.id) + assert data is not None + assert data["status"] == "stopped" + assert "search_text" in data + + +async def test_restore_all_sessions_skips_closed(tmp_data_dirs): + """closed_at set → keep on disk for /history; closed_at None → restore.""" + _save_session("active", { + "id": "active", "name": "alive", "model": "sonnet", "mode": "agent", + "status": "running", "messages": [], "closed_at": None, + }) + _save_session("closed", { + "id": "closed", "name": "dead", "model": "sonnet", "mode": "agent", + "status": "stopped", "messages": [], "closed_at": "2026-04-30T00:00:00", + }) + + mgr = AgentManager() + await mgr.restore_all_sessions() + + assert "active" in mgr.sessions + assert "closed" not in mgr.sessions + # The active one's status is normalized stopped+ file removed. + assert mgr.sessions["active"].status == "stopped" + assert _load_session_data("active") is None + assert _load_session_data("closed") is not None + + +async def test_restore_all_sessions_skips_corrupt_file(tmp_data_dirs, caplog): + """A corrupt session file must NOT abort restore — log + skip + move on.""" + # Garbage payload that AgentSession can't validate + _save_session("good", { + "id": "good", "name": "ok", "model": "sonnet", "mode": "agent", + "status": "stopped", "messages": [], + }) + _save_session("bad", {"this": "is not a session"}) + + mgr = AgentManager() + await mgr.restore_all_sessions() + + assert "good" in mgr.sessions + assert "bad" not in mgr.sessions + + +# --------------------------------------------------------------------------- +# Lifecycle: send_message cross-provider fork +# --------------------------------------------------------------------------- + + +async def test_send_message_marks_needs_fork_on_cross_provider_switch(tmp_data_dirs, stub_agent_loop): + mgr = AgentManager() + s = _seed_session(mgr, model="sonnet") # api=anthropic + # Switch to gpt-5.4-mini (api=codex) — different api_type → fork. + await mgr.send_message(s.id, "ping", model="gpt-5.4-mini") + + assert s.needs_fork is True + assert s.model == "gpt-5.4-mini" + + +async def test_send_message_same_api_no_fork(tmp_data_dirs, stub_agent_loop): + mgr = AgentManager() + s = _seed_session(mgr, model="sonnet") + s.needs_fork = False + # opus is also anthropic → no fork required. + await mgr.send_message(s.id, "ping", model="opus") + + assert s.needs_fork is False + assert s.model == "opus" + + +async def test_send_message_unknown_session_raises(tmp_data_dirs): + with pytest.raises(ValueError): + await AgentManager().send_message("nope", "hi") + + +async def test_send_message_skips_when_task_already_running(tmp_data_dirs, stub_agent_loop): + """If a task for this session is in flight, send_message must + early-return without appending the user message twice.""" + mgr = AgentManager() + s = _seed_session(mgr) + + # Plant a never-resolving task in the registry. + async def _hang(): + await asyncio.Event().wait() + + task = asyncio.create_task(_hang()) + mgr.tasks[s.id] = task + try: + before = len(s.messages) + await mgr.send_message(s.id, "second prompt") + assert len(s.messages) == before # nothing appended + finally: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass diff --git a/backend/tests/test_agents_lifespan_integration.py b/backend/tests/test_agents_lifespan_integration.py new file mode 100644 index 00000000..3c11c7b8 --- /dev/null +++ b/backend/tests/test_agents_lifespan_integration.py @@ -0,0 +1,377 @@ +"""Integration tests for `backend.apps.agents.agents` (router + lifespan). + +Existing `test_api_agents.py` covers the happy paths of each REST +endpoint. This file fills in branches that those don't: + + - `agents_lifespan` startup runs `reconcile_on_startup` + + `restore_all_sessions`; shutdown stops in-flight tasks + + persists every active session. + - `POST /sessions/{sid}/message` schedules `mcp_preflight.run_preflight` + in a background task; with patched run_preflight returning + suggestions we observe the `agent:mcp_suggestions` event reach + `ws_manager`. + - Preflight raising → message still returns 200 (fail-open contract). + - `POST /sessions/{sid}/warm-cache` returns 200 even if the manager + raises, and is wired to `agent_manager.warm_prompt_cache`. + - `GET /api/agents/models` four logical branches: + a) no creds + no 9Router → empty Anthropic block + b) anthropic_api_key only → Anthropic group emitted + c) openswarm-pro + claude sub → both "OpenSwarm Pro" and + "Anthropic" groups + d) openswarm-pro alone → only "OpenSwarm Pro" +""" + +from __future__ import annotations + +import asyncio +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from backend.apps.agents import agents as agents_mod +from backend.apps.agents.agent_manager import ( + _save_session, + agent_manager, +) +from backend.apps.agents.models import AgentSession + + +# --------------------------------------------------------------------------- +# agents_lifespan startup + shutdown +# --------------------------------------------------------------------------- + + +async def test_agents_lifespan_startup_runs_reconcile_and_restore(tmp_data_dirs): + """Drop a stale-running session on disk, drive lifespan, assert the + session was reconciled (status flipped to stopped) AND restored + into memory (closed_at=None on disk record).""" + _save_session("active-startup", { + "id": "active-startup", "name": "alive", "model": "sonnet", + "mode": "agent", "status": "running", "messages": [], + "closed_at": None, + }) + + # Reset in-memory state so the lifespan's restore is observable + agent_manager.sessions.clear() + agent_manager.tasks.clear() + + async with agents_mod.agents_lifespan(): + # Inside the lifespan: session is restored into memory and + # status was reconciled to stopped (since we marked it running). + assert "active-startup" in agent_manager.sessions + assert agent_manager.sessions["active-startup"].status == "stopped" + + +async def test_agents_lifespan_shutdown_stops_running_tasks_and_persists(tmp_data_dirs): + """Plant an in-flight task in the manager, drive shutdown, assert + the task is cancelled and the session is persisted to disk.""" + agent_manager.sessions.clear() + agent_manager.tasks.clear() + + sess = AgentSession(name="ToShutdown", model="sonnet") + agent_manager.sessions[sess.id] = sess + + async def _hang(): + await asyncio.Event().wait() + + task = asyncio.create_task(_hang()) + agent_manager.tasks[sess.id] = task + + try: + async with agents_mod.agents_lifespan(): + pass + finally: + if not task.done(): + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + # After shutdown: in-memory clear + persisted to disk. + assert sess.id not in agent_manager.sessions + from backend.apps.agents.agent_manager import _load_session_data + data = _load_session_data(sess.id) + assert data is not None + assert data["status"] == "stopped" + + +# --------------------------------------------------------------------------- +# POST /sessions/{sid}/message — preflight branch +# --------------------------------------------------------------------------- + + +async def _drain_pending_tasks(loop_iters: int = 5) -> None: + """Yield to the event loop a few times so a fire-and-forget + `asyncio.create_task(_emit_preflight())` finishes before we assert.""" + for _ in range(loop_iters): + await asyncio.sleep(0) + + +def test_send_message_fires_preflight_emits_mcp_suggestions(client, stub_agent_loop): + """Patched run_preflight returns suggestions → `ws_manager` should + see an `agent:mcp_suggestions` event sent to the session.""" + captured: list[tuple[str, str, dict]] = [] + + async def _capture_send(session_id: str, event: str, payload: dict): + captured.append((session_id, event, payload)) + + suggestions = [{ + "id": "Slack", "title": "Slack", "description": "x", + "reason": "user mentioned channel", + }] + + # Launch a session first so message has something to land on. + r = client.post("/api/agents/launch", json={ + "name": "test", "model": "sonnet", "mode": "agent", + }) + assert r.status_code == 200 + sid = r.json()["session_id"] + + async def _fake_preflight(prompt: str, timeout_s: float = 2.0): + return {"is_vague": True, "suggestions": suggestions} + + # The send_message handler does `from .ws_manager import ws_manager as _ws` + # at call time, so we have to patch the source module's attribute, not + # the alias on `agents_mod`. + from backend.apps.agents import ws_manager as _ws_mod + + fake_ws = MagicMock() + fake_ws.send_to_session = AsyncMock(side_effect=_capture_send) + + with patch.object(_ws_mod, "ws_manager", fake_ws), \ + patch("backend.apps.agents.mcp_preflight.run_preflight", _fake_preflight): + r = client.post(f"/api/agents/sessions/{sid}/message", + json={"prompt": "send an update to the team channel"}) + assert r.status_code == 200 + + # The preflight emit is fired in a background task. Drain by + # giving the event loop time to run pending tasks. TestClient's + # `requests`-shaped surface returns synchronously after the + # endpoint coroutine completes, so the background task may not + # have run yet — but the FastAPI/Starlette runner shares its + # loop across consecutive sync calls. A second tiny request + # forces a loop step. + for _ in range(10): + client.get("/api/health") + if captured: + break + + suggestion_events = [c for c in captured if c[1] == "agent:mcp_suggestions"] + assert len(suggestion_events) >= 1 + payload = suggestion_events[0][2] + assert payload["is_vague"] is True + assert payload["suggestions"] == suggestions + + +def test_send_message_preflight_raises_still_returns_ok(client, stub_agent_loop): + """Preflight is best-effort: any exception inside the classifier + must be swallowed so the agent still proceeds.""" + r = client.post("/api/agents/launch", json={ + "name": "fail-open", "model": "sonnet", "mode": "agent", + }) + sid = r.json()["session_id"] + + async def _broken_preflight(prompt, timeout_s=2.0): + raise RuntimeError("preflight kaboom") + + with patch("backend.apps.agents.mcp_preflight.run_preflight", _broken_preflight): + r = client.post(f"/api/agents/sessions/{sid}/message", + json={"prompt": "do the thing"}) + assert r.status_code == 200 + assert r.json() == {"ok": True} + + +# --------------------------------------------------------------------------- +# POST /sessions/{sid}/warm-cache +# --------------------------------------------------------------------------- + + +def test_warm_cache_calls_agent_manager(client): + """Endpoint should invoke `agent_manager.warm_prompt_cache(session_id)`.""" + called: list[str] = [] + + async def _fake_warm(session_id: str): + called.append(session_id) + + with patch.object(agent_manager, "warm_prompt_cache", side_effect=_fake_warm): + r = client.post("/api/agents/sessions/sess-x/warm-cache") + assert r.status_code == 200 + assert r.json() == {"ok": True} + assert called == ["sess-x"] + + +def test_warm_cache_swallows_manager_exceptions(client): + """Best-effort: an exception from warm_prompt_cache must NOT bubble.""" + async def _explode(session_id: str): + raise RuntimeError("kaboom") + + with patch.object(agent_manager, "warm_prompt_cache", side_effect=_explode): + r = client.post("/api/agents/sessions/sess-x/warm-cache") + assert r.status_code == 200 + assert r.json() == {"ok": True} + + +# --------------------------------------------------------------------------- +# GET /api/agents/models +# --------------------------------------------------------------------------- + + +def _patch_settings(monkeypatch, **overrides): + """Override `load_settings` returns to a tweaked AppSettings.""" + from backend.apps.settings.models import AppSettings + s = AppSettings(**overrides) + monkeypatch.setattr("backend.apps.settings.settings.load_settings", + lambda: s) + + +def _patch_9router(monkeypatch, *, running: bool, providers: list[dict]): + monkeypatch.setattr("backend.apps.nine_router.is_running", lambda: running) + + async def _fake_get_providers(): + return providers + + monkeypatch.setattr("backend.apps.nine_router.get_providers", _fake_get_providers) + + +def test_models_no_creds_no_9router_returns_empty_anthropic(client, monkeypatch): + """No api keys, no 9Router → no Anthropic group at all (own_key + branch only emits when has_api_key OR has_claude_sub).""" + _patch_settings(monkeypatch) + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + assert r.status_code == 200 + models = r.json()["models"] + assert "Anthropic" not in models + assert "OpenSwarm Pro" not in models + + +def test_models_anthropic_api_key_only(client, monkeypatch): + """anthropic_api_key set, own_key mode → adaptive Anthropic group + surfaces under 'Anthropic'.""" + _patch_settings(monkeypatch, anthropic_api_key="sk-test") + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "Anthropic" in models + assert any(m["value"] == "sonnet" for m in models["Anthropic"]) + assert "OpenSwarm Pro" not in models + + +def test_models_openswarm_pro_only(client, monkeypatch): + """openswarm-pro + bearer, no claude sub → only 'OpenSwarm Pro' group.""" + _patch_settings(monkeypatch, + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x") + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "OpenSwarm Pro" in models + assert "Anthropic" not in models + # Adaptive variants only (no -cc / -api suffix) + values = {m["value"] for m in models["OpenSwarm Pro"]} + assert "sonnet" in values + assert "sonnet-cc" not in values + assert "sonnet-api" not in values + + +def test_models_openswarm_pro_plus_claude_sub_emits_both(client, monkeypatch): + """openswarm-pro + 9Router claude sub → BOTH 'OpenSwarm Pro' (adaptive) + AND 'Anthropic' (-cc variants for personal sub routing).""" + _patch_settings(monkeypatch, + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x") + _patch_9router(monkeypatch, running=True, + providers=[{"provider": "claude", "isActive": True}]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "OpenSwarm Pro" in models + assert "Anthropic" in models + # The Anthropic group uses -cc variants + cc_values = {m["value"] for m in models["Anthropic"]} + assert "sonnet-cc" in cc_values + + +def test_models_subscription_only_models_gated_by_9router(client, monkeypatch): + """OpenAI/Codex models are subscription_only — they only surface + when 9Router has the codex provider connected.""" + _patch_settings(monkeypatch, anthropic_api_key="sk-test") + _patch_9router(monkeypatch, running=True, + providers=[{"provider": "codex", "isActive": True}]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "OpenAI" in models + assert any(m["value"] == "gpt-5.4" for m in models["OpenAI"]) + + +def test_models_openai_api_key_surfaces_pinned_api_variants(client, monkeypatch): + """openai_api_key set → -api variants surface under 'OpenAI' group + even without a Codex subscription.""" + _patch_settings(monkeypatch, + anthropic_api_key="sk-test", + openai_api_key="sk-openai") + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "OpenAI" in models + values = {m["value"] for m in models["OpenAI"]} + assert "gpt-5.4-api" in values + assert "gpt-5.4" not in values # subscription one is hidden + + +def test_models_google_api_key_surfaces_pinned_api_variants(client, monkeypatch): + _patch_settings(monkeypatch, + anthropic_api_key="sk-test", + google_api_key="AIza-test") + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + models = r.json()["models"] + assert "Google" in models + values = {m["value"] for m in models["Google"]} + assert "gemini-3-pro-api" in values + assert "gemini-3-pro" not in values # subscription-only hidden + + +def test_models_response_shape_includes_reasoning_and_context(client, monkeypatch): + _patch_settings(monkeypatch, anthropic_api_key="sk-test") + _patch_9router(monkeypatch, running=False, providers=[]) + r = client.get("/api/agents/models") + body = r.json() + assert "models" in body and "notes" in body + sonnet = next(m for m in body["models"]["Anthropic"] if m["value"] == "sonnet") + assert sonnet["context_window"] == 1_000_000 + assert sonnet["reasoning"] is True + assert "label" in sonnet + + +def test_models_9router_provider_fetch_failure_falls_back_to_unconnected(client, monkeypatch): + """If 9Router probe raises, log + treat as no providers connected.""" + _patch_settings(monkeypatch, + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x") + + monkeypatch.setattr("backend.apps.nine_router.is_running", lambda: True) + + async def _broken_get_providers(): + raise RuntimeError("9Router exploded") + + monkeypatch.setattr("backend.apps.nine_router.get_providers", _broken_get_providers) + + r = client.get("/api/agents/models") + models = r.json()["models"] + # Only the Pro group remains (claude sub treated as missing). + assert "OpenSwarm Pro" in models + assert "Anthropic" not in models + + +# --------------------------------------------------------------------------- +# Smoke: existing endpoints still work +# --------------------------------------------------------------------------- + + +def test_warm_cache_unknown_session_still_returns_ok(client): + """warm-cache is best-effort. Even if the session doesn't exist + (manager will raise ValueError), the endpoint returns 200.""" + r = client.post("/api/agents/sessions/does-not-exist/warm-cache") + assert r.status_code == 200 diff --git a/backend/tests/test_anthropic_proxy.py b/backend/tests/test_anthropic_proxy.py new file mode 100644 index 00000000..9874380d --- /dev/null +++ b/backend/tests/test_anthropic_proxy.py @@ -0,0 +1,411 @@ +"""Tests for `backend.apps.agents.anthropic_proxy`. + +The proxy splits Anthropic-format traffic between the OpenSwarm Pro +cloud proxy (for Claude models on a Pro subscription) and 9Router +(everything else). + +Coverage targets: + - `_is_claude_model` parameterized over expected matches/non-matches + - `_pick_upstream`: + - openswarm-pro + Claude → Pro proxy with bearer token + - non-Claude model → 9Router with `x-api-key: 9router` + - Claude + own_key → 9Router (Pro fallthrough) + - `_healthcheck` returns 200 + - non-streaming proxy: + - body round-trips + - JSON content-type bodies parsed; non-JSON wrapped as `{"raw": ...}` + - hop headers / x-api-key / authorization stripped before forward + - timeout → 504, generic exception → 502 + - streaming proxy: + - returns StreamingResponse with chunks from upstream + - `stream:true` honored +""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from backend.apps.agents import anthropic_proxy as proxy_mod +from backend.apps.agents.anthropic_proxy import ( + _HOP_HEADERS, + _is_claude_model, + _pick_upstream, +) +from backend.apps.settings.models import AppSettings + + +# --------------------------------------------------------------------------- +# _is_claude_model +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "model,expected", + [ + ("claude-sonnet-4-6", True), + ("claude/claude-3", True), + ("claude-opus-4-6", True), + ("sonnet", True), + ("opus", True), + ("haiku", True), + ("cc/claude-sonnet-4-6", True), + ("CLAUDE-haiku-4-5", True), # case insensitive + (" sonnet ", True), # whitespace trimmed + ("cx/gpt-5.4", False), + ("gc/gemini-3-pro-preview", False), + ("gpt-5.4", False), + ("gemini-2.5-pro", False), + ("", False), + ], +) +def test_is_claude_model(model: str, expected: bool): + assert _is_claude_model(model) is expected + + +# --------------------------------------------------------------------------- +# _pick_upstream +# --------------------------------------------------------------------------- + + +def test_pick_upstream_claude_with_openswarm_pro_returns_proxy_with_bearer(): + s = AppSettings( + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x", + openswarm_proxy_url="https://api.openswarm.com", + ) + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, headers = _pick_upstream("claude-sonnet-4-6") + assert base == "https://api.openswarm.com" + assert headers == {"Authorization": "Bearer bearer-x"} + + +def test_pick_upstream_claude_pro_strips_trailing_slash(): + s = AppSettings( + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x", + openswarm_proxy_url="https://api.openswarm.com/", + ) + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, _ = _pick_upstream("sonnet") + assert base == "https://api.openswarm.com" # no trailing / + + +def test_pick_upstream_non_claude_returns_9router(): + s = AppSettings(connection_mode="openswarm-pro", openswarm_bearer_token="bearer-x") + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, headers = _pick_upstream("cx/gpt-5.4") + assert base == "http://127.0.0.1:20128" + assert headers == {"x-api-key": "9router"} + + +def test_pick_upstream_claude_own_key_falls_back_to_9router(): + """When connection_mode != openswarm-pro, a Claude model still routes + through 9Router (the user might have a Claude subscription wired).""" + s = AppSettings(connection_mode="own_key", anthropic_api_key="sk-foo") + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, headers = _pick_upstream("claude-sonnet-4-6") + assert base == "http://127.0.0.1:20128" + assert headers == {"x-api-key": "9router"} + + +def test_pick_upstream_pro_without_bearer_falls_back_to_9router(): + """openswarm-pro mode but no bearer token → can't reach the Pro proxy, + fall through to 9Router.""" + s = AppSettings(connection_mode="openswarm-pro") # no bearer + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, headers = _pick_upstream("claude-sonnet-4-6") + assert base == "http://127.0.0.1:20128" + assert headers == {"x-api-key": "9router"} + + +def test_pick_upstream_pro_default_proxy_url(): + """Empty `openswarm_proxy_url` defaults to https://api.openswarm.com.""" + s = AppSettings( + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x", + ) + with patch("backend.apps.settings.settings.load_settings", return_value=s): + base, _ = _pick_upstream("claude-sonnet-4-6") + assert base == "https://api.openswarm.com" + + +# --------------------------------------------------------------------------- +# Routes +# --------------------------------------------------------------------------- + + +def test_healthcheck_via_test_client(client): + """GET on the proxy root must return 200 (CLI healthcheck path).""" + r = client.get("/api/anthropic-proxy") + assert r.status_code in (200, 307, 308) + if r.status_code == 200: + assert r.json() == {"ok": True} + + +def test_healthcheck_via_test_client_trailing_slash(client): + r = client.get("/api/anthropic-proxy/") + assert r.status_code == 200 + assert r.json() == {"ok": True} + + +# --------------------------------------------------------------------------- +# Non-streaming proxy +# --------------------------------------------------------------------------- + + +def _make_async_client_mock(response): + """Build a context-managed async client whose `request` returns + `response`. Mirrors `httpx.AsyncClient(...)` ergonomics.""" + inst = MagicMock() + inst.request = AsyncMock(return_value=response) + inst.__aenter__ = AsyncMock(return_value=inst) + inst.__aexit__ = AsyncMock(return_value=False) + return inst + + +def test_proxy_non_streaming_routes_claude_to_pro_proxy(client): + """Claude model + Pro mode → upstream URL is the Pro proxy with + bearer auth. Body and headers round-trip; x-api-key is stripped.""" + s = AppSettings( + connection_mode="openswarm-pro", + openswarm_bearer_token="bearer-x", + openswarm_proxy_url="https://api.openswarm.com", + ) + upstream_resp = httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=b'{"id": "msg_1", "model": "claude-sonnet-4-6"}', + ) + captured: dict[str, Any] = {} + + async def fake_request(method, url, content=None, headers=None, params=None): + captured["method"] = method + captured["url"] = url + captured["body"] = content + captured["headers"] = headers or {} + return upstream_resp + + fake_client = MagicMock() + fake_client.request = AsyncMock(side_effect=fake_request) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "claude-sonnet-4-6", "messages": [{"role": "user", "content": "hi"}]}, + headers={"x-api-key": "should-not-leak", "x-extra": "passthrough"}, + ) + + assert r.status_code == 200 + assert r.json() == {"id": "msg_1", "model": "claude-sonnet-4-6"} + # Routed to Pro proxy + assert captured["url"] == "https://api.openswarm.com/v1/messages" + # Bearer auth attached + assert captured["headers"].get("Authorization") == "Bearer bearer-x" + # x-api-key NEVER reaches upstream + keys = {k.lower() for k in captured["headers"].keys()} + assert "x-api-key" not in keys + # Hop headers stripped + for hop in ("host", "content-length", "connection"): + assert hop not in keys + # Custom headers passed through + assert captured["headers"].get("x-extra") == "passthrough" + # Body round-tripped + assert json.loads(captured["body"]) == { + "model": "claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hi"}], + } + + +def test_proxy_non_streaming_routes_non_claude_to_9router(client): + """Non-Claude model → 9Router with x-api-key=9router header.""" + s = AppSettings(connection_mode="openswarm-pro", openswarm_bearer_token="bearer-x") + upstream_resp = httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=b'{"id": "msg_2"}', + ) + captured: dict[str, Any] = {} + + async def fake_request(method, url, content=None, headers=None, params=None): + captured["url"] = url + captured["headers"] = headers or {} + return upstream_resp + + fake_client = MagicMock() + fake_client.request = AsyncMock(side_effect=fake_request) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4"}, + ) + + assert r.status_code == 200 + assert captured["url"] == "http://127.0.0.1:20128/v1/messages" + assert captured["headers"].get("x-api-key") == "9router" + assert "Authorization" not in captured["headers"] + + +def test_proxy_non_streaming_non_json_body_wrapped_in_raw(client): + """Upstream returning text/plain → body wrapped as {"raw": "..."}.""" + s = AppSettings() + upstream_resp = httpx.Response( + status_code=200, + headers={"content-type": "text/plain"}, + content=b"hello", + ) + fake_client = MagicMock() + fake_client.request = AsyncMock(return_value=upstream_resp) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4"}, + ) + assert r.status_code == 200 + assert r.json() == {"raw": "hello"} + + +def test_proxy_non_streaming_timeout_returns_504(client): + s = AppSettings() + fake_client = MagicMock() + fake_client.request = AsyncMock(side_effect=httpx.TimeoutException("timed out")) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4"}, + ) + assert r.status_code == 504 + assert r.json()["error"] == "upstream timeout" + + +def test_proxy_non_streaming_generic_exception_returns_502(client): + s = AppSettings() + fake_client = MagicMock() + fake_client.request = AsyncMock(side_effect=RuntimeError("kaboom")) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4"}, + ) + assert r.status_code == 502 + assert "kaboom" in r.json()["error"] + + +def test_proxy_non_streaming_non_json_body_doesnt_break_routing(client): + """If the body isn't valid JSON, model is "" → routes to 9Router + via the non-Claude branch. Must not crash.""" + s = AppSettings() + upstream_resp = httpx.Response(status_code=200, content=b"{}", + headers={"content-type": "application/json"}) + fake_client = MagicMock() + fake_client.request = AsyncMock(return_value=upstream_resp) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + content=b"not-json-at-all", + headers={"content-type": "application/json"}, + ) + assert r.status_code == 200 + + +# --------------------------------------------------------------------------- +# Streaming proxy +# --------------------------------------------------------------------------- + + +def test_proxy_streaming_returns_chunks(client): + """`stream:true` body → StreamingResponse with raw chunks from + upstream's aiter_raw.""" + s = AppSettings() + + upstream = MagicMock() + upstream.status_code = 200 + upstream.headers = {"content-type": "text/event-stream"} + + async def aiter_raw(): + yield b"event: message\n" + yield b"data: hello\n\n" + + upstream.aiter_raw = aiter_raw + upstream.aclose = AsyncMock() + + fake_client = MagicMock() + fake_client.build_request = MagicMock(return_value=MagicMock()) + fake_client.send = AsyncMock(return_value=upstream) + fake_client.aclose = AsyncMock() + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4", "stream": True}, + ) + + assert r.status_code == 200 + body = r.content + assert b"event: message" in body + assert b"data: hello" in body + fake_client.send.assert_awaited_once() + + +def test_proxy_streaming_stream_false_takes_non_streaming_path(client): + """Body with `stream:false` → goes through the non-streaming branch + (no StreamingResponse).""" + s = AppSettings() + upstream_resp = httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=b'{"ok": true}', + ) + fake_client = MagicMock() + fake_client.request = AsyncMock(return_value=upstream_resp) + fake_client.send = AsyncMock(side_effect=AssertionError("send must NOT be called")) + fake_client.__aenter__ = AsyncMock(return_value=fake_client) + fake_client.__aexit__ = AsyncMock(return_value=False) + + with patch("backend.apps.settings.settings.load_settings", return_value=s), \ + patch.object(proxy_mod.httpx, "AsyncClient", return_value=fake_client): + r = client.post( + "/api/anthropic-proxy/v1/messages", + json={"model": "cx/gpt-5.4", "stream": False}, + ) + assert r.status_code == 200 + + +# --------------------------------------------------------------------------- +# Constants sanity +# --------------------------------------------------------------------------- + + +def test_hop_headers_lowercase(): + """Hop-by-hop headers must be lowercased so the comparison in + proxy() works (`k.lower() in _HOP_HEADERS`).""" + for h in _HOP_HEADERS: + assert h == h.lower() diff --git a/backend/tests/test_mcp_preflight.py b/backend/tests/test_mcp_preflight.py new file mode 100644 index 00000000..dcb5d197 --- /dev/null +++ b/backend/tests/test_mcp_preflight.py @@ -0,0 +1,391 @@ +"""Tests for `backend.apps.agents.mcp_preflight`. + +Currently 0% covered. Drives the public entry point `run_preflight` and +the helpers behind it. The aux-model call is mocked at the function +boundary (`_call_classifier`); no real Anthropic / 9Router traffic. + +Coverage targets: + - `_is_obviously_local`: short / shell-prefixed / single-path / + normal prompts + - `_build_available_shortlist`: enabled tools removed, dismissed + entries removed + - `_decorate`: known id → full payload, unknown id → None, + `reason` truncated to 200 chars + - `run_preflight`: + - empty prompt → default + - obviously-local prompt → default (no LLM call) + - happy path → returns classifier JSON, suggestions decorated + - is_vague=False zeros suggestions (concrete-prompt guard) + - 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 + +import asyncio +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +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, +) +from backend.apps.settings.models import AppSettings + + +# --------------------------------------------------------------------------- +# _is_obviously_local +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "prompt,expected", + [ + ("", True), # < 8 chars + ("hi", True), + ("ok thx", True), + ("$ git status", True), # shell prefix + ("! ls", True), + ("/clear", True), + ("./src/foo.ts", True), # single path-like token + ("README.md", True), # extension match + ("/tmp/whatever.json", True), + ("write me an email about the demo", False), + ("summarize my notes from the meeting yesterday", False), + ("./src/foo.ts has a bug, fix it", False), # multi-token + ], +) +def test_is_obviously_local(prompt: str, expected: bool): + assert _is_obviously_local(prompt) is expected + + +# --------------------------------------------------------------------------- +# _build_available_shortlist +# --------------------------------------------------------------------------- + + +def _fake_tool(name: str, *, enabled: bool = True) -> object: + return SimpleNamespace(name=name, enabled=enabled) + + +def test_build_available_shortlist_enabled_tools_removed(): + """If a curated entry's `id` matches an enabled tool, it must NOT + appear in the shortlist (already connected → nothing to suggest).""" + settings = AppSettings() + with patch.object(pf, "load_all_tools", return_value=[ + _fake_tool("Slack", enabled=True), + _fake_tool("Notion", enabled=False), # disabled doesn't filter + ]): + out = _build_available_shortlist(settings) + ids = {e["id"] for e in out} + assert "Slack" not in ids + assert "Notion" in ids + + +def test_build_available_shortlist_dismissed_entries_filtered(): + """User-dismissed suggestions are suppressed on subsequent launches.""" + settings = AppSettings(dismissed_mcp_suggestions={"Reddit": "2026-04-30T00:00:00"}) + with patch.object(pf, "load_all_tools", return_value=[]): + out = _build_available_shortlist(settings) + ids = {e["id"] for e in out} + assert "Reddit" not in ids + + +def test_build_available_shortlist_handles_load_tools_exception(): + """If load_all_tools raises, the helper falls back to "no enabled + tools" rather than crashing.""" + settings = AppSettings() + with patch.object(pf, "load_all_tools", side_effect=RuntimeError("disk gone")): + out = _build_available_shortlist(settings) + # No enabled / no dismissed → entire curated shortlist returned + assert len(out) == len(CURATED_SHORTLIST) + + +# --------------------------------------------------------------------------- +# _decorate +# --------------------------------------------------------------------------- + + +def test_decorate_known_id_returns_full_shape(): + available = list(CURATED_SHORTLIST) + out = _decorate({"id": "Slack", "reason": "user mentioned channel"}, available) + assert out is not None + assert out["id"] == "Slack" + assert out["title"] == "Slack" + assert "Search channels" in out["description"] + assert out["reason"] == "user mentioned channel" + + +def test_decorate_unknown_id_returns_none(): + out = _decorate({"id": "DefinitelyNotReal", "reason": "x"}, list(CURATED_SHORTLIST)) + assert out is None + + +def test_decorate_truncates_reason_to_200_chars(): + long_reason = "x" * 500 + out = _decorate({"id": "Slack", "reason": long_reason}, list(CURATED_SHORTLIST)) + assert out is not None + assert len(out["reason"]) == 200 + + +def test_decorate_missing_reason_becomes_empty_string(): + out = _decorate({"id": "Slack"}, list(CURATED_SHORTLIST)) + assert out is not None + assert out["reason"] == "" + + +# --------------------------------------------------------------------------- +# run_preflight — early returns +# --------------------------------------------------------------------------- + + +async def test_run_preflight_empty_prompt_returns_default(): + out = await run_preflight("") + assert out == {"is_vague": False, "suggestions": []} + + +async def test_run_preflight_whitespace_only_returns_default(): + out = await run_preflight(" \n\t ") + assert out == {"is_vague": False, "suggestions": []} + + +async def test_run_preflight_obviously_local_skips_classifier(): + """Local prompts must short-circuit BEFORE any classifier call.""" + classifier = AsyncMock() + with patch.object(pf, "_call_classifier", classifier): + out = await run_preflight("./src/foo.ts") + assert out == {"is_vague": False, "suggestions": []} + classifier.assert_not_called() + + +# --------------------------------------------------------------------------- +# run_preflight — happy path +# --------------------------------------------------------------------------- + + +async def test_run_preflight_happy_path_decorates_suggestions(): + classifier_result = { + "is_vague": True, + "suggestions": [ + {"id": "Slack", "reason": "user mentioned channel"}, + {"id": "Notion", "reason": "wants to update notes"}, + ], + } + with patch.object(pf, "_call_classifier", AsyncMock(return_value=classifier_result)): + out = await run_preflight("send a status update to the team channel") + + assert out["is_vague"] is True + ids = {s["id"] for s in out["suggestions"]} + assert ids == {"Slack", "Notion"} + # Decorated to full shape + slack = next(s for s in out["suggestions"] if s["id"] == "Slack") + assert slack["title"] == "Slack" + assert slack["reason"] == "user mentioned channel" + + +async def test_run_preflight_concrete_prompt_zeros_suggestions(): + """is_vague=False MUST suppress all suggestions, even if the + classifier returned some — concrete tasks shouldn't be interrupted + with a connect-mcp modal.""" + classifier_result = { + "is_vague": False, + "suggestions": [{"id": "Slack", "reason": "x"}], + } + with patch.object(pf, "_call_classifier", AsyncMock(return_value=classifier_result)): + out = await run_preflight("refactor foo.ts to use async/await") + assert out["is_vague"] is False + assert out["suggestions"] == [] + + +async def test_run_preflight_drops_hallucinated_ids(): + """The classifier may invent an id — preflight must filter against + `CURATED_SHORTLIST` so the frontend never sees a phantom.""" + classifier_result = { + "is_vague": True, + "suggestions": [ + {"id": "Slack", "reason": "ok"}, + {"id": "PhantomService", "reason": "made up"}, + {"id": "AnotherFake", "reason": "also made up"}, + ], + } + with patch.object(pf, "_call_classifier", AsyncMock(return_value=classifier_result)): + out = await run_preflight("write me an email summarizing the call") + + ids = {s["id"] for s in out["suggestions"]} + assert "PhantomService" not in ids + assert "AnotherFake" not in ids + assert "Slack" in ids + + +async def test_run_preflight_drops_already_enabled_after_classifier(): + """If the user enables an MCP between preflight start and classifier + return, the suggestion should be dropped (matches `available` is + None in `_decorate`).""" + settings = AppSettings() + classifier_result = { + "is_vague": True, + "suggestions": [ + {"id": "Slack", "reason": "channel"}, # will be enabled (filtered out) + {"id": "Notion", "reason": "notes"}, + ], + } + + with patch.object(pf, "load_all_tools", return_value=[ + _fake_tool("Slack", enabled=True), # enabled mid-flight + ]), patch.object(pf, "load_settings", return_value=settings), \ + patch.object(pf, "_call_classifier", AsyncMock(return_value=classifier_result)): + out = await run_preflight("ping the team in our channel and update the doc") + + ids = {s["id"] for s in out["suggestions"]} + assert "Slack" not in ids + assert "Notion" in ids + + +async def test_run_preflight_non_dict_suggestion_filtered(): + """Defensive: classifier might return non-dict items in the + suggestions list (e.g. a bare string). They must be silently + dropped, not crash decoration.""" + classifier_result = { + "is_vague": True, + "suggestions": [ + "not a dict", + {"id": "Slack", "reason": "real one"}, + ], + } + with patch.object(pf, "_call_classifier", AsyncMock(return_value=classifier_result)): + out = await run_preflight("send a status update to the team channel") + assert [s["id"] for s in out["suggestions"]] == ["Slack"] + + +# --------------------------------------------------------------------------- +# run_preflight — fail-open contract +# --------------------------------------------------------------------------- + + +async def test_run_preflight_classifier_timeout_returns_default(): + """asyncio.TimeoutError → default. Real path: aux model is slow.""" + async def _slow(*_args, **_kw): + await asyncio.sleep(10) + return {} + + with patch.object(pf, "_call_classifier", _slow): + out = await run_preflight("write me an email summarizing the call", timeout_s=0.05) + assert out == {"is_vague": False, "suggestions": []} + + +async def test_run_preflight_classifier_exception_returns_default(): + """Any other exception (network, bad JSON, ValueError) must fail + open with the default response.""" + with patch.object(pf, "_call_classifier", AsyncMock(side_effect=RuntimeError("boom"))): + out = await run_preflight("write me an email summarizing the call") + assert out == {"is_vague": False, "suggestions": []} + + +async def test_run_preflight_no_provider_classifier_value_error_returns_default(): + """resolve_aux_model raises ValueError when no provider is wired — + that surfaces inside the classifier and must fail open.""" + with patch.object(pf, "_call_classifier", + AsyncMock(side_effect=ValueError("no provider"))): + out = await run_preflight("send a status update to the team channel") + 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"] diff --git a/backend/tests/test_mcp_servers_unit.py b/backend/tests/test_mcp_servers_unit.py new file mode 100644 index 00000000..c421c597 --- /dev/null +++ b/backend/tests/test_mcp_servers_unit.py @@ -0,0 +1,623 @@ +"""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"} diff --git a/backend/tests/test_providers_anthropic_extra.py b/backend/tests/test_providers_anthropic_extra.py new file mode 100644 index 00000000..e730df2b --- /dev/null +++ b/backend/tests/test_providers_anthropic_extra.py @@ -0,0 +1,530 @@ +"""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 == {} diff --git a/backend/tests/test_providers_openai_compat.py b/backend/tests/test_providers_openai_compat.py new file mode 100644 index 00000000..a38901c0 --- /dev/null +++ b/backend/tests/test_providers_openai_compat.py @@ -0,0 +1,676 @@ +"""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} diff --git a/backend/tests/test_providers_registry.py b/backend/tests/test_providers_registry.py new file mode 100644 index 00000000..3d240835 --- /dev/null +++ b/backend/tests/test_providers_registry.py @@ -0,0 +1,616 @@ +"""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/` + - gemini-cli + Antigravity active (mock httpx 200) → `ag/` + - gemini-cli fallthrough → `gc/` + - 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/.""" + 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 diff --git a/backend/tests/test_tools_unit.py b/backend/tests/test_tools_unit.py new file mode 100644 index 00000000..3f132e4c --- /dev/null +++ b/backend/tests/test_tools_unit.py @@ -0,0 +1,696 @@ +"""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'' + ) + 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="") + 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 = "

Hello world

" + 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 "