156 lines
6.3 KiB
Python
156 lines
6.3 KiB
Python
"""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
|