397 lines
17 KiB
Python
397 lines
17 KiB
Python
"""xAI Web Search — search-only provider backed by Grok's server-side ``web_search``
|
|
tool on the Responses API. Grok searches/browses server-side; we ask for structured
|
|
JSON so results match the ``{title, url, description, position}`` rows every other
|
|
Hermes web provider produces. Reference: https://docs.x.ai/developers/tools/web-search
|
|
|
|
Config: ``web.search_backend`` / ``web.backend: "xai"``. Optional ``web.xai``:
|
|
``model`` (reasoning model, default grok-build-0.1), ``allowed_domains`` /
|
|
``excluded_domains`` (max 5, mutually exclusive), ``timeout`` (seconds, default 90).
|
|
Auth: :func:`tools.xai_http.resolve_xai_http_credentials` (Grok OAuth via
|
|
``hermes auth``, else ``XAI_API_KEY``).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import re
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from plugins.web._common import BaseWebSearchProvider, search_fail as _fail, search_ok
|
|
from tools.xai_http import (
|
|
has_xai_credentials,
|
|
hermes_xai_user_agent,
|
|
resolve_xai_http_credentials,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
DEFAULT_MODEL = "grok-build-0.1"
|
|
DEFAULT_TIMEOUT = 90
|
|
_MAX_DOMAIN_FILTERS = 5 # xAI hard cap on allowed_domains / excluded_domains
|
|
|
|
# The JSON object Grok is asked to emit; tolerates leading/trailing prose since
|
|
# reasoning models occasionally narrate before the JSON block.
|
|
_JSON_BLOCK_RE = re.compile(r"\{[\s\S]*\}", re.MULTILINE)
|
|
|
|
|
|
def _load_xai_web_config() -> Dict[str, Any]:
|
|
"""Read ``web.xai`` from config.yaml (returns {} on miss)."""
|
|
try:
|
|
from hermes_cli.config import load_config
|
|
|
|
cfg = load_config()
|
|
web_section = cfg.get("web") if isinstance(cfg, dict) else None
|
|
xai_section = web_section.get("xai") if isinstance(web_section, dict) else None
|
|
return xai_section if isinstance(xai_section, dict) else {}
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.debug("Could not load web.xai config: %s", exc)
|
|
return {}
|
|
|
|
|
|
def _coerce_domain_list(value: Any) -> List[str]:
|
|
"""Coerce a config value to a clean list of <=5 domain strings."""
|
|
if not isinstance(value, list):
|
|
return []
|
|
cleaned = [item.strip() for item in value if isinstance(item, str) and item.strip()]
|
|
return cleaned[:_MAX_DOMAIN_FILTERS]
|
|
|
|
|
|
def _coerce(cast, value: Any, default: Any) -> Any:
|
|
try:
|
|
return cast(value)
|
|
except (TypeError, ValueError):
|
|
return default
|
|
|
|
|
|
def _row(title: str, url: str, description: str, position: int) -> Dict[str, Any]:
|
|
# Key order (title first) is this provider's wire shape; keep it distinct from web_hit.
|
|
return {"title": title, "url": url, "description": description, "position": position}
|
|
|
|
|
|
class XAIWebSearchProvider(BaseWebSearchProvider):
|
|
"""Search-only provider backed by xAI's agentic Web Search tool.
|
|
|
|
Sends a structured prompt with ``tools=[{"type": "web_search"}]`` and parses the
|
|
JSON Grok returns; falls back to message annotations, then the ``citations``
|
|
list, if Grok ignores the schema. No extract capability.
|
|
|
|
Trust model: unlike index-backed providers, Grok *generates* the URLs, titles
|
|
and descriptions and is steerable by the query text itself — treat returned
|
|
URLs like any model-generated link and validate before fetching.
|
|
"""
|
|
|
|
NAME = "xai"
|
|
DISPLAY_NAME = "xAI Web Search (Grok)"
|
|
|
|
def is_available(self) -> bool:
|
|
"""Cheap probe (env var OR auth-store tokens). Deliberately NOT
|
|
``resolve_xai_http_credentials``: must never refresh tokens or take the
|
|
auth-store lock, since this runs on every ``hermes tools`` repaint."""
|
|
return has_xai_credentials()
|
|
|
|
# -- Search -----------------------------------------------------------
|
|
|
|
def search(self, query: str, limit: int = 5) -> Dict[str, Any]:
|
|
"""Grok-backed web search → ``{"success": True, "data": {"web": [...]}}``
|
|
or ``{"success": False, "error": str}``."""
|
|
try:
|
|
from tools.interrupt import is_interrupted
|
|
|
|
if is_interrupted():
|
|
return _fail("Interrupted")
|
|
except Exception: # noqa: BLE001 — interrupt module is best-effort
|
|
pass
|
|
|
|
creds = resolve_xai_http_credentials()
|
|
api_key = str(creds.get("api_key") or "").strip()
|
|
base_url = str(creds.get("base_url") or "https://api.x.ai/v1").strip().rstrip("/")
|
|
if not api_key:
|
|
return _fail(
|
|
"No xAI credentials found. Run `hermes auth` to sign in with "
|
|
"xAI Grok OAuth, or set XAI_API_KEY."
|
|
)
|
|
|
|
# Same clamp range as web_search_tool so explicit limits aren't downgraded;
|
|
# cost scales with the count via reasoning tokens, but that's the caller's call.
|
|
limit = max(1, min(_coerce(int, limit, 5), 100))
|
|
|
|
cfg = _load_xai_web_config()
|
|
model = cfg.get("model") if isinstance(cfg.get("model"), str) else DEFAULT_MODEL
|
|
model = model.strip() or DEFAULT_MODEL
|
|
timeout = _coerce(float, cfg.get("timeout", DEFAULT_TIMEOUT), DEFAULT_TIMEOUT)
|
|
|
|
web_search_tool = self._web_search_tool(cfg)
|
|
if web_search_tool is None:
|
|
# xAI rejects this combo — surface a clear error rather than an API 400.
|
|
return _fail(
|
|
"web.xai.allowed_domains and web.xai.excluded_domains "
|
|
"cannot both be set (xAI restriction)."
|
|
)
|
|
|
|
payload: Dict[str, Any] = {
|
|
"model": model,
|
|
"input": [{"role": "user", "content": self._build_prompt(query, limit)}],
|
|
"tools": [web_search_tool],
|
|
# Keep the JSON block clean; URLs are read from annotations/citations.
|
|
"include": ["no_inline_citations"],
|
|
}
|
|
|
|
try:
|
|
import httpx # noqa: F401 — availability probe
|
|
except ImportError:
|
|
return _fail("httpx is not installed (required for xAI web search)")
|
|
|
|
logger.info("xAI web search via %s: '%s' (limit=%d, model=%s)", base_url, query, limit, model)
|
|
|
|
data, error = self._post_responses(
|
|
base_url, payload, api_key, timeout,
|
|
is_oauth_path=(creds.get("provider") == "xai-oauth"),
|
|
)
|
|
if error:
|
|
return error
|
|
|
|
# xAI sometimes returns HTTP 200 with an error envelope (overloaded, refusal);
|
|
# without this check we'd report success-with-no-rows and mask a real failure.
|
|
api_error = data.get("error") if isinstance(data, dict) else None
|
|
if isinstance(api_error, dict):
|
|
err_msg = api_error.get("message") or api_error.get("code") or "unknown error"
|
|
logger.warning("xAI web search returned error envelope: %s", err_msg)
|
|
return _fail(f"xAI returned an error: {err_msg}")
|
|
|
|
# Empty list on 0 hits is a success (matches brave-free / exa) so the model
|
|
# can decide whether to retry.
|
|
return search_ok(self._extract_results(data, limit=limit))
|
|
|
|
@staticmethod
|
|
def _web_search_tool(cfg: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
|
"""Build the ``web_search`` tool spec with optional domain filters; None when
|
|
both allowed and excluded are set (xAI rejects the combination)."""
|
|
allowed = _coerce_domain_list(cfg.get("allowed_domains"))
|
|
excluded = _coerce_domain_list(cfg.get("excluded_domains"))
|
|
if allowed and excluded:
|
|
return None
|
|
tool: Dict[str, Any] = {"type": "web_search"}
|
|
if allowed:
|
|
tool["filters"] = {"allowed_domains": allowed}
|
|
elif excluded:
|
|
tool["filters"] = {"excluded_domains": excluded}
|
|
return tool
|
|
|
|
@staticmethod
|
|
def _post_responses(
|
|
base_url: str,
|
|
payload: Dict[str, Any],
|
|
api_key: str,
|
|
timeout: float,
|
|
*,
|
|
is_oauth_path: bool,
|
|
) -> tuple[Any, Optional[Dict[str, Any]]]:
|
|
"""POST to ``/responses`` → ``(parsed_json, None)`` or ``(None, failure_envelope)``
|
|
on transport/HTTP/JSON failure.
|
|
|
|
Two attempts: on a first-call 401 with OAuth creds, force-refresh once and
|
|
retry. Covers opaque (non-JWT) tokens the resolver can't pre-check and
|
|
mid-window revocation/rotation. XAI_API_KEY creds can't be refreshed, so
|
|
they skip the retry rather than burn quota.
|
|
"""
|
|
import httpx
|
|
|
|
headers = {
|
|
"Authorization": f"Bearer {api_key}",
|
|
"Content-Type": "application/json",
|
|
"User-Agent": hermes_xai_user_agent(),
|
|
}
|
|
resp = None
|
|
for attempt in range(2):
|
|
try:
|
|
resp = httpx.post(f"{base_url}/responses", headers=headers, json=payload, timeout=timeout)
|
|
resp.raise_for_status()
|
|
break
|
|
except httpx.HTTPStatusError as exc:
|
|
status = exc.response.status_code if exc.response is not None else 0
|
|
if status == 401 and attempt == 0 and is_oauth_path:
|
|
logger.info(
|
|
"xAI web search got 401 on first attempt; forcing OAuth "
|
|
"refresh and retrying once.",
|
|
)
|
|
try:
|
|
refreshed = resolve_xai_http_credentials(force_refresh=True, api_key_hint=api_key)
|
|
refreshed_key = str(refreshed.get("api_key") or "").strip()
|
|
if refreshed_key and refreshed_key != api_key:
|
|
api_key = refreshed_key
|
|
headers["Authorization"] = f"Bearer {api_key}"
|
|
continue
|
|
# Same/empty token back — retrying is pointless; fall through.
|
|
except Exception as refresh_exc: # noqa: BLE001
|
|
logger.warning("xAI web search OAuth refresh after 401 failed: %s", refresh_exc)
|
|
try:
|
|
body = exc.response.text[:300] if exc.response is not None else ""
|
|
except Exception:
|
|
body = ""
|
|
logger.warning("xAI web search HTTP %d: %s", status, body)
|
|
return None, _fail(f"xAI web search returned HTTP {status}: {body}".rstrip())
|
|
except httpx.RequestError as exc:
|
|
logger.warning("xAI web search request error: %s", exc)
|
|
return None, _fail(f"Could not reach xAI: {exc}")
|
|
|
|
if resp is None:
|
|
return None, _fail("xAI web search produced no response")
|
|
|
|
try:
|
|
return resp.json(), None
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("xAI web search bad JSON: %s", exc)
|
|
return None, _fail("Could not parse xAI Responses API reply as JSON")
|
|
|
|
# -- Prompt + parsing -------------------------------------------------
|
|
|
|
@staticmethod
|
|
def _build_prompt(query: str, limit: int) -> str:
|
|
"""Ask Grok for a JSON *object* (cheap to match with ``_JSON_BLOCK_RE``) and
|
|
forbid prose/fences/inline citations to keep the payload parseable."""
|
|
return (
|
|
"Use the web_search tool to find current information for the query below, "
|
|
"then respond with ONLY a single JSON object — no prose, no markdown "
|
|
"fences, no inline citation links — matching this exact schema:\n\n"
|
|
'{"results": [{"title": "string", "url": "string", '
|
|
'"description": "1-2 sentence summary"}]}\n\n'
|
|
f'Return at most {limit} results, ordered by relevance, with absolute '
|
|
"https:// URLs. If no usable results exist, return "
|
|
'{"results": []}.\n\n'
|
|
f"Query: {query}"
|
|
)
|
|
|
|
@classmethod
|
|
def _extract_results(cls, response_data: Dict[str, Any], *, limit: int) -> List[Dict[str, Any]]:
|
|
"""Result rows from a Responses-API reply, in order of preference:
|
|
(1) the JSON object in ``output_text`` blocks, (2) ``url_citation``
|
|
annotations paired with surrounding text, (3) the raw ``citations`` list.
|
|
Only short-circuit on (2) when it yields rows, so future annotation types
|
|
don't mask real data in ``citations``."""
|
|
text_blocks, annotations = cls._collect_output_text(response_data)
|
|
|
|
for block in text_blocks:
|
|
parsed = cls._try_parse_json_results(block, limit=limit)
|
|
if parsed:
|
|
return parsed
|
|
|
|
if annotations:
|
|
annotation_results = cls._results_from_annotations(annotations, "\n".join(text_blocks), limit=limit)
|
|
if annotation_results:
|
|
return annotation_results
|
|
|
|
citations = response_data.get("citations") or []
|
|
if isinstance(citations, list):
|
|
return [
|
|
_row("", str(u), "", i + 1)
|
|
for i, u in enumerate(citations[:limit])
|
|
if isinstance(u, str) and u.strip()
|
|
]
|
|
|
|
return []
|
|
|
|
@staticmethod
|
|
def _collect_output_text(response_data: Dict[str, Any]) -> tuple[List[str], List[Dict[str, Any]]]:
|
|
"""Return (text_blocks, annotations) from ``response.output`` message chunks."""
|
|
text_blocks: List[str] = []
|
|
annotations: List[Dict[str, Any]] = []
|
|
output = response_data.get("output")
|
|
if not isinstance(output, list):
|
|
return text_blocks, annotations
|
|
|
|
for item in output:
|
|
if not isinstance(item, dict) or item.get("type") != "message":
|
|
continue
|
|
content = item.get("content")
|
|
if not isinstance(content, list):
|
|
continue
|
|
for chunk in content:
|
|
if not isinstance(chunk, dict) or chunk.get("type") != "output_text":
|
|
continue
|
|
text = chunk.get("text")
|
|
if isinstance(text, str) and text.strip():
|
|
text_blocks.append(text)
|
|
chunk_annotations = chunk.get("annotations")
|
|
if isinstance(chunk_annotations, list):
|
|
annotations.extend(a for a in chunk_annotations if isinstance(a, dict))
|
|
return text_blocks, annotations
|
|
|
|
@staticmethod
|
|
def _try_parse_json_results(text: str, *, limit: int) -> Optional[List[Dict[str, Any]]]:
|
|
"""Parse a JSON object with a ``results`` array out of ``text``; None when
|
|
absent. Tries the whole string first, then the regex-matched block, since
|
|
reasoning models sometimes prefix narration."""
|
|
candidates = [text]
|
|
match = _JSON_BLOCK_RE.search(text)
|
|
if match and match.group(0) != text:
|
|
candidates.append(match.group(0))
|
|
|
|
for candidate in candidates:
|
|
try:
|
|
parsed = json.loads(candidate)
|
|
except (json.JSONDecodeError, ValueError):
|
|
continue
|
|
if not isinstance(parsed, dict):
|
|
continue
|
|
results = parsed.get("results")
|
|
if not isinstance(results, list):
|
|
continue
|
|
normalized: List[Dict[str, Any]] = []
|
|
for row in results[:limit]:
|
|
if not isinstance(row, dict):
|
|
continue
|
|
url = str(row.get("url", "")).strip()
|
|
if not url:
|
|
continue
|
|
# Renumber from kept rows so a dropped malformed row leaves no gap.
|
|
normalized.append(_row(
|
|
str(row.get("title", "")).strip(), url,
|
|
str(row.get("description", "")).strip(), len(normalized) + 1,
|
|
))
|
|
if normalized:
|
|
return normalized
|
|
return None
|
|
|
|
@staticmethod
|
|
def _results_from_annotations(
|
|
annotations: List[Dict[str, Any]], joined_text: str, *, limit: int
|
|
) -> List[Dict[str, Any]]:
|
|
"""Fallback rows from ``url_citation`` annotations: URL plus ~200 chars of
|
|
preceding text as the description (the annotation title is just a number)."""
|
|
seen: set[str] = set()
|
|
results: List[Dict[str, Any]] = []
|
|
for ann in annotations:
|
|
if ann.get("type") != "url_citation":
|
|
continue
|
|
url = str(ann.get("url", "")).strip()
|
|
if not url or url in seen:
|
|
continue
|
|
seen.add(url)
|
|
|
|
description = ""
|
|
start = ann.get("start_index")
|
|
end = ann.get("end_index")
|
|
if isinstance(start, int) and isinstance(end, int) and 0 <= start < end <= len(joined_text):
|
|
description = joined_text[max(0, start - 200):start].strip()
|
|
if len(description) > 200:
|
|
description = description[-200:].strip()
|
|
|
|
results.append(_row("", url, description, len(results) + 1))
|
|
if len(results) >= limit:
|
|
break
|
|
return results
|
|
|
|
# -- Setup picker -----------------------------------------------------
|
|
|
|
def get_setup_schema(self) -> Dict[str, Any]:
|
|
# Auth resolution is delegated to the shared ``xai_grok`` post_setup hook
|
|
# (same one image_gen.xai / tts.xai use) for a consistent OAuth-or-key prompt.
|
|
return {
|
|
"name": "xAI Web Search (Grok)",
|
|
"badge": "paid",
|
|
"tag": "Agentic web search via Grok's web_search tool — uses xAI Grok OAuth or XAI_API_KEY.",
|
|
"env_vars": [],
|
|
"post_setup": "xai_grok",
|
|
}
|