diff --git a/plugins/video_gen/deepinfra/__init__.py b/plugins/video_gen/deepinfra/__init__.py index b37d256714..c94b1ef2ab 100644 --- a/plugins/video_gen/deepinfra/__init__.py +++ b/plugins/video_gen/deepinfra/__init__.py @@ -1,10 +1,9 @@ """DeepInfra video generation backend. -DeepInfra serves video over the OpenAI-compatible ``/v1/openai/videos`` endpoint (async job: -``create`` → poll → ``download_content``), so all SDK plumbing lives in -:class:`agent.video_gen_provider.OpenAICompatibleVideoGenProvider`. This plugin only declares -identity, credentials, and live model discovery — no hardcoded model ids, so retired models drop -out without a patch. Mirrors ``plugins/image_gen/deepinfra``. +DeepInfra serves video over the OpenAI-compatible ``/v1/openai/videos`` endpoint (``create`` → poll → +``download_content``), so all SDK plumbing lives in :class:`agent.video_gen_provider.OpenAICompatibleVideoGenProvider`. +This plugin only declares identity, credentials, and live model discovery — no hardcoded model ids, so retired +models drop out without a patch. Mirrors ``plugins/image_gen/deepinfra``. """ from __future__ import annotations @@ -21,47 +20,30 @@ class DeepInfraVideoGenProvider(OpenAICompatibleVideoGenProvider): """Text-to-video and image-to-video via DeepInfra's OpenAI-compatible API.""" name = "deepinfra" + display_name = "DeepInfra" _env_key = "DEEPINFRA_API_KEY" _default_base_url = "https://api.deepinfra.com/v1/openai" - @property - def display_name(self) -> str: - return "DeepInfra" - def list_models(self) -> List[Dict[str, Any]]: - """``video-gen``-tagged models from the live catalog; empty when it is - unreachable so the picker shows nothing rather than a retired model.""" + """``video-gen``-tagged models from the live catalog; empty when unreachable (nothing beats a retired model).""" try: from hermes_cli.models import _fetch_deepinfra_models_by_tag except Exception as exc: # noqa: BLE001 — never break the picker logger.debug("Cannot import _fetch_deepinfra_models_by_tag: %s", exc) return [] - return [ - {"id": item["id"], "display": item["id"].split("/")[-1], - "strengths": ((item.get("metadata", {}) or {}).get("description") or "")[:80]} - for item in (_fetch_deepinfra_models_by_tag("video-gen") or []) if item.get("id") - ] + return [{"id": item["id"], "display": item["id"].split("/")[-1], + "strengths": ((item.get("metadata", {}) or {}).get("description") or "")[:80]} + for item in (_fetch_deepinfra_models_by_tag("video-gen") or []) if item.get("id")] def capabilities(self) -> Dict[str, Any]: - return { - "modalities": ["text", "image"], - "aspect_ratios": ["16:9", "9:16", "1:1"], - "resolutions": ["480p", "720p", "1080p"], - "max_duration": 10, "min_duration": 1, - "supports_audio": False, "supports_negative_prompt": True, - "supports_seed": True, "supports_upscale": False, - "max_reference_images": 0, - } + return {"modalities": ["text", "image"], "aspect_ratios": ["16:9", "9:16", "1:1"], "resolutions": ["480p", "720p", "1080p"], + "max_duration": 10, "min_duration": 1, "supports_audio": False, "supports_negative_prompt": True, + "supports_seed": True, "supports_upscale": False, "max_reference_images": 0} def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "DeepInfra", - "badge": "paid", - "tag": "Wan, p-video, … — live catalog from api.deepinfra.com; text-to-video & image-to-video", - "env_vars": [ - {"key": "DEEPINFRA_API_KEY", "prompt": "DeepInfra API key", "url": "https://deepinfra.com/dash/api_keys"}, - ], - } + return {"name": "DeepInfra", "badge": "paid", + "tag": "Wan, p-video, … — live catalog from api.deepinfra.com; text-to-video & image-to-video", + "env_vars": [{"key": "DEEPINFRA_API_KEY", "prompt": "DeepInfra API key", "url": "https://deepinfra.com/dash/api_keys"}]} def register(ctx) -> None: diff --git a/plugins/video_gen/fal/__init__.py b/plugins/video_gen/fal/__init__.py index 6172da9f4c..2bf7b48b2b 100644 --- a/plugins/video_gen/fal/__init__.py +++ b/plugins/video_gen/fal/__init__.py @@ -1,11 +1,9 @@ """FAL.ai video generation backend. -The user picks a **model family** (e.g. "Pixverse v6", "Veo 3.1"); the plugin routes to the -family's text-to-video endpoint when called without ``image_url`` and to its image-to-video -endpoint otherwise (gemini-omni-flash is i2v only). Active-family precedence: tool ``model=`` -arg → ``FAL_VIDEO_MODEL`` env → ``video_gen.fal.model`` → ``video_gen.model`` (a family id or -an endpoint path containing one) → ``DEFAULT_MODEL``. Auth via ``FAL_KEY`` or the managed Nous -gateway. Output is an HTTPS URL from FAL's CDN; the gateway downloads it. +The user picks a **model family** (e.g. "Pixverse v6"); the plugin routes to its text-to-video endpoint without +``image_url`` and to its image-to-video endpoint otherwise (gemini-omni-flash is i2v only). Active-family precedence: +tool ``model=`` → ``FAL_VIDEO_MODEL`` env → ``video_gen.fal.model`` → ``video_gen.model`` (family id or an endpoint +path containing one) → ``DEFAULT_MODEL``. Auth via ``FAL_KEY`` or the managed Nous gateway; output is an HTTPS URL. """ from __future__ import annotations @@ -19,182 +17,106 @@ from agent.video_gen_provider import VideoGenProvider, error_response, success_r logger = logging.getLogger(__name__) - -# Family catalog. Capability flags gate which keys reach the payload — keys a family -# doesn't advertise are never sent (the managed gateway forwards everything verbatim). -# ``_family`` defaults every enum to None and every flag to False. Per-family keys: -# aspect_ratios / resolutions : supported enums (None = endpoint decides) -# durations : enum tuple OR (min, max) range (2 ints with gap > 1) -# audio / audio_native : generate_audio toggle / audio always on (description line only) -# negative / seed : negative_prompt / seed accepted -# duration_int / duration_suffix : send duration as JSON int (default: queue-API string) / "4s" suffix -# image_param_key / image_drop_keys : i2v image key when not `image_url` / keys the i2v endpoint rejects -# resolution_aliases / static_payload : tool-style resolution → endpoint enum / constants always required +# Family catalog. Capability flags gate which keys reach the payload — keys a family doesn't advertise are never sent (the +# managed gateway forwards everything verbatim). Enums default to None (endpoint decides), flags to False. ``durations`` is an +# enum tuple OR a ``(min, max)`` range (2 ints with gap > 1). Extras: audio_native (always on; description line only), +# duration_int (JSON int, default queue-API string), duration_suffix ("4s"), image_param_key (i2v key when not `image_url`), +# image_drop_keys (i2v endpoint rejects), resolution_aliases (tool value → endpoint enum), static_payload (always required). def _family(display: str, speed: str, tier: str, strengths: str, text: Optional[str], image: str, **caps: Any) -> Dict[str, Any]: - return { - "display": display, "speed": speed, "price": tier, "tier": tier, "strengths": strengths, - "text_endpoint": text, "image_endpoint": image, - "aspect_ratios": None, "resolutions": None, "durations": None, "audio": False, "negative": False, "seed": False, - **caps, - } + return {"display": display, "speed": speed, "price": tier, "tier": tier, "strengths": strengths, "text_endpoint": text, "image_endpoint": image, + "aspect_ratios": None, "resolutions": None, "durations": None, "audio": False, "negative": False, "seed": False, **caps} _SIX_ASPECTS = ("21:9", "16:9", "4:3", "1:1", "3:4", "9:16") -# MiniMax H3 uses capitalized/2K-style resolution enums; aliases map the tool's usual values. +# MiniMax H3 uses capitalized/2K-style resolution enums; aliases map the tool's usual values. Max tops out at 768P. _H3_ALIASES = {"480p": "768P", "540p": "768P", "720p": "768P", "768p": "768P", "1080p": "2K", "2k": "2K", "4k": "4K", "2160p": "4K"} _H3_MAX_ALIASES = {"480p": "480P", "540p": "480P", "720p": "768P", "768p": "768P", "1080p": "768P", "2k": "768P", "4k": "768P", "2160p": "768P"} FAL_FAMILIES: Dict[str, Dict[str, Any]] = { # ─── Cheap / fast tier ───────────────────────────────────────────── - "ltx-2.3": _family( # LTX docs expose no duration/aspect/resolution enums. - "LTX 2.3 (22B)", "~30-60s", "cheap", "22B model with native audio generation. Affordable.", - "fal-ai/ltx-2.3-22b/text-to-video", "fal-ai/ltx-2.3-22b/image-to-video", audio=True, negative=True, seed=True, - ), - "pixverse-v6": _family( - "Pixverse v6", "~30-90s", "cheap", "Affordable. Negative prompts. 1-15s durations.", - "fal-ai/pixverse/v6/text-to-video", "fal-ai/pixverse/v6/image-to-video", - resolutions=("360p", "540p", "720p", "1080p"), durations=(1, 15), audio=True, negative=True, seed=True, - ), - "seedance-2.0-mini": _family( - "Seedance 2.0 Mini", "~30-90s", "cheap", "ByteDance. Faster/cheaper Seedance tier, audio + lip-sync, 4-15s.", - "bytedance/seedance-2.0/mini/text-to-video", "bytedance/seedance-2.0/mini/image-to-video", - aspect_ratios=_SIX_ASPECTS, resolutions=("480p", "720p"), durations=(4, 15), audio=True, - ), + "ltx-2.3": _family("LTX 2.3 (22B)", "~30-60s", "cheap", "22B model with native audio generation. Affordable.", # docs expose no enums + "fal-ai/ltx-2.3-22b/text-to-video", "fal-ai/ltx-2.3-22b/image-to-video", audio=True, negative=True, seed=True), + "pixverse-v6": _family("Pixverse v6", "~30-90s", "cheap", "Affordable. Negative prompts. 1-15s durations.", "fal-ai/pixverse/v6/text-to-video", + "fal-ai/pixverse/v6/image-to-video", resolutions=("360p", "540p", "720p", "1080p"), durations=(1, 15), audio=True, negative=True, seed=True), + "seedance-2.0-mini": _family("Seedance 2.0 Mini", "~30-90s", "cheap", "ByteDance. Faster/cheaper Seedance tier, audio + lip-sync, 4-15s.", + "bytedance/seedance-2.0/mini/text-to-video", "bytedance/seedance-2.0/mini/image-to-video", aspect_ratios=_SIX_ASPECTS, + resolutions=("480p", "720p"), durations=(4, 15), audio=True), # ─── Expensive / premium tier ────────────────────────────────────── - "veo3.1": _family( - "Veo 3.1", "~60-120s", "premium", "Google DeepMind. Cinematic, native audio, strong prompt adherence.", - "fal-ai/veo3.1", "fal-ai/veo3.1/image-to-video", - aspect_ratios=("16:9", "9:16"), resolutions=("720p", "1080p", "4k"), - durations=(4, 6, 8), duration_suffix="s", # wants "4s" not "4" - audio=True, negative=True, seed=True, - ), - "seedance-2.0": _family( - "Seedance 2.0", "~60-120s", "premium", "ByteDance. Cinematic, synchronized audio + lip-sync, 4-15s.", - "bytedance/seedance-2.0/text-to-video", "bytedance/seedance-2.0/image-to-video", - # "auto" aspect is deliberately omitted so the agent can't pass it; input schema has no `seed`. - aspect_ratios=_SIX_ASPECTS, resolutions=("480p", "720p", "1080p"), durations=(4, 15), audio=True, - ), - "seedance-2.5": _family( - "Seedance 2.5", "~60-180s", "premium", - "ByteDance flagship. Native 30s single-pass, audio in the same latent space, lip-sync.", - "bytedance/seedance-2.5/text-to-video", "bytedance/seedance-2.5/image-to-video", - image_drop_keys=("aspect_ratio",), # i2v accepts only "auto" aspect (follows the input image) - aspect_ratios=_SIX_ASPECTS, resolutions=("480p", "720p"), durations=(4, 30), audio=True, - ), - "minimax-h3": _family( - "MiniMax H3", "~60-180s", "premium", "MiniMax frontier. Native 2K (up to 4K), 5-15s, seven aspect ratios.", - "minimax/h3/text-to-video", "minimax/h3/image-to-video", - duration_int=True, image_drop_keys=("aspect_ratio",), # i2v derives aspect from the input image - aspect_ratios=_SIX_ASPECTS, resolutions=("768P", "2K", "4K"), resolution_aliases=_H3_ALIASES, - durations=(5, 15), audio_native=True, - ), - "minimax-h3-max": _family( - "MiniMax H3 Max (fal post-train)", "~5-30s", "premium", - "fal's post-trained MiniMax H3. Top-ranked quality/prompt adherence/aesthetics, 768p in seconds, 5-15s.", - "minimax/h3-max/text-to-video", "minimax/h3-max/image-to-video", - duration_int=True, image_drop_keys=("aspect_ratio",), # i2v schema doesn't declare aspect_ratio - # Max tops out at 768P (no 2K/4K tiers like base H3). - aspect_ratios=_SIX_ASPECTS, resolutions=("480P", "768P"), resolution_aliases=_H3_MAX_ALIASES, - durations=(5, 15), - static_payload={"prompt_expansion_mode": "balanced"}, # in the schema's required array - audio_native=True, seed=True, # unlike base H3, Max declares `seed` on both endpoints - ), - "flux-3": _family( - "FLUX 3 (via FAL)", "~60-120s", "premium", "Black Forest Labs frontier video. Native audio, 5-20s, 8 aspect ratios.", - "blackforestlabs/flux-3/text-to-video", "blackforestlabs/flux-3/image-to-video", - duration_int=True, # enum is "auto" | 5..20 as JSON integers - aspect_ratios=("21:9", "2:1", "16:9", "4:3", "1:1", "3:4", "9:16"), resolutions=("720p", "1080p"), - durations=(5, 20), audio=True, - ), - "grok-imagine-1.5": _family( - "Grok Imagine 1.5 (via FAL)", "~30-90s", "premium", "xAI. Fast stylized video with audio, 1-15s, cheap per second.", - "xai/grok-imagine-video/v1.5/text-to-video", "xai/grok-imagine-video/v1.5/image-to-video", - duration_int=True, image_drop_keys=("aspect_ratio",), # t2v-only key; i2v follows the image - aspect_ratios=("16:9", "4:3", "3:2", "1:1", "2:3", "3:4", "9:16"), resolutions=("480p", "720p", "1080p"), - durations=(1, 15), audio_native=True, - ), - "gemini-omni-flash": _family( - "Gemini Omni Flash (via FAL)", "~60-120s", "premium", "Google. Image-to-video with audio, physics-grounded motion, 3-10s.", - None, "google/gemini-omni-flash/image-to-video", # image/reference only on FAL - duration_int=True, aspect_ratios=("16:9", "9:16"), durations=(3, 10), audio_native=True, - ), - "kling-v3-4k": _family( - "Kling v3 4K", "~120-300s", "premium", "4K output, native audio (Chinese/English), 3-15s.", - "fal-ai/kling-video/v3/4k/text-to-video", "fal-ai/kling-video/v3/4k/image-to-video", - image_param_key="start_image_url", aspect_ratios=("16:9", "9:16", "1:1"), durations=(3, 15), - audio=True, negative=True, seed=True, - ), - "happy-horse": _family( - "Happy Horse 1.0", "~60-120s", "premium", "Alibaba. New model, sparse public docs — conservative defaults.", - "alibaba/happy-horse/text-to-video", "alibaba/happy-horse/image-to-video", audio_native=True, seed=True, - ), + "veo3.1": _family("Veo 3.1", "~60-120s", "premium", "Google DeepMind. Cinematic, native audio, strong prompt adherence.", "fal-ai/veo3.1", + "fal-ai/veo3.1/image-to-video", aspect_ratios=("16:9", "9:16"), resolutions=("720p", "1080p", "4k"), durations=(4, 6, 8), + duration_suffix="s", audio=True, negative=True, seed=True), # wants "4s" not "4" + "seedance-2.0": _family("Seedance 2.0", "~60-120s", "premium", "ByteDance. Cinematic, synchronized audio + lip-sync, 4-15s.", # no "auto" aspect, no `seed` + "bytedance/seedance-2.0/text-to-video", "bytedance/seedance-2.0/image-to-video", aspect_ratios=_SIX_ASPECTS, + resolutions=("480p", "720p", "1080p"), durations=(4, 15), audio=True), + "seedance-2.5": _family("Seedance 2.5", "~60-180s", "premium", "ByteDance flagship. Native 30s single-pass, audio in the same latent space, lip-sync.", + "bytedance/seedance-2.5/text-to-video", "bytedance/seedance-2.5/image-to-video", aspect_ratios=_SIX_ASPECTS, + image_drop_keys=("aspect_ratio",), resolutions=("480p", "720p"), durations=(4, 30), audio=True), # i2v aspect is "auto" only + "minimax-h3": _family("MiniMax H3", "~60-180s", "premium", "MiniMax frontier. Native 2K (up to 4K), 5-15s, seven aspect ratios.", + "minimax/h3/text-to-video", "minimax/h3/image-to-video", duration_int=True, image_drop_keys=("aspect_ratio",), # i2v follows image + aspect_ratios=_SIX_ASPECTS, resolutions=("768P", "2K", "4K"), resolution_aliases=_H3_ALIASES, durations=(5, 15), audio_native=True), + # i2v schema doesn't declare aspect_ratio; unlike base H3, Max declares `seed` on both endpoints; static key is in the schema's required array. + "minimax-h3-max": _family("MiniMax H3 Max (fal post-train)", "~5-30s", "premium", "fal's post-trained MiniMax H3. Top-ranked quality/prompt " + "adherence/aesthetics, 768p in seconds, 5-15s.", "minimax/h3-max/text-to-video", "minimax/h3-max/image-to-video", + duration_int=True, image_drop_keys=("aspect_ratio",), aspect_ratios=_SIX_ASPECTS, resolutions=("480P", "768P"), + resolution_aliases=_H3_MAX_ALIASES, durations=(5, 15), static_payload={"prompt_expansion_mode": "balanced"}, audio_native=True, seed=True), + "flux-3": _family("FLUX 3 (via FAL)", "~60-120s", "premium", "Black Forest Labs frontier video. Native audio, 5-20s, 8 aspect ratios.", + "blackforestlabs/flux-3/text-to-video", "blackforestlabs/flux-3/image-to-video", duration_int=True, # enum "auto" | 5..20 ints + aspect_ratios=("21:9", "2:1", "16:9", "4:3", "1:1", "3:4", "9:16"), resolutions=("720p", "1080p"), durations=(5, 20), audio=True), + "grok-imagine-1.5": _family("Grok Imagine 1.5 (via FAL)", "~30-90s", "premium", "xAI. Fast stylized video with audio, 1-15s, cheap per second.", + "xai/grok-imagine-video/v1.5/text-to-video", "xai/grok-imagine-video/v1.5/image-to-video", duration_int=True, + image_drop_keys=("aspect_ratio",), aspect_ratios=("16:9", "4:3", "3:2", "1:1", "2:3", "3:4", "9:16"), # aspect is t2v-only + resolutions=("480p", "720p", "1080p"), durations=(1, 15), audio_native=True), + "gemini-omni-flash": _family("Gemini Omni Flash (via FAL)", "~60-120s", "premium", "Google. Image-to-video with audio, physics-grounded motion, 3-10s.", + None, "google/gemini-omni-flash/image-to-video", duration_int=True, aspect_ratios=("16:9", "9:16"), durations=(3, 10), audio_native=True), + "kling-v3-4k": _family("Kling v3 4K", "~120-300s", "premium", "4K output, native audio (Chinese/English), 3-15s.", "fal-ai/kling-video/v3/4k/text-to-video", + "fal-ai/kling-video/v3/4k/image-to-video", image_param_key="start_image_url", aspect_ratios=("16:9", "9:16", "1:1"), + durations=(3, 15), audio=True, negative=True, seed=True), + "happy-horse": _family("Happy Horse 1.0", "~60-120s", "premium", "Alibaba. New model, sparse public docs — conservative defaults.", + "alibaba/happy-horse/text-to-video", "alibaba/happy-horse/image-to-video", audio_native=True, seed=True), } DEFAULT_MODEL = "pixverse-v6" # cheap, both modalities, sane defaults -_ENDPOINT_MODALITY_LEAVES = frozenset({"text-to-video", "image-to-video"}) - -def _is_duration_range(durations: Tuple[int, ...]) -> bool: - """Heuristic: a 2-tuple of ints with a gap > 1 is treated as ``(min, max)``.""" - return len(durations) == 2 and all(isinstance(d, int) for d in durations) and durations[1] - durations[0] > 1 - - -def _duration_bounds(durations: Tuple[int, ...]) -> Tuple[int, int]: - """``(lo, hi)`` for a non-empty durations spec (range or enum).""" - return (durations[0], durations[1]) if _is_duration_range(durations) else (min(durations), max(durations)) +def _clamp_duration(durations: Tuple[int, ...], duration: Optional[int]) -> Optional[int]: + """Clamp into a ``(min, max)`` range (None stays None: endpoint default) or snap to the nearest enum entry (None → first). + Range heuristic: a 2-tuple of ints with a gap > 1.""" + if len(durations) == 2 and all(isinstance(d, int) for d in durations) and durations[1] - durations[0] > 1: + return None if duration is None else max(durations[0], min(durations[1], duration)) + return durations[0] if duration is None else min(durations, key=lambda d: abs(d - duration)) def _modalities(meta: Dict[str, Any]) -> List[str]: return [m for m in ("text", "image") if meta[f"{m}_endpoint"]] -def _clamp_duration(durations: Tuple[int, ...], duration: Optional[int]) -> Optional[int]: - """Clamp into a range, or snap to the nearest enum entry. ``None`` stays None for - range families (the endpoint applies its own default) but becomes the first enum entry.""" - is_range = _is_duration_range(durations) - if duration is None: - return None if is_range else durations[0] - if is_range: - return max(durations[0], min(durations[1], duration)) - return min(durations, key=lambda d: abs(d - duration)) - - def _normalize_family_key(c: str) -> Optional[str]: - """Extract a known family ID from a bare id, full endpoint path, - truncated endpoint stem (``minimax/h3``) or provider-prefixed name.""" + """Known family ID from a bare id, full endpoint path, truncated stem (``minimax/h3``) or provider-prefixed name.""" c = c.strip() - if not c: - return None - if c in FAL_FAMILIES: - return c + if not c or c in FAL_FAMILIES: + return c or None endpoints = [(fid, ep) for fid, meta in FAL_FAMILIES.items() for ep in (meta["text_endpoint"], meta["image_endpoint"]) if isinstance(ep, str)] # Exact declared endpoint beats any segment scan (which would see "seedance-2.0" inside ".../seedance-2.0/mini/..."). - exact = [fid for fid, ep in endpoints if c == ep] # Truncated stem: the segment after ``c`` must be a modality leaf so "bytedance/seedance-2.0" skips Mini's deeper path. - stem = [fid for fid, ep in endpoints - if ep.startswith(c + "/") and ep[len(c) + 1:].split("/", 1)[0] in _ENDPOINT_MODALITY_LEAVES] - if exact or stem: - return (exact or stem)[0] - # Longest family-id path-segment match (prefers seedance-2.0-mini over seedance-2.0 when both appear). - hits = [fid for fid in FAL_FAMILIES if fid in c.split("/")] - return max(hits, key=len) if hits else None + # Last resort: longest family-id path-segment match (prefers seedance-2.0-mini over seedance-2.0 when both appear). + hit = ([fid for fid, ep in endpoints if c == ep] + or [fid for fid, ep in endpoints if ep.startswith(c + "/") and ep[len(c) + 1:].split("/", 1)[0] in ("text-to-video", "image-to-video")] + or sorted((fid for fid in FAL_FAMILIES if fid in c.split("/")), key=len, reverse=True)) + return hit[0] if hit else None def _resolve_family(explicit: Optional[str]) -> Tuple[str, Dict[str, Any]]: """Decide which FAL family to use. Returns ``(family_id, meta)``.""" import os - try: from hermes_cli.config import load_config - cfg = load_config() - cfg = cfg.get("video_gen") if isinstance(cfg, dict) else None - cfg = cfg if isinstance(cfg, dict) else {} except Exception as exc: logger.debug("Could not load video_gen config: %s", exc) - cfg = {} + cfg = None + cfg = cfg.get("video_gen") if isinstance(cfg, dict) else None + cfg = cfg if isinstance(cfg, dict) else {} fal_cfg = cfg.get("fal") if isinstance(cfg.get("fal"), dict) else {} for c in (explicit, os.environ.get("FAL_VIDEO_MODEL"), fal_cfg.get("model"), cfg.get("model")): fid = _normalize_family_key(c) if isinstance(c, str) else None @@ -203,38 +125,26 @@ def _resolve_family(explicit: Optional[str]) -> Tuple[str, Dict[str, Any]]: return DEFAULT_MODEL, FAL_FAMILIES[DEFAULT_MODEL] -def _build_payload( - family: Dict[str, Any], *, prompt: str, image_url: Optional[str], duration: Optional[int], aspect_ratio: str, - resolution: str, negative_prompt: Optional[str], audio: Optional[bool], seed: Optional[int], -) -> Dict[str, Any]: - """Build a family-specific payload, dropping keys the family doesn't declare.""" - payload: Dict[str, Any] = {} - if prompt: - payload["prompt"] = prompt - if image_url: - payload[family.get("image_param_key") or "image_url"] = image_url - # Newer endpoints declare no `seed` and the managed gateway forwards whatever we send — gate on the family. - if seed is not None and family.get("seed", True): - payload["seed"] = seed - # Unsupported aspect/resolution values are dropped so the endpoint defaults. - if family["aspect_ratios"] and aspect_ratio in family["aspect_ratios"]: - payload["aspect_ratio"] = aspect_ratio +def _build_payload(family: Dict[str, Any], *, prompt: str, image_url: Optional[str], duration: Optional[int], aspect_ratio: str, + resolution: str, negative_prompt: Optional[str], audio: Optional[bool], seed: Optional[int]) -> Dict[str, Any]: + """Build a family-specific payload, dropping keys the family doesn't declare (unsupported enums → endpoint default).""" resolved = (family.get("resolution_aliases") or {}).get((resolution or "").lower(), resolution) - if family["resolutions"] and resolved in family["resolutions"]: - payload["resolution"] = resolved clamped = _clamp_duration(family["durations"], duration) if family["durations"] else None - if clamped is not None: - # FAL's queue API types duration as a string ("8" not 8) unless the family says int; - # some families (veo3.1) also need a unit suffix ("4s" not "4"). - payload["duration"] = clamped if family.get("duration_int") else f"{clamped}{family.get('duration_suffix', '')}" - if family["audio"] and audio is not None: - payload["generate_audio"] = bool(audio) - if family["negative"] and negative_prompt: - payload["negative_prompt"] = negative_prompt - # Keys the i2v endpoint rejects outright, then constants it always requires. - for key in family.get("image_drop_keys", ()) if image_url else (): + payload: Dict[str, Any] = {key: value for ok, key, value in ( + (prompt, "prompt", prompt), + (image_url, family.get("image_param_key") or "image_url", image_url), + # Newer endpoints declare no `seed` and the managed gateway forwards whatever we send — gate on the family. + (seed is not None and family.get("seed", True), "seed", seed), + (family["aspect_ratios"] and aspect_ratio in family["aspect_ratios"], "aspect_ratio", aspect_ratio), + (family["resolutions"] and resolved in family["resolutions"], "resolution", resolved), + # FAL's queue API types duration as a string ("8" not 8) unless the family says int; veo3.1 also wants a unit suffix. + (clamped is not None, "duration", clamped if family.get("duration_int") else f"{clamped}{family.get('duration_suffix', '')}"), + (family["audio"] and audio is not None, "generate_audio", bool(audio)), + (family["negative"] and negative_prompt, "negative_prompt", negative_prompt), + ) if ok} + for key in family.get("image_drop_keys", ()) if image_url else (): # keys the i2v endpoint rejects outright payload.pop(key, None) - for key, value in (family.get("static_payload") or {}).items(): + for key, value in (family.get("static_payload") or {}).items(): # constants the endpoint always requires payload.setdefault(key, value) return payload @@ -246,8 +156,6 @@ def _video_url_from_result(result: Any) -> Tuple[Any, Optional[str]]: return video, url or None -# ---- fal_client lazy import + managed FAL gateway (Nous Subscription) --------- - _fal_client: Any = None _fal_client_lock = threading.Lock() @@ -262,27 +170,20 @@ def _load_fal_client() -> Any: with _fal_client_lock: if _fal_client is None: from tools.fal_common import import_fal_client - _fal_client = import_fal_client() return _fal_client def _resolve_managed_fal_video_gateway(): - """Resolve the FAL video route from the stored ``video_gen`` selection. - - ``"nous"`` → managed only (unentitled ⇒ selection-naming error); any other stored provider → - direct only (missing FAL_KEY ⇒ selection-naming error); never-configured → legacy autodetect. - """ + """Resolve the FAL video route from the stored ``video_gen`` selection: ``"nous"`` → managed only (unentitled ⇒ + selection-naming error); other stored provider → direct only (missing FAL_KEY ⇒ error); never-configured → autodetect.""" from tools.managed_tool_gateway import resolve_managed_tool_gateway from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER, fal_key_is_configured, read_selection, selection_error - selected = read_selection("video_gen") if selected == NOUS_MANAGED_PROVIDER: gateway = resolve_managed_tool_gateway("fal-queue") if gateway is None: - raise ValueError(selection_error( - "video_gen", NOUS_MANAGED_PROVIDER, "the Nous Tool Gateway is not available (not entitled or unreachable)", - )) + raise ValueError(selection_error("video_gen", NOUS_MANAGED_PROVIDER, "the Nous Tool Gateway is not available (not entitled or unreachable)")) return gateway if selected is not None: if not fal_key_is_configured(): @@ -291,28 +192,21 @@ def _resolve_managed_fal_video_gateway(): return None if fal_key_is_configured() else resolve_managed_tool_gateway("fal-queue") -def _check_fal_video_available() -> bool: - """True if the selected (or, never-configured, any) FAL backend is reachable. Never raises on a - stored-but-broken selection — the honest selection-naming error surfaces at call time.""" +def _fal_video_available() -> bool: + """True if the selected (or, never-configured, any) FAL backend is reachable; raises on a stored-but-broken selection.""" from tools.tool_backend_helpers import fal_key_is_configured - - try: - return _resolve_managed_fal_video_gateway() is not None or fal_key_is_configured() - except ValueError: - return False + return _resolve_managed_fal_video_gateway() is not None or fal_key_is_configured() def _get_managed_fal_video_client(managed_gateway): """Reuse the managed FAL client so its internal httpx.Client is not leaked per call.""" global _managed_fal_video_client, _managed_fal_video_client_config from tools.fal_common import _ManagedFalSyncClient - client_config = (managed_gateway.gateway_origin.rstrip("/"), managed_gateway.nous_user_token) with _managed_fal_video_client_lock: if _managed_fal_video_client is None or _managed_fal_video_client_config != client_config: - _managed_fal_video_client = _ManagedFalSyncClient( - _load_fal_client(), key=managed_gateway.nous_user_token, queue_run_origin=managed_gateway.gateway_origin, - ) + _managed_fal_video_client = _ManagedFalSyncClient(_load_fal_client(), key=managed_gateway.nous_user_token, + queue_run_origin=managed_gateway.gateway_origin) _managed_fal_video_client_config = client_config return _managed_fal_video_client @@ -328,15 +222,11 @@ def _submit_fal_video_request(endpoint: str, arguments: Dict[str, Any]): return _get_managed_fal_video_client(managed_gateway).submit(endpoint, arguments=arguments, headers=headers) except Exception as exc: from tools.fal_common import _extract_http_status - status = _extract_http_status(exc) if status is not None and 400 <= status < 500: - raise ValueError( - f"Nous Subscription gateway rejected endpoint '{endpoint}' (HTTP {status}). This model may not yet " - f"be enabled on the Nous Portal's FAL proxy. Either:\n" - f" • Set FAL_KEY in your environment to use FAL.ai directly, or\n" - f" • Pick a different model via `hermes tools` → Video Generation." - ) from exc + raise ValueError(f"Nous Subscription gateway rejected endpoint '{endpoint}' (HTTP {status}). This model may not yet be enabled " + f"on the Nous Portal's FAL proxy. Either:\n • Set FAL_KEY in your environment to use FAL.ai directly, or\n" + f" • Pick a different model via `hermes tools` → Video Generation.") from exc raise @@ -364,17 +254,11 @@ def _upscale_video(video_url: str, source_request_id: Optional[str] = None) -> O return url -# ---- Provider --------------------------------------------------------------- - -_NO_BACKEND_MSG = ( - "No FAL backend available. Either set FAL_KEY (run `hermes tools` → Video Generation → FAL to configure) " - "or sign in to Nous (`hermes setup`) for managed gateway access." -) +_NO_BACKEND_MSG = ("No FAL backend available. Either set FAL_KEY (run `hermes tools` → Video Generation → FAL to configure) " + "or sign in to Nous (`hermes setup`) for managed gateway access.") _MODALITY_MISSING_MSG = { - "image": "FAL family {fid} has no image-to-video endpoint. Pick a family with image-to-video support via " - "`hermes tools` → Video Generation.", - "text": "FAL family {fid} has no text-to-video endpoint. Pass an image_url to use its image-to-video endpoint, " - "or pick a different family.", + "image": "FAL family {fid} has no image-to-video endpoint. Pick a family with image-to-video support via `hermes tools` → Video Generation.", + "text": "FAL family {fid} has no text-to-video endpoint. Pass an image_url to use its image-to-video endpoint, or pick a different family.", } @@ -385,104 +269,72 @@ def _fal_error(error: str, error_type: str, prompt: str, model: str = "", aspect class FALVideoGenProvider(VideoGenProvider): """FAL.ai multi-family backend; routes t2v/i2v on ``image_url`` presence.""" - @property - def name(self) -> str: - return "fal" - - @property - def display_name(self) -> str: - return "FAL" + name = "fal" + display_name = "FAL" def is_available(self) -> bool: + # A stored-but-broken selection raises the selection-naming ValueError; report unavailable, never break the picker. try: - return _check_fal_video_available() - except Exception: # noqa: BLE001 — never break the picker + return _fal_video_available() + except Exception: # noqa: BLE001 return False def list_models(self) -> List[Dict[str, Any]]: - out: List[Dict[str, Any]] = [] - for fid, meta in FAL_FAMILIES.items(): - entry: Dict[str, Any] = {"id": fid, **{k: meta[k] for k in ("display", "speed", "strengths", "price", "tier")}, - "modalities": _modalities(meta)} - if meta["durations"]: - entry["min_duration"], entry["max_duration"] = _duration_bounds(meta["durations"]) - out.append(entry) - return out + return [{"id": fid, **{k: meta[k] for k in ("display", "speed", "strengths", "price", "tier")}, "modalities": _modalities(meta), + **({"min_duration": min(d), "max_duration": max(d)} if (d := meta["durations"]) else {})} + for fid, meta in FAL_FAMILIES.items()] def default_model(self) -> Optional[str]: return DEFAULT_MODEL def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "FAL", "badge": "paid", - "tag": "LTX, Pixverse, Seedance 2.0/2.5/Mini, Veo 3.1, MiniMax H3, FLUX 3, Kling 4K, Happy Horse, Grok Imagine, " - "Gemini Omni — text-to-video & image-to-video", - "env_vars": [{"key": "FAL_KEY", "prompt": "FAL.ai API key", "url": "https://fal.ai/dashboard/keys"}], - } + return {"name": "FAL", "badge": "paid", "env_vars": [{"key": "FAL_KEY", "prompt": "FAL.ai API key", "url": "https://fal.ai/dashboard/keys"}], + "tag": "LTX, Pixverse, Seedance 2.0/2.5/Mini, Veo 3.1, MiniMax H3, FLUX 3, Kling 4K, Happy Horse, Grok Imagine, " + "Gemini Omni — text-to-video & image-to-video"} def capabilities(self) -> Dict[str, Any]: - # Report the RESOLVED family's surface so the dynamic tool schema gates params on what - # the selected model honors; fall back to the cross-family union if resolution fails (never raises). + # RESOLVED family's surface so the dynamic tool schema gates params on what the selected model honors; union fallback (never raises). try: _family_id, family = _resolve_family(None) except Exception: # noqa: BLE001 family = None if family: - lo, hi = _duration_bounds(family["durations"] or (1, 1)) - return { - "modalities": _modalities(family) or ["text"], - "aspect_ratios": list(family["aspect_ratios"] or []), "resolutions": list(family["resolutions"] or []), - "max_duration": hi, "min_duration": lo, "supports_audio": bool(family["audio"]), - "audio_always_on": bool(family.get("audio_native")), # no toggle: description line, not a param - "supports_negative_prompt": bool(family["negative"]), - "supports_seed": bool(family["seed"]), - "supports_upscale": True, # SeedVR chains for any family - "max_reference_images": 0, - } - bounds = [_duration_bounds(m["durations"]) for m in FAL_FAMILIES.values() if m["durations"]] - return { - "modalities": ["text", "image"], "aspect_ratios": ["16:9", "9:16", "1:1"], - "resolutions": ["360p", "540p", "720p", "1080p"], "max_duration": max([1] + [hi for _lo, hi in bounds]), - "min_duration": min([lo for lo, _hi in bounds], default=1), "supports_audio": True, - "supports_negative_prompt": True, "supports_seed": True, "supports_upscale": True, "max_reference_images": 0, - } + durations = family["durations"] or (1, 1) + return {"modalities": _modalities(family) or ["text"], "aspect_ratios": list(family["aspect_ratios"] or []), + "resolutions": list(family["resolutions"] or []), "max_duration": max(durations), "min_duration": min(durations), + "supports_audio": bool(family["audio"]), "audio_always_on": bool(family.get("audio_native")), # no toggle: description only + "supports_negative_prompt": bool(family["negative"]), "supports_seed": bool(family["seed"]), + "supports_upscale": True, "max_reference_images": 0} # SeedVR chains for any family + spans = [m["durations"] for m in FAL_FAMILIES.values() if m["durations"]] + return {"modalities": ["text", "image"], "aspect_ratios": ["16:9", "9:16", "1:1"], "resolutions": ["360p", "540p", "720p", "1080p"], + "max_duration": max([1] + [max(d) for d in spans]), "min_duration": min([min(d) for d in spans], default=1), + "supports_audio": True, "supports_negative_prompt": True, "supports_seed": True, "supports_upscale": True, "max_reference_images": 0} def generate( - self, prompt: str, *, model: Optional[str] = None, image_url: Optional[str] = None, - reference_image_urls: Optional[List[str]] = None, duration: Optional[int] = None, - aspect_ratio: str = "16:9", resolution: str = "720p", negative_prompt: Optional[str] = None, + self, prompt: str, *, model: Optional[str] = None, image_url: Optional[str] = None, reference_image_urls: Optional[List[str]] = None, + duration: Optional[int] = None, aspect_ratio: str = "16:9", resolution: str = "720p", negative_prompt: Optional[str] = None, audio: Optional[bool] = None, seed: Optional[int] = None, upscale: Optional[bool] = None, **kwargs: Any, ) -> Dict[str, Any]: - if not _check_fal_video_available(): - from tools.tool_backend_helpers import read_selection - - if read_selection("video_gen") is not None: - # A stored selection that cannot run gets the honest selection-naming error from the strict resolver. - try: - _resolve_managed_fal_video_gateway() - except ValueError as exc: - return _fal_error(str(exc), "auth_required", prompt) - return _fal_error(_NO_BACKEND_MSG, "auth_required", prompt) + try: # a stored selection that cannot run gets the honest selection-naming error from the strict resolver + if not _fal_video_available(): + return _fal_error(_NO_BACKEND_MSG, "auth_required", prompt) + except ValueError as exc: + return _fal_error(str(exc), "auth_required", prompt) try: _load_fal_client() except ImportError: return _fal_error("fal_client Python package not installed (pip install fal-client)", "missing_dependency", prompt) - prompt = (prompt or "").strip() family_id, family = _resolve_family(model) image_url_norm = (image_url or "").strip() or None modality_used = "image" if image_url_norm else "text" # routes to the i2v vs t2v endpoint endpoint = family[f"{modality_used}_endpoint"] if not endpoint: - msg = _MODALITY_MISSING_MSG[modality_used].format(fid=family_id) - return _fal_error(msg, "modality_unsupported", prompt, model=family_id) + return _fal_error(_MODALITY_MISSING_MSG[modality_used].format(fid=family_id), "modality_unsupported", prompt, model=family_id) if not prompt: return _fal_error("prompt is required.", "missing_prompt", prompt, model=family_id) - - payload = _build_payload( - family, prompt=prompt, image_url=image_url_norm, duration=duration, aspect_ratio=aspect_ratio, - resolution=resolution, negative_prompt=negative_prompt, audio=audio, seed=seed, - ) + payload = _build_payload(family, prompt=prompt, image_url=image_url_norm, duration=duration, aspect_ratio=aspect_ratio, + resolution=resolution, negative_prompt=negative_prompt, audio=audio, seed=seed) try: handle = _submit_fal_video_request(endpoint, payload) source_request_id = getattr(handle, "request_id", None) @@ -492,24 +344,19 @@ class FALVideoGenProvider(VideoGenProvider): return _fal_error(f"FAL video generation failed: {exc}", "api_error", prompt, model=family_id, aspect_ratio=aspect_ratio) if not url: return _fal_error("FAL returned no video URL in response", "empty_response", prompt, model=family_id) - # Optional SeedVR2 pass — explicit opt-in, best-effort: failure falls back to the native video. upscaled_url = _upscale_video(url, source_request_id) if upscale else None upscaled = bool(upscaled_url) if upscale and not upscaled: logger.warning("Video upscale pass failed — returning native-resolution video") - extra: Dict[str, Any] = {"endpoint": endpoint, "upscaled": upscaled} - if upscaled: - url, extra["upscale_factor"] = upscaled_url, UPSCALER_FACTOR - if isinstance(video, dict): - extra.update({k: video[k] for k in ("file_size", "content_type") if video.get(k)}) - if upscaled: - extra.pop("file_size", None) # native-resolution size no longer applies + url = upscaled_url or url + extra: Dict[str, Any] = {"endpoint": endpoint, "upscaled": upscaled, **({"upscale_factor": UPSCALER_FACTOR} if upscaled else {})} + if isinstance(video, dict): # native-resolution file_size no longer applies after an upscale + extra.update({k: video[k] for k in (("content_type",) if upscaled else ("file_size", "content_type")) if video.get(k)}) return success_response( - video=url, model=family_id, prompt=prompt, modality=modality_used, + video=url, model=family_id, prompt=prompt, modality=modality_used, provider="fal", extra=extra, aspect_ratio=aspect_ratio if "aspect_ratio" in payload else "", duration=int("".join(c for c in str(payload["duration"]) if c.isdigit()) or "0") if "duration" in payload else 0, - provider="fal", extra=extra, ) diff --git a/plugins/video_gen/xai/__init__.py b/plugins/video_gen/xai/__init__.py index bd0baac3e9..7e59fe66f8 100644 --- a/plugins/video_gen/xai/__init__.py +++ b/plugins/video_gen/xai/__init__.py @@ -1,13 +1,10 @@ """xAI Grok-Imagine video generation backend. -Surface: text-, image- and reference-to-video through the unified video provider; xAI -edit/extend are exposed by ``tools.xai_video_tools`` via ``run_xai_video_edit`` / ``run_xai_video_extend``. - -Authentication: xAI Grok OAuth tokens (preferred — billed to the user's SuperGrok / X Premium+ -subscription) or ``XAI_API_KEY``, both via ``tools.xai_http.resolve_xai_http_credentials`` so one -login covers chat + TTS + image gen + video gen + transcription. When xAI storage is enabled, the -primary ``video`` / ``public_url`` fields are the stored files-cdn HTTPS link; pass that public MP4 -URL as ``video_url`` for edit/extend (sent to xAI as ``video.url``). +Text-, image- and reference-to-video through the unified video provider; edit/extend are exposed by +``tools.xai_video_tools`` via ``run_xai_video_edit`` / ``run_xai_video_extend``. Auth: xAI Grok OAuth tokens +(preferred — billed to the user's SuperGrok / X Premium+ subscription) or ``XAI_API_KEY``, both via +``tools.xai_http.resolve_xai_http_credentials``. With xAI storage enabled the primary ``video`` / ``public_url`` +fields are the stored files-cdn HTTPS link; pass it as ``video_url`` for edit/extend (sent as ``video.url``). """ from __future__ import annotations @@ -18,8 +15,9 @@ import logging import mimetypes import os import uuid +from contextlib import closing from pathlib import Path -from typing import Any, Callable, Coroutine, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple import httpx @@ -27,17 +25,12 @@ from agent.video_gen_provider import VideoGenProvider, error_response, success_r logger = logging.getLogger(__name__) - DEFAULT_XAI_BASE_URL = "https://api.x.ai/v1" -DEFAULT_TEXT_TO_VIDEO_MODEL = "grok-imagine-video" +DEFAULT_TEXT_TO_VIDEO_MODEL = DEFAULT_MODEL = "grok-imagine-video" DEFAULT_IMAGE_TO_VIDEO_MODEL = "grok-imagine-video-1.5" -DEFAULT_MODEL = DEFAULT_TEXT_TO_VIDEO_MODEL -DEFAULT_DURATION = 8 -DEFAULT_ASPECT_RATIO = "16:9" -DEFAULT_RESOLUTION = "720p" -DEFAULT_TIMEOUT_SECONDS = 240 -DEFAULT_POLL_INTERVAL_SECONDS = 5 -DEFAULT_EXTEND_DURATION = 6 +DEFAULT_DURATION, DEFAULT_EXTEND_DURATION = 8, 6 +DEFAULT_ASPECT_RATIO, DEFAULT_RESOLUTION = "16:9", "720p" +DEFAULT_TIMEOUT_SECONDS, DEFAULT_POLL_INTERVAL_SECONDS = 240, 5 VALID_ASPECT_RATIOS = {"1:1", "16:9", "9:16", "4:3", "3:4", "3:2", "2:3"} VALID_RESOLUTIONS = {"480p", "720p"} @@ -45,7 +38,10 @@ MAX_REFERENCE_IMAGES = 7 _REMOTE_PREFIXES = ("http://", "https://") _TERMINAL_POLL_STATUSES = {"done", "failed", "error", "expired", "cancelled"} - +_IMAGE_TO_VIDEO_COMPAT_MODEL_IDS = {"grok-imagine-video-1.5-preview", "grok-imagine-video-1.5-2026-05-30"} +_AUTH_REQUIRED_MSG = ("No xAI credentials found. Sign in via `hermes auth add xai-oauth` " + "(SuperGrok / Premium+) or set XAI_API_KEY from https://console.x.ai/.") +_PUBLIC_URL_HINT = "(e.g. the `image`/`public_url` from a prior Imagine result)" _MODELS: Dict[str, Dict[str, Any]] = { "grok-imagine-video": { "display": "Grok Imagine Video", "speed": "~60-240s", "strengths": "Text-to-video; legacy image-to-video fallback.", @@ -57,65 +53,41 @@ _MODELS: Dict[str, Dict[str, Any]] = { }, } -_IMAGE_TO_VIDEO_COMPAT_MODEL_IDS = {"grok-imagine-video-1.5-preview", "grok-imagine-video-1.5-2026-05-30"} -_AUTH_REQUIRED_MSG = ( - "No xAI credentials found. Sign in via `hermes auth add xai-oauth` " - "(SuperGrok / Premium+) or set XAI_API_KEY from https://console.x.ai/." -) -_PUBLIC_URL_HINT = "(e.g. the `image`/`public_url` from a prior Imagine result)" - - -# ---- Credentials / HTTP helpers ------------------------------------------- +def _xai_http(helper: str, fallback: Any, *args: Any, log: Optional[str] = None) -> Any: + """``tools.xai_http.(*args)``, or ``fallback`` when it is unavailable or raises (never breaks video gen).""" + try: + import tools.xai_http as xai_http + return getattr(xai_http, helper)(*args) + except Exception as exc: + if log: + logger.debug(log, exc) + return fallback def _resolve_xai_credentials() -> Tuple[str, str]: - """Return ``(api_key, base_url)`` from the shared xAI credential resolver. - - Order: runtime provider (xai-oauth pool entry) → singleton ``auth.json`` OAuth tokens → - ``XAI_API_KEY`` env var. ``api_key`` is empty when no source is available; callers must check. - """ - try: - from tools.xai_http import resolve_xai_http_credentials - - creds = resolve_xai_http_credentials() or {} - except Exception as exc: - logger.debug("xAI credential resolver failed: %s", exc) - creds = {} + """``(api_key, base_url)``: runtime xai-oauth pool entry → ``auth.json`` OAuth tokens → ``XAI_API_KEY`` (empty key = none; callers check).""" + creds = _xai_http("resolve_xai_http_credentials", {}, log="xAI credential resolver failed: %s") or {} base_url = str(creds.get("base_url") or os.getenv("XAI_BASE_URL") or DEFAULT_XAI_BASE_URL) return str(creds.get("api_key") or os.getenv("XAI_API_KEY", "")).strip(), base_url.strip().rstrip("/") def _xai_headers(api_key: str) -> Dict[str, str]: - try: - from tools.xai_http import hermes_xai_user_agent - - user_agent = hermes_xai_user_agent() - except Exception: - user_agent = "hermes-agent/video_gen" - return {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json", "User-Agent": user_agent} + return {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json", + "User-Agent": _xai_http("hermes_xai_user_agent", "hermes-agent/video_gen")} def _xai_error(error: str, error_type: str, prompt: str, model: str = "", aspect_ratio: str = "") -> Dict[str, Any]: return error_response(error=error, error_type=error_type, provider="xai", model=model, prompt=prompt, aspect_ratio=aspect_ratio) -# ---- Input normalization -------------------------------------------------- - - def _media_ref_to_xai_url(value: str, *, kind: str, fallback_mime: str) -> str: - """Return a URL/data URI accepted by xAI for ``kind`` (``image``/``video``) inputs. - - Remote URLs and matching data URIs pass through; a readable local file of the right MIME - class is inlined as base64; anything else is returned as-is so the caller's URL check rejects - it clearly. Local reads go through Hermes' read deny-list (same credential-store guard as the - image providers), which fails open if its machinery is unavailable. - """ + """URL/data URI accepted by xAI for ``kind`` (``image``/``video``) inputs: remote URLs and matching data URIs pass + through; a readable local file of the right MIME class is inlined as base64 (after Hermes' read deny-list / + credential-store guard, which fails open if unavailable); anything else is returned as-is so the caller rejects it.""" ref = (value or "").strip() - if not ref or ref.lower().startswith(_REMOTE_PREFIXES + (f"data:{kind}/",)): - return ref path = Path(ref).expanduser() - if not path.is_file(): + if not ref or ref.lower().startswith(_REMOTE_PREFIXES + (f"data:{kind}/",)) or not path.is_file(): return ref try: from agent.file_safety import raise_if_read_blocked @@ -124,9 +96,7 @@ def _media_ref_to_xai_url(value: str, *, kind: str, fallback_mime: str) -> str: else: raise_if_read_blocked(ref) mime = mimetypes.guess_type(path.name)[0] or fallback_mime - if not mime.startswith(f"{kind}/"): - return ref - return f"data:{mime};base64,{base64.b64encode(path.read_bytes()).decode('ascii')}" + return f"data:{mime};base64,{base64.b64encode(path.read_bytes()).decode('ascii')}" if mime.startswith(f"{kind}/") else ref def _image_ref_to_xai_input(value: str) -> Optional[Dict[str, str]]: @@ -138,47 +108,33 @@ async def _video_input_from_public_url(value: str, *, api_key: str, base_url: st """Build xAI ``video`` input using a public HTTPS URL (``url`` field only).""" ref = (value or "").strip() if ref and Path(ref).expanduser().is_file(): - ref = _media_ref_to_xai_url(ref, kind="video", fallback_mime="video/mp4") - return {"url": ref} if ref else None + return {"url": _media_ref_to_xai_url(ref, kind="video", fallback_mime="video/mp4")} return {"url": ref} if ref.lower().startswith(_REMOTE_PREFIXES) else None def _clamp_duration(duration: Optional[int], *, has_reference_images: bool = False, max_seconds: int = 15, default: int = DEFAULT_DURATION) -> int: """Clamp to ``[1, max_seconds]``; reference-to-video additionally caps at 10s.""" - value = max(1, min(max_seconds, duration if duration is not None else default)) - return min(value, 10) if has_reference_images else value + return min(max(1, min(max_seconds, duration if duration is not None else default)), 10 if has_reference_images else max_seconds) def _resolve_model_for_modality(model: Optional[str], *, modality: str, explicit_model: bool) -> str: - """Select xAI's text/video model without treating config as a prompt override. - - ``grok-imagine-video-1.5`` rejects text-only generation but is the desired image-to-video - backend. Explicit tool ``model=`` still wins for users who intentionally request another model. - """ + """Select xAI's text/video model without treating config as a prompt override: ``grok-imagine-video-1.5`` + rejects text-only generation but is the desired image-to-video backend; explicit tool ``model=`` still wins.""" requested = (model or "").strip() if explicit_model and requested: return requested if modality == "image": return DEFAULT_IMAGE_TO_VIDEO_MODEL - if requested == DEFAULT_IMAGE_TO_VIDEO_MODEL or requested in _IMAGE_TO_VIDEO_COMPAT_MODEL_IDS: - return DEFAULT_TEXT_TO_VIDEO_MODEL - return requested or DEFAULT_TEXT_TO_VIDEO_MODEL - - -# ---- Provider --------------------------------------------------------------- + is_i2v_id = requested == DEFAULT_IMAGE_TO_VIDEO_MODEL or requested in _IMAGE_TO_VIDEO_COMPAT_MODEL_IDS + return DEFAULT_TEXT_TO_VIDEO_MODEL if is_i2v_id or not requested else requested class XAIVideoGenProvider(VideoGenProvider): """xAI Grok Imagine video backend.""" - @property - def name(self) -> str: - return "xai" - - @property - def display_name(self) -> str: - return "xAI" + name = "xai" + display_name = "xAI" def is_available(self) -> bool: return has_xai_video_credentials() @@ -190,187 +146,118 @@ class XAIVideoGenProvider(VideoGenProvider): return DEFAULT_MODEL def get_setup_schema(self) -> Dict[str, Any]: - # Auth resolution lives in the shared ``xai_grok`` post_setup hook (hermes_cli/tools_config.py) so the - # picker doesn't prompt for an API key when already signed in via xAI Grok OAuth; the hook offers an - # OAuth-vs-API-key choice when neither is configured. - try: - from tools.xai_http import xai_storage_notice_text - - storage_notice = xai_storage_notice_text("video_gen") - except Exception: - storage_notice = "" - tag = ( - "grok-imagine-video for text/reference; grok-imagine-video-1.5 for image-to-video; edit/extend: pass " - "the stored public HTTPS MP4 (`video` / `public_url` from a prior Imagine result); uses xAI Grok OAuth " - "or XAI_API_KEY" - ) - if storage_notice: - tag += f". {storage_notice}" + # Auth resolution lives in the shared ``xai_grok`` post_setup hook (hermes_cli/tools_config.py): no API-key + # prompt when already signed in via xAI Grok OAuth; OAuth-vs-API-key choice when neither is configured. + storage_notice = _xai_http("xai_storage_notice_text", "", "video_gen") + tag = ("grok-imagine-video for text/reference; grok-imagine-video-1.5 for image-to-video; edit/extend: pass the stored public " + "HTTPS MP4 (`video` / `public_url` from a prior Imagine result); uses xAI Grok OAuth or XAI_API_KEY" + ) + (f". {storage_notice}" if storage_notice else "") return {"name": "xAI Grok Imagine", "badge": "paid", "tag": tag, "env_vars": [], "post_setup": "xai_grok"} def capabilities(self) -> Dict[str, Any]: - return { - "modalities": ["text", "image"], "aspect_ratios": sorted(VALID_ASPECT_RATIOS), - "resolutions": sorted(VALID_RESOLUTIONS), "max_duration": 15, "min_duration": 1, - "supports_audio": False, "supports_negative_prompt": False, "supports_seed": True, - "supports_upscale": False, "max_reference_images": MAX_REFERENCE_IMAGES, - } + return {"modalities": ["text", "image"], "aspect_ratios": sorted(VALID_ASPECT_RATIOS), "resolutions": sorted(VALID_RESOLUTIONS), + "max_duration": 15, "min_duration": 1, "supports_audio": False, "supports_negative_prompt": False, "supports_seed": True, + "supports_upscale": False, "max_reference_images": MAX_REFERENCE_IMAGES} def generate( self, prompt: str, *, model: Optional[str] = None, image_url: Optional[str] = None, - reference_image_urls: Optional[List[str]] = None, duration: Optional[int] = None, - aspect_ratio: str = DEFAULT_ASPECT_RATIO, resolution: str = DEFAULT_RESOLUTION, - negative_prompt: Optional[str] = None, audio: Optional[bool] = None, seed: Optional[int] = None, - **kwargs: Any, + reference_image_urls: Optional[List[str]] = None, duration: Optional[int] = None, aspect_ratio: str = DEFAULT_ASPECT_RATIO, + resolution: str = DEFAULT_RESOLUTION, negative_prompt: Optional[str] = None, audio: Optional[bool] = None, + seed: Optional[int] = None, **kwargs: Any, ) -> Dict[str, Any]: - return _run_xai_video_coroutine( - lambda api_key, base_url: _generate_xai_video_async( - api_key=api_key, base_url=base_url, prompt=prompt, model=model, - explicit_model=bool(kwargs.get("_model_override_explicit")), image_url=image_url, - reference_image_urls=reference_image_urls, duration=duration, aspect_ratio=aspect_ratio, - resolution=resolution, - ), - operation_label="generation", model=model, prompt=prompt, aspect_ratio=aspect_ratio, + return _run_xai_video( + "generation", _generate_xai_video_async, prompt=prompt, model=model, + explicit_model=bool(kwargs.get("_model_override_explicit")), image_url=image_url, + reference_image_urls=reference_image_urls, duration=duration, aspect_ratio=aspect_ratio, resolution=resolution, ) -# ---- Sync entry points (provider + tools.xai_video_tools) ------------------- - - def has_xai_video_credentials() -> bool: return bool(_resolve_xai_credentials()[0]) def run_xai_video_edit(*, prompt: str, video_url: str, model: Optional[str] = None) -> Dict[str, Any]: - return _run_xai_video_mutation(prompt, video_url, model, endpoint="edits", operation="edit", duration=DEFAULT_DURATION) + return _run_xai_video("edit", _mutate_xai_video_async, prompt=prompt, video_url=video_url, model=model, + endpoint="edits", operation="edit", duration=DEFAULT_DURATION) -def run_xai_video_extend(*, prompt: str, video_url: str, duration: Optional[int] = None, - model: Optional[str] = None) -> Dict[str, Any]: - return _run_xai_video_mutation( - prompt, video_url, model, endpoint="extensions", operation="extend", - duration=_clamp_duration(duration, max_seconds=10, default=DEFAULT_EXTEND_DURATION), - ) +def run_xai_video_extend(*, prompt: str, video_url: str, duration: Optional[int] = None, model: Optional[str] = None) -> Dict[str, Any]: + return _run_xai_video("extend", _mutate_xai_video_async, prompt=prompt, video_url=video_url, model=model, + endpoint="extensions", operation="extend", + duration=_clamp_duration(duration, max_seconds=10, default=DEFAULT_EXTEND_DURATION)) -def _run_xai_video_mutation(prompt: str, video_url: str, model: Optional[str], *, endpoint: str, operation: str, - duration: int) -> Dict[str, Any]: - return _run_xai_video_coroutine( - lambda api_key, base_url: _mutate_xai_video_async( - api_key=api_key, base_url=base_url, prompt=prompt, video_url=video_url, model=model, - endpoint=endpoint, operation=operation, duration=duration, - ), - operation_label=operation, model=model, prompt=prompt, aspect_ratio=DEFAULT_ASPECT_RATIO, - ) - - -def _run_xai_video_coroutine( - start: Callable[[str, str], Coroutine[Any, Any, Dict[str, Any]]], *, operation_label: str, - model: Optional[str], prompt: str, aspect_ratio: str, -) -> Dict[str, Any]: - """Resolve credentials, then drive ``start(api_key, base_url)`` on a fresh event loop; - any escaped exception → api_error response.""" +def _run_xai_video(label: str, flow, /, **kwargs: Any) -> Dict[str, Any]: + """Resolve credentials, then drive ``flow(api_key=, base_url=, **kwargs)`` on a fresh event loop; escaped exception → api_error.""" + prompt, model = kwargs["prompt"], kwargs["model"] api_key, base_url = _resolve_xai_credentials() if not api_key: return _xai_error(_AUTH_REQUIRED_MSG, "auth_required", prompt) try: - loop = asyncio.new_event_loop() - try: - return loop.run_until_complete(start(api_key, base_url)) - finally: - loop.close() + with closing(asyncio.new_event_loop()) as loop: + return loop.run_until_complete(flow(api_key=api_key, base_url=base_url, **kwargs)) except Exception as exc: - logger.warning("xAI video %s unexpected failure: %s", operation_label, exc, exc_info=True) - return _xai_error( - f"xAI video {operation_label} failed: {exc}", "api_error", prompt, - model=model or DEFAULT_MODEL, aspect_ratio=aspect_ratio, - ) - - -# ---- Async flows ------------------------------------------------------------ + logger.warning("xAI video %s unexpected failure: %s", label, exc, exc_info=True) + return _xai_error(f"xAI video {label} failed: {exc}", "api_error", prompt, + model=model or DEFAULT_MODEL, aspect_ratio=kwargs.get("aspect_ratio", DEFAULT_ASPECT_RATIO)) async def _generate_xai_video_async( *, api_key: str, base_url: str, prompt: str, model: Optional[str], explicit_model: bool, image_url: Optional[str], - reference_image_urls: Optional[List[str]], duration: Optional[int], aspect_ratio: str, resolution: str, -) -> Dict[str, Any]: - prompt = (prompt or "").strip() - image_input = _image_ref_to_xai_input(image_url) if (image_url or "").strip() else None - if (image_url or "").strip() and not image_input: - return _xai_error(f"image_url must be a public HTTPS URL or data URI {_PUBLIC_URL_HINT}", "invalid_image_url", prompt) - aspect_ratio = (aspect_ratio or DEFAULT_ASPECT_RATIO).strip() - resolution = (resolution or DEFAULT_RESOLUTION).strip().lower() + reference_image_urls: Optional[List[str]], duration: Optional[int], aspect_ratio: str, resolution: str) -> Dict[str, Any]: + prompt, image_url = (prompt or "").strip(), (image_url or "").strip() + image_input = _image_ref_to_xai_input(image_url) if image_url else None refs = [_image_ref_to_xai_input(url.strip()) for url in reference_image_urls or [] if (url or "").strip()] - if not all(refs): - return _xai_error( - f"reference_image_urls must be public HTTPS URLs or data URIs {_PUBLIC_URL_HINT}", "invalid_reference_image_urls", prompt, - ) - if not prompt: - return _xai_error("prompt is required for xAI video generation", "missing_prompt", prompt) - if len(refs) > MAX_REFERENCE_IMAGES: - return _xai_error( - f"reference_image_urls supports at most {MAX_REFERENCE_IMAGES} images on xAI", "too_many_references", prompt, - ) - if image_input and refs: - return _xai_error("image_url and reference_image_urls cannot be combined on xAI", "conflicting_inputs", prompt) - + for bad, message, error_type in ( # validation order is part of the contract + (image_url and not image_input, f"image_url must be a public HTTPS URL or data URI {_PUBLIC_URL_HINT}", "invalid_image_url"), + (not all(refs), f"reference_image_urls must be public HTTPS URLs or data URIs {_PUBLIC_URL_HINT}", "invalid_reference_image_urls"), + (not prompt, "prompt is required for xAI video generation", "missing_prompt"), + (len(refs) > MAX_REFERENCE_IMAGES, f"reference_image_urls supports at most {MAX_REFERENCE_IMAGES} images on xAI", "too_many_references"), + (image_input and refs, "image_url and reference_image_urls cannot be combined on xAI", "conflicting_inputs"), + ): + if bad: + return _xai_error(message, error_type, prompt) # Unsupported values silently fall back to defaults rather than erroring. + aspect_ratio = (aspect_ratio or DEFAULT_ASPECT_RATIO).strip() aspect_ratio = aspect_ratio if aspect_ratio in VALID_ASPECT_RATIOS else DEFAULT_ASPECT_RATIO + resolution = (resolution or DEFAULT_RESOLUTION).strip().lower() resolution = resolution if resolution in VALID_RESOLUTIONS else DEFAULT_RESOLUTION - modality_used = "reference" if refs else ("image" if image_input else "text") resolved_model = _resolve_model_for_modality(model, modality=modality_used, explicit_model=explicit_model) # Reference-to-video only exists on the text model: explicit other model = error, implicit (config) = corrected. if refs and resolved_model != DEFAULT_TEXT_TO_VIDEO_MODEL: if explicit_model: - return _xai_error( - f"xAI reference-to-video requires {DEFAULT_TEXT_TO_VIDEO_MODEL}; got {resolved_model}", - "unsupported_model", prompt, model=resolved_model, - ) + return _xai_error(f"xAI reference-to-video requires {DEFAULT_TEXT_TO_VIDEO_MODEL}; got {resolved_model}", + "unsupported_model", prompt, model=resolved_model) resolved_model = DEFAULT_TEXT_TO_VIDEO_MODEL - clamped_duration = _clamp_duration(duration, has_reference_images=bool(refs)) - payload = {"model": resolved_model, "prompt": prompt, "duration": clamped_duration, "aspect_ratio": aspect_ratio, - "resolution": resolution} - payload.update({k: v for k, v in (("image", image_input), ("reference_images", refs)) if v}) - return await _submit_xai_video_payload( - api_key=api_key, base_url=base_url, endpoint="generations", payload=payload, - prompt=prompt, resolved_model=resolved_model, modality=modality_used, - aspect_ratio=aspect_ratio, duration=clamped_duration, operation="generate", resolution=resolution, - ) + payload = {"model": resolved_model, "prompt": prompt, "duration": clamped_duration, "aspect_ratio": aspect_ratio, "resolution": resolution, + **{k: v for k, v in (("image", image_input), ("reference_images", refs)) if v}} + return await _submit_xai_video_payload(api_key, base_url, "generations", payload, modality=modality_used, operation="generate", + aspect_ratio=aspect_ratio, duration=clamped_duration, resolution=resolution) -async def _mutate_xai_video_async( - *, api_key: str, base_url: str, prompt: str, video_url: str, model: Optional[str], endpoint: str, operation: str, - duration: int, -) -> Dict[str, Any]: +async def _mutate_xai_video_async(*, api_key: str, base_url: str, prompt: str, video_url: str, model: Optional[str], endpoint: str, + operation: str, duration: int) -> Dict[str, Any]: """Edit or extend using a public HTTPS ``video_url`` input (``url`` on the wire).""" prompt = (prompt or "").strip() video_input = await _video_input_from_public_url(video_url or "", api_key=api_key, base_url=base_url) if not prompt: return _xai_error("prompt is required for xAI video edit/extend", "missing_prompt", prompt) if not video_input: - msg = "video_url must be a public HTTPS MP4 URL (the `video`/`public_url` from a prior Imagine result)" - return _xai_error(msg, "missing_video", prompt) - resolved_model = _resolve_model_for_modality(model, modality="text", explicit_model=bool(model)) - payload: Dict[str, Any] = {"model": resolved_model, "prompt": prompt, "video": video_input} - if endpoint == "extensions": - payload["duration"] = duration - return await _submit_xai_video_payload( - api_key=api_key, base_url=base_url, endpoint=endpoint, payload=payload, - prompt=prompt, resolved_model=resolved_model, modality=operation, - aspect_ratio=DEFAULT_ASPECT_RATIO, duration=duration, operation=operation, - ) + return _xai_error("video_url must be a public HTTPS MP4 URL (the `video`/`public_url` from a prior Imagine result)", "missing_video", prompt) + payload: Dict[str, Any] = {"model": _resolve_model_for_modality(model, modality="text", explicit_model=bool(model)), "prompt": prompt, + "video": video_input, **({"duration": duration} if endpoint == "extensions" else {})} + return await _submit_xai_video_payload(api_key, base_url, endpoint, payload, modality=operation, operation=operation, + aspect_ratio=DEFAULT_ASPECT_RATIO, duration=duration) -async def _submit_xai_video_payload( - *, api_key: str, base_url: str, endpoint: str, payload: Dict[str, Any], prompt: str, resolved_model: str, - modality: str, aspect_ratio: str, duration: int, operation: str, resolution: Optional[str] = None, -) -> Dict[str, Any]: +async def _submit_xai_video_payload(api_key: str, base_url: str, endpoint: str, payload: Dict[str, Any], *, modality: str, operation: str, + aspect_ratio: str, duration: int, resolution: Optional[str] = None) -> Dict[str, Any]: """POST ``payload`` to ``/videos/{endpoint}``, poll ``/videos/{request_id}`` to a terminal status, shape the response.""" + prompt, resolved_model = payload["prompt"], payload["model"] try: from tools.xai_http import build_xai_storage_options, maybe_mark_xai_storage_notice_seen, read_xai_imagine_storage_config - storage_options = build_xai_storage_options("video_gen", filename_prefix="hermes-xai-video", extension="mp4") storage_notice = maybe_mark_xai_storage_notice_seen("video_gen") storage_cfg = read_xai_imagine_storage_config("video_gen") @@ -378,30 +265,22 @@ async def _submit_xai_video_payload( storage_options, storage_notice, storage_cfg = None, None, {"enabled": False} if storage_options is not None: payload["storage_options"] = storage_options - headers = _xai_headers(api_key) async with httpx.AsyncClient() as client: try: - response = await client.post( - f"{base_url}/videos/{endpoint}", headers={**headers, "x-idempotency-key": str(uuid.uuid4())}, - json=payload, timeout=60, - ) + response = await client.post(f"{base_url}/videos/{endpoint}", headers={**headers, "x-idempotency-key": str(uuid.uuid4())}, + json=payload, timeout=60) response.raise_for_status() except httpx.HTTPStatusError as exc: - detail = "" try: detail = exc.response.text[:500] except Exception: - pass - return _xai_error( - f"xAI submit failed ({exc.response.status_code}): {detail or exc}", "api_error", prompt, model=resolved_model, - ) + detail = "" + return _xai_error(f"xAI submit failed ({exc.response.status_code}): {detail or exc}", "api_error", prompt, model=resolved_model) request_id = response.json().get("request_id") if not request_id: raise RuntimeError("xAI video response did not include request_id") - - elapsed = 0.0 - status, body = "queued", {} + elapsed, status, body = 0.0, "queued", {} while elapsed < DEFAULT_TIMEOUT_SECONDS: response = await client.get(f"{base_url}/videos/{request_id}", headers=headers, timeout=30) response.raise_for_status() @@ -412,48 +291,26 @@ async def _submit_xai_video_payload( await asyncio.sleep(DEFAULT_POLL_INTERVAL_SECONDS) elapsed += DEFAULT_POLL_INTERVAL_SECONDS else: - return _xai_error( - f"Timed out waiting for xAI video request after {DEFAULT_TIMEOUT_SECONDS}s", "timeout", prompt, - model=resolved_model, - ) - + return _xai_error(f"Timed out waiting for xAI video request after {DEFAULT_TIMEOUT_SECONDS}s", "timeout", prompt, + model=resolved_model) if status != "done": message = (body.get("error", {}) or {}).get("message") or body.get("message") return _xai_error(message or f"xAI video request ended with status '{status}'", f"xai_{status}", prompt, model=resolved_model) - video = body.get("video") if isinstance(body.get("video"), dict) else {} - # Primary URL is the stored files-cdn HTTPS MP4 (``public_url``) when storage is enabled, else xAI's - # temporary ``video.url``; pass it as ``video_url`` for edit/extend chaining. The temporary URL is only - # reported when it differs from the stored one. - file_output = video.get("file_output") - file_output = file_output if isinstance(file_output, dict) else {} - stored_public = file_output.get("public_url") - stored_public = stored_public.strip() if isinstance(stored_public, str) else None - temporary = video.get("url") - temporary = temporary.strip() if isinstance(temporary, str) else None + # Primary URL is the stored files-cdn HTTPS MP4 (``public_url``) when storage is enabled, else xAI's temporary + # ``video.url``; pass it as ``video_url`` for edit/extend chaining. The temporary URL is only reported when it differs. + file_output = video.get("file_output") if isinstance(video.get("file_output"), dict) else {} + stored_public, temporary = (v.strip() if isinstance(v, str) else None for v in (file_output.get("public_url"), video.get("url"))) public_video_url = stored_public or temporary or "" if not public_video_url: - return _xai_error( - "xAI video request completed without a video URL", "empty_response", prompt, - model=body.get("model") or resolved_model, - ) + return _xai_error("xAI video request completed without a video URL", "empty_response", prompt, model=body.get("model") or resolved_model) extra: Dict[str, Any] = {"request_id": request_id, "operation": operation, "storage_enabled": bool(storage_cfg.get("enabled"))} - if resolution: - extra["resolution"] = resolution - if storage_notice: - extra["storage_notice"] = storage_notice - if stored_public: - extra["public_url"] = stored_public - if temporary and temporary != stored_public: - extra["temporary_url"] = temporary + extra.update({k: v for k, v in (("resolution", resolution), ("storage_notice", storage_notice), ("public_url", stored_public), + ("temporary_url", stored_public and temporary != stored_public and temporary)) if v}) extra.update({k: file_output[k] for k in ("filename", "expires_at", "public_url_expires_at", "public_url_error", "storage_error") - if k in file_output}) - if body.get("usage"): - extra["usage"] = body["usage"] - return success_response( - video=public_video_url, model=body.get("model") or resolved_model, prompt=prompt, modality=modality, - aspect_ratio=aspect_ratio, duration=video.get("duration") or duration, provider="xai", extra=extra, - ) + if k in file_output}, **({"usage": body["usage"]} if body.get("usage") else {})) + return success_response(video=public_video_url, model=body.get("model") or resolved_model, prompt=prompt, modality=modality, + aspect_ratio=aspect_ratio, duration=video.get("duration") or duration, provider="xai", extra=extra) def register(ctx) -> None: diff --git a/plugins/web/__init__.py b/plugins/web/__init__.py index ad557e1774..def8570ddd 100644 --- a/plugins/web/__init__.py +++ b/plugins/web/__init__.py @@ -1,7 +1,3 @@ -# Bundled web search providers — plugins/web/. -# -# Each subdirectory follows the image_gen plugin layout: -# plugins/web//{plugin.yaml, __init__.py, provider.py} -# -# They auto-load via kind: backend and register via +# Bundled web search providers: plugins/web//{plugin.yaml, __init__.py, provider.py} +# (image_gen plugin layout). They auto-load via kind: backend and register via # ctx.register_web_search_provider() into agent.web_search_registry. diff --git a/plugins/web/_common.py b/plugins/web/_common.py index f5fba3b161..3576c399bc 100644 --- a/plugins/web/_common.py +++ b/plugins/web/_common.py @@ -20,25 +20,20 @@ SEARCH_LIMIT_CAP = 20 # every vendor here caps max_results at 20 server-side def provider_env(name: str) -> str: """Config-aware env lookup (os.environ, then ~/.hermes/.env).""" from agent.web_search_provider import get_provider_env - return get_provider_env(name) def use_keyless(name: str, api_key: str) -> bool: from plugins.web.keyless_mcp import use_keyless as _use_keyless - return _use_keyless(name, api_key) def _interrupted() -> bool: from tools.interrupt import is_interrupted - return is_interrupted() # --- Result shapes (key order is part of the contract — it reaches the model as JSON) --- - - def search_ok(web_results: List[Dict[str, Any]]) -> Dict[str, Any]: return {"success": True, "data": {"web": web_results}} @@ -51,33 +46,45 @@ def web_hit(url: str, title: str, description: str, position: int) -> Dict[str, return {"url": url, "title": title, "description": description, "position": position} +def title_hit(title: str, url: str, description: str, position: int) -> Dict[str, Any]: + """Title-first row — the historical wire shape of brave/searxng/ddgs/tavily/xai.""" + return {"title": title, "url": url, "description": description, "position": position} + + def document(url: str, title: str, content: str, *, source_url: Optional[str] = None) -> Dict[str, Any]: """Successful extract entry; ``raw_content`` mirrors ``content`` for the legacy pipeline.""" return { - "url": url, - "title": title, - "content": content, - "raw_content": content, + "url": url, "title": title, "content": content, "raw_content": content, "metadata": {"sourceURL": url if source_url is None else source_url, "title": title}, } +def page_error(url: str, error: str) -> Dict[str, Any]: + return {"url": url, "title": "", "content": "", "error": error} + + def extract_fail(urls: List[str], error: str) -> List[Dict[str, Any]]: - return [{"url": u, "title": "", "content": "", "error": error} for u in urls] + return [page_error(u, error) for u in urls] + + +# --- Keyless ring hand-off (shared by exa / parallel / keenable) --------------- +def keyless_search(display: str, name: str, query: str, limit: int, logger: logging.Logger) -> Dict[str, Any]: + from plugins.web.keyless_mcp import search_with_failover + logger.info("%s keyless search: '%s' (limit=%d)", display, query, limit) + return search_with_failover(name, query, limit) + + +def keyless_extract(display: str, name: str, urls: List[str], logger: logging.Logger) -> List[Dict[str, Any]]: + from plugins.web.keyless_mcp import extract_with_failover + logger.info("%s keyless extract: %d URL(s)", display, len(urls)) + return extract_with_failover(name, list(urls)) # --- Guarded execution: interrupt check + uniform failure classification --- - - -def _failure_message( - vendor: str, kind: str, exc: Exception, logger: logging.Logger, *, sdk: bool, verbatim_value_error: bool -) -> str: - """Map an exception to the user-facing error string. - - ``verbatim_value_error``: ValueError carries a pre-formatted message (missing - key, HTTP body) and is returned as-is. ``sdk``: ImportError means the lazily - installed vendor SDK is missing. Anything else is logged and wrapped. - """ +def _failure_message(vendor: str, kind: str, exc: Exception, logger: logging.Logger, *, sdk: bool, verbatim_value_error: bool) -> str: + """``verbatim_value_error``: ValueError carries a pre-formatted message (missing key, + HTTP body) and is returned as-is. ``sdk``: ImportError means the lazily installed + vendor SDK is missing. Anything else is logged and wrapped.""" if verbatim_value_error and isinstance(exc, ValueError): return str(exc) if sdk and isinstance(exc, ImportError): @@ -86,15 +93,17 @@ def _failure_message( return f"{vendor} {kind} failed: {exc}" -def run_search( - vendor: str, logger: logging.Logger, body: Callable[[], Dict[str, Any]], *, sdk: bool = False, verbatim_value_error: bool = True -) -> Dict[str, Any]: +def _guarded(vendor: str, kind: str, logger: logging.Logger, body: Callable[[], Any], interrupted: Any, fail: Callable[[str], Any], sdk: bool, vve: bool) -> Any: try: if _interrupted(): - return search_fail("Interrupted") + return interrupted return body() - except Exception as exc: # noqa: BLE001 — surface as failure dict - return search_fail(_failure_message(vendor, "search", exc, logger, sdk=sdk, verbatim_value_error=verbatim_value_error)) + except Exception as exc: # noqa: BLE001 — surface as failure shape + return fail(_failure_message(vendor, kind, exc, logger, sdk=sdk, verbatim_value_error=vve)) + + +def run_search(vendor: str, logger: logging.Logger, body: Callable[[], Dict[str, Any]], *, sdk: bool = False, verbatim_value_error: bool = True) -> Dict[str, Any]: + return _guarded(vendor, "search", logger, body, search_fail("Interrupted"), search_fail, sdk, verbatim_value_error) def _extract_interrupted(urls: List[str]) -> List[Dict[str, Any]]: @@ -106,18 +115,14 @@ def run_extract( *, sdk: bool = False, verbatim_value_error: bool = True, ) -> List[Dict[str, Any]]: """Per-URL failures are returned as entries with ``error`` — never raised.""" - try: - if _interrupted(): - return _extract_interrupted(urls) - return body() - except Exception as exc: # noqa: BLE001 - return extract_fail(urls, _failure_message(vendor, "extract", exc, logger, sdk=sdk, verbatim_value_error=verbatim_value_error)) + return _guarded(vendor, "extract", logger, body, _extract_interrupted(urls), lambda m: extract_fail(urls, m), sdk, verbatim_value_error) async def run_extract_async( vendor: str, logger: logging.Logger, urls: List[str], body: Callable[[], Awaitable[List[Dict[str, Any]]]], *, sdk: bool = False, verbatim_value_error: bool = True, ) -> List[Dict[str, Any]]: + """Async twin of :func:`run_extract` (``body`` is awaited inside the guard).""" try: if _interrupted(): return _extract_interrupted(urls) @@ -127,8 +132,6 @@ async def run_extract_async( # --- HTTP + SDK client helpers --- - - def http_status_detail(response: Any) -> str: """Response body text for a >=400 reply, or ``HTTP `` when the body is empty.""" return (response.text or "").strip() or f"HTTP {response.status_code}" @@ -138,11 +141,8 @@ def http_get_json( label: str, url: str, *, params: Dict[str, Any], headers: Dict[str, str], timeout: int, logger: logging.Logger, reach_target: Optional[str] = None, ) -> tuple[Any, Optional[Dict[str, Any]]]: - """GET ``url`` and parse JSON. Returns ``(data, None)`` or ``(None, failure_dict)``. - - ``label`` names the vendor in error strings; ``reach_target`` overrides the - "Could not reach ..." subject (SearXNG includes its instance URL). - """ + """GET ``url`` and parse JSON → ``(data, None)`` or ``(None, failure_dict)``. + ``reach_target`` overrides the "Could not reach ..." subject (SearXNG includes its URL).""" try: resp = httpx.get(url, params=params, headers=headers, timeout=timeout) resp.raise_for_status() @@ -159,51 +159,49 @@ def http_get_json( return None, search_fail(f"Could not parse {label} response as JSON") -def cached_sdk_client(slot: str, env_var: str, missing_key_error: str, feature: str, factory: Callable[[str], Any]) -> Any: - """Lazy-build + cache a vendor SDK client on ``tools.web_tools.``. +def titled_rows(raw_results: List[Dict[str, Any]], description_key: str) -> List[Dict[str, Any]]: + """Brave/SearXNG row normalizer: ``str()`` every field, 1-based positions.""" + return [ + title_hit(str(r.get("title", "")), str(r.get("url", "")), str(r.get(description_key, "")), i + 1) + for i, r in enumerate(raw_results) + ] - The cache slot lives on ``tools.web_tools`` so tests that reset - ``tools.web_tools.__client = None`` between cases see fresh state. - Raises ValueError when the key is unset; lazy_deps install hints are - re-raised as ImportError (its own ImportError is benign and swallowed). - """ - import tools.web_tools as _wt - - cached = getattr(_wt, slot, None) - if cached is not None: - return cached - - api_key = provider_env(env_var) - if not api_key: - raise ValueError(missing_key_error) +def lazy_ensure(feature: str) -> None: + """Best-effort ``tools.lazy_deps.ensure``: its own ImportError is benign and swallowed; + an install hint (any other error) is re-raised as ImportError.""" try: from tools.lazy_deps import ensure as _lazy_ensure - _lazy_ensure(feature, prompt=False) except ImportError: pass except Exception as exc: # noqa: BLE001 raise ImportError(str(exc)) + +def cached_sdk_client(slot: str, env_var: str, missing_key_error: str, feature: str, factory: Callable[[str], Any]) -> Any: + """Lazy-build + cache a vendor SDK client on ``tools.web_tools.`` (so tests that + reset ``tools.web_tools.__client = None`` see fresh state). Raises ValueError + when the key is unset.""" + import tools.web_tools as _wt + cached = getattr(_wt, slot, None) + if cached is not None: + return cached + api_key = provider_env(env_var) + if not api_key: + raise ValueError(missing_key_error) + lazy_ensure(feature) client = factory(api_key) setattr(_wt, slot, client) return client # --- Provider base --- - - class BaseWebSearchProvider(WebSearchProvider): - """Common surface for the bundled providers. - - Subclasses set ``NAME`` / ``DISPLAY_NAME`` / ``KEY_ENV`` and flip - ``EXTRACT`` / ``KEYLESS``. ``is_available`` deliberately ignores the - keyless tier: otherwise the legacy preference walk would route users - holding a key for a lower-priority backend onto this vendor's free tier. - ``is_keyless_available`` is True for ring members and opt-in keyless - vendors unless the user pinned ``web.provider_tier.: paid``. - """ + """Subclasses set ``NAME`` / ``DISPLAY_NAME`` / ``KEY_ENV`` and flip ``EXTRACT`` / ``KEYLESS``. + ``is_available`` deliberately ignores the keyless tier: otherwise the legacy preference walk + would route users holding a key for a lower-priority backend onto this vendor's free tier. + ``is_keyless_available`` is True for keyless vendors unless pinned ``web.provider_tier.: paid``.""" NAME: str = "" DISPLAY_NAME: str = "" @@ -211,20 +209,14 @@ class BaseWebSearchProvider(WebSearchProvider): EXTRACT: bool = False KEYLESS: bool = False - @property - def name(self) -> str: - return self.NAME - - @property - def display_name(self) -> str: - return self.DISPLAY_NAME + name = property(lambda self: self.NAME) + display_name = property(lambda self: self.DISPLAY_NAME) def is_available(self) -> bool: return bool(provider_env(self.KEY_ENV)) def is_keyless_available(self) -> bool: from plugins.web.keyless_mcp import keyless_enabled, provider_tier - return self.KEYLESS and keyless_enabled() and provider_tier(self.NAME) != "paid" def supports_search(self) -> bool: @@ -234,21 +226,13 @@ class BaseWebSearchProvider(WebSearchProvider): return self.EXTRACT +def setup_schema(name: str, badge: str, tag: str, key_env: str = "", prompt: str = "", url: str = "", **extra: Any) -> Dict[str, Any]: + """``hermes tools`` picker entry; ``env_vars`` is empty when ``key_env`` is blank.""" + env_vars = [{"key": key_env, "prompt": prompt, "url": url}] if key_env else [] + return {"name": name, "badge": badge, "tag": tag, "env_vars": env_vars, **extra} + + def keyless_variant_schema(display: str, key_env: str, key_url: str, *, free_tag: str, paid_tag: str) -> Dict[str, Any]: - """``hermes tools`` picker entry for a keyless-ring vendor with a paid variant.""" - return { - "name": f"{display} · Free (keyless)", - "badge": "free · no key", - "tag": free_tag, - "env_vars": [], - "web_tier": "free", - "variants": [ - { - "name": f"{display} · Paid (API key)", - "badge": "paid", - "tag": paid_tag, - "env_vars": [{"key": key_env, "prompt": f"{display} API key", "url": key_url}], - "web_tier": "paid", - }, - ], - } + """Picker entry for a keyless-ring vendor with a paid variant.""" + paid = setup_schema(f"{display} · Paid (API key)", "paid", paid_tag, key_env, f"{display} API key", key_url, web_tier="paid") + return setup_schema(f"{display} · Free (keyless)", "free · no key", free_tag, web_tier="free", variants=[paid]) diff --git a/plugins/web/brave_free/__init__.py b/plugins/web/brave_free/__init__.py index abd05807b9..e908a6e8d2 100644 --- a/plugins/web/brave_free/__init__.py +++ b/plugins/web/brave_free/__init__.py @@ -1,7 +1,5 @@ """Brave Search (free tier) plugin — bundled, auto-loaded.""" - from __future__ import annotations - from plugins.web.brave_free.provider import BraveFreeWebSearchProvider diff --git a/plugins/web/brave_free/provider.py b/plugins/web/brave_free/provider.py index 1ffd6a6c15..924763d4ba 100644 --- a/plugins/web/brave_free/provider.py +++ b/plugins/web/brave_free/provider.py @@ -10,7 +10,7 @@ from __future__ import annotations import logging from typing import Any, Dict -from plugins.web._common import BaseWebSearchProvider, http_get_json, provider_env, search_fail, search_ok +from plugins.web._common import BaseWebSearchProvider, http_get_json, provider_env, search_fail, search_ok, setup_schema, titled_rows logger = logging.getLogger(__name__) @@ -28,44 +28,21 @@ class BraveFreeWebSearchProvider(BaseWebSearchProvider): api_key = provider_env("BRAVE_SEARCH_API_KEY") if not api_key: return search_fail("BRAVE_SEARCH_API_KEY is not set") - data, failure = http_get_json( - "Brave Search", - _BRAVE_ENDPOINT, + "Brave Search", _BRAVE_ENDPOINT, params={"q": query, "count": max(1, min(int(limit), 20))}, # Brave caps count at 20 headers={"X-Subscription-Token": api_key, "Accept": "application/json"}, - timeout=15, - logger=logger, + timeout=15, logger=logger, ) if failure is not None: return failure - raw_results = (data.get("web") or {}).get("results", []) or [] - web_results = [ - { - "title": str(r.get("title", "")), - "url": str(r.get("url", "")), - "description": str(r.get("description", "")), - "position": i + 1, - } - for i, r in enumerate(raw_results[:limit]) - ] - logger.info( - "Brave Search '%s': %d results (from %d raw, limit %d)", - query, len(web_results), len(raw_results), limit, - ) + web_results = titled_rows(raw_results[:limit], "description") + logger.info("Brave Search '%s': %d results (from %d raw, limit %d)", query, len(web_results), len(raw_results), limit) return search_ok(web_results) def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "Brave Search (Free)", - "badge": "free", - "tag": "Free-tier API key — 2k queries/mo, search only.", - "env_vars": [ - { - "key": "BRAVE_SEARCH_API_KEY", - "prompt": "Brave Search API key (free tier)", - "url": "https://brave.com/search/api/", - }, - ], - } + return setup_schema( + "Brave Search (Free)", "free", "Free-tier API key — 2k queries/mo, search only.", + "BRAVE_SEARCH_API_KEY", "Brave Search API key (free tier)", "https://brave.com/search/api/", + ) diff --git a/plugins/web/ddgs/__init__.py b/plugins/web/ddgs/__init__.py index ee14d558fc..f6ad740c9f 100644 --- a/plugins/web/ddgs/__init__.py +++ b/plugins/web/ddgs/__init__.py @@ -1,7 +1,5 @@ """DuckDuckGo search plugin (``ddgs`` package, optional dep) — bundled, auto-loaded.""" - from __future__ import annotations - from plugins.web.ddgs.provider import DDGSWebSearchProvider diff --git a/plugins/web/ddgs/_search_worker.py b/plugins/web/ddgs/_search_worker.py index 9ac3795160..f4243d46f8 100644 --- a/plugins/web/ddgs/_search_worker.py +++ b/plugins/web/ddgs/_search_worker.py @@ -2,8 +2,8 @@ Reads one JSON request ``{"query": str, "safe_limit": int}`` from stdin, writes one envelope ``{"ok": true, "results": [...]}`` / ``{"ok": false, "error": str}`` to -stdout, exits. Test hooks (``"test_hook": "sleep"|"gil"|"success"|"error"|"empty"``) -are honored only when ``HERMES_DDGS_ALLOW_TEST_HOOKS=1``. +stdout, exits. Test hooks (``"test_hook": "sleep"|"gil"|"empty"``) are honored only +when ``HERMES_DDGS_ALLOW_TEST_HOOKS=1``. """ from __future__ import annotations @@ -18,42 +18,23 @@ def _hold_gil(secs: int) -> None: """Block in a foreign call that keeps the GIL — mirrors native ``primp``. ``ctypes.PyDLL`` (unlike ``CDLL``/``WinDLL``) does not release the GIL.""" import ctypes - if sys.platform == "win32": - lib = ctypes.PyDLL("kernel32") - lib.Sleep.argtypes = [ctypes.c_uint] - lib.Sleep(int(secs * 1000)) - return - - lib = ctypes.PyDLL(None) - try: - sleep = lib.sleep - except AttributeError: # pragma: no cover — macOS libSystem fallback - sleep = ctypes.PyDLL("/usr/lib/libSystem.B.dylib").sleep + sleep, secs = ctypes.PyDLL("kernel32").Sleep, secs * 1000 + else: + try: + sleep = ctypes.PyDLL(None).sleep + except AttributeError: # pragma: no cover — macOS libSystem fallback + sleep = ctypes.PyDLL("/usr/lib/libSystem.B.dylib").sleep sleep.argtypes = [ctypes.c_uint] sleep(int(secs)) -def _hook_sleep() -> dict: - time.sleep(30) - return {"ok": False, "error": "sleep hook returned unexpectedly"} +def _hang(block, name: str) -> dict: + block(30) + return {"ok": False, "error": f"{name} hook returned unexpectedly"} -def _hook_gil() -> dict: - _hold_gil(30) - return {"ok": False, "error": "gil hook returned unexpectedly"} - - -_TEST_HOOKS = { - "sleep": _hook_sleep, - "gil": _hook_gil, - "success": lambda: { - "ok": True, - "results": [{"title": "Hit", "url": "https://example.com", "description": "body", "position": 1}], - }, - "empty": lambda: {"ok": True, "results": []}, - "error": lambda: {"ok": False, "error": "RuntimeError: boom"}, -} +_TEST_HOOKS = {"sleep": lambda: _hang(time.sleep, "sleep"), "gil": lambda: _hang(_hold_gil, "gil"), "empty": lambda: {"ok": True, "results": []}} def _write_envelope(envelope: dict) -> None: @@ -61,34 +42,31 @@ def _write_envelope(envelope: dict) -> None: sys.stdout.flush() +def _fail(error: str, code: int) -> int: + _write_envelope({"ok": False, "error": error}) + return code + + def main() -> int: try: request = json.load(sys.stdin) except Exception as exc: # noqa: BLE001 - _write_envelope({"ok": False, "error": f"invalid request: {exc}"}) - return 2 - + return _fail(f"invalid request: {exc}", 2) hook = request.get("test_hook") if hook: if os.environ.get("HERMES_DDGS_ALLOW_TEST_HOOKS") != "1": - _write_envelope({"ok": False, "error": "test_hook refused (hooks not enabled)"}) - return 3 + return _fail("test_hook refused (hooks not enabled)", 3) fn = _TEST_HOOKS.get(str(hook)) envelope = fn() if fn else {"ok": False, "error": f"unknown test_hook: {hook!r}"} _write_envelope(envelope) return 0 if envelope.get("ok") else 1 - - query = str(request.get("query") or "") - safe_limit = max(1, int(request.get("safe_limit") or 1)) + query, safe_limit = str(request.get("query") or ""), max(1, int(request.get("safe_limit") or 1)) try: from plugins.web.ddgs.provider import _run_ddgs_search # lazy: light startup, patchable - - results = _run_ddgs_search(query, safe_limit) - _write_envelope({"ok": True, "results": results}) + _write_envelope({"ok": True, "results": _run_ddgs_search(query, safe_limit)}) return 0 except Exception as exc: # noqa: BLE001 - _write_envelope({"ok": False, "error": f"{type(exc).__name__}: {exc}"}) - return 1 + return _fail(f"{type(exc).__name__}: {exc}", 1) if __name__ == "__main__": diff --git a/plugins/web/ddgs/provider.py b/plugins/web/ddgs/provider.py index cca68b18ca..0a4053114f 100644 --- a/plugins/web/ddgs/provider.py +++ b/plugins/web/ddgs/provider.py @@ -1,12 +1,8 @@ -"""DuckDuckGo search via the optional ``ddgs`` package (search only, no key). - -``is_available()`` reflects package importability; the plugin registers either way -so ``hermes tools`` can offer to install it. - -Isolation: ``ddgs``/``primp`` can block inside native code while holding the GIL, so -a thread-pool ``future.result(timeout=…)`` cap can never fire and Ctrl+C/SIGTERM -freeze the whole process. Each search therefore runs in a disposable child process -the parent can terminate/kill. +"""DuckDuckGo search via the optional ``ddgs`` package (search only, no key). ``is_available()`` +reflects package importability; the plugin registers either way so ``hermes tools`` can offer to +install it. Isolation: ``ddgs``/``primp`` can block inside native code while holding the GIL, so a +thread-pool ``future.result(timeout=…)`` cap can never fire and Ctrl+C/SIGTERM freeze the process — +each search runs in a disposable child process the parent can terminate/kill. """ from __future__ import annotations @@ -20,22 +16,18 @@ import sys import time from typing import Any, Dict, Optional -from plugins.web._common import BaseWebSearchProvider, search_fail, search_ok +from plugins.web._common import BaseWebSearchProvider, search_fail, search_ok, setup_schema, title_hit logger = logging.getLogger(__name__) -# Hard wall-clock cap per search. ``DDGS(timeout=…)`` only bounds individual HTTP -# requests; ddgs's multi-engine retry loop has no overall cap, so a rate-limited -# response could otherwise hang the shared agent loop indefinitely. +# Hard wall-clock cap per search: ``DDGS(timeout=…)`` only bounds individual HTTP requests; +# ddgs's multi-engine retry loop has no overall cap, so a rate-limited response could +# otherwise hang the shared agent loop indefinitely. _SEARCH_TIMEOUT_SECS = 30 _POLL_INTERVAL_SECS = 0.1 # parent stdout / interrupt-flag poll cadence _TERMINATE_GRACE_SECS = 1.0 # wait after terminate() before escalating to kill() - -# Test-only hook name forwarded to the child (see _search_worker.py); never set in production. -_test_hook: Optional[str] = None - -# Last worker Popen started by ``_run_ddgs_search_bounded`` (test reap checks). -_last_worker_proc: Optional[subprocess.Popen] = None +_test_hook: Optional[str] = None # test-only hook forwarded to the child (see _search_worker.py) +_last_worker_proc: Optional[subprocess.Popen] = None # last worker Popen (test reap checks) class _SearchInterrupted(Exception): @@ -43,98 +35,71 @@ class _SearchInterrupted(Exception): def _run_ddgs_search(query: str, safe_limit: int) -> list[dict[str, Any]]: - """Blocking ddgs query → normalized hits. Module-level so the child worker can - import it and tests can patch it for in-process runs.""" + """Blocking ddgs query → normalized hits (module-level: the child worker imports it, + tests patch it for in-process runs).""" from ddgs import DDGS # type: ignore - results: list[dict[str, Any]] = [] with DDGS(timeout=10) as client: for i, hit in enumerate(client.text(query, max_results=safe_limit)): if i >= safe_limit: break - results.append( - { - "title": str(hit.get("title", "")), - "url": str(hit.get("href") or hit.get("url") or ""), - "description": str(hit.get("body", "")), - "position": i + 1, - } - ) + results.append(title_hit(str(hit.get("title", "")), str(hit.get("href") or hit.get("url") or ""), str(hit.get("body", "")), i + 1)) return results def _plugins_path_entry() -> str: - """``sys.path`` entry that makes ``import plugins`` work in the child (prefers - the live package location; correct for source checkouts and site-packages).""" + """``sys.path`` entry that makes ``import plugins`` work in the child (live package + location first; correct for source checkouts and site-packages).""" try: import plugins as plugins_pkg - - pkg_file = getattr(plugins_pkg, "__file__", None) - if pkg_file: + if pkg_file := getattr(plugins_pkg, "__file__", None): return os.path.dirname(os.path.dirname(os.path.abspath(pkg_file))) except Exception: # noqa: BLE001 — fall through to path-walk fallback pass - here = os.path.abspath(__file__) - for _ in range(4): - here = os.path.dirname(here) - return here + return os.path.abspath(os.path.join(__file__, *([os.pardir] * 4))) def _terminate_and_reap(proc: Optional[subprocess.Popen], *, grace: float = _TERMINATE_GRACE_SECS) -> None: - """Terminate a worker, escalate to kill, and wait so no orphan remains. - - Does not close the parent's pipe ends — closing stdout while another thread is - blocked in ``read()`` deadlocks on some platforms; the caller drains first. - """ + """Terminate a worker, escalate to kill, and wait so no orphan remains. Does not close + the parent's pipe ends — closing stdout while another thread is blocked in ``read()`` + deadlocks on some platforms; the caller drains first.""" if proc is None: return + alive = False - def _wait_until_dead(seconds: float) -> bool: - deadline = time.monotonic() + seconds + def _wait_until_dead() -> bool: + deadline = time.monotonic() + grace while proc.poll() is None and time.monotonic() < deadline: time.sleep(0.05) return proc.poll() is not None try: - if proc.poll() is None: - proc.terminate() - _wait_until_dead(grace) - if proc.poll() is None: - proc.kill() - if not _wait_until_dead(grace): - logger.warning("DDGS worker pid=%s did not exit after kill", proc.pid) + for escalate in (proc.terminate, proc.kill): + if proc.poll() is None: + escalate() + alive = not _wait_until_dead() + if alive: + logger.warning("DDGS worker pid=%s did not exit after kill", proc.pid) except Exception as exc: # noqa: BLE001 — best-effort cleanup logger.debug("DDGS worker reap error: %s", exc) def _spawn_worker(env: dict[str, str]) -> subprocess.Popen: - """Start ``_search_worker.py`` as a script with ``plugins`` importable.""" - # Running the worker as a script puts ``plugins/web/ddgs/`` on ``sys.path[0]``, - # which breaks ``import plugins...``; prepend the real package location. + """Start ``_search_worker.py`` as a script with ``plugins`` importable. Running as a + script puts ``plugins/web/ddgs/`` on ``sys.path[0]``, breaking ``import plugins...``, + so the real package location is prepended to PYTHONPATH.""" child_pythonpath = env.get("PYTHONPATH", "") path_entry = _plugins_path_entry() if path_entry and path_entry not in child_pythonpath.split(os.pathsep): env["PYTHONPATH"] = path_entry + os.pathsep + child_pythonpath if child_pythonpath else path_entry - worker_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "_search_worker.py") - # stdin/stdout/stderr must stay explicit keyword args on the Popen call so - # scripts/check_subprocess_stdin.py can see them (TUI gateway inherits stdin). - extra_kwargs: dict[str, Any] = ( - {"creationflags": subprocess.CREATE_NEW_PROCESS_GROUP} # so terminate/kill reach the worker - if sys.platform == "win32" - else {"start_new_session": True} # own session so a hung primp grandchild is reaped too - ) + # Own process group/session so terminate/kill also reach a hung primp grandchild. + extra_kwargs: dict[str, Any] = {"creationflags": subprocess.CREATE_NEW_PROCESS_GROUP} if sys.platform == "win32" else {"start_new_session": True} + # stdin/stdout/stderr stay explicit keyword args so scripts/check_subprocess_stdin.py sees them + # (TUI gateway inherits stdin). stderr=DEVNULL: a chatty child would deadlock a stdout-only drain. return subprocess.Popen( - [sys.executable, worker_path], - stdin=subprocess.PIPE, - stdout=subprocess.PIPE, - # DEVNULL: a chatty child filling the stderr pipe while we only drain stdout would deadlock. - stderr=subprocess.DEVNULL, - env=env, - text=True, - encoding="utf-8", - errors="replace", - **extra_kwargs, + [sys.executable, worker_path], stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, + env=env, text=True, encoding="utf-8", errors="replace", **extra_kwargs, ) @@ -149,72 +114,52 @@ def _parse_envelope(raw: str, proc: subprocess.Popen) -> list[dict[str, Any]]: raise RuntimeError(f"DDGS worker returned invalid JSON: {raw[:200]!r}") from exc if not isinstance(envelope, dict): raise RuntimeError(f"DDGS worker returned an invalid envelope: {envelope!r}") - if envelope.get("ok"): - results = envelope.get("results") or [] - if not isinstance(results, list): - raise RuntimeError("DDGS worker returned non-list results") - return results - raise RuntimeError(str(envelope.get("error") or "DDGS worker failed")) + if not envelope.get("ok"): + raise RuntimeError(str(envelope.get("error") or "DDGS worker failed")) + results = envelope.get("results") or [] + if not isinstance(results, list): + raise RuntimeError("DDGS worker returned non-list results") + return results def _run_ddgs_search_bounded(query: str, safe_limit: int) -> list[dict[str, Any]]: - """Run ``_run_ddgs_search`` in a disposable process with a hard deadline. - - The parent never joins a child that may be in native code holding *its* GIL — - it polls a communicator thread and, on timeout/interrupt, kills the OS process. - Raises ``TimeoutError``, ``_SearchInterrupted``, or ``RuntimeError``. - """ + """Run ``_run_ddgs_search`` in a disposable process with a hard deadline. The parent + never joins a child that may be in native code holding *its* GIL — it polls a + communicator thread and, on timeout/interrupt, kills the OS process. + Raises ``TimeoutError``, ``_SearchInterrupted``, or ``RuntimeError``.""" from tools.interrupt import is_interrupted # lazy: keep plugin import light from tools.environments.local import _sanitize_subprocess_env - global _last_worker_proc - request: dict[str, Any] = {"query": query, "safe_limit": safe_limit} env = _sanitize_subprocess_env(dict(os.environ)) if _test_hook: request["test_hook"] = _test_hook env["HERMES_DDGS_ALLOW_TEST_HOOKS"] = "1" - - proc = _spawn_worker(env) - _last_worker_proc = proc - + proc = _last_worker_proc = _spawn_worker(env) # ``communicate`` runs in a side thread so the parent can poll interrupt / # deadline without blocking; killing the child unblocks it. pool = cf.ThreadPoolExecutor(max_workers=1) fut = pool.submit(proc.communicate, json.dumps(request)) - timed_out = interrupted = False - raw = "" + interrupted, done, raw = False, False, "" try: deadline = time.monotonic() + _SEARCH_TIMEOUT_SECS - while True: - if is_interrupted(): - interrupted = True - break - remaining = deadline - time.monotonic() - if remaining <= 0: - timed_out = True - break + while not done and not (interrupted := is_interrupted()) and (remaining := deadline - time.monotonic()) > 0: try: - out, _err = fut.result(timeout=min(_POLL_INTERVAL_SECS, remaining)) - raw = out or "" - break + raw, done = fut.result(timeout=min(_POLL_INTERVAL_SECS, remaining))[0] or "", True except cf.TimeoutError: - continue + pass finally: _terminate_and_reap(proc) # After kill, communicate should return promptly; don't block forever. if not fut.done(): try: - out, _err = fut.result(timeout=_TERMINATE_GRACE_SECS) - if not raw: - raw = out or "" + raw = raw or fut.result(timeout=_TERMINATE_GRACE_SECS)[0] or "" except Exception: # noqa: BLE001 pass pool.shutdown(wait=False, cancel_futures=True) - if interrupted: raise _SearchInterrupted("DuckDuckGo search interrupted") - if timed_out: + if not done: raise TimeoutError(f"DuckDuckGo search timed out after {_SEARCH_TIMEOUT_SECS}s") return _parse_envelope(raw, proc) @@ -231,7 +176,6 @@ class DDGSWebSearchProvider(BaseWebSearchProvider): tool-registration time and on every ``hermes tools`` paint.""" try: import ddgs # noqa: F401 - return True except ImportError: return False @@ -241,18 +185,14 @@ class DDGSWebSearchProvider(BaseWebSearchProvider): hung native ``primp`` call cannot freeze the Hermes process.""" if not self.is_available(): return search_fail("ddgs package is not installed — run `pip install ddgs`") - - # Defensive cap in case the package ignores its max_results hint. - safe_limit = max(1, int(limit)) - try: - web_results = _run_ddgs_search_bounded(query, safe_limit) + # max(1, …): defensive cap in case the package ignores its max_results hint. + web_results = _run_ddgs_search_bounded(query, max(1, int(limit))) except TimeoutError: logger.warning("DDGS search timed out after %ds for query: %r", _SEARCH_TIMEOUT_SECS, query) return search_fail( f"DuckDuckGo search timed out after {_SEARCH_TIMEOUT_SECS}s — " - "DuckDuckGo may be rate-limiting or slow. Try again later " - "or switch to a different search provider." + "DuckDuckGo may be rate-limiting or slow. Try again later or switch to a different search provider." ) except _SearchInterrupted: logger.info("DDGS search interrupted for query: %r", query) @@ -260,16 +200,12 @@ class DDGSWebSearchProvider(BaseWebSearchProvider): except Exception as exc: # noqa: BLE001 — ddgs raises its own exceptions logger.warning("DDGS search error: %s", exc) return search_fail(f"DuckDuckGo search failed: {exc}") - logger.info("DDGS search '%s': %d results (limit %d)", query, len(web_results), limit) return search_ok(web_results) def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "DuckDuckGo (ddgs)", - "badge": "free · no key · search only", - "tag": "Search via the ddgs Python package — no API key (pair with any extract provider)", - "env_vars": [], - # Triggers `_run_post_setup("ddgs")` so the package gets pip-installed on first pick. - "post_setup": "ddgs", - } + # post_setup triggers `_run_post_setup("ddgs")` so the package gets pip-installed on first pick. + return setup_schema( + "DuckDuckGo (ddgs)", "free · no key · search only", + "Search via the ddgs Python package — no API key (pair with any extract provider)", post_setup="ddgs", + ) diff --git a/plugins/web/exa/__init__.py b/plugins/web/exa/__init__.py index 47a989fb59..fd9305d6d2 100644 --- a/plugins/web/exa/__init__.py +++ b/plugins/web/exa/__init__.py @@ -1,7 +1,5 @@ """Exa web search + extract plugin — bundled, auto-loaded (sync SDK; dispatcher wraps async callers).""" - from __future__ import annotations - from plugins.web.exa.provider import ExaWebSearchProvider diff --git a/plugins/web/exa/provider.py b/plugins/web/exa/provider.py index 562e117230..982598fc19 100644 --- a/plugins/web/exa/provider.py +++ b/plugins/web/exa/provider.py @@ -10,16 +10,8 @@ import logging from typing import Any, Dict, List from plugins.web._common import ( - BaseWebSearchProvider, - cached_sdk_client, - document, - keyless_variant_schema, - provider_env, - run_extract, - run_search, - search_ok, - use_keyless, - web_hit, + BaseWebSearchProvider, cached_sdk_client, document, keyless_extract, keyless_search, keyless_variant_schema, + provider_env, run_extract, run_search, search_ok, use_keyless, web_hit, ) logger = logging.getLogger(__name__) @@ -30,7 +22,6 @@ _MISSING_KEY = "EXA_API_KEY environment variable not set. Get your API key at ht def _get_exa_client() -> Any: def _factory(api_key: str) -> Any: from exa_py import Exa # deliberately lazy - client = Exa(api_key=api_key) client.headers["x-exa-integration"] = "hermes-agent" return client @@ -49,12 +40,8 @@ class ExaWebSearchProvider(BaseWebSearchProvider): def search(self, query: str, limit: int = 5) -> Dict[str, Any]: def _body() -> Dict[str, Any]: - from plugins.web.keyless_mcp import search_with_failover - if use_keyless("exa", provider_env("EXA_API_KEY")): - logger.info("Exa keyless search: '%s' (limit=%d)", query, limit) - return search_with_failover("exa", query, limit) - + return keyless_search("Exa", "exa", query, limit, logger) logger.info("Exa search: '%s' (limit=%d)", query, limit) response = _get_exa_client().search(query, num_results=limit, contents={"highlights": True}) return search_ok([ @@ -66,12 +53,8 @@ class ExaWebSearchProvider(BaseWebSearchProvider): def extract(self, urls: List[str], **kwargs: Any) -> List[Dict[str, Any]]: def _body() -> List[Dict[str, Any]]: - from plugins.web.keyless_mcp import extract_with_failover - if use_keyless("exa", provider_env("EXA_API_KEY")): - logger.info("Exa keyless extract: %d URL(s)", len(urls)) - return extract_with_failover("exa", list(urls)) - + return keyless_extract("Exa", "exa", urls, logger) logger.info("Exa extract: %d URL(s)", len(urls)) response = _get_exa_client().get_contents(urls, text=True) return [document(r.url or "", r.title or "", r.text or "") for r in response.results or []] diff --git a/plugins/web/firecrawl/__init__.py b/plugins/web/firecrawl/__init__.py index d52314d909..9551e0ccc7 100644 --- a/plugins/web/firecrawl/__init__.py +++ b/plugins/web/firecrawl/__init__.py @@ -1,7 +1,5 @@ """Firecrawl web search + extract plugin — bundled, auto-loaded.""" - from __future__ import annotations - from plugins.web.firecrawl.provider import FirecrawlWebSearchProvider diff --git a/plugins/web/firecrawl/provider.py b/plugins/web/firecrawl/provider.py index f35c50e3a1..a1e42503b2 100644 --- a/plugins/web/firecrawl/provider.py +++ b/plugins/web/firecrawl/provider.py @@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional import httpx -from plugins.web._common import BaseWebSearchProvider, search_fail, search_ok +from plugins.web._common import BaseWebSearchProvider, keyless_extract, keyless_search, lazy_ensure, search_fail, search_ok, setup_schema from tools.url_safety import is_safe_url # Module-level (cheap import) so tests can monkeypatch the policy gate on this module. from tools.website_policy import check_website_access @@ -20,13 +20,9 @@ from tools.website_policy import check_website_access logger = logging.getLogger(__name__) _FIRECRAWL_CLOUD_API_URL = "https://api.firecrawl.dev" -_SELECTION_KEYS = ("backend", "search_backend", "extract_backend") - -# --- Lazy Firecrawl SDK proxy ------------------------------------------------- # The SDK costs ~200ms of imports on a cold CLI; defer to first use. tools.web_tools # re-exports ``Firecrawl`` so ``patch("tools.web_tools.Firecrawl")`` keeps working. - _FIRECRAWL_CLS_CACHE: Optional[type] = None @@ -34,16 +30,8 @@ def _load_firecrawl_cls() -> type: """Import and cache ``firecrawl.Firecrawl`` (lazy_deps install hint → ImportError).""" global _FIRECRAWL_CLS_CACHE if _FIRECRAWL_CLS_CACHE is None: - try: - from tools.lazy_deps import ensure as _lazy_ensure - - _lazy_ensure("search.firecrawl", prompt=False) - except ImportError: - pass - except Exception as exc: # noqa: BLE001 — surface install hint - raise ImportError(str(exc)) + lazy_ensure("search.firecrawl") from firecrawl import Firecrawl as _cls - _FIRECRAWL_CLS_CACHE = _cls return _FIRECRAWL_CLS_CACHE @@ -65,365 +53,225 @@ class _FirecrawlProxy: Firecrawl = _FirecrawlProxy() - # --- Client construction (direct vs managed-gateway) --------------------------- -# Client cache slots and gateway/token helpers are read through tools.web_tools so -# tests that reset ``tools.web_tools._firecrawl_client`` or patch -# ``tools.web_tools._peek_nous_access_token`` see their changes. +def _wt(): + """Client cache slots and gateway/token helpers are read through tools.web_tools so tests + that reset ``_firecrawl_client`` or patch ``_peek_nous_access_token`` there see their changes.""" + import tools.web_tools as _mod + return _mod def _env(name: str) -> str: from hermes_cli.config import get_env_value - return (get_env_value(name) or "").strip() def _get_direct_firecrawl_config() -> Optional[tuple]: - """Return direct Firecrawl ``(mode, kwargs, cache_key)`` or None. - - ``mode`` is ``"sdk"`` (keyed / self-hosted) or ``"keyless"`` (explicit - Firecrawl selection with no credentials — public cloud API, anonymous - rate-limited). Keyless requires the explicit selection so an unconfigured - install never silently routes to it. - """ + """Direct Firecrawl ``(mode, kwargs, cache_key)`` or None. ``mode`` is ``"sdk"`` (keyed / self-hosted) or + ``"keyless"`` (explicit selection + no credentials → anonymous public cloud; the explicit selection is + required so an unconfigured install never silently routes to it).""" api_key = _env("FIRECRAWL_API_KEY") api_url = _env("FIRECRAWL_API_URL").rstrip("/") - - if not api_key and not api_url: - if _is_explicit_firecrawl_selection(): - return "keyless", {"api_url": _FIRECRAWL_CLOUD_API_URL}, ("direct-keyless", _FIRECRAWL_CLOUD_API_URL, None) - return None - - kwargs = {k: v for k, v in (("api_key", api_key), ("api_url", api_url)) if v} - return "sdk", kwargs, ("direct", api_url or None, api_key or None) + if api_key or api_url: + return "sdk", {k: v for k, v in (("api_key", api_key), ("api_url", api_url)) if v}, ("direct", api_url or None, api_key or None) + if _is_explicit_firecrawl_selection(): + return "keyless", {"api_url": _FIRECRAWL_CLOUD_API_URL}, ("direct-keyless", _FIRECRAWL_CLOUD_API_URL, None) + return None def _is_explicit_firecrawl_selection() -> bool: - """True when config explicitly selects Firecrawl for web tools.""" - import tools.web_tools as _wt - - cfg = _wt._load_web_config() - return any((cfg.get(key) or "").lower().strip() == "firecrawl" for key in _SELECTION_KEYS) + from plugins.web.keyless_mcp import _web_config_selects + return _web_config_selects("firecrawl") def _use_keyless_ring() -> bool: - """True when Firecrawl calls should route via the keyless ring. - - Only when there are no direct credentials, the managed Nous gateway isn't - the selected path, and the keyless tier isn't disabled or pinned paid. - """ + """Route via the keyless ring only with no direct credentials, when the managed Nous + gateway isn't the selected path, and the keyless tier isn't disabled or pinned paid.""" if _env("FIRECRAWL_API_KEY") or _env("FIRECRAWL_API_URL"): return False - import tools.web_tools as _wt from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER, read_selection - - try: - if read_selection("web") == NOUS_MANAGED_PROVIDER: - return False - except Exception: # noqa: BLE001 — selection helpers optional - pass - try: - if _wt._is_tool_gateway_ready() and not _is_explicit_firecrawl_selection(): - return False - except Exception: # noqa: BLE001 — probe optional - pass from plugins.web.keyless_mcp import use_keyless - + # Both probes are optional layers: a failing probe never blocks the ring. + for probe in (lambda: read_selection("web") == NOUS_MANAGED_PROVIDER, lambda: _wt()._is_tool_gateway_ready() and not _is_explicit_firecrawl_selection()): + try: + if probe(): + return False + except Exception: # noqa: BLE001 + pass return use_keyless("firecrawl", "") class _KeylessFirecrawlClient: - """Minimal REST client for Firecrawl's keyless cloud mode. - - Duck-types the SDK's ``search`` / ``scrape``; never sends an Authorization header. - """ + """Minimal REST client for Firecrawl's keyless cloud mode; duck-types the SDK's + ``search`` / ``scrape`` and never sends an Authorization header.""" def __init__(self, api_url: str = _FIRECRAWL_CLOUD_API_URL): self.api_url = api_url.rstrip("/") def _post(self, path: str, payload: Dict[str, Any]) -> Dict[str, Any]: - response = httpx.post( - f"{self.api_url}{path}", - json=payload, - headers={"Content-Type": "application/json"}, - timeout=60.0, - ) + response = httpx.post(f"{self.api_url}{path}", json=payload, headers={"Content-Type": "application/json"}, timeout=60.0) response.raise_for_status() return response.json() - def search(self, *, query: str, limit: int = 5) -> Dict[str, Any]: - return self._post("/v2/search", {"query": query, "limit": limit}) - - def scrape(self, *, url: str, formats: List[str]) -> Dict[str, Any]: - return self._post("/v2/scrape", {"url": url, "formats": formats}) + search = lambda self, *, query, limit=5: self._post("/v2/search", {"query": query, "limit": limit}) # noqa: E731 + scrape = lambda self, *, url, formats: self._post("/v2/scrape", {"url": url, "formats": formats}) # noqa: E731 def _get_firecrawl_gateway_url() -> str: - """Return the configured Firecrawl gateway URL.""" - import tools.web_tools as _wt - - return _wt.build_vendor_gateway_url("firecrawl") + return _wt().build_vendor_gateway_url("firecrawl") def _is_tool_gateway_ready() -> bool: """True when gateway URL + Nous Subscriber token are available.""" - import tools.web_tools as _wt - - return _wt.resolve_managed_tool_gateway("firecrawl", token_reader=_wt._peek_nous_access_token) is not None + return _wt().resolve_managed_tool_gateway("firecrawl", token_reader=_wt()._peek_nous_access_token) is not None def check_firecrawl_api_key() -> bool: - """True when the Firecrawl route selected via ``hermes tools`` (or, on a - never-configured install, either route) is usable. Re-exported by tools.web_tools.""" + """True when the route selected via ``hermes tools`` (or, on a never-configured + install, either route) is usable. Re-exported by tools.web_tools.""" from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER, read_selection - selected = read_selection("web") if selected == NOUS_MANAGED_PROVIDER: return _is_tool_gateway_ready() - has_direct = _get_direct_firecrawl_config() is not None - if selected is not None: - return has_direct - return has_direct or _is_tool_gateway_ready() + return _get_direct_firecrawl_config() is not None or (selected is None and _is_tool_gateway_ready()) def _firecrawl_backend_help_suffix() -> str: - """Return optional managed-gateway guidance for Firecrawl help text.""" - import tools.web_tools as _wt - - if not _wt.managed_nous_tools_enabled(): - return "" - return ", or use the Nous Tool Gateway via your subscription (FIRECRAWL_GATEWAY_URL or TOOL_GATEWAY_DOMAIN)" + return ", or use the Nous Tool Gateway via your subscription (FIRECRAWL_GATEWAY_URL or TOOL_GATEWAY_DOMAIN)" if _wt().managed_nous_tools_enabled() else "" def _get_firecrawl_client() -> Any: - """Get or create the cached Firecrawl client. - - Strict selection semantics on the stored ``web`` selection: ``"nous"`` → - managed Tool Gateway ONLY; any other stored backend → direct Firecrawl ONLY - (never a silent managed fallback billed to Nous); never-configured → direct - when present, else managed. Raises ValueError when the resolved path is unusable. - """ - import tools.web_tools as _wt - from tools.tool_backend_helpers import ( - NOUS_MANAGED_PROVIDER, - read_selection, - selection_error, - selection_exists, - ) - + """Get or create the cached Firecrawl client. Strict selection semantics on the stored ``web`` selection: + ``"nous"`` → managed Tool Gateway ONLY; any other stored backend → direct Firecrawl ONLY (never a silent + managed fallback billed to Nous); never-configured → direct when present, else managed. Raises ValueError + when the resolved path is unusable.""" + wt = _wt() + from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER, read_selection, selection_error, selection_exists selected = read_selection("web") direct_config = _get_direct_firecrawl_config() - def _managed_kwargs(): - gw = _wt.resolve_managed_tool_gateway("firecrawl", token_reader=_wt._read_nous_access_token) + def _managed(): + gw = wt.resolve_managed_tool_gateway("firecrawl", token_reader=wt._read_nous_access_token) if gw is None: return None - kwargs = {"api_key": gw.nous_user_token, "api_url": gw.gateway_origin} - return kwargs, ("tool-gateway", kwargs["api_url"], gw.nous_user_token) + return "sdk", {"api_key": gw.nous_user_token, "api_url": gw.gateway_origin}, ("tool-gateway", gw.gateway_origin, gw.nous_user_token) - client_mode = "sdk" + def _unconfigured_message() -> str: + message = "Web tools are not configured. Set FIRECRAWL_API_KEY for cloud Firecrawl or set FIRECRAWL_API_URL for a self-hosted Firecrawl instance." + if wt.managed_nous_tools_enabled(): + return message + " With your Nous subscription you can also use the Tool Gateway. run `hermes tools` and select Nous Subscription as the web provider." + return message + " " + wt.nous_tool_gateway_unavailable_message("managed Firecrawl web tools") + + # (resolved config, log detail, error message) per selection state; the message is built lazily. if selected == NOUS_MANAGED_PROVIDER: - managed = _managed_kwargs() - if managed is None: - logger.error( - "Firecrawl client initialization failed: the Nous " - "Subscription web selection is stored but the tool gateway " - "is unavailable." - ) - raise ValueError(selection_error( - "web", NOUS_MANAGED_PROVIDER, - "the Nous Tool Gateway is not available (not entitled or unreachable)", - )) - kwargs, client_config = managed + resolved, log, message = _managed(), "the Nous Subscription web selection is stored but the tool gateway is unavailable.", lambda: selection_error( + "web", NOUS_MANAGED_PROVIDER, "the Nous Tool Gateway is not available (not entitled or unreachable)") elif selected is not None or selection_exists("web"): - # Stored vendor selection: direct Firecrawl only. With no credentials the - # explicit selection unlocks keyless cloud mode instead of erroring. - if direct_config is None: - logger.error( - "Firecrawl client initialization failed: direct Firecrawl " - "selected but FIRECRAWL_API_KEY/FIRECRAWL_API_URL is not set." - ) - raise ValueError(selection_error( - "web", selected or "firecrawl", "neither FIRECRAWL_API_KEY nor FIRECRAWL_API_URL is set", - )) - client_mode, kwargs, client_config = direct_config + # Stored vendor selection: direct only (no credentials → explicit selection unlocks keyless cloud mode). + resolved, log, message = direct_config, "direct Firecrawl selected but FIRECRAWL_API_KEY/FIRECRAWL_API_URL is not set.", lambda: selection_error( + "web", selected or "firecrawl", "neither FIRECRAWL_API_KEY nor FIRECRAWL_API_URL is set") elif direct_config is not None: - client_mode, kwargs, client_config = direct_config - else: - # Never-configured web section: legacy managed fallback. - managed = _managed_kwargs() - if managed is None: - logger.error("Firecrawl client initialization failed: missing direct config and tool-gateway auth.") - message = ( - "Web tools are not configured. " - "Set FIRECRAWL_API_KEY for cloud Firecrawl or set FIRECRAWL_API_URL " - "for a self-hosted Firecrawl instance." - ) - if _wt.managed_nous_tools_enabled(): - message += ( - " With your Nous subscription you can also use the Tool Gateway. " - "run `hermes tools` and select Nous Subscription as the web provider." - ) - else: - message += " " + _wt.nous_tool_gateway_unavailable_message("managed Firecrawl web tools") - raise ValueError(message) - kwargs, client_config = managed - - cached = getattr(_wt, "_firecrawl_client", None) - if cached is not None and getattr(_wt, "_firecrawl_client_config", None) == client_config: + resolved = direct_config + else: # never-configured web section: legacy managed fallback + resolved, log, message = _managed(), "missing direct config and tool-gateway auth.", _unconfigured_message + if resolved is None: + logger.error("Firecrawl client initialization failed: %s", log) + raise ValueError(message()) + client_mode, kwargs, client_config = resolved + cached = getattr(wt, "_firecrawl_client", None) + if cached is not None and getattr(wt, "_firecrawl_client_config", None) == client_config: return cached - - if client_mode == "keyless": - _wt._firecrawl_client = _KeylessFirecrawlClient(api_url=kwargs["api_url"]) - else: - _wt._firecrawl_client = _wt.Firecrawl(**kwargs) - _wt._firecrawl_client_config = client_config - return _wt._firecrawl_client + wt._firecrawl_client = _KeylessFirecrawlClient(api_url=kwargs["api_url"]) if client_mode == "keyless" else wt.Firecrawl(**kwargs) + wt._firecrawl_client_config = client_config + return wt._firecrawl_client # --- Response shape normalization (SDK / direct / gateway differ) -------------- - - def _to_plain_object(value: Any) -> Any: - """Convert SDK objects (pydantic ``model_dump`` / ``__dict__``) to plain data when possible.""" + """SDK objects (pydantic ``model_dump`` / ``__dict__``) → plain data when possible.""" if value is None or isinstance(value, (dict, list, str, int, float, bool)): return value - if hasattr(value, "model_dump"): - try: - return value.model_dump() - except Exception: # noqa: BLE001 - pass - if hasattr(value, "__dict__"): - try: - return {k: v for k, v in value.__dict__.items() if not k.startswith("_")} - except Exception: # noqa: BLE001 - pass + for attr, convert in (("model_dump", lambda v: v.model_dump()), ("__dict__", lambda v: {k: x for k, x in v.__dict__.items() if not k.startswith("_")})): + if hasattr(value, attr): + try: + return convert(value) + except Exception: # noqa: BLE001 + pass return value def _normalize_result_list(values: Any) -> List[Dict[str, Any]]: - """Normalize mixed SDK/list payloads into a list of dicts.""" - if not isinstance(values, list): - return [] - plain = (_to_plain_object(item) for item in values) - return [p for p in plain if isinstance(p, dict)] + return [p for p in map(_to_plain_object, values) if isinstance(p, dict)] if isinstance(values, list) else [] def _extract_web_search_results(response: Any) -> List[Dict[str, Any]]: - """Extract Firecrawl search results across SDK/direct/gateway response shapes.""" - response_plain = _to_plain_object(response) - - if isinstance(response_plain, dict): - data = response_plain.get("data") + """Search results across SDK/direct/gateway response shapes.""" + plain = _to_plain_object(response) + if isinstance(plain, dict): + data = plain.get("data") if isinstance(data, list): return _normalize_result_list(data) - candidates = [] - if isinstance(data, dict): - candidates += [data.get("web"), data.get("results")] - candidates += [response_plain.get("web"), response_plain.get("results")] - for candidate in candidates: + candidates = [data.get("web"), data.get("results")] if isinstance(data, dict) else [] + for candidate in candidates + [plain.get("web"), plain.get("results")]: normalized = _normalize_result_list(candidate) if normalized: return normalized - if hasattr(response, "web"): return _normalize_result_list(getattr(response, "web", [])) return [] def _extract_scrape_payload(scrape_result: Any) -> Dict[str, Any]: - """Normalize Firecrawl scrape payload shape across SDK and gateway variants.""" - result_plain = _to_plain_object(scrape_result) - if not isinstance(result_plain, dict): + plain = _to_plain_object(scrape_result) + if not isinstance(plain, dict): return {} - nested = result_plain.get("data") - return nested if isinstance(nested, dict) else result_plain + return plain["data"] if isinstance(plain.get("data"), dict) else plain -def _error_entry( - url: str, - error: str, - *, - title: str = "", - raw: bool = False, - blocked: Optional[Dict[str, Any]] = None, -) -> Dict[str, Any]: - """Per-URL extract failure item. ``raw`` adds the ``raw_content`` key (post-scrape - failures carry it, pre-scrape ones don't); ``blocked`` adds ``blocked_by_policy``.""" - entry: Dict[str, Any] = {"url": url, "title": title, "content": ""} - if raw: - entry["raw_content"] = "" - entry["error"] = error - if blocked: - entry["blocked_by_policy"] = {k: blocked[k] for k in ("host", "rule", "source")} - return entry +def _error_entry(url: str, error: str, *, title: str = "", raw: bool = False, blocked: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: + """Per-URL extract failure. ``raw`` adds ``raw_content`` (post-scrape failures carry + it, pre-scrape ones don't); ``blocked`` adds ``blocked_by_policy``.""" + policy = {"blocked_by_policy": {k: blocked[k] for k in ("host", "rule", "source")}} if blocked else {} + return {"url": url, "title": title, "content": "", **({"raw_content": ""} if raw else {}), "error": error, **policy} -_SCRAPE_TIMEOUT_MSG = ( - "Scrape timed out after 60s — page may be too large " - "or unresponsive. Try browser_navigate instead." -) +_SCRAPE_TIMEOUT_MSG = "Scrape timed out after 60s — page may be too large or unresponsive. Try browser_navigate instead." _UNSAFE_REDIRECT_MSG = "Blocked: URL targets a private or internal network address" async def _scrape_one(url: str, formats: List[str], format: Optional[str]) -> Dict[str, Any]: """Scrape one URL (60s timeout) and re-check SSRF + website policy against the post-redirect URL. Never raises for scrape errors; returns an error entry instead.""" - blocked = check_website_access(url) - if blocked: + if blocked := check_website_access(url): logger.info("Blocked web_extract for %s by rule %s", blocked["host"], blocked["rule"]) return _error_entry(url, blocked["message"], blocked=blocked) - try: logger.info("Firecrawl scraping: %s", url) try: - scrape_result = await asyncio.wait_for( - asyncio.to_thread(_get_firecrawl_client().scrape, url=url, formats=formats), - timeout=60, - ) + scrape_result = await asyncio.wait_for(asyncio.to_thread(_get_firecrawl_client().scrape, url=url, formats=formats), timeout=60) except asyncio.TimeoutError: logger.warning("Firecrawl scrape timed out for %s", url) return _error_entry(url, _SCRAPE_TIMEOUT_MSG) - - scrape_payload = _extract_scrape_payload(scrape_result) - metadata = scrape_payload.get("metadata", {}) - content_markdown = scrape_payload.get("markdown") - content_html = scrape_payload.get("html") - + payload = _extract_scrape_payload(scrape_result) + metadata = payload.get("metadata", {}) # SDK may return a typed object for metadata (raw __dict__ here, unlike _to_plain_object). if not isinstance(metadata, dict): - if hasattr(metadata, "model_dump"): - metadata = metadata.model_dump() - elif hasattr(metadata, "__dict__"): - metadata = metadata.__dict__ - else: - metadata = {} - - title = metadata.get("title", "") - final_url = metadata.get("sourceURL", url) - + metadata = metadata.model_dump() if hasattr(metadata, "model_dump") else getattr(metadata, "__dict__", {}) + title, final_url = metadata.get("title", ""), metadata.get("sourceURL", url) if not is_safe_url(final_url): logger.info("Blocked redirected web_extract for unsafe final URL: %s", final_url) return _error_entry(final_url, _UNSAFE_REDIRECT_MSG, title=title, raw=True) - - final_blocked = check_website_access(final_url) - if final_blocked: + if final_blocked := check_website_access(final_url): logger.info("Blocked redirected web_extract for %s by rule %s", final_blocked["host"], final_blocked["rule"]) return _error_entry(final_url, final_blocked["message"], title=title, raw=True, blocked=final_blocked) - - if format == "markdown" or (format is None and content_markdown): - chosen_content = content_markdown - else: - chosen_content = content_html or content_markdown or "" - return {"url": final_url, "title": title, "content": chosen_content, "raw_content": chosen_content, "metadata": metadata} + markdown, html = payload.get("markdown"), payload.get("html") + content = markdown if format == "markdown" or (format is None and markdown) else html or markdown or "" + return {"url": final_url, "title": title, "content": content, "raw_content": content, "metadata": metadata} except Exception as scrape_err: # noqa: BLE001 logger.debug("Firecrawl scrape failed for %s: %s", url, scrape_err) return _error_entry(url, str(scrape_err), raw=True) -# --- Provider class ------------------------------------------------------------ - - class FirecrawlWebSearchProvider(BaseWebSearchProvider): """Firecrawl search + extract provider with dual auth paths.""" @@ -436,25 +284,17 @@ class FirecrawlWebSearchProvider(BaseWebSearchProvider): return check_firecrawl_api_key() def search(self, query: str, limit: int = 5) -> Dict[str, Any]: - """Sync search. Pre-flight errors (ValueError / ImportError) propagate so the - dispatcher emits the legacy ``tool_error`` envelope; in-flight errors are - returned as ``{"success": False, "error": ...}``.""" + """Pre-flight errors (ValueError / ImportError) propagate so the dispatcher emits + the legacy ``tool_error`` envelope; in-flight errors become failure dicts.""" from tools.interrupt import is_interrupted - if is_interrupted(): return search_fail("Interrupted") - if _use_keyless_ring(): - from plugins.web.keyless_mcp import search_with_failover - - logger.info("Firecrawl keyless search: '%s' (limit=%d)", query, limit) - return search_with_failover("firecrawl", query, limit) - + return keyless_search("Firecrawl", "firecrawl", query, limit, logger) logger.info("Firecrawl search: '%s' (limit=%d)", query, limit) client = _get_firecrawl_client() try: - response = client.search(query=query, limit=limit) - web_results = _extract_web_search_results(response) + web_results = _extract_web_search_results(client.search(query=query, limit=limit)) logger.info("Firecrawl: found %d search results", len(web_results)) return search_ok(web_results) except Exception as exc: # noqa: BLE001 @@ -462,37 +302,23 @@ class FirecrawlWebSearchProvider(BaseWebSearchProvider): return search_fail(f"Firecrawl search failed: {exc}") async def extract(self, urls: List[str], **kwargs: Any) -> List[Dict[str, Any]]: - """Per-URL scrape via :func:`_scrape_one`; failures become items with an - ``error`` field. ``format``: "markdown" | "html" | both (markdown preferred).""" + """Per-URL scrape; failures become items with an ``error`` field. + ``format``: "markdown" | "html" | both (markdown preferred).""" from tools.interrupt import is_interrupted as _is_interrupted - if _is_interrupted(): return [{"url": u, "error": "Interrupted", "title": ""} for u in urls] - if _use_keyless_ring(): - from plugins.web.keyless_mcp import extract_with_failover - - logger.info("Firecrawl keyless extract: %d URL(s)", len(urls)) - return await asyncio.to_thread(extract_with_failover, "firecrawl", list(urls)) - + return await asyncio.to_thread(keyless_extract, "Firecrawl", "firecrawl", urls, logger) format = kwargs.get("format") formats = [format] if format in ("markdown", "html") else ["markdown", "html"] - return [ {"url": url, "error": "Interrupted", "title": ""} if _is_interrupted() else await _scrape_one(url, formats, format) for url in urls ] def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "Firecrawl", - "badge": "keyless/paid · optional gateway", - "tag": "Full search + extract; supports keyless cloud, direct API, and Nous tool-gateway routing.", - "env_vars": [ - { - "key": "FIRECRAWL_API_KEY", - "prompt": "Firecrawl API key (optional; blank = keyless cloud or self-hosted)", - "url": "https://docs.firecrawl.dev/introduction", - }, - ], - } + return setup_schema( + "Firecrawl", "keyless/paid · optional gateway", + "Full search + extract; supports keyless cloud, direct API, and Nous tool-gateway routing.", + "FIRECRAWL_API_KEY", "Firecrawl API key (optional; blank = keyless cloud or self-hosted)", "https://docs.firecrawl.dev/introduction", + ) diff --git a/plugins/web/keenable/__init__.py b/plugins/web/keenable/__init__.py index ffb4affe2e..da904743e7 100644 --- a/plugins/web/keenable/__init__.py +++ b/plugins/web/keenable/__init__.py @@ -1,7 +1,5 @@ """Keenable web search + extract plugin — bundled, auto-loaded; keyless-ring member.""" - from __future__ import annotations - from plugins.web.keenable.provider import KeenableWebSearchProvider diff --git a/plugins/web/keenable/provider.py b/plugins/web/keenable/provider.py index c61c885d41..df704cb84a 100644 --- a/plugins/web/keenable/provider.py +++ b/plugins/web/keenable/provider.py @@ -10,18 +10,8 @@ import logging from typing import Any, Dict, List from plugins.web._common import ( - SEARCH_LIMIT_CAP, - BaseWebSearchProvider, - document, - http_status_detail, - keyless_variant_schema, - provider_env, - run_extract, - run_search, - search_fail, - search_ok, - use_keyless, - web_hit, + SEARCH_LIMIT_CAP, BaseWebSearchProvider, document, http_status_detail, keyless_extract, keyless_search, + keyless_variant_schema, page_error, provider_env, run_extract, run_search, search_fail, search_ok, use_keyless, web_hit, ) logger = logging.getLogger(__name__) @@ -48,21 +38,15 @@ class KeenableWebSearchProvider(BaseWebSearchProvider): def search(self, query: str, limit: int = 5) -> Dict[str, Any]: def _body() -> Dict[str, Any]: - from plugins.web.keyless_mcp import search_with_failover - api_key = provider_env("KEENABLE_API_KEY") if use_keyless("keenable", api_key): - logger.info("Keenable keyless search: '%s' (limit=%d)", query, limit) - return search_with_failover("keenable", query, limit) - + return keyless_search("Keenable", "keenable", query, limit, logger) import requests - logger.info("Keenable search: '%s' (limit=%d)", query, limit) response = requests.post( f"{_KEENABLE_API_URL}/v1/search", json={"query": query, "max_results": min(max(1, int(limit)), SEARCH_LIMIT_CAP)}, - headers=_keenable_headers(api_key), - timeout=30, + headers=_keenable_headers(api_key), timeout=30, ) if response.status_code >= 400: return search_fail(f"Keenable search failed: {http_status_detail(response)}") @@ -75,30 +59,21 @@ class KeenableWebSearchProvider(BaseWebSearchProvider): def extract(self, urls: List[str], **kwargs: Any) -> List[Dict[str, Any]]: def _body() -> List[Dict[str, Any]]: - from plugins.web.keyless_mcp import extract_with_failover - api_key = provider_env("KEENABLE_API_KEY") if use_keyless("keenable", api_key): - logger.info("Keenable keyless extract: %d URL(s)", len(urls)) - return extract_with_failover("keenable", list(urls)) - + return keyless_extract("Keenable", "keenable", urls, logger) import requests - logger.info("Keenable extract: %d URL(s)", len(urls)) results: List[Dict[str, Any]] = [] for url in urls: try: - response = requests.get( - f"{_KEENABLE_API_URL}/v1/fetch", params={"url": url}, headers=_keenable_headers(api_key), timeout=30 - ) + response = requests.get(f"{_KEENABLE_API_URL}/v1/fetch", params={"url": url}, headers=_keenable_headers(api_key), timeout=30) if response.status_code >= 400: raise ValueError(http_status_detail(response)) data = response.json() - results.append( - document(data.get("url") or url, data.get("title") or "", data.get("content") or "", source_url=url) - ) + results.append(document(data.get("url") or url, data.get("title") or "", data.get("content") or "", source_url=url)) except Exception as exc: # noqa: BLE001 — per-URL error entry - results.append({"url": url, "title": "", "content": "", "error": f"Keenable extract failed: {exc}"}) + results.append(page_error(url, f"Keenable extract failed: {exc}")) return results return run_extract("Keenable", logger, urls, _body, verbatim_value_error=False) diff --git a/plugins/web/keyless_mcp.py b/plugins/web/keyless_mcp.py index 98af33012e..1855c8897c 100644 --- a/plugins/web/keyless_mcp.py +++ b/plugins/web/keyless_mcp.py @@ -1,10 +1,8 @@ """Keyless web search/extract via public free-tier endpoints (Exa, Parallel, Firecrawl, Keenable). - -Resolved strictly LAST — after every keyed backend, the managed gateway, ddgs and -custom plugin providers — so it never pre-empts a deliberate setup. Privacy: no user -identifiers are sent; Parallel gets a random per-process ``session_id`` (rate -limiting only) and its optional ``model_name`` analytics field is deliberately -omitted. Disable the tier with ``web.keyless_fallback: false``. +Resolved strictly LAST — after every keyed backend, the managed gateway, ddgs and custom plugin +providers — so it never pre-empts a deliberate setup. Privacy: no user identifiers are sent; +Parallel gets a random per-process ``session_id`` (rate limiting only) and its optional +``model_name`` analytics field is deliberately omitted. Disable with ``web.keyless_fallback: false``. """ from __future__ import annotations @@ -15,7 +13,7 @@ import threading import uuid from typing import Any, Callable, Dict, List, Optional -from plugins.web._common import document as _page, search_fail, search_ok, web_hit as _row +from plugins.web._common import document as _page, page_error as _page_error, search_fail, search_ok, web_hit as _row logger = logging.getLogger(__name__) @@ -38,10 +36,8 @@ _RATE_LIMIT_MARKERS = ("rate limit", "rate-limit", "ratelimit", "too many reques # vendor -> (display label, env key, signup URL) for the standard failure hint. _VENDOR_HINTS = { - "exa": ("Exa", "EXA_API_KEY", "https://exa.ai"), - "parallel": ("Parallel", "PARALLEL_API_KEY", "https://parallel.ai"), - "firecrawl": ("Firecrawl", "FIRECRAWL_API_KEY", "https://firecrawl.dev"), - "keenable": ("Keenable", "KEENABLE_API_KEY", "https://keenable.ai"), + "exa": ("Exa", "EXA_API_KEY", "https://exa.ai"), "parallel": ("Parallel", "PARALLEL_API_KEY", "https://parallel.ai"), + "firecrawl": ("Firecrawl", "FIRECRAWL_API_KEY", "https://firecrawl.dev"), "keenable": ("Keenable", "KEENABLE_API_KEY", "https://keenable.ai"), } @@ -56,33 +52,54 @@ def _fail_msg(vendor: str, kind: str, exc: Any, *, other_backends: bool = True) return f"Keyless {label} {kind} failed: {exc}. Set {env_key} ({site}){alt} for reliable service." -def _page_error(url: str, message: str) -> Dict[str, Any]: - return {"url": url, "title": "", "content": "", "error": message} +def _search(vendor: str, rows: Callable[[], List[Dict[str, Any]]], catch: Any = (), fmt: Optional[Callable[[Exception], str]] = None) -> Dict[str, Any]: + """``search_ok(rows())``; :class:`KeylessMCPError` → standard vendor hint, exception + types in ``catch`` → ``fmt(exc)``; anything else propagates.""" + try: + return search_ok(rows()) + except KeylessMCPError as exc: + return search_fail(_fail_msg(vendor, "search", exc)) + except catch as exc: + return search_fail(fmt(exc)) + + +def _per_url(urls: List[str], fetch: Callable[[str], Dict[str, Any]], vendor: str, catch: Any = Exception, hint: bool = False) -> List[Dict[str, Any]]: + """Per-URL extract loop: a ``catch`` failure becomes an error entry (``hint`` adds the ``hermes tools`` hint).""" + def _one(url: str) -> Dict[str, Any]: + try: + return fetch(url) + except catch as exc: # noqa: BLE001 — per-URL error entry + return _page_error(url, _fail_msg(vendor, "extract", exc, other_backends=hint)) + + return [_one(u) for u in urls] # --- Tier / config ------------------------------------------------------------ - - def keyless_enabled() -> bool: - """True when the keyless tier is enabled. Delegates to the registry so the - ``web.keyless_fallback`` (default on) chokepoint lives with backend resolution.""" + """Delegates to the registry so the ``web.keyless_fallback`` (default on) chokepoint lives with backend resolution.""" try: from agent.web_search_registry import _keyless_tier_enabled - return _keyless_tier_enabled() except Exception as exc: # noqa: BLE001 — resolver optional in stripped envs logger.debug("keyless_enabled(): registry helper unavailable: %s", exc) return True +_BACKEND_KEYS = ("backend", "search_backend", "extract_backend") + + +def _web_config_selects(name: str) -> bool: + """True when any ``web.backend`` / ``search_backend`` / ``extract_backend`` names *name*.""" + import tools.web_tools as _wt + web_cfg = _wt._load_web_config() + return any((web_cfg.get(key) or "").lower().strip() == name for key in _BACKEND_KEYS) + + def provider_tier(name: str) -> str: - """Return ``web.provider_tier.`` (set by the ``hermes tools`` Free/Paid rows): - ``free``, ``paid``, or ``auto`` for anything else including unset.""" + """``web.provider_tier.`` (``hermes tools`` Free/Paid rows): ``free``, ``paid``, or ``auto`` (anything else/unset).""" try: from hermes_cli.config import load_config - - web_cfg = load_config().get("web") or {} - tiers = web_cfg.get("provider_tier") or {} + tiers = (load_config().get("web") or {}).get("provider_tier") or {} value = str(tiers.get(name, "") or "").lower().strip() return value if value in ("free", "paid") else "auto" except Exception as exc: # noqa: BLE001 — config layer optional @@ -91,28 +108,19 @@ def provider_tier(name: str) -> str: def use_keyless(name: str, api_key: str) -> bool: - """Single chokepoint for search + extract so tier semantics can't drift: - ``free`` → keyless even with a key; ``paid`` → keyed even without one (the keyed - path raises its usual missing-key error); ``auto`` → keyless only when no key - and the tier is enabled.""" + """Single chokepoint for search + extract: ``free`` → keyless even with a key; ``paid`` → keyed even + without one (the keyed path raises its usual missing-key error); ``auto`` → keyless only when no key + tier enabled.""" tier = provider_tier(name) - if tier == "free": - return True - if tier == "paid": - return False + if tier in ("free", "paid"): + return tier == "free" return not api_key and keyless_enabled() # --- MCP transport ------------------------------------------------------------ - - def _parse_mcp_body(body: str) -> str: - """Extract the first text content item from an MCP tools/call response. - - Handles plain-JSON bodies (Parallel) and SSE ``data: {...}`` lines (Exa). Raises - :class:`KeylessMCPError` for JSON-RPC errors and ``isError`` tool results (e.g. - Exa's free-tier rate-limit message). - """ + """First text content item from an MCP tools/call response — plain-JSON bodies + (Parallel) or SSE ``data: {...}`` lines (Exa). Raises :class:`KeylessMCPError` for + JSON-RPC errors and ``isError`` tool results (e.g. Exa's free-tier rate limit).""" def _from_payload(payload: str) -> Optional[str]: payload = payload.strip() @@ -123,51 +131,30 @@ def _parse_mcp_body(body: str) -> str: if err: raise KeylessMCPError(str(err.get("message") or err)) result = data.get("result") or {} - content = result.get("content") or [] + texts = [c.get("text", "") for c in result.get("content") or [] if isinstance(c, dict)] if result.get("isError"): - texts = [c.get("text", "") for c in content if isinstance(c, dict)] raise KeylessMCPError(" ".join(t for t in texts if t) or "MCP tool call failed") - for item in content: - if isinstance(item, dict) and item.get("text"): - return str(item["text"]) - return None + return next((str(t) for t in texts if t), None) stripped = body.strip() - if stripped.startswith("{"): + candidates = [stripped] if stripped.startswith("{") else [] + candidates += [line[len("data: "):] for line in body.splitlines() if line.startswith("data: ")] + for candidate in candidates: try: - text = _from_payload(stripped) - if text is not None: - return text - except json.JSONDecodeError: - pass - - for line in body.splitlines(): - if not line.startswith("data: "): - continue - try: - text = _from_payload(line[len("data: "):]) + text = _from_payload(candidate) except json.JSONDecodeError: continue if text is not None: return text - raise KeylessMCPError("Unrecognized MCP response shape") def mcp_call(url: str, tool: str, arguments: Dict[str, Any], timeout: int = _TIMEOUT_SECONDS) -> str: - """POST a JSON-RPC ``tools/call`` to *url* and return the text payload. - - Raises :class:`KeylessMCPError` on transport failures, non-2xx statuses, - JSON-RPC errors, and error-shaped tool results. - """ + """POST a JSON-RPC ``tools/call`` and return the text payload. Raises + :class:`KeylessMCPError` on transport failures, non-2xx, JSON-RPC and tool errors.""" import requests - payload = {"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": tool, "arguments": arguments}} - headers = { - "Content-Type": "application/json", - "Accept": "application/json, text/event-stream", - "User-Agent": "hermes-agent", - } + headers = {"Content-Type": "application/json", "Accept": "application/json, text/event-stream", "User-Agent": "hermes-agent"} try: response = requests.post(url, json=payload, headers=headers, timeout=timeout) except requests.RequestException as exc: @@ -178,62 +165,43 @@ def mcp_call(url: str, tool: str, arguments: Dict[str, Any], timeout: int = _TIM # --- Parallel (search.parallel.ai) — JSON text payloads ----------------------- - - def parallel_search_keyless(query: str, limit: int = 5) -> Dict[str, Any]: - """Keyless Parallel web search → legacy search response shape.""" - try: - text = mcp_call( - PARALLEL_MCP_URL, "web_search", - {"objective": query, "search_queries": [query], "session_id": _SESSION_ID}, - ) - data = json.loads(text) - web_results = [] - for i, result in enumerate(data.get("results") or []): - if limit and i >= limit: - break - web_results.append(_row( - result.get("url") or "", result.get("title") or "", - " ".join(result.get("excerpts") or []), i + 1, - )) - return search_ok(web_results) - except KeylessMCPError as exc: - return search_fail(_fail_msg("parallel", "search", exc)) - except (json.JSONDecodeError, TypeError, KeyError) as exc: - return search_fail(f"Keyless Parallel search returned an unexpected payload: {exc}") + def _rows() -> List[Dict[str, Any]]: + text = mcp_call(PARALLEL_MCP_URL, "web_search", {"objective": query, "search_queries": [query], "session_id": _SESSION_ID}) + results = json.loads(text).get("results") or [] + return [ + _row(r.get("url") or "", r.get("title") or "", " ".join(r.get("excerpts") or []), i + 1) + for i, r in enumerate(results[:max(limit, 0)] if limit else results) + ] + + return _search("parallel", _rows, (json.JSONDecodeError, TypeError, KeyError), lambda exc: f"Keyless Parallel search returned an unexpected payload: {exc}") def parallel_extract_keyless(urls: List[str]) -> List[Dict[str, Any]]: - """Keyless Parallel web fetch → legacy extract result list.""" try: - text = mcp_call( - PARALLEL_MCP_URL, "web_fetch", - {"urls": list(urls), "objective": "Full page content", "session_id": _SESSION_ID}, - ) - data = json.loads(text) + data = json.loads(mcp_call(PARALLEL_MCP_URL, "web_fetch", {"urls": list(urls), "objective": "Full page content", "session_id": _SESSION_ID})) except (KeylessMCPError, json.JSONDecodeError, TypeError) as exc: message = _fail_msg("parallel", "extract", exc) return [_page_error(u, message) for u in urls] - - results: List[Dict[str, Any]] = [] - seen = set() - for result in data.get("results") or []: - url = result.get("url") or "" - content = result.get("full_content") or result.get("content") or "\n\n".join(result.get("excerpts") or []) - seen.add(url) - results.append(_page(url, result.get("title") or "", content)) + results = [ + _page(r.get("url") or "", r.get("title") or "", r.get("full_content") or r.get("content") or "\n\n".join(r.get("excerpts") or [])) + for r in data.get("results") or [] + ] for error in data.get("errors") or []: url = error.get("url") or "" - seen.add(url) - entry = _page_error(url, str(error.get("content") or error.get("error_type") or "extraction failed")) - entry["metadata"] = {"sourceURL": url} - results.append(entry) + results.append({**_page_error(url, str(error.get("content") or error.get("error_type") or "extraction failed")), "metadata": {"sourceURL": url}}) # URLs the endpoint silently dropped still get an error entry (per-URL contract). + seen = {r["url"] for r in results} results.extend(_page_error(u, "no content returned") for u in urls if u not in seen) return results # --- Exa (mcp.exa.ai) — formatted plain-text payloads ------------------------- +def _after(line: str, prefix: str) -> str: + return line[len(prefix):].strip() + + +_EXA_LABELS = ("Title:", "URL:", "Highlights:", "Published:", "Author:") def _parse_exa_search_text(text: str, limit: int) -> List[Dict[str, Any]]: @@ -243,20 +211,16 @@ def _parse_exa_search_text(text: str, limit: int) -> List[Dict[str, Any]]: title = url = "" highlight_lines: List[str] = [] in_highlights = False - for line in block.splitlines(): - stripped = line.strip() + for stripped in map(str.strip, block.splitlines()): if stripped.startswith("Title:"): - title = stripped[len("Title:"):].strip() - in_highlights = False + title = _after(stripped, "Title:") elif stripped.startswith("URL:"): - url = stripped[len("URL:"):].strip() - in_highlights = False - elif stripped.startswith("Highlights:"): - in_highlights = True - elif stripped.startswith(("Published:", "Author:")): - in_highlights = False - elif in_highlights and stripped: + url = _after(stripped, "URL:") + elif in_highlights and stripped and not stripped.startswith(_EXA_LABELS): highlight_lines.append(stripped) + # Highlights run until the next labelled field. + if stripped.startswith(_EXA_LABELS): + in_highlights = stripped.startswith("Highlights:") if url: results.append(_row(url, title, " ".join(highlight_lines), len(results) + 1)) if limit and len(results) >= limit: @@ -265,77 +229,44 @@ def _parse_exa_search_text(text: str, limit: int) -> List[Dict[str, Any]]: def exa_search_keyless(query: str, limit: int = 5) -> Dict[str, Any]: - """Keyless Exa web search → legacy search response shape.""" - try: - text = mcp_call(EXA_MCP_URL, "web_search_exa", {"query": query, "numResults": max(1, int(limit))}) - except KeylessMCPError as exc: - return search_fail(_fail_msg("exa", "search", exc)) - return search_ok(_parse_exa_search_text(text, limit)) + return _search("exa", lambda: _parse_exa_search_text(mcp_call(EXA_MCP_URL, "web_search_exa", {"query": query, "numResults": max(1, int(limit))}), limit)) def exa_extract_keyless(urls: List[str]) -> List[Dict[str, Any]]: - """Keyless Exa web fetch → legacy extract result list (called per-URL; the - tool returns one combined text payload).""" - results: List[Dict[str, Any]] = [] - for url in urls: - try: - text = mcp_call(EXA_MCP_URL, "web_fetch_exa", {"urls": [url]}) - except KeylessMCPError as exc: - results.append(_page_error(url, _fail_msg("exa", "extract", exc))) - continue - title = "" - for line in text.splitlines(): - stripped = line.strip() - if stripped.startswith("# "): - title = stripped[2:].strip() - break - if stripped.startswith("Title:"): - title = stripped[len("Title:"):].strip() - break - results.append(_page(url, title, text)) - return results + """Called per-URL; the tool returns one combined text payload.""" + def _fetch(url: str) -> Dict[str, Any]: + text = mcp_call(EXA_MCP_URL, "web_fetch_exa", {"urls": [url]}) + # Title: first markdown H1 or ``Title:`` line, whichever comes first. + titles = (_after(s, "# " if s.startswith("# ") else "Title:") for s in map(str.strip, text.splitlines()) if s.startswith(("# ", "Title:"))) + return _page(url, next(titles, ""), text) + + return _per_url(urls, _fetch, "exa", catch=KeylessMCPError, hint=True) # --- Firecrawl keyless (public cloud API, no auth header) --------------------- - - def firecrawl_search_keyless(query: str, limit: int = 5) -> Dict[str, Any]: - """Keyless Firecrawl cloud search → legacy search response shape.""" from plugins.web.firecrawl.provider import _KeylessFirecrawlClient, _extract_web_search_results - - try: - response = _KeylessFirecrawlClient().search(query=query, limit=limit) - return search_ok(_extract_web_search_results(response)) - except Exception as exc: # noqa: BLE001 — normalized below - return search_fail(_fail_msg("firecrawl", "search", exc)) + rows = lambda: _extract_web_search_results(_KeylessFirecrawlClient().search(query=query, limit=limit)) # noqa: E731 + return _search("firecrawl", rows, Exception, lambda exc: _fail_msg("firecrawl", "search", exc)) def firecrawl_extract_keyless(urls: List[str]) -> List[Dict[str, Any]]: - """Keyless Firecrawl cloud scrape → legacy extract result list.""" from plugins.web.firecrawl.provider import _KeylessFirecrawlClient, _extract_scrape_payload - client = _KeylessFirecrawlClient() - results: List[Dict[str, Any]] = [] - for url in urls: - try: - payload = _extract_scrape_payload(client.scrape(url=url, formats=["markdown"])) or {} - metadata = payload.get("metadata") or {} - if not isinstance(metadata, dict): - metadata = {} - content = payload.get("markdown") or payload.get("html") or "" - results.append(_page(url, metadata.get("title") or "", content)) - except Exception as exc: # noqa: BLE001 — per-URL error entry - results.append(_page_error(url, _fail_msg("firecrawl", "extract", exc, other_backends=False))) - return results + + def _fetch(url: str) -> Dict[str, Any]: + payload = _extract_scrape_payload(client.scrape(url=url, formats=["markdown"])) or {} + metadata = payload.get("metadata") or {} + title = metadata.get("title") if isinstance(metadata, dict) else None + return _page(url, title or "", payload.get("markdown") or payload.get("html") or "") + + return _per_url(urls, _fetch, "firecrawl") # --- Keenable keyless (api.keenable.ai public endpoints) ---------------------- - - def _keenable_request(method: str, path: str, **kwargs: Any) -> Dict[str, Any]: """Call a Keenable public endpoint with the mandatory X-Keenable-Title app id.""" import requests - headers = {"X-Keenable-Title": _KEENABLE_TITLE} if method == "post": headers["Content-Type"] = "application/json" @@ -346,43 +277,31 @@ def _keenable_request(method: str, path: str, **kwargs: Any) -> Dict[str, Any]: def keenable_search_keyless(query: str, limit: int = 5) -> Dict[str, Any]: - """Keyless Keenable search (POST /v1/search/public) → legacy search response shape.""" - try: + def _rows() -> List[Dict[str, Any]]: data = _keenable_request("post", "/v1/search/public", json={"query": query, "max_results": max(1, int(limit))}) - except KeylessMCPError as exc: - return search_fail(_fail_msg("keenable", "search", exc)) - except Exception as exc: # noqa: BLE001 — transport/JSON errors - return search_fail(f"Keyless Keenable search failed: {exc}.") - return search_ok([ - _row(r.get("url") or "", r.get("title") or "", r.get("snippet") or r.get("description") or "", i + 1) - for i, r in enumerate(data.get("results") or []) - ]) + return [ + _row(r.get("url") or "", r.get("title") or "", r.get("snippet") or r.get("description") or "", i + 1) + for i, r in enumerate(data.get("results") or []) + ] + + return _search("keenable", _rows, Exception, lambda exc: f"Keyless Keenable search failed: {exc}.") def keenable_extract_keyless(urls: List[str]) -> List[Dict[str, Any]]: - """Keyless Keenable page fetch (GET /v1/fetch/public, per-URL) → legacy extract list.""" - results: List[Dict[str, Any]] = [] - for url in urls: - try: - data = _keenable_request("get", "/v1/fetch/public", params={"url": url}) - results.append(_page(data.get("url") or url, data.get("title") or "", data.get("content") or "", source_url=url)) - except Exception as exc: # noqa: BLE001 — per-URL error entry - results.append(_page_error(url, _fail_msg("keenable", "extract", exc, other_backends=False))) - return results + def _fetch(url: str) -> Dict[str, Any]: + data = _keenable_request("get", "/v1/fetch/public", params={"url": url}) + return _page(data.get("url") or url, data.get("title") or "", data.get("content") or "", source_url=url) + + return _per_url(urls, _fetch, "keenable") # --- Round-robin ring + next-in-line failover (rate-limited free tiers) ------- - _KEYLESS_RING = ("exa", "parallel", "firecrawl", "keenable") # Late-bound lookups (not bare references) so ``patch.object(keyless_mcp, "_search_keyless")`` # is honored at call time. Tests also ``setitem`` these dicts directly. -_KEYLESS_SEARCHERS: Dict[str, Callable[[str, int], Dict[str, Any]]] = { - v: (lambda query, limit, _v=v: globals()[f"{_v}_search_keyless"](query, limit)) for v in _KEYLESS_RING -} -_KEYLESS_EXTRACTORS: Dict[str, Callable[[List[str]], List[Dict[str, Any]]]] = { - v: (lambda urls, _v=v: globals()[f"{_v}_extract_keyless"](urls)) for v in _KEYLESS_RING -} +_KEYLESS_SEARCHERS: Dict[str, Callable[[str, int], Dict[str, Any]]] = {v: (lambda query, limit, _v=v: globals()[f"{_v}_search_keyless"](query, limit)) for v in _KEYLESS_RING} +_KEYLESS_EXTRACTORS: Dict[str, Callable[[List[str]], List[Dict[str, Any]]]] = {v: (lambda urls, _v=v: globals()[f"{_v}_extract_keyless"](urls)) for v in _KEYLESS_RING} # Per-process round-robin cursor, seeded by the random session id so the fleet # spreads across vendors; advances once per unpinned keyless request. @@ -396,22 +315,16 @@ def _vendor_pinned(name: str) -> bool: if provider_tier(name) == "free": return True try: - import tools.web_tools as _wt - - web_cfg = _wt._load_web_config() - return any( - (web_cfg.get(key) or "").lower().strip() == name - for key in ("backend", "search_backend", "extract_backend") - ) + return _web_config_selects(name) except Exception as exc: # noqa: BLE001 — config layer optional logger.debug("_vendor_pinned(%r) config read failed: %s", name, exc) return False def _ring_order(name: str) -> List[str]: - """Vendor walk order: pinned → start at *name* (its ring position fixes the - failover succession); else round-robin from the cursor, advancing it per request. - Vendors pinned ``paid`` are excluded (explicit paid opts their free endpoint out).""" + """Vendor walk order: pinned → start at *name* (its ring position fixes the failover + succession); else round-robin from the cursor, advancing it per request. Vendors + pinned ``paid`` are excluded (explicit paid opts their free endpoint out).""" global _ring_cursor if _vendor_pinned(name): start = _KEYLESS_RING.index(name) if name in _KEYLESS_RING else 0 @@ -427,12 +340,11 @@ _ALL_PAID_MSG = "All keyless web providers are pinned to paid tiers." def _walk_ring(name: str, kind: str, call, throttled) -> tuple: - """Shared ring walk: call each vendor from :func:`_ring_order` until a result is - not ``throttled``. Returns ``(order, vendor, result, exhausted)``; ``order`` is - empty (and result None) when every vendor is pinned paid.""" + """Call each vendor from :func:`_ring_order` until a result is not ``throttled``. + Returns ``(order, vendor, result, exhausted)``; ``order`` is empty (result None) + when every vendor is pinned paid.""" order = _ring_order(name) - vendor = None - result: Any = None + vendor, result = None, None for i, vendor in enumerate(order): result = call(vendor) if not throttled(result): @@ -443,16 +355,14 @@ def _walk_ring(name: str, kind: str, call, throttled) -> tuple: def search_with_failover(name: str, query: str, limit: int = 5) -> Dict[str, Any]: - """Keyless search across the ring; rate-limit-shaped errors advance to the - next vendor, other errors stop the walk (a malformed query fails everywhere). - ``data.served_by`` is set when the serving vendor differs from *name*.""" + """Rate-limit-shaped errors advance to the next vendor, other errors stop the walk + (a malformed query fails everywhere). ``data.served_by`` is set when the serving + vendor differs from *name*.""" def _throttled(result: Dict[str, Any]) -> bool: return not result.get("success") and _is_rate_limitish(result.get("error", "")) - order, vendor, result, exhausted = _walk_ring( - name, "search", lambda v: _KEYLESS_SEARCHERS[v](query, limit), _throttled - ) + order, vendor, result, exhausted = _walk_ring(name, "search", lambda v: _KEYLESS_SEARCHERS[v](query, limit), _throttled) if not order: return search_fail(_ALL_PAID_MSG) if exhausted: @@ -463,16 +373,13 @@ def search_with_failover(name: str, query: str, limit: int = 5) -> Dict[str, Any def extract_with_failover(name: str, urls: List[str]) -> List[Dict[str, Any]]: - """Keyless extract across the ring; fails over only when EVERY url in a batch - is rate-limit-shaped (partial failures are page problems, returned as-is).""" + """Fails over only when EVERY url in a batch is rate-limit-shaped (partial failures + are page problems, returned as-is).""" def _all_throttled(results: List[Dict[str, Any]]) -> bool: - errors = [r.get("error", "") for r in results] - return bool(results) and all(e and _is_rate_limitish(e) for e in errors) + return bool(results) and all(r.get("error", "") and _is_rate_limitish(r.get("error", "")) for r in results) - order, _vendor, results, _exhausted = _walk_ring( - name, "extract", lambda v: _KEYLESS_EXTRACTORS[v](list(urls)), _all_throttled - ) + order, _vendor, results, _exhausted = _walk_ring(name, "extract", lambda v: _KEYLESS_EXTRACTORS[v](list(urls)), _all_throttled) if not order: return [_page_error(u, _ALL_PAID_MSG) for u in urls] return results diff --git a/plugins/web/parallel/__init__.py b/plugins/web/parallel/__init__.py index 8e2d1c91bf..4eb2095a28 100644 --- a/plugins/web/parallel/__init__.py +++ b/plugins/web/parallel/__init__.py @@ -1,7 +1,5 @@ """Parallel.ai web search + extract plugin — bundled, auto-loaded; async-native ``extract``.""" - from __future__ import annotations - from plugins.web.parallel.provider import ParallelWebSearchProvider diff --git a/plugins/web/parallel/provider.py b/plugins/web/parallel/provider.py index 05be452496..c01d8665f6 100644 --- a/plugins/web/parallel/provider.py +++ b/plugins/web/parallel/provider.py @@ -12,17 +12,8 @@ import os from typing import Any, Dict, List from plugins.web._common import ( - SEARCH_LIMIT_CAP, - BaseWebSearchProvider, - cached_sdk_client, - document, - keyless_variant_schema, - provider_env, - run_extract_async, - run_search, - search_ok, - use_keyless, - web_hit, + SEARCH_LIMIT_CAP, BaseWebSearchProvider, cached_sdk_client, document, keyless_extract, keyless_search, + keyless_variant_schema, page_error, provider_env, run_extract_async, run_search, search_ok, use_keyless, web_hit, ) logger = logging.getLogger(__name__) @@ -33,7 +24,6 @@ _MISSING_KEY = "PARALLEL_API_KEY environment variable not set. Get your API key def _client(slot: str, cls_name: str) -> Any: def _factory(api_key: str) -> Any: import parallel # deliberately lazy - return getattr(parallel, cls_name)(api_key=api_key) return cached_sdk_client(slot, "PARALLEL_API_KEY", _MISSING_KEY, "search.parallel", _factory) @@ -48,8 +38,7 @@ def _get_async_client() -> Any: # Names re-exported by tools.web_tools for existing tests/callers. -_get_parallel_client = _get_sync_client -_get_async_parallel_client = _get_async_client +_get_parallel_client, _get_async_parallel_client = _get_sync_client, _get_async_client def _resolve_search_mode() -> str: @@ -68,17 +57,11 @@ class ParallelWebSearchProvider(BaseWebSearchProvider): def search(self, query: str, limit: int = 5) -> Dict[str, Any]: def _body() -> Dict[str, Any]: - from plugins.web.keyless_mcp import search_with_failover - if use_keyless("parallel", provider_env("PARALLEL_API_KEY")): - logger.info("Parallel keyless search: '%s' (limit=%d)", query, limit) - return search_with_failover("parallel", query, limit) - + return keyless_search("Parallel", "parallel", query, limit, logger) mode = _resolve_search_mode() logger.info("Parallel search: '%s' (mode=%s, limit=%d)", query, mode, limit) - response = _get_sync_client().beta.search( - search_queries=[query], objective=query, mode=mode, max_results=min(limit, SEARCH_LIMIT_CAP) - ) + response = _get_sync_client().beta.search(search_queries=[query], objective=query, mode=mode, max_results=min(limit, SEARCH_LIMIT_CAP)) return search_ok([ web_hit(r.url or "", r.title or "", " ".join(r.excerpts or []), i + 1) for i, r in enumerate(response.results or []) @@ -88,27 +71,16 @@ class ParallelWebSearchProvider(BaseWebSearchProvider): async def extract(self, urls: List[str], **kwargs: Any) -> List[Dict[str, Any]]: async def _body() -> List[Dict[str, Any]]: - from plugins.web.keyless_mcp import extract_with_failover - if use_keyless("parallel", provider_env("PARALLEL_API_KEY")): # Keyless ring is blocking HTTP — hop off the event loop. - logger.info("Parallel keyless extract: %d URL(s)", len(urls)) - return await asyncio.to_thread(extract_with_failover, "parallel", list(urls)) - + return await asyncio.to_thread(keyless_extract, "Parallel", "parallel", urls, logger) logger.info("Parallel extract: %d URL(s)", len(urls)) response = await _get_async_client().beta.extract(urls=urls, full_content=True) - - results = [ - document(r.url or "", r.title or "", r.full_content or "\n\n".join(r.excerpts or [])) - for r in response.results or [] + results = [document(r.url or "", r.title or "", r.full_content or "\n\n".join(r.excerpts or [])) for r in response.results or []] + return results + [ + {**page_error(e.url or "", e.content or e.error_type or "extraction failed"), "metadata": {"sourceURL": e.url or ""}} + for e in response.errors or [] ] - for error in response.errors or []: - results.append({ - "url": error.url or "", "title": "", "content": "", - "error": error.content or error.error_type or "extraction failed", - "metadata": {"sourceURL": error.url or ""}, - }) - return results return await run_extract_async("Parallel", logger, urls, _body, sdk=True) diff --git a/plugins/web/searxng/__init__.py b/plugins/web/searxng/__init__.py index 5ed1d8c73d..a690f007d0 100644 --- a/plugins/web/searxng/__init__.py +++ b/plugins/web/searxng/__init__.py @@ -1,7 +1,5 @@ """SearXNG search-only plugin — bundled, auto-loaded (``SEARXNG_URL``).""" - from __future__ import annotations - from plugins.web.searxng.provider import SearXNGWebSearchProvider diff --git a/plugins/web/searxng/provider.py b/plugins/web/searxng/provider.py index 35e6e22ff1..14f29a4630 100644 --- a/plugins/web/searxng/provider.py +++ b/plugins/web/searxng/provider.py @@ -9,7 +9,7 @@ from __future__ import annotations import logging from typing import Any, Dict -from plugins.web._common import BaseWebSearchProvider, http_get_json, provider_env, search_fail, search_ok +from plugins.web._common import BaseWebSearchProvider, http_get_json, provider_env, search_fail, search_ok, setup_schema, titled_rows logger = logging.getLogger(__name__) @@ -25,47 +25,21 @@ class SearXNGWebSearchProvider(BaseWebSearchProvider): base_url = provider_env("SEARXNG_URL").rstrip("/") if not base_url: return search_fail("SEARXNG_URL is not set") - data, failure = http_get_json( - "SearXNG", - f"{base_url}/search", - params={"q": query, "format": "json", "pageno": 1}, - headers={"Accept": "application/json"}, - timeout=15, - logger=logger, - reach_target=f"SearXNG at {base_url}", + "SearXNG", f"{base_url}/search", params={"q": query, "format": "json", "pageno": 1}, + headers={"Accept": "application/json"}, timeout=15, logger=logger, reach_target=f"SearXNG at {base_url}", ) if failure is not None: return failure - raw_results = data.get("results", []) # SearXNG may return a score field; sort descending and cap to limit. sorted_results = sorted(raw_results, key=lambda r: float(r.get("score", 0)), reverse=True)[:limit] - web_results = [ - { - "title": str(r.get("title", "")), - "url": str(r.get("url", "")), - "description": str(r.get("content", "")), - "position": i + 1, - } - for i, r in enumerate(sorted_results) - ] - logger.info( - "SearXNG search '%s': %d results (from %d raw, limit %d)", - query, len(web_results), len(raw_results), limit, - ) + web_results = titled_rows(sorted_results, "content") + logger.info("SearXNG search '%s': %d results (from %d raw, limit %d)", query, len(web_results), len(raw_results), limit) return search_ok(web_results) def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "SearXNG", - "badge": "free · self-hosted", - "tag": "Free, privacy-respecting metasearch. Point SEARXNG_URL at your instance.", - "env_vars": [ - { - "key": "SEARXNG_URL", - "prompt": "SearXNG instance URL (e.g. http://localhost:8080)", - "url": "https://searx.space/", - }, - ], - } + return setup_schema( + "SearXNG", "free · self-hosted", "Free, privacy-respecting metasearch. Point SEARXNG_URL at your instance.", + "SEARXNG_URL", "SearXNG instance URL (e.g. http://localhost:8080)", "https://searx.space/", + ) diff --git a/plugins/web/tavily/__init__.py b/plugins/web/tavily/__init__.py index 71da624b6b..18b3b7030f 100644 --- a/plugins/web/tavily/__init__.py +++ b/plugins/web/tavily/__init__.py @@ -1,7 +1,5 @@ """Tavily web search + extract plugin — bundled, auto-loaded.""" - from __future__ import annotations - from plugins.web.tavily.provider import TavilyWebSearchProvider diff --git a/plugins/web/tavily/provider.py b/plugins/web/tavily/provider.py index e29f16de0b..db31fc8566 100644 --- a/plugins/web/tavily/provider.py +++ b/plugins/web/tavily/provider.py @@ -14,17 +14,8 @@ from typing import Any, Dict, List, Optional import httpx from plugins.web._common import ( - SEARCH_LIMIT_CAP, - BaseWebSearchProvider, - document, - extract_fail, - http_status_detail, - provider_env, - run_extract, - run_search, - search_fail, - search_ok, - use_keyless, + SEARCH_LIMIT_CAP, BaseWebSearchProvider, document, extract_fail, http_status_detail, provider_env, run_extract, + run_search, search_fail, search_ok, setup_schema, title_hit, use_keyless, ) logger = logging.getLogger(__name__) @@ -44,18 +35,15 @@ def _tavily_headers(api_key: str) -> Dict[str, str]: def _tavily_request(endpoint: str, payload: Dict[str, Any], *, api_key: Optional[str] = None) -> Dict[str, Any]: - """POST to Tavily and return parsed JSON. - - ``api_key=None`` reads ``TAVILY_API_KEY``; pass ``""`` to force the keyless - header even when a key exists (``web.provider_tier.tavily: free``). Non-2xx - raises ValueError with the body so Tavily's rate-limit/upgrade text reaches the model. - """ + """POST to Tavily and return parsed JSON. ``api_key=None`` reads ``TAVILY_API_KEY``; + pass ``""`` to force the keyless header even when a key exists + (``web.provider_tier.tavily: free``). Non-2xx raises ValueError with the body so + Tavily's rate-limit/upgrade text reaches the model.""" if api_key is None: api_key = provider_env("TAVILY_API_KEY") base_url = provider_env("TAVILY_BASE_URL") or "https://api.tavily.com" url = f"{base_url}/{endpoint.lstrip('/')}" logger.info("Tavily %s request to %s", endpoint, url) - response = httpx.post(url, json=payload, timeout=60, headers=_tavily_headers(api_key)) if response.status_code >= 400: raise ValueError(http_status_detail(response)) @@ -63,25 +51,20 @@ def _tavily_request(endpoint: str, payload: Dict[str, Any], *, api_key: Optional def _normalize_tavily_search_results(response: Dict[str, Any]) -> Dict[str, Any]: - # title-first key order is Tavily's historical wire shape (differs from web_hit). return search_ok([ - {"title": r.get("title", ""), "url": r.get("url", ""), "description": r.get("content", ""), "position": i + 1} + title_hit(r.get("title", ""), r.get("url", ""), r.get("content", ""), i + 1) for i, r in enumerate(response.get("results", [])) ]) def _normalize_tavily_documents(response: Dict[str, Any], fallback_url: str = "") -> List[Dict[str, Any]]: """Map ``/extract`` to documents; ``failed_results`` / ``failed_urls`` become ``error`` entries.""" - documents: List[Dict[str, Any]] = [] - for result in response.get("results", []): - url = result.get("url", fallback_url) - raw = result.get("raw_content", "") or result.get("content", "") - documents.append(document(url, result.get("title", ""), raw)) - for fail in response.get("failed_results", []): - url = fail.get("url", fallback_url) - documents.append(_failed_document(url, fail.get("error", "extraction failed"))) - for fail_url in response.get("failed_urls", []): - documents.append(_failed_document(str(fail_url), "extraction failed")) + documents = [ + document(r.get("url", fallback_url), r.get("title", ""), r.get("raw_content", "") or r.get("content", "")) + for r in response.get("results", []) + ] + documents += [_failed_document(f.get("url", fallback_url), f.get("error", "extraction failed")) for f in response.get("failed_results", [])] + documents += [_failed_document(str(u), "extraction failed") for u in response.get("failed_urls", [])] return documents @@ -90,10 +73,17 @@ def _failed_document(url: str, error: str) -> Dict[str, Any]: def _missing_key_error(action: str) -> str: - return ( - "TAVILY_API_KEY is not set. Get a key at https://app.tavily.com/home " - f"or select Tavily in `hermes tools` for opt-in keyless {action}." - ) + return f"TAVILY_API_KEY is not set. Get a key at https://app.tavily.com/home or select Tavily in `hermes tools` for opt-in keyless {action}." + + +def _auth(action: str) -> tuple[Optional[str], Optional[str], str]: + """``(request_key, missing_key_error, log_prefix)``: request key is ``""`` when forcing + keyless, ``None`` when neither key nor keyless applies (``missing_key_error`` set).""" + api_key = provider_env("TAVILY_API_KEY") + force_keyless = use_keyless("tavily", api_key) + if not force_keyless and not api_key: + return None, _missing_key_error(action), "" + return "" if force_keyless else api_key, None, "keyless " if force_keyless else "" class TavilyWebSearchProvider(BaseWebSearchProvider): @@ -107,42 +97,28 @@ class TavilyWebSearchProvider(BaseWebSearchProvider): def search(self, query: str, limit: int = 5) -> Dict[str, Any]: def _body() -> Dict[str, Any]: - api_key = provider_env("TAVILY_API_KEY") - force_keyless = use_keyless("tavily", api_key) - if not force_keyless and not api_key: - return search_fail(_missing_key_error("search")) - - logger.info("Tavily %ssearch: '%s' (limit=%d)", "keyless " if force_keyless else "", query, limit) + key, missing, prefix = _auth("search") + if missing: + return search_fail(missing) + logger.info("Tavily %ssearch: '%s' (limit=%d)", prefix, query, limit) payload = {"query": query, "max_results": min(limit, SEARCH_LIMIT_CAP), **_SEARCH_PAYLOAD} - return _normalize_tavily_search_results( - _tavily_request("search", payload, api_key="" if force_keyless else api_key) - ) + return _normalize_tavily_search_results(_tavily_request("search", payload, api_key=key)) return run_search("Tavily", logger, _body) def extract(self, urls: List[str], **kwargs: Any) -> List[Dict[str, Any]]: def _body() -> List[Dict[str, Any]]: - api_key = provider_env("TAVILY_API_KEY") - force_keyless = use_keyless("tavily", api_key) - if not force_keyless and not api_key: - return extract_fail(urls, _missing_key_error("extract")) - - logger.info("Tavily %sextract: %d URL(s)", "keyless " if force_keyless else "", len(urls)) - raw = _tavily_request("extract", {"urls": urls, "include_images": False}, api_key="" if force_keyless else api_key) + key, missing, prefix = _auth("extract") + if missing: + return extract_fail(urls, missing) + logger.info("Tavily %sextract: %d URL(s)", prefix, len(urls)) + raw = _tavily_request("extract", {"urls": urls, "include_images": False}, api_key=key) return _normalize_tavily_documents(raw, fallback_url=urls[0] if urls else "") return run_extract("Tavily", logger, urls, _body) def get_setup_schema(self) -> Dict[str, Any]: - return { - "name": "Tavily", - "badge": "free · key optional", - "tag": "Search + extract. Opt-in keyless; set TAVILY_API_KEY for higher limits.", - "env_vars": [ - { - "key": "TAVILY_API_KEY", - "prompt": "Tavily API key (optional — keyless works when Tavily is selected)", - "url": "https://app.tavily.com/home", - }, - ], - } + return setup_schema( + "Tavily", "free · key optional", "Search + extract. Opt-in keyless; set TAVILY_API_KEY for higher limits.", + "TAVILY_API_KEY", "Tavily API key (optional — keyless works when Tavily is selected)", "https://app.tavily.com/home", + ) diff --git a/plugins/web/xai/__init__.py b/plugins/web/xai/__init__.py index a3aed4577d..662077818f 100644 --- a/plugins/web/xai/__init__.py +++ b/plugins/web/xai/__init__.py @@ -1,7 +1,5 @@ """xAI web search plugin — bundled, auto-loaded.""" - from __future__ import annotations - from plugins.web.xai.provider import XAIWebSearchProvider diff --git a/plugins/web/xai/provider.py b/plugins/web/xai/provider.py index 8620aa2345..18b12fc09d 100644 --- a/plugins/web/xai/provider.py +++ b/plugins/web/xai/provider.py @@ -1,13 +1,8 @@ -"""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``). +"""xAI Web Search — search-only provider backed by Grok's server-side ``web_search`` tool on the +Responses API (https://docs.x.ai/developers/tools/web-search); Grok is asked for structured JSON +so rows match every other Hermes web provider. Config: ``web.backend: "xai"``; optional ``web.xai``: +``model`` (default grok-build-0.1), ``allowed_domains`` / ``excluded_domains`` (max 5, mutually +exclusive), ``timeout`` (default 90s). Auth: Grok OAuth via ``hermes auth``, else XAI_API_KEY. """ from __future__ import annotations @@ -17,12 +12,8 @@ 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, -) +from plugins.web._common import BaseWebSearchProvider, search_fail as _fail, search_ok, setup_schema, title_hit as _row +from tools.xai_http import has_xai_credentials, hermes_xai_user_agent, resolve_xai_http_credentials logger = logging.getLogger(__name__) @@ -30,8 +21,7 @@ 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. +# Tolerates leading/trailing prose — reasoning models occasionally narrate before the JSON block. _JSON_BLOCK_RE = re.compile(r"\{[\s\S]*\}", re.MULTILINE) @@ -39,22 +29,17 @@ 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 {} + for key in ("web", "xai"): + cfg = cfg.get(key) if isinstance(cfg, dict) else None + return cfg if isinstance(cfg, 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] + return [item.strip() for item in value if isinstance(item, str) and item.strip()][:_MAX_DOMAIN_FILTERS] if isinstance(value, list) else [] def _coerce(cast, value: Any, default: Any) -> Any: @@ -64,22 +49,10 @@ def _coerce(cast, value: Any, default: Any) -> Any: 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. - """ + """Sends a structured prompt with ``tools=[{"type": "web_search"}]`` and parses the JSON Grok + returns; falls back to message annotations, then ``citations``. Trust model: Grok *generates* + the URLs/titles/descriptions and is steerable by the query text — validate before fetching.""" NAME = "xai" DISPLAY_NAME = "xAI Web Search (Grok)" @@ -90,67 +63,39 @@ class XAIWebSearchProvider(BaseWebSearchProvider): 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. + 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. 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) - + model = (cfg["model"].strip() if isinstance(cfg.get("model"), str) else "") or DEFAULT_MODEL 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"], - } - + return _fail("web.xai.allowed_domains and web.xai.excluded_domains cannot both be set (xAI restriction).") + # include=no_inline_citations keeps the JSON block clean; URLs come from annotations/citations. + payload: Dict[str, Any] = {"model": model, "input": [{"role": "user", "content": self._build_prompt(query, limit)}], "tools": [web_search_tool], "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, + base_url, payload, api_key, _coerce(float, cfg.get("timeout", DEFAULT_TIMEOUT), DEFAULT_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 @@ -158,50 +103,38 @@ class XAIWebSearchProvider(BaseWebSearchProvider): 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. + # Empty list on 0 hits is a success (matches brave-free / exa). 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: + """``web_search`` tool spec with optional domain filters; None when both + allowed and excluded are set (xAI rejects the combination).""" + filters = {k: _coerce_domain_list(cfg.get(k)) for k in ("allowed_domains", "excluded_domains")} + filters = {k: v for k, v in filters.items() if v} + if len(filters) == 2: 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 + return {"type": "web_search", "filters": filters} if filters else {"type": "web_search"} @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. + 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 ``/responses`` → ``(parsed_json, None)`` or ``(None, failure_envelope)``. - 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. + Two attempts: on a first-call 401 with OAuth creds, force-refresh once and retry + (opaque tokens the resolver can't pre-check; 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()} + def _refreshed_key() -> str: + """New bearer after a 401, or "" when refresh fails / returns the same token (retry would be pointless).""" + try: + key = str(resolve_xai_http_credentials(force_refresh=True, api_key_hint=api_key).get("api_key") or "").strip() + return key if key != api_key else "" + except Exception as refresh_exc: # noqa: BLE001 + logger.warning("xAI web search OAuth refresh after 401 failed: %s", refresh_exc) + return "" - headers = { - "Authorization": f"Bearer {api_key}", - "Content-Type": "application/json", - "User-Agent": hermes_xai_user_agent(), - } resp = None for attempt in range(2): try: @@ -211,20 +144,10 @@ class XAIWebSearchProvider(BaseWebSearchProvider): 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) + logger.info("xAI web search got 401 on first attempt; forcing OAuth refresh and retrying once.") + if new_key := _refreshed_key(): + api_key, headers["Authorization"] = new_key, f"Bearer {new_key}" + continue try: body = exc.response.text[:300] if exc.response is not None else "" except Exception: @@ -234,163 +157,104 @@ class XAIWebSearchProvider(BaseWebSearchProvider): 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.""" + """Ask 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 " + "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``.""" + """Rows 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. (2) only short-circuits 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 - + parsed = next((p for p in (cls._try_parse_json_results(b, limit=limit) for b in text_blocks) if p), None) + if parsed or (annotations and (parsed := cls._results_from_annotations(annotations, "\n".join(text_blocks), limit=limit))): + return parsed 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 [] + return [_row("", str(u), "", i + 1) for i, u in enumerate(citations[:limit]) if isinstance(u, str) and u.strip()] if isinstance(citations, list) else [] @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]] = [] + """(text_blocks, annotations) from ``response.output`` message chunks.""" 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)) + chunks = [ + chunk + for item in (output if isinstance(output, list) else []) + if isinstance(item, dict) and item.get("type") == "message" and isinstance(item.get("content"), list) + for chunk in item["content"] + if isinstance(chunk, dict) and chunk.get("type") == "output_text" + ] + text_blocks = [c["text"] for c in chunks if isinstance(c.get("text"), str) and c["text"].strip()] + annotations = [a for c in chunks if isinstance(c.get("annotations"), list) for a in c["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] + """Parse a JSON object with a ``results`` array out of ``text``; None when absent. + Whole string first, then the regex-matched block (reasoning models prefix narration).""" match = _JSON_BLOCK_RE.search(text) - if match and match.group(0) != text: - candidates.append(match.group(0)) - - for candidate in candidates: + for candidate in [text] + ([match.group(0)] if match and match.group(0) != text else []): try: parsed = json.loads(candidate) except (json.JSONDecodeError, ValueError): continue - if not isinstance(parsed, dict): - continue - results = parsed.get("results") + results = parsed.get("results") if isinstance(parsed, dict) else None 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, - )) + url = str(row.get("url", "")).strip() if isinstance(row, dict) else "" + if url: + # 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]]: + 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() + url = str(ann.get("url", "")).strip() if ann.get("type") == "url_citation" else "" if not url or url in seen: continue seen.add(url) - description = "" - start = ann.get("start_index") - end = ann.get("end_index") + start, end = ann.get("start_index"), 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", - } + return setup_schema( + "xAI Web Search (Grok)", "paid", + "Agentic web search via Grok's web_search tool — uses xAI Grok OAuth or XAI_API_KEY.", post_setup="xai_grok", + )