mirror of
https://github.com/usestrix/strix.git
synced 2026-09-07 08:25:56 +00:00
feat(web_search): add Exa as a web search provider alongside Perplexity (#1270)
This commit is contained in:
parent
f6d9790ecb
commit
afa7c4a77f
9 changed files with 785 additions and 63 deletions
|
|
@ -273,7 +273,6 @@ export LLM_API_KEY="your-api-key"
|
|||
|
||||
# Optional
|
||||
export LLM_API_BASE="your-api-base-url" # if using a local model, e.g. Ollama, LMStudio
|
||||
export PERPLEXITY_API_KEY="your-api-key" # for search capabilities
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
|
|
|
|||
|
|
@ -80,6 +80,22 @@ affecting the agents that do the actual testing.
|
|||
API key for Perplexity AI. Enables real-time web search during scans for OSINT and vulnerability research.
|
||||
</ParamField>
|
||||
|
||||
<ParamField path="EXA_API_KEY" type="string">
|
||||
API key for Exa. Enables real-time web search through the Exa `/search` endpoint. Exa also powers the `web_get_contents` tool, which fetches the full text of a page through the Exa `/contents` endpoint. This is the preferred web search provider.
|
||||
</ParamField>
|
||||
|
||||
<ParamField path="STRIX_WEB_SEARCH_PROVIDER" default="auto" type="string">
|
||||
Web search provider: `auto`, `perplexity`, or `exa`. With `auto`, Strix uses Exa when `EXA_API_KEY` is set, and Perplexity otherwise. Set an explicit provider to pin one when you configure both keys.
|
||||
</ParamField>
|
||||
|
||||
<ParamField path="STRIX_EXA_SEARCH_TYPE" default="auto" type="string">
|
||||
Exa search mode: `auto`, `fast`, `instant`, `deep-lite`, `deep`, or `deep-reasoning`. Lower modes return results faster. Higher modes plan across more steps and take more time. This setting applies only to the Exa provider.
|
||||
</ParamField>
|
||||
|
||||
<ParamField path="STRIX_EXA_NUM_RESULTS" default="5" type="integer">
|
||||
Number of Exa results to return, from `1` to `100`. Each result includes a title, a URL, and a short security-focused summary. To read a full page, the agent calls `web_get_contents` with the result URL. This setting applies only to the Exa provider.
|
||||
</ParamField>
|
||||
|
||||
<ParamField path="POSTMAN_API_KEY" type="string">
|
||||
Postman API key (`PMAK-…`). Enables fetching Postman collections by id as a target (`postman://<collection-uid>`), and Postman environments (`postman://<collection-uid>?env=<environment-uid>`) to resolve collection variables. Not needed when passing a local collection export file.
|
||||
</ParamField>
|
||||
|
|
@ -159,7 +175,8 @@ strix --target ./app --config /path/to/config.json
|
|||
export STRIX_LLM="openrouter/z-ai/glm-5.3"
|
||||
export LLM_API_KEY="sk-..."
|
||||
|
||||
# Optional: Enable web search
|
||||
# Optional: Enable web search (Exa preferred, Perplexity supported)
|
||||
export EXA_API_KEY="..."
|
||||
export PERPLEXITY_API_KEY="pplx-..."
|
||||
|
||||
# Optional: Custom timeouts
|
||||
|
|
|
|||
|
|
@ -28,6 +28,6 @@ Strix agents use specialized tools to test your applications like a real penetra
|
|||
| -------------- | ---------------------------------------- |
|
||||
| Python Runtime | Write and execute custom exploit scripts |
|
||||
| File Editor | Read and modify source code |
|
||||
| Web Search | Real-time OSINT via Perplexity |
|
||||
| Web Search | Real-time OSINT with Exa or Perplexity |
|
||||
| Notes | Document findings during the scan |
|
||||
| Reporting | Generate vulnerability reports with PoCs |
|
||||
|
|
|
|||
|
|
@ -69,7 +69,7 @@ from strix.tools.todo.tools import (
|
|||
mark_todo_pending,
|
||||
update_todo,
|
||||
)
|
||||
from strix.tools.web_search.tool import web_search
|
||||
from strix.tools.web_search.tool import web_get_contents, web_search
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -579,6 +579,7 @@ _BASE_TOOLS: tuple[Tool, ...] = (
|
|||
save_threat_model,
|
||||
amend_threat_model,
|
||||
web_search,
|
||||
web_get_contents,
|
||||
create_vulnerability_report,
|
||||
create_dependency_report,
|
||||
update_vulnerability_report,
|
||||
|
|
|
|||
|
|
@ -120,6 +120,10 @@ class TelemetrySettings(BaseSettings):
|
|||
enabled: bool = Field(default=True, alias="STRIX_TELEMETRY")
|
||||
|
||||
|
||||
WebSearchProvider = Literal["auto", "perplexity", "exa"]
|
||||
ExaSearchType = Literal["auto", "fast", "instant", "deep-lite", "deep", "deep-reasoning"]
|
||||
|
||||
|
||||
class IntegrationSettings(BaseSettings):
|
||||
model_config = _BASE_CONFIG
|
||||
|
||||
|
|
@ -128,6 +132,25 @@ class IntegrationSettings(BaseSettings):
|
|||
alias="PERPLEXITY_API_KEY",
|
||||
repr=False,
|
||||
)
|
||||
exa_api_key: str | None = Field(
|
||||
default=None,
|
||||
alias="EXA_API_KEY",
|
||||
repr=False,
|
||||
)
|
||||
web_search_provider: WebSearchProvider = Field(
|
||||
default="auto",
|
||||
alias="STRIX_WEB_SEARCH_PROVIDER",
|
||||
)
|
||||
exa_search_type: ExaSearchType = Field(
|
||||
default="auto",
|
||||
alias="STRIX_EXA_SEARCH_TYPE",
|
||||
)
|
||||
exa_num_results: int = Field(
|
||||
default=5,
|
||||
ge=1,
|
||||
le=100,
|
||||
alias="STRIX_EXA_NUM_RESULTS",
|
||||
)
|
||||
postman_api_key: str | None = Field(
|
||||
default=None,
|
||||
alias="POSTMAN_API_KEY",
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from rich.console import Console
|
|||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from strix.config import codex, load_settings
|
||||
from strix.config import IntegrationSettings, codex, load_settings
|
||||
from strix.interface.utils import (
|
||||
check_docker_connection,
|
||||
image_exists,
|
||||
|
|
@ -19,6 +19,17 @@ from strix.interface.utils import (
|
|||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _missing_web_search_vars(integrations: IntegrationSettings) -> list[str]:
|
||||
"""Mirror the web_search provider rules: which key(s) the selected provider needs."""
|
||||
if integrations.web_search_provider == "exa":
|
||||
return [] if integrations.exa_api_key else ["EXA_API_KEY"]
|
||||
if integrations.web_search_provider == "perplexity":
|
||||
return [] if integrations.perplexity_api_key else ["PERPLEXITY_API_KEY"]
|
||||
if integrations.exa_api_key or integrations.perplexity_api_key:
|
||||
return []
|
||||
return ["EXA_API_KEY", "PERPLEXITY_API_KEY"]
|
||||
|
||||
|
||||
def validate_environment() -> None:
|
||||
logger.info("Validating environment")
|
||||
console = Console()
|
||||
|
|
@ -46,8 +57,7 @@ def validate_environment() -> None:
|
|||
if not settings.llm.api_base:
|
||||
missing_optional_vars.append("LLM_API_BASE")
|
||||
|
||||
if not settings.integrations.perplexity_api_key:
|
||||
missing_optional_vars.append("PERPLEXITY_API_KEY")
|
||||
missing_optional_vars.extend(_missing_web_search_vars(settings.integrations))
|
||||
|
||||
if missing_required_vars:
|
||||
error_text = Text()
|
||||
|
|
@ -89,7 +99,14 @@ def validate_environment() -> None:
|
|||
error_text.append("• ", style="white")
|
||||
error_text.append("PERPLEXITY_API_KEY", style="bold cyan")
|
||||
error_text.append(
|
||||
" - API key for Perplexity AI web search (enables real-time research)\n",
|
||||
" - API key for Perplexity AI web search (alternative to Exa)\n",
|
||||
style="white",
|
||||
)
|
||||
elif var == "EXA_API_KEY":
|
||||
error_text.append("• ", style="white")
|
||||
error_text.append("EXA_API_KEY", style="bold cyan")
|
||||
error_text.append(
|
||||
" - API key for Exa web search (enables real-time research)\n",
|
||||
style="white",
|
||||
)
|
||||
elif var == "STRIX_REASONING_EFFORT":
|
||||
|
|
@ -116,6 +133,8 @@ def validate_environment() -> None:
|
|||
error_text.append(
|
||||
"export PERPLEXITY_API_KEY='your-perplexity-key-here'\n", style="dim white"
|
||||
)
|
||||
elif var == "EXA_API_KEY":
|
||||
error_text.append("export EXA_API_KEY='your-exa-key-here'\n", style="dim white")
|
||||
elif var == "STRIX_REASONING_EFFORT":
|
||||
error_text.append(
|
||||
"export STRIX_REASONING_EFFORT='high'\n",
|
||||
|
|
|
|||
|
|
@ -1,11 +1,12 @@
|
|||
"""``web_search`` — Perplexity-backed security-focused web search."""
|
||||
"""Security-focused web research tools (Exa or Perplexity)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
import requests
|
||||
from agents import RunContextWrapper, function_tool
|
||||
|
|
@ -13,6 +14,10 @@ from agents import RunContextWrapper, function_tool
|
|||
from strix.config import load_settings
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
|
|
@ -41,22 +46,7 @@ Structure your response to be comprehensive yet concise, emphasizing the most cr
|
|||
security implications and details."""
|
||||
|
||||
|
||||
def _do_search(query: str) -> dict[str, Any]: # noqa: PLR0911 - each error class needs its own sanitized return
|
||||
if not query or not query.strip():
|
||||
return {"success": False, "error": "Query cannot be empty"}
|
||||
|
||||
api_key = load_settings().integrations.perplexity_api_key
|
||||
if not api_key:
|
||||
logger.warning("web_search invoked without PERPLEXITY_API_KEY configured")
|
||||
return {
|
||||
"success": False,
|
||||
"error": (
|
||||
"Web search is not configured for this scan "
|
||||
"(operator needs to set PERPLEXITY_API_KEY). Proceed without it"
|
||||
),
|
||||
}
|
||||
logger.info("web_search query (len=%d): %s", len(query), query[:120])
|
||||
|
||||
def _perplexity_content(api_key: str, query: str) -> str:
|
||||
url = "https://api.perplexity.ai/chat/completions"
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
payload = {
|
||||
|
|
@ -66,61 +56,269 @@ def _do_search(query: str) -> dict[str, Any]: # noqa: PLR0911 - each error clas
|
|||
{"role": "user", "content": query},
|
||||
],
|
||||
}
|
||||
with requests.post(url, headers=headers, json=payload, timeout=300) as response:
|
||||
response.raise_for_status()
|
||||
return str(response.json()["choices"][0]["message"]["content"])
|
||||
|
||||
|
||||
_EXA_PAGE_MAX_CHARS = 20000
|
||||
_EXA_MAX_CONTENT_URLS = 10
|
||||
_EXA_SUMMARY_PROMPT = (
|
||||
"Summarize this page for a penetration tester. Keep concrete technical detail: "
|
||||
"affected products and exact versions, CVE and CWE identifiers, CVSS scores, "
|
||||
"exploitation preconditions, payloads or commands, and mitigations. "
|
||||
"Leave out marketing copy and navigation text."
|
||||
)
|
||||
|
||||
|
||||
def _exa_result_block(result: dict[str, Any]) -> str | None:
|
||||
result_url = str(result.get("url") or result.get("id") or "")
|
||||
if not result_url:
|
||||
return None
|
||||
title = str(result.get("title") or result_url)
|
||||
parts = [f"### {title}\n{result_url}"]
|
||||
summary = str(result.get("summary") or "").strip()
|
||||
if summary:
|
||||
parts.append(summary)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _exa_page_block(result: dict[str, Any]) -> str | None:
|
||||
result_url = str(result.get("url") or result.get("id") or "")
|
||||
text = str(result.get("text") or "").strip()
|
||||
if not result_url or not text:
|
||||
return None
|
||||
if len(text) > _EXA_PAGE_MAX_CHARS:
|
||||
text = f"{text[:_EXA_PAGE_MAX_CHARS]}\n[truncated at {_EXA_PAGE_MAX_CHARS} characters]"
|
||||
title = str(result.get("title") or result_url)
|
||||
return f"### {title}\n{result_url}\n\n{text}"
|
||||
|
||||
|
||||
def _exa_blocks(
|
||||
results: list[Any],
|
||||
render: Callable[[dict[str, Any]], str | None],
|
||||
) -> list[str]:
|
||||
blocks: list[str] = []
|
||||
for result in results:
|
||||
if not isinstance(result, dict):
|
||||
continue
|
||||
block = render(cast("dict[str, Any]", result))
|
||||
if block:
|
||||
blocks.append(block)
|
||||
return blocks
|
||||
|
||||
|
||||
def _exa_post(api_key: str, endpoint: str, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
headers = {"x-api-key": api_key, "Content-Type": "application/json"}
|
||||
with requests.post(endpoint, headers=headers, json=payload, timeout=300) as response:
|
||||
response.raise_for_status()
|
||||
body: dict[str, Any] = response.json()
|
||||
return body
|
||||
|
||||
|
||||
def _exa_content(api_key: str, query: str, search_type: str, num_results: int) -> str:
|
||||
body = _exa_post(
|
||||
api_key,
|
||||
"https://api.exa.ai/search",
|
||||
{
|
||||
"query": f"{_SYSTEM_PROMPT}\n\n{query}",
|
||||
"type": search_type,
|
||||
"numResults": num_results,
|
||||
"contents": {"summary": {"query": _EXA_SUMMARY_PROMPT}},
|
||||
},
|
||||
)
|
||||
blocks = _exa_blocks(body.get("results") or [], _exa_result_block)
|
||||
if not blocks:
|
||||
raise ValueError("Exa response has no results")
|
||||
return "\n\n".join(blocks)
|
||||
|
||||
|
||||
def _normalize_url(url: str) -> str:
|
||||
"""Canonical form for matching: case-fold scheme and host only, drop a trailing slash."""
|
||||
parts = urlsplit(url.strip())
|
||||
return urlunsplit(
|
||||
(parts.scheme.lower(), parts.netloc.lower(), parts.path.rstrip("/"), parts.query, "")
|
||||
)
|
||||
|
||||
|
||||
def _exa_page_text(api_key: str, urls: list[str]) -> tuple[str, set[str]]:
|
||||
"""Fetch page text and report which of the requested URLs Exa returned."""
|
||||
body = _exa_post(api_key, "https://api.exa.ai/contents", {"urls": urls, "text": True})
|
||||
blocks: list[str] = []
|
||||
fetched: set[str] = set()
|
||||
results: list[Any] = body.get("results") or []
|
||||
for result in results:
|
||||
if not isinstance(result, dict):
|
||||
continue
|
||||
page = cast("dict[str, Any]", result)
|
||||
block = _exa_page_block(page)
|
||||
if not block:
|
||||
continue
|
||||
blocks.append(block)
|
||||
fetched.add(_normalize_url(str(page.get("url") or page.get("id") or "")))
|
||||
if not blocks:
|
||||
raise ValueError("Exa returned no page contents")
|
||||
return "\n\n".join(blocks), fetched
|
||||
|
||||
|
||||
def _resolve_provider( # noqa: PLR0911 - each provider/missing-key case needs its own return
|
||||
integrations: Any,
|
||||
) -> tuple[str, str] | dict[str, Any]:
|
||||
"""Pick the search provider and its key, or return a sanitized error dict."""
|
||||
provider = integrations.web_search_provider
|
||||
perplexity_key = integrations.perplexity_api_key
|
||||
exa_key = integrations.exa_api_key
|
||||
|
||||
if provider == "perplexity":
|
||||
if not perplexity_key:
|
||||
return _not_configured_error("PERPLEXITY_API_KEY")
|
||||
return ("perplexity", perplexity_key)
|
||||
if provider == "exa":
|
||||
if not exa_key:
|
||||
return _not_configured_error("EXA_API_KEY")
|
||||
return ("exa", exa_key)
|
||||
|
||||
if exa_key:
|
||||
return ("exa", exa_key)
|
||||
if perplexity_key:
|
||||
return ("perplexity", perplexity_key)
|
||||
return _not_configured_error("EXA_API_KEY or PERPLEXITY_API_KEY")
|
||||
|
||||
|
||||
def _not_configured_error(missing: str) -> dict[str, Any]:
|
||||
logger.warning("web_search invoked without %s configured", missing)
|
||||
return {
|
||||
"success": False,
|
||||
"error": (
|
||||
"Web search is not configured for this scan "
|
||||
f"(operator needs to set {missing}). Proceed without it"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _guarded_call[T]( # noqa: PLR0911 - each error class needs its own sanitized return
|
||||
tool: str,
|
||||
rejected_hint: str,
|
||||
fetch: Callable[[], T],
|
||||
) -> T | dict[str, Any]:
|
||||
"""Run a provider call and translate any failure into a sanitized error dict."""
|
||||
try:
|
||||
with requests.post(url, headers=headers, json=payload, timeout=300) as response:
|
||||
response.raise_for_status()
|
||||
content = response.json()["choices"][0]["message"]["content"]
|
||||
return fetch()
|
||||
except requests.exceptions.Timeout:
|
||||
logger.warning("web_search timed out")
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Web search timed out. Try again or shorten the query",
|
||||
}
|
||||
logger.warning("%s timed out", tool)
|
||||
return {"success": False, "error": f"{tool} timed out. Try again or narrow the request"}
|
||||
except requests.exceptions.HTTPError as exc:
|
||||
status = exc.response.status_code if exc.response is not None else None
|
||||
logger.exception("web_search HTTP error status=%s", status)
|
||||
logger.exception("%s HTTP error status=%s", tool, status)
|
||||
if status is not None and 400 <= status < 500:
|
||||
return {
|
||||
"success": False,
|
||||
"error": (
|
||||
"Web search rejected the query. Refine it "
|
||||
"(more specific, shorter, no unusual characters) and retry"
|
||||
),
|
||||
}
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Web search service is unavailable. Try again later",
|
||||
}
|
||||
return {"success": False, "error": rejected_hint}
|
||||
return {"success": False, "error": f"{tool} service is unavailable. Try again later"}
|
||||
except requests.exceptions.RequestException:
|
||||
logger.exception("web_search network error")
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Web search network error. Try again later",
|
||||
}
|
||||
logger.exception("%s network error", tool)
|
||||
return {"success": False, "error": f"{tool} network error. Try again later"}
|
||||
except (KeyError, IndexError, ValueError):
|
||||
logger.exception("web_search response shape unexpected")
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Web search returned an unexpected response. Try again",
|
||||
}
|
||||
logger.exception("%s response shape unexpected", tool)
|
||||
return {"success": False, "error": f"{tool} returned an unexpected response. Try again"}
|
||||
except Exception:
|
||||
logger.exception("web_search failed")
|
||||
logger.exception("%s failed", tool)
|
||||
return {"success": False, "error": f"{tool} failed unexpectedly"}
|
||||
|
||||
|
||||
def _do_search(query: str) -> dict[str, Any]:
|
||||
if not query or not query.strip():
|
||||
return {"success": False, "error": "Query cannot be empty"}
|
||||
|
||||
integrations = load_settings().integrations
|
||||
resolved = _resolve_provider(integrations)
|
||||
if isinstance(resolved, dict):
|
||||
return resolved
|
||||
provider, api_key = resolved
|
||||
logger.info("web_search provider=%s query (len=%d): %s", provider, len(query), query[:120])
|
||||
|
||||
def fetch() -> str:
|
||||
if provider == "exa":
|
||||
return _exa_content(
|
||||
api_key,
|
||||
query,
|
||||
integrations.exa_search_type,
|
||||
integrations.exa_num_results,
|
||||
)
|
||||
return _perplexity_content(api_key, query)
|
||||
|
||||
outcome = _guarded_call(
|
||||
"Web search",
|
||||
(
|
||||
"Web search rejected the query. Refine it "
|
||||
"(more specific, shorter, no unusual characters) and retry"
|
||||
),
|
||||
fetch,
|
||||
)
|
||||
if isinstance(outcome, dict):
|
||||
return outcome
|
||||
return {
|
||||
"success": True,
|
||||
"query": query,
|
||||
"provider": provider,
|
||||
"content": outcome,
|
||||
}
|
||||
|
||||
|
||||
def _do_get_contents(urls: list[str]) -> dict[str, Any]:
|
||||
cleaned = [url.strip() for url in urls if url and url.strip()]
|
||||
if not cleaned:
|
||||
return {"success": False, "error": "Provide at least one URL"}
|
||||
if len(cleaned) > _EXA_MAX_CONTENT_URLS:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Web search failed unexpectedly",
|
||||
"error": f"Too many URLs. Pass at most {_EXA_MAX_CONTENT_URLS} per call",
|
||||
}
|
||||
else:
|
||||
|
||||
integrations = load_settings().integrations
|
||||
api_key = integrations.exa_api_key
|
||||
if not api_key:
|
||||
return _not_configured_error("EXA_API_KEY")
|
||||
if integrations.web_search_provider == "perplexity":
|
||||
logger.warning("web_get_contents invoked while the provider is pinned to Perplexity")
|
||||
return {
|
||||
"success": True,
|
||||
"query": query,
|
||||
"content": content,
|
||||
"success": False,
|
||||
"error": (
|
||||
"Page fetching needs the Exa provider "
|
||||
"(operator pinned STRIX_WEB_SEARCH_PROVIDER to perplexity). "
|
||||
"Use web_search instead"
|
||||
),
|
||||
}
|
||||
|
||||
logger.info("web_get_contents urls=%d", len(cleaned))
|
||||
outcome = _guarded_call(
|
||||
"Page fetch",
|
||||
"Page fetch was rejected. Check the URLs are complete, public, and correctly formed",
|
||||
lambda: _exa_page_text(api_key, cleaned),
|
||||
)
|
||||
if isinstance(outcome, dict):
|
||||
return outcome
|
||||
content, fetched = outcome
|
||||
missing = [url for url in cleaned if _normalize_url(url) not in fetched]
|
||||
result: dict[str, Any] = {
|
||||
"success": True,
|
||||
"urls": [url for url in cleaned if url not in missing],
|
||||
"provider": "exa",
|
||||
"content": content,
|
||||
}
|
||||
if missing:
|
||||
logger.warning(
|
||||
"web_get_contents returned %d of %d pages", len(cleaned) - len(missing), len(cleaned)
|
||||
)
|
||||
result["failed_urls"] = missing
|
||||
result["warning"] = (
|
||||
f"Exa returned no content for {len(missing)} of {len(cleaned)} requested URLs. "
|
||||
"Those pages are missing from the content below"
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@function_tool(timeout=330)
|
||||
async def web_search(ctx: RunContextWrapper, query: str) -> str:
|
||||
"""Real-time web search via Perplexity — your primary research tool.
|
||||
"""Real-time web search (Exa or Perplexity) — your primary research tool.
|
||||
|
||||
Use it liberally for anything that's not in your training data:
|
||||
|
||||
|
|
@ -150,6 +348,12 @@ async def web_search(ctx: RunContextWrapper, query: str) -> str:
|
|||
exploits, Kali-compatible tooling, and concrete code/command
|
||||
examples.
|
||||
|
||||
With the Exa provider you get a ranked list of results, each with a
|
||||
title, URL, and a short security-focused summary. Read the result
|
||||
you need, then call ``web_get_contents`` with its URL to pull the
|
||||
full page text when a summary is not enough. With Perplexity you get
|
||||
a single synthesized cited answer.
|
||||
|
||||
**Good example queries** (each is a full sentence, names a
|
||||
version/product, and asks one concrete thing):
|
||||
|
||||
|
|
@ -177,3 +381,33 @@ async def web_search(ctx: RunContextWrapper, query: str) -> str:
|
|||
"""
|
||||
result = await asyncio.to_thread(_do_search, query)
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
@function_tool(timeout=330)
|
||||
async def web_get_contents(ctx: RunContextWrapper, urls: list[str]) -> str:
|
||||
"""Fetch the full, cleaned text of specific web pages (Exa only).
|
||||
|
||||
Use this as the drill-down step after ``web_search``: when a result's
|
||||
summary is not enough, pass that result's URL here to read the whole
|
||||
page. Good for reading a full advisory, a CVE writeup,
|
||||
an exploit proof-of-concept, or vendor documentation end to end.
|
||||
|
||||
Prefer ``web_search`` first to find the right pages, then fetch only
|
||||
the few URLs worth reading in full — each page can be large, so avoid
|
||||
fetching many pages you do not need.
|
||||
|
||||
This tool needs the Exa provider (``EXA_API_KEY``). When the operator
|
||||
pins the provider to Perplexity, it returns an error and you should
|
||||
use ``web_search`` instead.
|
||||
|
||||
Some pages block extraction. When a page returns no content, the
|
||||
result lists it under ``failed_urls`` and the ``content`` field holds
|
||||
only the pages that came back. Check ``failed_urls`` before you
|
||||
conclude that a page had nothing useful.
|
||||
|
||||
Args:
|
||||
urls: The page URLs to fetch, at most 10 per call. Use complete,
|
||||
public URLs (for example the ones returned by ``web_search``).
|
||||
"""
|
||||
result = await asyncio.to_thread(_do_get_contents, urls)
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
|
|
|||
|
|
@ -30,6 +30,8 @@ _LLM_ENV_KEYS = [
|
|||
"STRIX_FORCE_REQUIRED_TOOL_CHOICE",
|
||||
"LLM_TIMEOUT",
|
||||
"PERPLEXITY_API_KEY",
|
||||
"EXA_API_KEY",
|
||||
"STRIX_WEB_SEARCH_PROVIDER",
|
||||
# RuntimeSettings
|
||||
"STRIX_IMAGE",
|
||||
"STRIX_RUNTIME_BACKEND",
|
||||
|
|
@ -80,6 +82,17 @@ def test_read_json_overrides_maps_to_nested_settings(tmp_path: Path) -> None:
|
|||
}
|
||||
|
||||
|
||||
def test_read_json_overrides_maps_exa_and_provider(tmp_path: Path) -> None:
|
||||
path = tmp_path / "cli-config.json"
|
||||
path.write_text(
|
||||
json.dumps({"env": {"EXA_API_KEY": "exa-key", "STRIX_WEB_SEARCH_PROVIDER": "exa"}}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
assert loader._read_json_overrides(path) == {
|
||||
"integrations": {"exa_api_key": "exa-key", "web_search_provider": "exa"},
|
||||
}
|
||||
|
||||
|
||||
def test_read_json_overrides_skips_keys_already_in_environ(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
|
|
|
|||
416
tests/test_web_search.py
Normal file
416
tests/test_web_search.py
Normal file
|
|
@ -0,0 +1,416 @@
|
|||
"""Tests for web_search/web_get_contents provider selection and the Exa backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from strix.config.settings import IntegrationSettings
|
||||
from strix.interface.environment import _missing_web_search_vars
|
||||
from strix.tools.web_search import tool
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing import Self
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, body: dict[str, Any]) -> None:
|
||||
self._body = body
|
||||
self.headers: dict[str, str] = {}
|
||||
|
||||
def __enter__(self) -> Self:
|
||||
return self
|
||||
|
||||
def __exit__(self, *_exc: object) -> None:
|
||||
return None
|
||||
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
def json(self) -> dict[str, Any]:
|
||||
return self._body
|
||||
|
||||
|
||||
def test_auto_prefers_exa_when_both_keys_set() -> None:
|
||||
integrations = IntegrationSettings(PERPLEXITY_API_KEY="pk", EXA_API_KEY="ek")
|
||||
assert tool._resolve_provider(integrations) == ("exa", "ek")
|
||||
|
||||
|
||||
def test_auto_falls_back_to_perplexity_when_only_perplexity_is_set() -> None:
|
||||
integrations = IntegrationSettings(PERPLEXITY_API_KEY="pk")
|
||||
assert tool._resolve_provider(integrations) == ("perplexity", "pk")
|
||||
|
||||
|
||||
def test_explicit_exa_ignores_a_configured_perplexity_key() -> None:
|
||||
integrations = IntegrationSettings(
|
||||
PERPLEXITY_API_KEY="pk",
|
||||
EXA_API_KEY="ek",
|
||||
STRIX_WEB_SEARCH_PROVIDER="exa",
|
||||
)
|
||||
assert tool._resolve_provider(integrations) == ("exa", "ek")
|
||||
|
||||
|
||||
def test_explicit_perplexity_ignores_a_configured_exa_key() -> None:
|
||||
integrations = IntegrationSettings(
|
||||
PERPLEXITY_API_KEY="pk",
|
||||
EXA_API_KEY="ek",
|
||||
STRIX_WEB_SEARCH_PROVIDER="perplexity",
|
||||
)
|
||||
assert tool._resolve_provider(integrations) == ("perplexity", "pk")
|
||||
|
||||
|
||||
def test_explicit_exa_without_a_key_names_only_exa() -> None:
|
||||
integrations = IntegrationSettings(
|
||||
PERPLEXITY_API_KEY="pk",
|
||||
STRIX_WEB_SEARCH_PROVIDER="exa",
|
||||
)
|
||||
resolved = tool._resolve_provider(integrations)
|
||||
assert isinstance(resolved, dict)
|
||||
assert resolved["success"] is False
|
||||
assert "EXA_API_KEY" in resolved["error"]
|
||||
assert "PERPLEXITY_API_KEY" not in resolved["error"]
|
||||
|
||||
|
||||
def test_no_keys_names_both_providers() -> None:
|
||||
resolved = tool._resolve_provider(IntegrationSettings())
|
||||
assert isinstance(resolved, dict)
|
||||
assert "EXA_API_KEY or PERPLEXITY_API_KEY" in resolved["error"]
|
||||
|
||||
|
||||
def test_exa_content_requests_summaries_and_renders_results(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def fake_post(url: str, **kwargs: Any) -> _FakeResponse:
|
||||
captured["url"] = url
|
||||
captured["headers"] = kwargs["headers"]
|
||||
captured["json"] = kwargs["json"]
|
||||
return _FakeResponse(
|
||||
{
|
||||
"results": [
|
||||
{
|
||||
"url": "https://nvd.example/cve",
|
||||
"title": "NVD entry",
|
||||
"summary": " CVE-2024-0001 is a heap overflow. ",
|
||||
},
|
||||
{"id": "https://blog.example/post"},
|
||||
"not-a-dict",
|
||||
{"title": "no url"},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
|
||||
content = tool._exa_content("ek", "OpenSSH 7.4 RCE?", "auto", 5)
|
||||
|
||||
assert captured["url"] == "https://api.exa.ai/search"
|
||||
assert captured["headers"]["x-api-key"] == "ek"
|
||||
assert "OpenSSH 7.4 RCE?" in captured["json"]["query"]
|
||||
assert captured["json"]["type"] == "auto"
|
||||
assert captured["json"]["numResults"] == 5
|
||||
assert captured["json"]["contents"] == {"summary": {"query": tool._EXA_SUMMARY_PROMPT}}
|
||||
assert content == (
|
||||
"### NVD entry\nhttps://nvd.example/cve\nCVE-2024-0001 is a heap overflow.\n\n"
|
||||
"### https://blog.example/post\nhttps://blog.example/post"
|
||||
)
|
||||
|
||||
|
||||
def test_exa_content_renders_a_result_without_contents(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
requests,
|
||||
"post",
|
||||
lambda *_a, **_kw: _FakeResponse(
|
||||
{"results": [{"url": "https://ex.example", "title": "Ex"}]}
|
||||
),
|
||||
)
|
||||
assert tool._exa_content("ek", "q", "auto", 5) == "### Ex\nhttps://ex.example"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", [{}, {"results": None}, {"results": []}, {"results": ["x"]}])
|
||||
def test_exa_content_rejects_empty_results(
|
||||
monkeypatch: pytest.MonkeyPatch, body: dict[str, Any]
|
||||
) -> None:
|
||||
monkeypatch.setattr(requests, "post", lambda *_a, **_kw: _FakeResponse(body))
|
||||
with pytest.raises(ValueError, match="no results"):
|
||||
tool._exa_content("ek", "q", "auto", 5)
|
||||
|
||||
|
||||
def test_do_search_reports_empty_exa_results_as_unexpected(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(requests, "post", lambda *_a, **_kw: _FakeResponse({}))
|
||||
|
||||
result = tool._do_search("q")
|
||||
|
||||
assert result["success"] is False
|
||||
assert "unexpected response" in result["error"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("env", "expected"),
|
||||
[
|
||||
({}, ["EXA_API_KEY", "PERPLEXITY_API_KEY"]),
|
||||
({"EXA_API_KEY": "ek"}, []),
|
||||
({"PERPLEXITY_API_KEY": "pk"}, []),
|
||||
({"STRIX_WEB_SEARCH_PROVIDER": "exa", "PERPLEXITY_API_KEY": "pk"}, ["EXA_API_KEY"]),
|
||||
({"STRIX_WEB_SEARCH_PROVIDER": "exa", "EXA_API_KEY": "ek"}, []),
|
||||
({"STRIX_WEB_SEARCH_PROVIDER": "perplexity", "EXA_API_KEY": "ek"}, ["PERPLEXITY_API_KEY"]),
|
||||
({"STRIX_WEB_SEARCH_PROVIDER": "perplexity", "PERPLEXITY_API_KEY": "pk"}, []),
|
||||
],
|
||||
)
|
||||
def test_environment_validation_follows_provider_rules(
|
||||
env: dict[str, str], expected: list[str]
|
||||
) -> None:
|
||||
integrations = IntegrationSettings.model_validate(env)
|
||||
assert _missing_web_search_vars(integrations) == expected
|
||||
|
||||
|
||||
def test_exa_search_type_and_num_results_are_configurable(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(
|
||||
EXA_API_KEY="ek",
|
||||
STRIX_EXA_SEARCH_TYPE="deep-reasoning",
|
||||
STRIX_EXA_NUM_RESULTS=3,
|
||||
)
|
||||
|
||||
def fake_post(_url: str, **kwargs: Any) -> _FakeResponse:
|
||||
captured["json"] = kwargs["json"]
|
||||
return _FakeResponse({"results": [{"url": "https://ex.example", "title": "Ex"}]})
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
|
||||
assert tool._do_search("q")["success"] is True
|
||||
assert captured["json"]["type"] == "deep-reasoning"
|
||||
assert captured["json"]["numResults"] == 3
|
||||
|
||||
|
||||
def test_exa_page_text_requests_full_text_and_renders_pages(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def fake_post(url: str, **kwargs: Any) -> _FakeResponse:
|
||||
captured["url"] = url
|
||||
captured["headers"] = kwargs["headers"]
|
||||
captured["json"] = kwargs["json"]
|
||||
return _FakeResponse(
|
||||
{
|
||||
"results": [
|
||||
{
|
||||
"url": "https://nvd.example/cve",
|
||||
"title": "NVD entry",
|
||||
"text": " Full advisory body. ",
|
||||
},
|
||||
{"url": "https://empty.example", "text": " "},
|
||||
"not-a-dict",
|
||||
{"text": "no url"},
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setattr(requests, "post", fake_post)
|
||||
|
||||
content, fetched = tool._exa_page_text("ek", ["https://nvd.example/cve"])
|
||||
|
||||
assert captured["url"] == "https://api.exa.ai/contents"
|
||||
assert captured["headers"]["x-api-key"] == "ek"
|
||||
assert captured["json"] == {"urls": ["https://nvd.example/cve"], "text": True}
|
||||
assert content == "### NVD entry\nhttps://nvd.example/cve\n\nFull advisory body."
|
||||
assert fetched == {"https://nvd.example/cve"}
|
||||
|
||||
|
||||
def test_exa_page_text_truncates_a_long_page(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
body = "A" * (tool._EXA_PAGE_MAX_CHARS + 500)
|
||||
monkeypatch.setattr(
|
||||
requests,
|
||||
"post",
|
||||
lambda *_a, **_kw: _FakeResponse(
|
||||
{"results": [{"url": "https://ex.example", "text": body}]}
|
||||
),
|
||||
)
|
||||
content, _fetched = tool._exa_page_text("ek", ["https://ex.example"])
|
||||
assert "truncated at" in content
|
||||
assert content.count("A") == tool._EXA_PAGE_MAX_CHARS
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", [{}, {"results": []}, {"results": [{"url": "u"}]}])
|
||||
def test_exa_page_text_rejects_pages_without_text(
|
||||
monkeypatch: pytest.MonkeyPatch, body: dict[str, Any]
|
||||
) -> None:
|
||||
monkeypatch.setattr(requests, "post", lambda *_a, **_kw: _FakeResponse(body))
|
||||
with pytest.raises(ValueError, match="no page contents"):
|
||||
tool._exa_page_text("ek", ["https://ex.example"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("urls", [[], ["", " "]])
|
||||
def test_do_get_contents_requires_a_url(urls: list[str]) -> None:
|
||||
result = tool._do_get_contents(urls)
|
||||
assert result["success"] is False
|
||||
assert "at least one URL" in result["error"]
|
||||
|
||||
|
||||
def test_do_get_contents_caps_the_url_count() -> None:
|
||||
urls = [f"https://ex{index}.example" for index in range(tool._EXA_MAX_CONTENT_URLS + 1)]
|
||||
result = tool._do_get_contents(urls)
|
||||
assert result["success"] is False
|
||||
assert "Too many URLs" in result["error"]
|
||||
|
||||
|
||||
def test_do_get_contents_needs_an_exa_key(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(PERPLEXITY_API_KEY="pk")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
|
||||
result = tool._do_get_contents(["https://ex.example"])
|
||||
|
||||
assert result["success"] is False
|
||||
assert "EXA_API_KEY" in result["error"]
|
||||
|
||||
|
||||
def test_do_get_contents_refuses_a_perplexity_pinned_provider(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(
|
||||
EXA_API_KEY="ek",
|
||||
PERPLEXITY_API_KEY="pk",
|
||||
STRIX_WEB_SEARCH_PROVIDER="perplexity",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
|
||||
result = tool._do_get_contents(["https://ex.example"])
|
||||
|
||||
assert result["success"] is False
|
||||
assert "web_search" in result["error"]
|
||||
|
||||
|
||||
def test_do_get_contents_returns_page_text(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(tool, "_exa_page_text", lambda *_a: ("page", {"https://ex.example"}))
|
||||
|
||||
result = tool._do_get_contents([" https://ex.example "])
|
||||
|
||||
assert result == {
|
||||
"success": True,
|
||||
"urls": ["https://ex.example"],
|
||||
"provider": "exa",
|
||||
"content": "page",
|
||||
}
|
||||
|
||||
|
||||
def test_do_get_contents_reports_urls_exa_did_not_return(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(
|
||||
requests,
|
||||
"post",
|
||||
lambda *_a, **_kw: _FakeResponse(
|
||||
{"results": [{"url": "https://ok.example/", "text": "Body."}]}
|
||||
),
|
||||
)
|
||||
|
||||
result = tool._do_get_contents(["https://ok.example", "https://blocked.example"])
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["urls"] == ["https://ok.example"]
|
||||
assert result["failed_urls"] == ["https://blocked.example"]
|
||||
assert "1 of 2" in result["warning"]
|
||||
assert "blocked.example" not in result["content"]
|
||||
|
||||
|
||||
def test_normalize_url_folds_only_scheme_and_host() -> None:
|
||||
assert tool._normalize_url("HTTPS://Ex.Example/Path/") == tool._normalize_url(
|
||||
"https://ex.example/Path"
|
||||
)
|
||||
assert tool._normalize_url("https://ex.example/Path") != tool._normalize_url(
|
||||
"https://ex.example/path"
|
||||
)
|
||||
assert tool._normalize_url("https://ex.example/p?Q=A") != tool._normalize_url(
|
||||
"https://ex.example/p?q=a"
|
||||
)
|
||||
|
||||
|
||||
def test_do_get_contents_omits_the_warning_when_every_page_returns(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(
|
||||
requests,
|
||||
"post",
|
||||
lambda *_a, **_kw: _FakeResponse(
|
||||
{
|
||||
"results": [
|
||||
{"url": "https://a.example", "text": "A."},
|
||||
{"url": "https://b.example", "text": "B."},
|
||||
]
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
result = tool._do_get_contents(["https://a.example", "https://b.example"])
|
||||
|
||||
assert result["urls"] == ["https://a.example", "https://b.example"]
|
||||
assert "failed_urls" not in result
|
||||
assert "warning" not in result
|
||||
|
||||
|
||||
def test_do_get_contents_sanitizes_a_network_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
def boom(*_args: Any, **_kwargs: Any) -> None:
|
||||
raise requests.exceptions.ConnectionError
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(requests, "post", boom)
|
||||
|
||||
result = tool._do_get_contents(["https://ex.example"])
|
||||
|
||||
assert result["success"] is False
|
||||
assert "network error" in result["error"]
|
||||
assert "ek" not in result["error"]
|
||||
|
||||
|
||||
def test_do_search_reports_the_provider_it_used(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
class _Settings:
|
||||
integrations = IntegrationSettings(EXA_API_KEY="ek")
|
||||
|
||||
monkeypatch.setattr(tool, "load_settings", _Settings)
|
||||
monkeypatch.setattr(tool, "_exa_content", lambda *_a: "answer")
|
||||
|
||||
result = tool._do_search("OpenSSH 7.4 RCE?")
|
||||
|
||||
assert result == {
|
||||
"success": True,
|
||||
"query": "OpenSSH 7.4 RCE?",
|
||||
"provider": "exa",
|
||||
"content": "answer",
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue