"""Unit tests for the keyed web-search providers (#92): Tavily, Brave, SerpAPI, Exa — plus provider ordering and the transient retry.""" import httpx import pytest import backend.tools.web as web from backend.tools.context import ToolContext, ToolError, ToolMode class FakeResponse: def __init__(self, json_data=None, status_code=200, content=b""): self._json = json_data self.status_code = status_code self.content = content self.encoding = "utf-8" self.headers = {} def json(self): return self._json def raise_for_status(self): if self.status_code >= 400: raise web.httpx.HTTPStatusError("boom", request=None, response=self) # type: ignore[arg-type] def _ctx() -> ToolContext: return ToolContext(user={"username": "tester", "vaults": []}, mode=ToolMode.IN_APP) def _no_fallback(monkeypatch): """Limit the chain to the provider under test (no searxng/ddg/bing noise).""" monkeypatch.setattr(web, "WEB_FALLBACK_ENABLED", False) monkeypatch.setattr(web, "SEARXNG_URL", "http://searxng.invalid") monkeypatch.setattr(web.httpx, "get", lambda *a, **kw: (_ for _ in ()).throw( httpx.ConnectError("offline"))) monkeypatch.setattr(web.httpx, "post", lambda *a, **kw: (_ for _ in ()).throw( httpx.ConnectError("offline"))) class TestKeyedProviderParsers: def test_tavily_maps_results(self, monkeypatch): captured = {} def fake_post(url, json=None, **kw): captured["url"] = url captured["payload"] = json return FakeResponse(json_data={"results": [ {"title": "T", "url": "https://a.dev", "content": "c" * 800}, ]}) monkeypatch.setattr(web.httpx, "post", fake_post) results, engines = web._search_tavily("q", web.WebSearchInput(query="q", max_results=3)) assert results[0]["title"] == "T" assert len(results[0]["snippet"]) <= 600 assert engines == [] assert captured["payload"]["api_key"] == "" assert captured["payload"]["max_results"] == 3 def test_brave_maps_results(self, monkeypatch): captured = {} def fake_get(url, params=None, headers=None, **kw): captured["url"] = url captured["headers"] = headers return FakeResponse(json_data={"web": {"results": [ {"title": "B", "url": "https://b.dev", "description": "d"}, ]}}) monkeypatch.setattr(web.httpx, "get", fake_get) results, _ = web._search_brave("q", web.WebSearchInput(query="q")) assert results[0]["title"] == "B" assert "api.search.brave.com" in str(captured["url"]) def test_serpapi_maps_results(self, monkeypatch): monkeypatch.setattr(web.httpx, "get", lambda *a, **kw: FakeResponse(json_data={ "organic_results": [{"title": "S", "link": "https://s.dev", "snippet": "sn"}], })) results, _ = web._search_serpapi("q", web.WebSearchInput(query="q")) assert results[0]["url"] == "https://s.dev" def test_exa_maps_results(self, monkeypatch): monkeypatch.setattr(web.httpx, "post", lambda *a, **kw: FakeResponse(json_data={ "results": [{"title": "E", "url": "https://e.dev", "text": "t" * 900}], })) results, _ = web._search_exa("q", web.WebSearchInput(query="q")) assert results[0]["title"] == "E" assert len(results[0]["snippet"]) <= 600 class TestProviderChain: def test_no_key_falls_back_to_searxng(self, monkeypatch): for var in ("OBSIGATE_TAVILY_API_KEY", "OBSIGATE_BRAVE_API_KEY", "OBSIGATE_SERPAPI_API_KEY", "OBSIGATE_EXA_API_KEY"): monkeypatch.delenv(var, raising=False) chain = [name for name, _ in web._provider_chain()] assert chain[0] == "searxng" assert "tavily" not in chain def test_keyed_provider_used_first_when_key_set(self, monkeypatch): monkeypatch.setenv("OBSIGATE_TAVILY_API_KEY", "k") chain = [name for name, _ in web._provider_chain()] assert chain[0] == "tavily" assert "brave" not in chain # no key → skipped def test_explicit_order_env(self, monkeypatch): monkeypatch.setenv("OBSIGATE_TAVILY_API_KEY", "k") monkeypatch.setenv("OBSIGATE_EXA_API_KEY", "k") monkeypatch.setenv("OBSIGATE_WEB_PROVIDERS", "exa,unknown,tavily") chain = [name for name, _ in web._provider_chain()] assert chain[:2] == ["exa", "tavily"] def test_search_uses_keyed_provider_first(self, monkeypatch): _no_fallback(monkeypatch) monkeypatch.setenv("OBSIGATE_BRAVE_API_KEY", "k") monkeypatch.setattr(web.httpx, "get", lambda *a, **kw: FakeResponse(json_data={ "web": {"results": [{"title": "B", "url": "https://b.dev", "description": "d"}]}, })) out = web.web_search(_ctx(), web.WebSearchInput(query="q")) assert out["provider"] == "brave" assert out["count"] == 1 class TestRetry: def test_transient_error_retried_then_succeeds(self, monkeypatch): monkeypatch.setattr(web, "WEB_RETRY_ATTEMPTS", 1) monkeypatch.setattr(web, "_provider_chain", lambda: [("searxng", web._search_searxng)]) calls = {"n": 0} def flaky_get(*a, **kw): calls["n"] += 1 if calls["n"] == 1: raise httpx.ConnectError("blip") return FakeResponse(json_data={"results": [ {"title": "A", "url": "https://a.dev", "content": "x"}]}) monkeypatch.setattr(web.httpx, "get", flaky_get) out = web.web_search(_ctx(), web.WebSearchInput(query="q")) assert out["provider"] == "searxng" assert calls["n"] == 2 def test_persistent_error_not_retried_forever(self, monkeypatch): monkeypatch.setattr(web, "WEB_RETRY_ATTEMPTS", 1) monkeypatch.setattr(web, "_provider_chain", lambda: [("searxng", web._search_searxng)]) calls = {"n": 0} def dead_get(*a, **kw): calls["n"] += 1 raise httpx.ConnectError("down") monkeypatch.setattr(web.httpx, "get", dead_get) with pytest.raises(ToolError) as ei: web.web_search(_ctx(), web.WebSearchInput(query="q")) assert ei.value.code == "web_search_unavailable" assert calls["n"] == 2 # initial + 1 retry, per provider