refactor(plugins/web,video_gen): unify web-provider/video_gen plugin scaffolding into _common helpers, drop dead code (3739->2755 LOC, request/result parity verified)
This commit is contained in:
@@ -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:
|
||||
|
||||
+136
-289
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
+118
-261
@@ -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.<helper>(*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:
|
||||
|
||||
@@ -1,7 +1,3 @@
|
||||
# Bundled web search providers — plugins/web/.
|
||||
#
|
||||
# Each subdirectory follows the image_gen plugin layout:
|
||||
# plugins/web/<name>/{plugin.yaml, __init__.py, provider.py}
|
||||
#
|
||||
# They auto-load via kind: backend and register via
|
||||
# Bundled web search providers: plugins/web/<name>/{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.
|
||||
|
||||
+77
-93
@@ -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 <code>`` 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.<slot>``.
|
||||
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._<vendor>_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.<slot>`` (so tests that
|
||||
reset ``tools.web_tools._<vendor>_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.<name>: 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.<name>: 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])
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""Brave Search (free tier) plugin — bundled, auto-loaded."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from plugins.web.brave_free.provider import BraveFreeWebSearchProvider
|
||||
|
||||
|
||||
|
||||
@@ -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/",
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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__":
|
||||
|
||||
+64
-128
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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 []]
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""Firecrawl web search + extract plugin — bundled, auto-loaded."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from plugins.web.firecrawl.provider import FirecrawlWebSearchProvider
|
||||
|
||||
|
||||
|
||||
+115
-289
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
+133
-226
@@ -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.<name>`` (set by the ``hermes tools`` Free/Paid rows):
|
||||
``free``, ``paid``, or ``auto`` for anything else including unset."""
|
||||
"""``web.provider_tier.<name>`` (``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, "<vendor>_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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""SearXNG search-only plugin — bundled, auto-loaded (``SEARXNG_URL``)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from plugins.web.searxng.provider import SearXNGWebSearchProvider
|
||||
|
||||
|
||||
|
||||
@@ -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/",
|
||||
)
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""Tavily web search + extract plugin — bundled, auto-loaded."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from plugins.web.tavily.provider import TavilyWebSearchProvider
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""xAI web search plugin — bundled, auto-loaded."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from plugins.web.xai.provider import XAIWebSearchProvider
|
||||
|
||||
|
||||
|
||||
+86
-222
@@ -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",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user