Merge branch 'simp/hclia-providers' into simp/integration
This commit is contained in:
+30
-121
@@ -2,25 +2,15 @@
|
||||
Browser Provider ABC
|
||||
====================
|
||||
|
||||
Defines the pluggable-backend interface for cloud browser providers
|
||||
(Browserbase, Browser Use, Firecrawl, …). Providers register instances via
|
||||
:meth:`PluginContext.register_browser_provider`; the active one (selected via
|
||||
Pluggable-backend interface for cloud browser providers (Browserbase, Browser
|
||||
Use, Firecrawl, …). Providers register via
|
||||
:meth:`PluginContext.register_browser_provider`; the active one (selected by
|
||||
``browser.cloud_provider`` in ``config.yaml``) services every cloud-mode
|
||||
``browser_*`` tool call.
|
||||
``browser_*`` tool call. Providers live in ``<repo>/plugins/browser/<name>/``
|
||||
(built-in) or ``~/.hermes/plugins/browser/<name>/`` (user, opt-in).
|
||||
|
||||
Providers live in ``<repo>/plugins/browser/<name>/`` (built-in, auto-loaded as
|
||||
``kind: backend``) or ``~/.hermes/plugins/browser/<name>/`` (user, opt-in via
|
||||
``plugins.enabled``).
|
||||
|
||||
This ABC mirrors :class:`agent.web_search_provider.WebSearchProvider` (PR
|
||||
#25182) — same shape, same registration flow, same picker integration. The
|
||||
legacy in-tree ``tools.browser_providers.base.CloudBrowserProvider`` ABC was
|
||||
deleted in PR #25214 (this work) along with the per-vendor inline modules in
|
||||
``tools/browser_providers/``; the lifecycle contract documented below is
|
||||
preserved bit-for-bit so the tool wrapper (:mod:`tools.browser_tool`) does
|
||||
not have to translate.
|
||||
|
||||
Session metadata contract (preserved from the legacy ``CloudBrowserProvider``)::
|
||||
Session metadata contract (preserved from the legacy ``CloudBrowserProvider``
|
||||
so :mod:`tools.browser_tool` needs no translation)::
|
||||
|
||||
{
|
||||
"session_name": str, # unique name for agent-browser --session
|
||||
@@ -31,79 +21,41 @@ Session metadata contract (preserved from the legacy ``CloudBrowserProvider``)::
|
||||
"external_call_id": str, # optional, managed-gateway billing key
|
||||
}
|
||||
|
||||
``bb_session_id`` is a legacy key name kept verbatim for backward compat with
|
||||
:mod:`tools.browser_tool` — it holds the provider's session ID regardless of
|
||||
which provider is in use.
|
||||
``bb_session_id`` is a legacy key name kept verbatim for backward compat — it
|
||||
holds the provider's session ID regardless of which provider is in use.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Dict
|
||||
|
||||
from agent.provider_base import ProviderBase
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BrowserProvider(abc.ABC):
|
||||
class BrowserProvider(ProviderBase):
|
||||
"""Abstract base class for a cloud browser backend.
|
||||
|
||||
Subclasses must implement :meth:`name`, :meth:`is_available`, and the
|
||||
three lifecycle methods: :meth:`create_session`, :meth:`close_session`,
|
||||
:meth:`emergency_cleanup`.
|
||||
|
||||
The lifecycle shape preserves the legacy ``CloudBrowserProvider`` contract
|
||||
bit-for-bit so the dispatcher in :mod:`tools.browser_tool` is a pure
|
||||
registry lookup — no per-provider conditionals, no shape translation.
|
||||
Subclasses implement :attr:`name` (the ``browser.cloud_provider`` value,
|
||||
e.g. ``browserbase``, ``browser-use``, ``firecrawl``), :meth:`is_available`,
|
||||
and the lifecycle trio :meth:`create_session` / :meth:`close_session` /
|
||||
:meth:`emergency_cleanup`. ``get_setup_schema`` may add ``"post_setup"``
|
||||
(e.g. ``"agent_browser"``) to trigger the install hook.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in the ``browser.cloud_provider``
|
||||
config key.
|
||||
|
||||
Lowercase, hyphens permitted to preserve existing user-visible names.
|
||||
Examples: ``browserbase``, ``browser-use``, ``firecrawl``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name``."""
|
||||
return self.name
|
||||
|
||||
@abc.abstractmethod
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
"""True when this provider can service calls.
|
||||
|
||||
Typically a cheap check (env var present, managed-gateway token
|
||||
readable, optional Python dep importable). Must NOT make network
|
||||
calls — this runs at tool-registration time and on every
|
||||
``hermes tools`` paint.
|
||||
|
||||
Mirrors the legacy ``CloudBrowserProvider.is_configured()`` method;
|
||||
renamed for parity with :class:`agent.web_search_provider.WebSearchProvider`.
|
||||
Cheap check only (env var present, managed-gateway token readable, dep
|
||||
importable) — must NOT make network calls; runs at tool-registration
|
||||
time and on every ``hermes tools`` paint.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def create_session(self, task_id: str) -> Dict[str, object]:
|
||||
"""Create a cloud browser session and return session metadata.
|
||||
|
||||
Must return a dict with at least::
|
||||
|
||||
{
|
||||
"session_name": str, # unique name for agent-browser --session
|
||||
"bb_session_id": str, # provider session ID (for close/cleanup)
|
||||
"cdp_url": str, # CDP websocket URL
|
||||
"expires_at": str, # optional provider-authoritative ISO timestamp
|
||||
"features": dict, # feature flags that were enabled
|
||||
}
|
||||
|
||||
``bb_session_id`` is a legacy key name kept for backward compat with
|
||||
the rest of :mod:`tools.browser_tool` — it holds the provider's
|
||||
session ID regardless of which provider is in use.
|
||||
"""Create a cloud browser session and return the session metadata dict
|
||||
described in the module docstring.
|
||||
|
||||
May raise ``ValueError`` (missing credentials) or ``RuntimeError``
|
||||
(network / API failure); the dispatcher surfaces these to the user.
|
||||
@@ -111,62 +63,19 @@ class BrowserProvider(abc.ABC):
|
||||
|
||||
@abc.abstractmethod
|
||||
def close_session(self, session_id: str) -> bool:
|
||||
"""Release / terminate a cloud session by its provider session ID.
|
||||
"""Release a cloud session by provider session ID.
|
||||
|
||||
Returns True on success, False on failure. Should not raise — log and
|
||||
return False on any exception so the dispatcher's cleanup loop keeps
|
||||
moving across sessions.
|
||||
return False so the dispatcher's cleanup loop keeps moving.
|
||||
"""
|
||||
|
||||
@abc.abstractmethod
|
||||
def emergency_cleanup(self, session_id: str) -> None:
|
||||
"""Best-effort session teardown during process exit.
|
||||
"""Best-effort teardown from atexit / signal handlers. Must tolerate
|
||||
missing credentials and network errors; must not raise."""
|
||||
|
||||
Called from atexit / signal handlers. Must tolerate missing
|
||||
credentials, network errors, etc. — log and move on. Must not raise.
|
||||
"""
|
||||
|
||||
def get_setup_schema(self) -> Optional[Dict[str, Any]]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by :mod:`hermes_cli.tools_config` to inject this provider as a
|
||||
row in the Browser Automation picker. Shape mirrors the existing
|
||||
hardcoded entries in ``TOOL_CATEGORIES["browser"]``::
|
||||
|
||||
{
|
||||
"name": "Browserbase",
|
||||
"badge": "paid",
|
||||
"tag": "Cloud browser with stealth and proxies",
|
||||
"env_vars": [
|
||||
{"key": "BROWSERBASE_API_KEY",
|
||||
"prompt": "Browserbase API key",
|
||||
"url": "https://browserbase.com"},
|
||||
],
|
||||
"post_setup": "agent_browser",
|
||||
}
|
||||
|
||||
Default: minimal entry derived from :attr:`display_name`. Override to
|
||||
expose API key prompts, badges, managed-Nous gating, and the
|
||||
``post_setup`` install hook.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Backward-compat shims for the legacy CloudBrowserProvider API
|
||||
# ------------------------------------------------------------------
|
||||
#
|
||||
# The pre-PR-#25214 ABC exposed ``is_configured()`` and ``provider_name()``;
|
||||
# ``tools.browser_tool`` has ~6 callers that still use those names. Rather
|
||||
# than churn every callsite (and break out-of-tree downstream code that
|
||||
# subclassed CloudBrowserProvider), we expose the old names as thin
|
||||
# delegations to the new API. Subclasses MUST implement :meth:`is_available`
|
||||
# and :attr:`name`; they may override ``is_configured`` / ``provider_name``
|
||||
# for compatibility with the legacy ABC but it is not required.
|
||||
# Legacy ``CloudBrowserProvider`` names still used by ``tools.browser_tool``
|
||||
# and out-of-tree subclasses; thin delegations to the current API.
|
||||
|
||||
def is_configured(self) -> bool:
|
||||
"""Backward-compat alias for :meth:`is_available`."""
|
||||
|
||||
+24
-172
@@ -37,117 +37,18 @@ job is purely selection, not capability routing.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from agent.browser_provider import BrowserProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
from agent.provider_registry import ProviderRegistry, is_available_safe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_providers: Dict[str, BrowserProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, BrowserProvider]] = {}
|
||||
_generation = 0
|
||||
_scoped_generations: Dict[str, int] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def register_provider(provider: BrowserProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a cloud browser provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and logs
|
||||
a debug message — makes hot-reload scenarios (tests, dev loops) behave
|
||||
predictably.
|
||||
"""
|
||||
if not isinstance(provider, BrowserProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a BrowserProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Browser provider .name must be a non-empty string")
|
||||
name = raw_name.strip()
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"Browser provider '%s' re-registered (was %r)",
|
||||
name, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered browser provider '%s' (%s)",
|
||||
name, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[BrowserProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[BrowserProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
with _lock:
|
||||
key = name.strip()
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[BrowserProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip())
|
||||
|
||||
|
||||
def registry_generation(*, scope: Optional[str] = None) -> tuple[int, int]:
|
||||
"""Return a cache fingerprint for the global base and one profile."""
|
||||
active_scope = scope or hermes_home_key()
|
||||
with _lock:
|
||||
return _generation, _scoped_generations.get(active_scope, 0)
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: BrowserProvider,
|
||||
previous: Optional[BrowserProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip()
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
_registry: ProviderRegistry[BrowserProvider] = ProviderRegistry(
|
||||
label="Browser", provider_cls=BrowserProvider, logger=logger,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -155,11 +56,9 @@ def restore_registration(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# Legacy auto-detect order — used when no ``browser.cloud_provider`` is set.
|
||||
# Matches the pre-migration walk in :func:`tools.browser_tool._get_cloud_provider`.
|
||||
# Firecrawl is intentionally absent so users with ``FIRECRAWL_API_KEY`` set
|
||||
# for web-extract don't get silently routed to a paid cloud browser. See
|
||||
# :func:`_resolve` for the full rationale.
|
||||
# Auto-detect order when ``browser.cloud_provider`` is unset (pre-migration
|
||||
# walk of :func:`tools.browser_tool._get_cloud_provider`); see :func:`_resolve`
|
||||
# for why Firecrawl is absent.
|
||||
_LEGACY_PREFERENCE = (
|
||||
"browser-use",
|
||||
"browserbase",
|
||||
@@ -167,60 +66,22 @@ _LEGACY_PREFERENCE = (
|
||||
|
||||
|
||||
def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]:
|
||||
"""Resolve the active browser provider.
|
||||
"""Resolve the active browser provider (rules in the module docstring).
|
||||
|
||||
Resolution rules (in order):
|
||||
|
||||
1. **Explicit "local".** Returns None — the dispatcher disables cloud
|
||||
mode entirely. Mirrors legacy short-circuit in
|
||||
:func:`tools.browser_tool._get_cloud_provider`.
|
||||
2. **Explicit config wins, ignoring availability.** If ``configured``
|
||||
names a registered provider, return it even if its
|
||||
:meth:`is_available` returns False — the dispatcher will surface a
|
||||
precise "X_API_KEY is not set" error instead of silently routing
|
||||
somewhere else.
|
||||
3. **Legacy preference walk, filtered by availability.** Walk
|
||||
:data:`_LEGACY_PREFERENCE` (``browser-use`` → ``browserbase``) looking
|
||||
for a provider whose ``is_available()`` is True.
|
||||
|
||||
There is intentionally NO "single-eligible shortcut" rule here (unlike
|
||||
:func:`agent.web_search_registry._resolve`). Pre-migration, the
|
||||
auto-detect branch in ``tools.browser_tool._get_cloud_provider`` only
|
||||
considered Browser Use and Browserbase; Firecrawl was reachable only
|
||||
via an explicit ``browser.cloud_provider: firecrawl`` config key.
|
||||
Preserving that gate matters because Firecrawl shares its API key with
|
||||
the *web* extract plugin (``plugins/web/firecrawl/``), so users who set
|
||||
``FIRECRAWL_API_KEY`` for web extract must NOT get silently routed to a
|
||||
paid cloud browser on a fresh install. Third-party browser-provider
|
||||
plugins added under ``~/.hermes/plugins/browser/<vendor>/`` are subject
|
||||
to the same gate — they must be explicitly configured to take effect.
|
||||
|
||||
Returns None when no provider is configured AND no available provider
|
||||
matches the legacy preference; the dispatcher then falls back to local
|
||||
browser mode.
|
||||
There is intentionally NO "single-eligible shortcut" (unlike
|
||||
:func:`agent.web_search_registry._resolve`): only ``_LEGACY_PREFERENCE``
|
||||
names are auto-eligible. Firecrawl shares its API key with the *web*
|
||||
extract plugin, so a user with ``FIRECRAWL_API_KEY`` must never be routed
|
||||
to a paid cloud browser without setting ``browser.cloud_provider``; the
|
||||
same gate applies to third-party browser-provider plugins.
|
||||
"""
|
||||
with _lock:
|
||||
snapshot = dict(_providers)
|
||||
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
|
||||
snapshot = _registry.merged()
|
||||
|
||||
def _is_available_safe(p: BrowserProvider) -> bool:
|
||||
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
|
||||
try:
|
||||
return bool(p.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning(
|
||||
"Browser provider %s.is_available() raised %s — treating as unavailable",
|
||||
p.name, exc, exc_info=True,
|
||||
)
|
||||
return False
|
||||
|
||||
# 1. Explicit "local" short-circuit.
|
||||
if configured == "local":
|
||||
return None
|
||||
|
||||
# 2. Explicit config wins — return regardless of is_available() so the
|
||||
# user gets a precise downstream error message rather than a silent
|
||||
# backend switch. Matches _get_cloud_provider() in browser_tool.py.
|
||||
# Explicit config wins regardless of is_available(): the dispatcher then
|
||||
# surfaces a precise "X_API_KEY is not set" error instead of a silent switch.
|
||||
if configured:
|
||||
provider = snapshot.get(configured)
|
||||
if provider is not None:
|
||||
@@ -231,23 +92,14 @@ def _resolve(configured: Optional[str]) -> Optional[BrowserProvider]:
|
||||
configured,
|
||||
)
|
||||
|
||||
# 3. Legacy preference walk — only providers in _LEGACY_PREFERENCE are
|
||||
# auto-eligible. Filtered by availability so we don't surface a
|
||||
# provider the user has no credentials for. See docstring for why
|
||||
# we do NOT fall back to "any single-eligible registered provider".
|
||||
for legacy in _LEGACY_PREFERENCE:
|
||||
provider = snapshot.get(legacy)
|
||||
if provider is not None and _is_available_safe(provider):
|
||||
if provider is not None and is_available_safe(
|
||||
provider, logger,
|
||||
"Browser provider %s.is_available() raised %s — treating as unavailable",
|
||||
level=logging.WARNING, exc_info=True,
|
||||
):
|
||||
return provider
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
global _generation
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
_scoped_generations.clear()
|
||||
_generation += 1
|
||||
|
||||
+85
-252
@@ -1,28 +1,14 @@
|
||||
"""Abstract base class for pluggable context engines.
|
||||
|
||||
A context engine controls how conversation context is managed when
|
||||
approaching the model's token limit. The built-in ContextCompressor
|
||||
is the default implementation. Third-party engines (e.g. LCM) can
|
||||
replace it via the plugin system or by being placed in the
|
||||
``plugins/context_engine/<name>/`` directory.
|
||||
A context engine decides when and how conversation context is compacted near
|
||||
the model's token limit, tracks token usage, and may expose tools. The
|
||||
built-in ContextCompressor is the default; ``context.engine`` in config.yaml
|
||||
selects a plugin engine (``plugins/context_engine/<name>/``). One engine is
|
||||
active at a time.
|
||||
|
||||
Selection is config-driven: ``context.engine`` in config.yaml.
|
||||
Default is ``"compressor"`` (the built-in). Only one engine is active.
|
||||
|
||||
The engine is responsible for:
|
||||
- Deciding when compaction should fire
|
||||
- Performing compaction (summarization, DAG construction, etc.)
|
||||
- Optionally exposing tools the agent can call (e.g. lcm_grep)
|
||||
- Tracking token usage from API responses
|
||||
|
||||
Lifecycle:
|
||||
1. Engine is instantiated and registered (plugin register() or default)
|
||||
2. on_session_start() called when a conversation begins
|
||||
3. update_from_response() called after each API response with usage data
|
||||
4. should_compress() checked after each turn
|
||||
5. compress() called when should_compress() returns True
|
||||
6. on_session_end() called at real session boundaries (CLI exit, /reset,
|
||||
gateway session expiry) — NOT per-turn
|
||||
Lifecycle: on_session_start() -> per API response update_from_response() ->
|
||||
per turn should_compress() / compress() -> on_session_end() at real session
|
||||
boundaries only (CLI exit, /reset, gateway expiry), never per-turn.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -60,12 +46,10 @@ def automatic_compaction_status_message(
|
||||
default_message: str,
|
||||
**context: Any,
|
||||
) -> str | None:
|
||||
"""Resolve host-visible status for an automatic compaction event.
|
||||
"""Host-visible status for an automatic compaction event; ``None`` = emit nothing.
|
||||
|
||||
Engines can suppress routine automatic status with
|
||||
``emit_automatic_compaction_status = False`` or customize it by defining
|
||||
``get_automatic_compaction_status_message(...)``. Empty strings and
|
||||
``None`` mean "do not emit a lifecycle status".
|
||||
Engines suppress via ``emit_automatic_compaction_status = False`` or
|
||||
customize via ``get_automatic_compaction_status_message(...)``.
|
||||
"""
|
||||
if not getattr(engine, "emit_automatic_compaction_status", True):
|
||||
return None
|
||||
@@ -96,9 +80,7 @@ class ContextEngine(ABC):
|
||||
def name(self) -> str:
|
||||
"""Short identifier (e.g. 'compressor', 'lcm')."""
|
||||
|
||||
# -- Token state (read by run_agent.py for display/logging) ------------
|
||||
#
|
||||
# Engines MUST maintain these. run_agent.py reads them directly.
|
||||
# -- Token state: engines MUST maintain these; run_agent.py reads them directly.
|
||||
|
||||
last_prompt_tokens: int = 0
|
||||
last_completion_tokens: int = 0
|
||||
@@ -107,39 +89,28 @@ class ContextEngine(ABC):
|
||||
context_length: int = 0
|
||||
compression_count: int = 0
|
||||
|
||||
# -- Compaction parameters (read by run_agent.py for preflight) --------
|
||||
#
|
||||
# These control the preflight compression check. Subclasses may
|
||||
# override via __init__ or property; defaults are sensible for most
|
||||
# engines.
|
||||
#
|
||||
# protect_first_n semantics (since PR #13754): count of non-system head
|
||||
# messages always preserved verbatim, IN ADDITION to the system prompt
|
||||
# which is always implicitly protected. Default 3 keeps the
|
||||
# historical "system + first 3 non-system messages" head shape.
|
||||
# -- Compaction parameters (read by run_agent.py for preflight). protect_first_n
|
||||
# counts non-system head messages kept verbatim IN ADDITION to the always-
|
||||
# protected system prompt (3 keeps the historical head shape).
|
||||
|
||||
threshold_percent: float = 0.75
|
||||
protect_first_n: int = 3
|
||||
protect_last_n: int = 6
|
||||
|
||||
# User-visible lifecycle status for automatic host-triggered compaction.
|
||||
# Alternative engines that treat compaction as routine background
|
||||
# maintenance can set this false to keep successful automatic passes silent;
|
||||
# warnings, errors, and explicit manual commands should still surface.
|
||||
# False keeps successful automatic compaction passes silent (routine
|
||||
# background maintenance); warnings, errors and manual /compress still surface.
|
||||
emit_automatic_compaction_status: bool = True
|
||||
|
||||
# -- Core interface ----------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
def update_from_response(self, usage: Dict[str, Any]) -> None:
|
||||
"""Update tracked token usage from an API response.
|
||||
"""Update tracked token usage after every LLM call.
|
||||
|
||||
Called after every LLM call with a normalized usage dict. The legacy
|
||||
keys ``prompt_tokens``, ``completion_tokens``, and ``total_tokens``
|
||||
are always present. Newer hosts also include canonical buckets:
|
||||
``input_tokens``, ``output_tokens``, ``cache_read_tokens``,
|
||||
``cache_write_tokens``, and ``reasoning_tokens``. Engines should
|
||||
treat those fields as optional for compatibility with older hosts.
|
||||
``prompt_tokens``/``completion_tokens``/``total_tokens`` are always
|
||||
present; the canonical buckets (``input_tokens``, ``output_tokens``,
|
||||
``cache_read_tokens``, ``cache_write_tokens``, ``reasoning_tokens``)
|
||||
are optional on older hosts.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -149,13 +120,9 @@ class ContextEngine(ABC):
|
||||
def should_compress_info(self, prompt_tokens: int = None) -> "tuple[bool, str | None]":
|
||||
"""Return ``(should_compress, reason)``.
|
||||
|
||||
The base implementation is backward-compatible: engines that only
|
||||
implement ``should_compress`` get ``(should_compress(prompt_tokens),
|
||||
None)``. Concrete engines with richer block reasons (e.g. a
|
||||
summary-LLM cooldown or an anti-thrashing guard) override this to
|
||||
surface a human-readable reason so callers can warn the user instead
|
||||
of silently skipping compression. Added for the silent-overflow
|
||||
warning fix (#62625) so plugin engines don't raise AttributeError.
|
||||
Engines with block reasons (summary-LLM cooldown, anti-thrashing guard)
|
||||
override this so callers can warn the user instead of silently skipping
|
||||
compression. The default keeps plugin engines from raising AttributeError.
|
||||
"""
|
||||
return self.should_compress(prompt_tokens), None
|
||||
|
||||
@@ -168,25 +135,14 @@ class ContextEngine(ABC):
|
||||
force: bool = False,
|
||||
memory_context: str = "",
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Compact the message list and return the new message list.
|
||||
"""Compact ``messages`` and return a valid OpenAI-format message list
|
||||
that fits the context budget (summarize, build a DAG, anything).
|
||||
|
||||
This is the main entry point. The engine receives the full message
|
||||
list and returns a (possibly shorter) list that fits within the
|
||||
context budget. The implementation is free to summarize, build a
|
||||
DAG, or do anything else — as long as the returned list is a valid
|
||||
OpenAI-format message sequence.
|
||||
|
||||
Args:
|
||||
focus_topic: Optional topic string from manual ``/compress <focus>``.
|
||||
Engines that support guided compression should prioritise
|
||||
preserving information related to this topic. Engines that
|
||||
don't support it may simply ignore this argument.
|
||||
force: Whether a user-requested compression should bypass an
|
||||
engine-owned cooldown. Engines without cooldowns may ignore it.
|
||||
memory_context: Text returned by memory providers immediately before
|
||||
compaction. Summarizing engines should include non-empty text in
|
||||
their handoff prompt. Older engines may omit this parameter; the
|
||||
host filters unsupported optional arguments by signature.
|
||||
``focus_topic`` comes from manual ``/compress <focus>`` (prioritise that
|
||||
topic); ``force`` asks to bypass an engine-owned cooldown;
|
||||
``memory_context`` is provider text to include in the handoff prompt.
|
||||
Older engines may omit optional parameters — the host filters them by
|
||||
signature.
|
||||
"""
|
||||
|
||||
# -- Optional: proactive tool-result prune -----------------------------
|
||||
@@ -199,14 +155,9 @@ class ContextEngine(ABC):
|
||||
"""Deterministically trim old tool-result payloads without an LLM call.
|
||||
|
||||
Runs on a low, cost-oriented trigger independent of ``should_compress``
|
||||
so large-window engines can reclaim re-sent tool output long before full
|
||||
compaction would fire. Returns ``(messages, n_pruned)``.
|
||||
|
||||
Default is a safe no-op: the list is returned unchanged with ``0``
|
||||
pruned. Engines that don't implement a cheap prune — and any engine that
|
||||
predates this hook — inherit this default, so the agent loop's
|
||||
post-tool-call prune path never raises ``AttributeError`` on them. The
|
||||
built-in ContextCompressor overrides this with the real implementation.
|
||||
so large-window engines reclaim re-sent tool output long before full
|
||||
compaction. Returns ``(messages, n_pruned)``; default is a no-op so
|
||||
engines predating this hook never raise in the post-tool-call prune path.
|
||||
"""
|
||||
return messages, 0
|
||||
|
||||
@@ -220,61 +171,27 @@ class ContextEngine(ABC):
|
||||
incoming_message: Dict[str, Any] = None,
|
||||
budget_tokens: int = 0,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Optionally choose/replace the context for THIS request, pre-generation.
|
||||
"""Optionally *select* (replace) the context for THIS request, pre-generation.
|
||||
|
||||
Called every turn after the request message list is assembled and
|
||||
before it is dispatched to the provider — independent of
|
||||
``should_compress()``. This lets an engine *select* which context
|
||||
enters the prompt (retrieval, topic routing, role/branch switching)
|
||||
rather than *shrink* context that is already there. The two verbs are
|
||||
orthogonal:
|
||||
Runs every provider request (so also on retries), independent of
|
||||
``should_compress()``: ``compress()`` shrinks context that is too long,
|
||||
``select_context()`` swaps in a different context (retrieval, topic
|
||||
routing, branch switching) without abusing ``compress()`` as a per-turn
|
||||
callback. Return ``None`` to leave the request unchanged.
|
||||
|
||||
- ``compress()`` : context is too long -> make it shorter.
|
||||
- ``select_context()``: this turn belongs to a different context
|
||||
-> use that one instead.
|
||||
The returned list is request-only — it MUST NOT be treated as persisted
|
||||
transcript state; the session DB history is untouched. Unlike the
|
||||
``pre_llm_call`` hook it may replace the list. The host runs it before
|
||||
prompt cache-control and before every request sanitizer, so a malformed
|
||||
replacement never reaches the provider and the default no-op keeps the
|
||||
request byte-identical (prompt-cache stability preserved). An engine that
|
||||
replaces the list changes its own cache prefix; breakpoints are
|
||||
re-derived on the selected list.
|
||||
|
||||
Without this hook, engines that need per-turn access to the message
|
||||
list have to force ``should_compress()`` to return ``True`` so that
|
||||
``compress()`` is invoked every turn purely as a callback — which
|
||||
conflates selection with compression and degrades behaviour when the
|
||||
engine's backend is unavailable. ``select_context()`` removes the need
|
||||
for that workaround.
|
||||
|
||||
The returned list is request-only: it replaces the messages sent to
|
||||
the provider for this single call and MUST NOT be treated as persisted
|
||||
transcript state. The conversation history in the session DB is left
|
||||
untouched, so nothing leaks across turns. Return ``None`` to leave the
|
||||
request unchanged.
|
||||
|
||||
Unlike the ``pre_llm_call`` plugin hook (which appends to the user
|
||||
message and intentionally never rewrites the list, to preserve the
|
||||
cache prefix), ``select_context()`` may *replace* the message list.
|
||||
|
||||
Ordering / cache contract: the host runs this hook **before** prompt
|
||||
cache-control and **before** every request sanitizer (orphaned-tool
|
||||
cleanup, thinking-only/role normalization, whitespace/JSON
|
||||
normalization). So (a) whatever the hook returns still passes through
|
||||
the same validation as any request — a malformed replacement cannot
|
||||
reach the provider — and (b) prompt-cache stability (an AGENTS.md
|
||||
invariant) is preserved: the default no-op leaves the request
|
||||
byte-identical, so cache behaviour is unchanged for the built-in
|
||||
compressor and any non-implementing engine. An engine that *does*
|
||||
replace the list changes its own cache prefix by definition; that is
|
||||
the engine's concern, and cache-control breakpoints are re-derived on
|
||||
the selected list. The hook is evaluated per provider request (so it
|
||||
re-runs on retries within a turn), consistent with "select the context
|
||||
for THIS request".
|
||||
|
||||
Args:
|
||||
request_messages: The assembled request message list (system
|
||||
prompt + history + any ephemeral prefill), in OpenAI format.
|
||||
conversation_messages: The unmodified persisted conversation
|
||||
history, for reference only (do not mutate).
|
||||
incoming_message: The current turn's user message, if available.
|
||||
budget_tokens: The active model's context length, or 0 if unknown.
|
||||
|
||||
Default returns ``None`` (no-op) — zero impact on the built-in
|
||||
compressor or any existing engine.
|
||||
``request_messages`` is the assembled request (system prompt + history +
|
||||
ephemeral prefill); ``conversation_messages`` is the persisted history
|
||||
for reference only (do not mutate); ``budget_tokens`` is the model's
|
||||
context length or 0 if unknown.
|
||||
"""
|
||||
return None
|
||||
|
||||
@@ -284,66 +201,29 @@ class ContextEngine(ABC):
|
||||
usage: Dict[str, Any] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""Observe a finished user turn (post-turn ingestion / observation).
|
||||
"""Observe a finished turn (complement of ``select_context()``) so the
|
||||
engine can ingest/index/update routing state for the next request.
|
||||
|
||||
Called from the standard turn-finalization path once the assistant/tool
|
||||
loop completes, with the finalized in-memory transcript snapshot. This
|
||||
is the complement to ``select_context()``: selection happens *before*
|
||||
the request, while observation happens *after* the turn. It lets an
|
||||
engine ingest, index, summarize, or update routing / topic / session
|
||||
state from what actually happened — so the next ``select_context()``
|
||||
can act on it.
|
||||
|
||||
Coverage: this fires from the normal finalization seam. Some abnormal
|
||||
early-return paths in the loop (e.g. a content-policy block or a
|
||||
provider terminal failure) persist and return without routing through
|
||||
finalization, and therefore do not currently emit this hook. Treat it
|
||||
as a best-effort post-turn observation for completed turns, not a
|
||||
guaranteed callback for every possible early exit; unifying all
|
||||
terminal paths behind one finalization seam is a separate follow-up.
|
||||
|
||||
Together the two hooks remove the need to abuse ``should_compress()`` /
|
||||
``compress()`` as a generic per-turn callback just to observe history,
|
||||
and they cover the case where a turn finishes and there may be no next
|
||||
request from which to infer the previous turn.
|
||||
|
||||
``messages`` is a shallow copy and should be treated as read-only:
|
||||
return values are ignored and this hook must not rely on transcript
|
||||
mutation for persistence. ``kwargs`` may include ``turn_id``,
|
||||
``task_id``, ``api_call_count``, ``interrupted``, ``failed``, and
|
||||
``turn_exit_reason``.
|
||||
|
||||
``usage`` carries the completed turn's canonical token usage (the same
|
||||
dict shape passed to ``update_from_response`` — ``prompt_tokens`` /
|
||||
``completion_tokens`` / ``total_tokens`` plus the canonical
|
||||
``input_tokens`` / ``output_tokens`` / ``cache_read_tokens`` /
|
||||
``cache_write_tokens`` / ``reasoning_tokens`` buckets) so an engine can
|
||||
weigh how large/expensive the selected context actually was when
|
||||
deciding the next ``select_context()``. It is ``None`` on finalized
|
||||
turns that never reached a provider response (e.g. interrupt); engines
|
||||
must treat it as optional.
|
||||
|
||||
Default is a no-op.
|
||||
Fires from the normal finalization seam only; some abnormal early
|
||||
returns (content-policy block, provider terminal failure) do not emit
|
||||
it — treat it as best-effort, not guaranteed. ``messages`` is a
|
||||
read-only shallow copy (return value ignored; never rely on transcript
|
||||
mutation). ``usage`` has the ``update_from_response`` dict shape and is
|
||||
``None`` when the turn never reached a provider response (interrupt).
|
||||
``kwargs`` may include ``turn_id``, ``task_id``, ``api_call_count``,
|
||||
``interrupted``, ``failed``, ``turn_exit_reason``.
|
||||
"""
|
||||
return None
|
||||
|
||||
# -- Optional: pre-flight check ----------------------------------------
|
||||
|
||||
def should_compress_preflight(self, messages: List[Dict[str, Any]]) -> bool:
|
||||
"""Quick rough check before the API call (no real token count yet).
|
||||
|
||||
Default returns False (skip pre-flight). Override if your engine
|
||||
can do a cheap estimate.
|
||||
"""
|
||||
"""Cheap rough check before the API call (no real token count yet); default skips."""
|
||||
return False
|
||||
|
||||
def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool:
|
||||
"""Return True when preflight should trust recent real usage instead.
|
||||
|
||||
Built-in compression uses this to avoid re-compacting from known-noisy
|
||||
rough estimates after a compressed request has already fit. Third-party
|
||||
engines can ignore it safely.
|
||||
"""
|
||||
"""True when preflight should trust recent real usage over the noisy rough
|
||||
estimate (avoids re-compacting after a compressed request already fit)."""
|
||||
return False
|
||||
|
||||
def get_automatic_compaction_status_message(
|
||||
@@ -353,15 +233,11 @@ class ContextEngine(ABC):
|
||||
default_message: str,
|
||||
**context: Any,
|
||||
) -> str | None:
|
||||
"""Return user-visible status for automatic host-triggered compaction.
|
||||
"""User-visible status for automatic compaction, or ``None`` to suppress it.
|
||||
|
||||
Return ``None`` to suppress successful automatic lifecycle status for
|
||||
this compaction event. ``phase`` identifies the host call site (for
|
||||
example ``"preflight"`` or ``"compress"``). ``context`` contains
|
||||
best-effort fields such as ``approx_tokens`` and ``threshold_tokens``.
|
||||
|
||||
This hook does not control warning/error messages or explicit manual
|
||||
commands such as ``/compress``.
|
||||
``phase`` is the host call site (``"preflight"`` / ``"compress"``);
|
||||
``context`` carries best-effort ``approx_tokens`` / ``threshold_tokens``.
|
||||
Warnings, errors and manual ``/compress`` are not governed by this hook.
|
||||
"""
|
||||
if not self.emit_automatic_compaction_status:
|
||||
return None
|
||||
@@ -370,39 +246,20 @@ class ContextEngine(ABC):
|
||||
# -- Optional: manual /compress preflight ------------------------------
|
||||
|
||||
def has_content_to_compress(self, messages: List[Dict[str, Any]]) -> bool:
|
||||
"""Quick check: is there anything in ``messages`` that can be compacted?
|
||||
|
||||
Used by the gateway ``/compress`` command as a preflight guard —
|
||||
returning False lets the gateway report "nothing to compress yet"
|
||||
without making an LLM call.
|
||||
|
||||
Default returns True (always attempt). Engines with a cheap way
|
||||
to introspect their own head/tail boundaries should override this
|
||||
to return False when the transcript is still entirely protected.
|
||||
"""
|
||||
"""Preflight guard for gateway ``/compress``: False reports "nothing to
|
||||
compress yet" without an LLM call (e.g. transcript entirely protected)."""
|
||||
return True
|
||||
|
||||
# -- Optional: session lifecycle ---------------------------------------
|
||||
|
||||
def on_session_start(self, session_id: str, **kwargs) -> None:
|
||||
"""Called when a new conversation session begins.
|
||||
|
||||
Use this to load persisted state (DAG, store) for the session.
|
||||
kwargs may include hermes_home, platform, model, etc.
|
||||
"""
|
||||
"""Session begins: load persisted state. kwargs may include hermes_home, platform, model."""
|
||||
|
||||
def on_session_end(self, session_id: str, messages: List[Dict[str, Any]]) -> None:
|
||||
"""Called at real session boundaries (CLI exit, /reset, gateway expiry).
|
||||
|
||||
Use this to flush state, close DB connections, etc.
|
||||
NOT called per-turn — only when the session truly ends.
|
||||
"""
|
||||
"""Real session boundary (CLI exit, /reset, gateway expiry) — never per-turn."""
|
||||
|
||||
def on_session_reset(self) -> None:
|
||||
"""Called on /new or /reset. Reset per-session state.
|
||||
|
||||
Default resets compression_count and token tracking.
|
||||
"""
|
||||
"""/new or /reset: reset per-session state (default: counters and token tracking)."""
|
||||
self.last_prompt_tokens = 0
|
||||
self.last_completion_tokens = 0
|
||||
self.last_total_tokens = 0
|
||||
@@ -411,36 +268,21 @@ class ContextEngine(ABC):
|
||||
# -- Optional: tools ---------------------------------------------------
|
||||
|
||||
def get_tool_schemas(self) -> List[Dict[str, Any]]:
|
||||
"""Return tool schemas this engine provides to the agent.
|
||||
|
||||
Default returns empty list (no tools). LCM would return schemas
|
||||
for lcm_grep, lcm_describe, lcm_expand here.
|
||||
"""
|
||||
"""Tool schemas this engine exposes to the agent (default: none)."""
|
||||
return []
|
||||
|
||||
def handle_tool_call(self, name: str, args: Dict[str, Any], **kwargs) -> str:
|
||||
"""Handle a tool call from the agent.
|
||||
|
||||
Only called for tool names returned by get_tool_schemas().
|
||||
Must return a JSON string.
|
||||
|
||||
kwargs may include:
|
||||
messages: the current in-memory message list (for live ingestion)
|
||||
"""
|
||||
"""Handle a call to one of this engine's tools; must return a JSON string.
|
||||
kwargs may include ``messages`` (live in-memory list)."""
|
||||
import json
|
||||
return json.dumps({"error": f"Unknown context engine tool: {name}"})
|
||||
|
||||
# -- Optional: status / display ----------------------------------------
|
||||
|
||||
def get_status(self) -> Dict[str, Any]:
|
||||
"""Return status dict for display/logging.
|
||||
|
||||
Default returns the standard fields run_agent.py expects.
|
||||
"""
|
||||
# Clamp the -1 "compression just ran, awaiting real usage" sentinel
|
||||
# (set by conversation_compression) to 0 so status readers don't see a
|
||||
# raw -1 or a negative usage_percent on the transitional turn. Mirrors
|
||||
# the CLI/gateway status-bar paths (cli.py, tui_gateway/server.py).
|
||||
"""Status dict with the standard fields run_agent.py expects."""
|
||||
# Clamp the -1 "compression just ran, awaiting real usage" sentinel to 0
|
||||
# so no reader sees a negative usage_percent on the transitional turn.
|
||||
last_prompt = self.last_prompt_tokens if self.last_prompt_tokens > 0 else 0
|
||||
return {
|
||||
"last_prompt_tokens": last_prompt,
|
||||
@@ -464,22 +306,13 @@ class ContextEngine(ABC):
|
||||
provider: str = "",
|
||||
api_mode: str = "",
|
||||
) -> None:
|
||||
"""Called when the user switches models or on fallback activation.
|
||||
|
||||
Default updates context_length and recalculates threshold_tokens
|
||||
from threshold_percent. Override if your engine needs more
|
||||
(e.g. recalculate DAG budgets, switch summary models).
|
||||
"""
|
||||
"""Model switch / fallback: recompute threshold_tokens (override for more)."""
|
||||
self.context_length = context_length
|
||||
# Apply per-model threshold overrides if set (longest substring match).
|
||||
# Falls back to _config_threshold_percent (the raw config value) when
|
||||
# no override matches. Plugin engines that override update_model() can
|
||||
# call resolve_model_threshold() for the same logic.
|
||||
# Per-model threshold override (longest substring match), else the raw
|
||||
# config percent. Snapshot that percent ONCE so repeated switches fall
|
||||
# back to the configured value, not the previous model's override.
|
||||
from agent.context_compressor import resolve_model_threshold
|
||||
if not hasattr(self, "_config_threshold_percent"):
|
||||
# Snapshot the pre-override percent ONCE so repeated model
|
||||
# switches fall back to the engine's configured value, not the
|
||||
# previous model's override.
|
||||
self._config_threshold_percent = self.threshold_percent
|
||||
self._base_threshold_percent = resolve_model_threshold(
|
||||
model, getattr(self, "model_thresholds", {}),
|
||||
|
||||
+47
-236
@@ -2,31 +2,18 @@
|
||||
Image Generation Provider ABC
|
||||
=============================
|
||||
|
||||
Defines the pluggable-backend interface for image generation. Providers register
|
||||
instances via ``PluginContext.register_image_gen_provider()``; the active one
|
||||
(selected via ``image_gen.provider`` in ``config.yaml``) services every
|
||||
``image_generate`` tool call.
|
||||
Pluggable-backend interface for image generation. Providers register via
|
||||
``PluginContext.register_image_gen_provider()``; the one selected by
|
||||
``image_gen.provider`` services every ``image_generate`` call. Providers live in
|
||||
``<repo>/plugins/image_gen/<name>/`` (built-in) or
|
||||
``~/.hermes/plugins/image_gen/<name>/`` (user, opt-in).
|
||||
|
||||
Providers live in ``<repo>/plugins/image_gen/<name>/`` (built-in, auto-loaded
|
||||
as ``kind: backend``) or ``~/.hermes/plugins/image_gen/<name>/`` (user, opt-in
|
||||
via ``plugins.enabled``).
|
||||
One tool covers text-to-image and image-to-image/editing: the presence of
|
||||
``image_url`` (and/or ``reference_image_urls``) routes to the provider's edit
|
||||
endpoint, otherwise text-to-image. Users pick one model; the provider picks the
|
||||
endpoint. Mirrors ``agent/video_gen_provider.py`` so the two stay learnable.
|
||||
|
||||
Unified surface
|
||||
---------------
|
||||
One tool — ``image_generate`` — covers **text-to-image** and
|
||||
**image-to-image / image editing**. The router is the presence of
|
||||
``image_url`` (and/or ``reference_image_urls``): if any source image is
|
||||
provided, the provider routes to its image-to-image / edit endpoint; if
|
||||
omitted, the provider routes to text-to-image. Users pick one **model**
|
||||
(e.g. nano-banana-pro, gpt-image-2, grok-imagine-image); the provider
|
||||
handles which underlying endpoint to hit. This mirrors the ``video_gen``
|
||||
provider design (``agent/video_gen_provider.py``) so the two surfaces
|
||||
stay learnable together.
|
||||
|
||||
Response shape
|
||||
--------------
|
||||
All providers return a dict that :func:`success_response` / :func:`error_response`
|
||||
produce. The tool wrapper JSON-serializes it. Keys:
|
||||
Response shape (built by :func:`success_response` / :func:`error_response`)::
|
||||
|
||||
success bool
|
||||
image str | None URL or absolute file path
|
||||
@@ -42,13 +29,13 @@ produce. The tool wrapper JSON-serializes it. Keys:
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
import base64
|
||||
import datetime
|
||||
import logging
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from agent import provider_media
|
||||
from agent.provider_base import CatalogProviderBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -56,106 +43,21 @@ VALID_ASPECT_RATIOS: Tuple[str, ...] = ("landscape", "square", "portrait")
|
||||
DEFAULT_ASPECT_RATIO = "landscape"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ImageGenProvider(abc.ABC):
|
||||
class ImageGenProvider(CatalogProviderBase):
|
||||
"""Abstract base class for an image generation backend.
|
||||
|
||||
Subclasses must implement :meth:`generate`. Everything else has sane
|
||||
defaults — override only what your provider needs.
|
||||
Subclasses must implement :attr:`name` and :meth:`generate`; everything else
|
||||
has defaults. ``list_models`` entries may add ``speed`` / ``strengths`` /
|
||||
``price`` for the picker.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in ``image_gen.provider`` config.
|
||||
|
||||
Lowercase, no spaces. Examples: ``fal``, ``openai``, ``replicate``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name.title()``."""
|
||||
return self.name.title()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
|
||||
Typically checks for a required API key. Default: True
|
||||
(providers with no external dependencies are always available).
|
||||
"""
|
||||
return True
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
"""Return catalog entries for ``hermes tools`` model picker.
|
||||
|
||||
Each entry::
|
||||
|
||||
{
|
||||
"id": "gpt-image-1.5", # required
|
||||
"display": "GPT Image 1.5", # optional; defaults to id
|
||||
"speed": "~10s", # optional
|
||||
"strengths": "...", # optional
|
||||
"price": "$...", # optional
|
||||
}
|
||||
|
||||
Default: empty list (provider has no user-selectable models).
|
||||
"""
|
||||
return []
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by ``tools_config.py`` to inject this provider as a row in
|
||||
the Image Generation provider list. Shape::
|
||||
|
||||
{
|
||||
"name": "OpenAI", # picker label
|
||||
"badge": "paid", # optional short tag
|
||||
"tag": "One-line description...", # optional subtitle
|
||||
"env_vars": [ # keys to prompt for
|
||||
{"key": "OPENAI_API_KEY",
|
||||
"prompt": "OpenAI API key",
|
||||
"url": "https://platform.openai.com/api-keys"},
|
||||
],
|
||||
}
|
||||
|
||||
Default: minimal entry derived from ``display_name``. Override to
|
||||
expose API key prompts and custom badges.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
"""Return the default model id, or None if not applicable."""
|
||||
models = self.list_models()
|
||||
if models:
|
||||
return models[0].get("id")
|
||||
return None
|
||||
|
||||
def capabilities(self) -> Dict[str, Any]:
|
||||
"""Return what this provider supports.
|
||||
"""What this provider supports: ``modalities`` (``"text"`` and/or
|
||||
``"image"``) and ``max_reference_images``.
|
||||
|
||||
Returned dict (all keys optional)::
|
||||
|
||||
{
|
||||
"modalities": ["text", "image"], # which inputs the backend accepts
|
||||
"max_reference_images": 9, # cap for reference_image_urls
|
||||
}
|
||||
|
||||
``modalities`` declares whether the active backend/model supports
|
||||
text-to-image (``"text"``), image-to-image / editing (``"image"``),
|
||||
or both. The tool layer surfaces this in the dynamic schema so the
|
||||
model knows when ``image_url`` is honored. Used by ``hermes tools``
|
||||
for the picker too. Default: text-only (backward compatible — a
|
||||
provider that doesn't override this advertises text-to-image only).
|
||||
The tool layer surfaces this in the dynamic schema so the model knows
|
||||
when ``image_url`` is honored. Default is text-only so a provider that
|
||||
doesn't override advertises only text-to-image (backward compatible).
|
||||
"""
|
||||
return {
|
||||
"modalities": ["text"],
|
||||
@@ -172,25 +74,15 @@ class ImageGenProvider(abc.ABC):
|
||||
reference_image_urls: Optional[List[str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Dict[str, Any]:
|
||||
"""Generate an image from a text prompt, or edit/transform a source image.
|
||||
"""Generate an image, or edit/transform a source image.
|
||||
|
||||
Routing: if ``image_url`` (or any ``reference_image_urls``) is
|
||||
provided, the provider should route to its image-to-image / edit
|
||||
endpoint; otherwise text-to-image. ``image_url`` is the primary
|
||||
source image to edit; ``reference_image_urls`` are additional
|
||||
style/composition references (provider clamps to its declared
|
||||
``max_reference_images``).
|
||||
|
||||
Implementations should return the dict from :func:`success_response`
|
||||
or :func:`error_response`. ``kwargs`` may contain forward-compat
|
||||
parameters future versions of the schema will expose —
|
||||
implementations MUST ignore unknown keys (no TypeError).
|
||||
|
||||
Known optional kwarg: ``upscale`` (bool) — when true, the caller
|
||||
requests a post-generation high-resolution pass through the
|
||||
backend's upscaler/enhancer. Providers without an upscaler simply
|
||||
ignore it; providers that honor it should report ``upscaled: True``
|
||||
in the response ``extra``.
|
||||
``image_url`` is the primary source to edit; ``reference_image_urls``
|
||||
are extra style/composition references (clamp to ``max_reference_images``).
|
||||
Any source image routes to the edit endpoint, otherwise text-to-image.
|
||||
Return :func:`success_response` / :func:`error_response`. Unknown
|
||||
``kwargs`` MUST be ignored (forward compat). Known optional kwarg:
|
||||
``upscale`` (bool) — a post-generation high-res pass; providers that
|
||||
honor it report ``upscaled: True`` in ``extra``.
|
||||
"""
|
||||
|
||||
|
||||
@@ -200,11 +92,8 @@ class ImageGenProvider(abc.ABC):
|
||||
|
||||
|
||||
def resolve_aspect_ratio(value: Optional[str]) -> str:
|
||||
"""Clamp an aspect_ratio value to the valid set, defaulting to landscape.
|
||||
|
||||
Invalid values are coerced rather than rejected so the tool surface is
|
||||
forgiving of agent mistakes.
|
||||
"""
|
||||
"""Clamp to :data:`VALID_ASPECT_RATIOS`; invalid values coerce to landscape so
|
||||
the tool surface forgives agent mistakes instead of rejecting them."""
|
||||
if not isinstance(value, str):
|
||||
return DEFAULT_ASPECT_RATIO
|
||||
v = value.strip().lower()
|
||||
@@ -214,12 +103,8 @@ def resolve_aspect_ratio(value: Optional[str]) -> str:
|
||||
|
||||
|
||||
def normalize_reference_images(value: Any) -> Optional[List[str]]:
|
||||
"""Coerce a reference-image argument into a clean list of URL/path strings.
|
||||
|
||||
Accepts a single string or a list; strips blanks and whitespace. Returns
|
||||
``None`` when nothing usable remains so providers can treat "no refs" as a
|
||||
single sentinel.
|
||||
"""
|
||||
"""Coerce a str or list into a clean list of non-blank strings; ``None`` when
|
||||
nothing usable remains so providers treat "no refs" as one sentinel."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
@@ -235,11 +120,7 @@ def normalize_reference_images(value: Any) -> Optional[List[str]]:
|
||||
|
||||
def _images_cache_dir() -> Path:
|
||||
"""Return ``$HERMES_HOME/cache/images/``, creating parents as needed."""
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
path = get_hermes_home() / "cache" / "images"
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
return path
|
||||
return provider_media.cache_dir("images")
|
||||
|
||||
|
||||
def save_b64_image(
|
||||
@@ -248,24 +129,10 @@ def save_b64_image(
|
||||
prefix: str = "image",
|
||||
extension: str = "png",
|
||||
) -> Path:
|
||||
"""Decode base64 image data and write it under ``$HERMES_HOME/cache/images/``.
|
||||
|
||||
Returns the absolute :class:`Path` to the saved file.
|
||||
|
||||
Filename format: ``<prefix>_<YYYYMMDD_HHMMSS>_<short-uuid>.<ext>``.
|
||||
"""
|
||||
raw = base64.b64decode(b64_data)
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
path = _images_cache_dir() / f"{prefix}_{ts}_{short}.{extension}"
|
||||
path.write_bytes(raw)
|
||||
return path
|
||||
"""Decode base64 image data into ``$HERMES_HOME/cache/images/``; return the path."""
|
||||
return provider_media.save_b64("images", b64_data, prefix=prefix, extension=extension)
|
||||
|
||||
|
||||
# Extension inference for save_url_image — keep small and explicit. We don't
|
||||
# want to import mimetypes for a handful of formats every image_gen provider
|
||||
# actually returns, and we never want to inherit a content-type that points
|
||||
# at HTML or JSON when the API gives us a degenerate response.
|
||||
_URL_IMAGE_CONTENT_TYPES = {
|
||||
"image/png": "png",
|
||||
"image/jpeg": "jpg",
|
||||
@@ -282,66 +149,17 @@ def save_url_image(
|
||||
timeout: float = 60.0,
|
||||
max_bytes: int = 25 * 1024 * 1024,
|
||||
) -> Path:
|
||||
"""Download an image URL and write it under ``$HERMES_HOME/cache/images/``.
|
||||
"""Download an (often ephemeral) image URL into ``$HERMES_HOME/cache/images/``.
|
||||
|
||||
Used by providers (xAI, fallback OpenAI) whose API returns an *ephemeral*
|
||||
URL instead of inline base64 — those URLs frequently expire before a
|
||||
downstream consumer (Telegram ``send_photo``, browser fetch) can resolve
|
||||
them, so we materialise the bytes locally at tool-completion time.
|
||||
Mirrors :func:`save_b64_image`'s shape so providers can swap in one line.
|
||||
|
||||
Returns the absolute :class:`Path` to the saved file. Raises on any
|
||||
network / HTTP / oversize / non-image-content-type error so callers can
|
||||
fall back to returning the bare URL with a clear error message.
|
||||
Raises on network / HTTP / oversize / empty errors so callers can fall back
|
||||
to returning the bare URL with a clear message. See :mod:`agent.provider_media`.
|
||||
"""
|
||||
import requests
|
||||
|
||||
response = requests.get(url, timeout=timeout, stream=True)
|
||||
response.raise_for_status()
|
||||
|
||||
# Infer extension from the response content-type, falling back to the
|
||||
# URL suffix when xAI / OpenAI omit a precise type (some CDNs return
|
||||
# ``application/octet-stream``). Defaults to ``png``.
|
||||
content_type = (response.headers.get("Content-Type") or "").split(";", 1)[0].strip().lower()
|
||||
extension = _URL_IMAGE_CONTENT_TYPES.get(content_type)
|
||||
if extension is None:
|
||||
url_path = url.split("?", 1)[0].lower()
|
||||
for ext in ("png", "jpg", "jpeg", "webp", "gif"):
|
||||
if url_path.endswith(f".{ext}"):
|
||||
extension = "jpg" if ext == "jpeg" else ext
|
||||
break
|
||||
if extension is None:
|
||||
extension = "png"
|
||||
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
path = _images_cache_dir() / f"{prefix}_{ts}_{short}.{extension}"
|
||||
|
||||
bytes_written = 0
|
||||
with path.open("wb") as fh:
|
||||
for chunk in response.iter_content(chunk_size=64 * 1024):
|
||||
if not chunk:
|
||||
continue
|
||||
bytes_written += len(chunk)
|
||||
if bytes_written > max_bytes:
|
||||
fh.close()
|
||||
try:
|
||||
path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
raise ValueError(
|
||||
f"Image at {url} exceeds {max_bytes // (1024 * 1024)}MB cap; refusing to cache."
|
||||
)
|
||||
fh.write(chunk)
|
||||
|
||||
if bytes_written == 0:
|
||||
try:
|
||||
path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
raise ValueError(f"Image at {url} returned 0 bytes; refusing to cache.")
|
||||
|
||||
return path
|
||||
return provider_media.save_url(
|
||||
"images", url, prefix=prefix, timeout=timeout, max_bytes=max_bytes,
|
||||
chunk_size=64 * 1024, content_types=_URL_IMAGE_CONTENT_TYPES,
|
||||
url_extensions=("png", "jpg", "jpeg", "webp", "gif"), default_extension="png",
|
||||
label="Image", empty_error="Image at {url} returned 0 bytes; refusing to cache.",
|
||||
)
|
||||
|
||||
|
||||
def success_response(
|
||||
@@ -354,14 +172,7 @@ def success_response(
|
||||
modality: str = "text",
|
||||
extra: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Build a uniform success response dict.
|
||||
|
||||
``image`` may be an HTTP URL or an absolute filesystem path (for b64
|
||||
providers like OpenAI). ``modality`` is ``"text"`` (text-to-image) or
|
||||
``"image"`` (image-to-image / editing) — indicates which endpoint was
|
||||
actually hit, useful for diagnostics. Callers that need to pass through
|
||||
additional backend-specific fields can supply ``extra``.
|
||||
"""
|
||||
"""Uniform success dict; ``extra`` keys are added without overriding standard ones."""
|
||||
payload: Dict[str, Any] = {
|
||||
"success": True,
|
||||
"image": image,
|
||||
|
||||
+19
-145
@@ -11,9 +11,9 @@ Active selection
|
||||
The active provider is chosen by ``image_gen.provider`` in ``config.yaml``.
|
||||
If unset, :func:`get_active_provider` applies fallback logic:
|
||||
|
||||
1. If exactly one provider is registered, use it.
|
||||
2. Otherwise if a provider named ``fal`` is registered, use it (legacy
|
||||
default — matches pre-plugin behavior).
|
||||
1. If exactly one *available* provider is registered, use it.
|
||||
2. Otherwise if a provider named ``fal`` is registered and available, use it
|
||||
(legacy default — matches pre-plugin behavior).
|
||||
3. Otherwise return ``None`` (the tool surfaces a helpful error pointing
|
||||
the user at ``hermes tools``).
|
||||
"""
|
||||
@@ -21,151 +21,35 @@ If unset, :func:`get_active_provider` applies fallback logic:
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from agent.image_gen_provider import ImageGenProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
from agent.provider_registry import ProviderRegistry, configured_provider_name, is_available_safe
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_providers: Dict[str, ImageGenProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, ImageGenProvider]] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def register_provider(provider: ImageGenProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register an image generation provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and logs
|
||||
a debug message — this makes hot-reload scenarios (tests, dev loops)
|
||||
behave predictably.
|
||||
"""
|
||||
if not isinstance(provider, ImageGenProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects an ImageGenProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Image gen provider .name must be a non-empty string")
|
||||
name = raw_name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if existing is not None:
|
||||
logger.debug("Image gen provider '%s' re-registered (was %r)", name, type(existing).__name__)
|
||||
else:
|
||||
logger.debug("Registered image gen provider '%s' (%s)", name, type(provider).__name__)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[ImageGenProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[ImageGenProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
with _lock:
|
||||
key = name.strip()
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[ImageGenProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip())
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: ImageGenProvider,
|
||||
previous: Optional[ImageGenProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is not None and not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
_registry: ProviderRegistry[ImageGenProvider] = ProviderRegistry(
|
||||
label="Image gen", provider_cls=ImageGenProvider, logger=logger,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
def get_active_provider() -> Optional[ImageGenProvider]:
|
||||
"""Resolve the currently-active provider.
|
||||
|
||||
Reads ``image_gen.provider`` from config.yaml; falls back per the
|
||||
module docstring.
|
||||
|
||||
**Availability semantics** (mirrors :mod:`agent.web_search_registry`):
|
||||
|
||||
- When ``image_gen.provider`` is explicitly set, the configured
|
||||
provider is returned even if :meth:`ImageGenProvider.is_available`
|
||||
reports False — the dispatcher surfaces a precise "X_API_KEY is not
|
||||
set" error rather than silently switching backends.
|
||||
- When ``image_gen.provider`` is unset, the fallback path (single-
|
||||
provider shortcut and the FAL legacy preference) is filtered by
|
||||
``is_available()`` so we don't pick a provider the user has no
|
||||
credentials for.
|
||||
an explicitly configured provider is returned even if ``is_available()``
|
||||
is False, so the dispatcher surfaces a precise "X_API_KEY is not set"
|
||||
error instead of silently switching backends. Only the unconfigured
|
||||
fallback path is filtered by availability.
|
||||
"""
|
||||
configured: Optional[str] = None
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
configured = configured_provider_name("image_gen", logger)
|
||||
snapshot = _registry.merged()
|
||||
|
||||
cfg = load_config_readonly()
|
||||
section = cfg.get("image_gen") if isinstance(cfg, dict) else None
|
||||
if isinstance(section, dict):
|
||||
raw = section.get("provider")
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
configured = raw.strip()
|
||||
except Exception as exc:
|
||||
logger.debug("Could not read image_gen.provider from config: %s", exc)
|
||||
def _available(p: ImageGenProvider) -> bool:
|
||||
return is_available_safe(p, logger, "image_gen provider %s.is_available() raised %s")
|
||||
|
||||
# The managed "Nous Subscription" selection is serviced by the FAL
|
||||
# plugin through the managed fal-queue gateway (the legacy FAL pipeline
|
||||
# routes managed when the stored selection is "nous").
|
||||
if configured:
|
||||
try:
|
||||
from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER
|
||||
|
||||
if configured.lower() == NOUS_MANAGED_PROVIDER:
|
||||
configured = "fal"
|
||||
except Exception: # pragma: no cover — helpers are in-repo
|
||||
pass
|
||||
|
||||
with _lock:
|
||||
snapshot = dict(_providers)
|
||||
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
|
||||
|
||||
def _is_available_safe(p: ImageGenProvider) -> bool:
|
||||
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
|
||||
try:
|
||||
return bool(p.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("image_gen provider %s.is_available() raised %s", p.name, exc)
|
||||
return False
|
||||
|
||||
# 1. Explicit config wins — return regardless of is_available() so the
|
||||
# user gets a precise downstream error message rather than a silent
|
||||
# backend switch.
|
||||
if configured:
|
||||
provider = snapshot.get(configured)
|
||||
if provider is not None:
|
||||
@@ -175,22 +59,12 @@ def get_active_provider() -> Optional[ImageGenProvider]:
|
||||
configured,
|
||||
)
|
||||
|
||||
# 2. Fallback: single registered provider — but only if it's actually
|
||||
# available (no credentials = don't surface it as "active").
|
||||
available = [p for p in snapshot.values() if _is_available_safe(p)]
|
||||
available = [p for p in snapshot.values() if _available(p)]
|
||||
if len(available) == 1:
|
||||
return available[0]
|
||||
|
||||
# 3. Fallback: prefer legacy FAL for backward compat, when available.
|
||||
fal = snapshot.get("fal")
|
||||
if fal is not None and _is_available_safe(fal):
|
||||
if fal is not None and _available(fal):
|
||||
return fal
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
|
||||
+286
-512
File diff suppressed because it is too large
Load Diff
+24
-51
@@ -1,30 +1,17 @@
|
||||
"""Language Server Protocol (LSP) integration for Hermes Agent.
|
||||
|
||||
Hermes runs full language servers (pyright, gopls, rust-analyzer,
|
||||
typescript-language-server, etc.) as subprocesses and pipes their
|
||||
``textDocument/publishDiagnostics`` output into the post-write lint
|
||||
delta filter used by ``write_file`` and ``patch``.
|
||||
Hermes runs real language servers (pyright, gopls, rust-analyzer, ...) as
|
||||
subprocesses and pipes their ``textDocument/publishDiagnostics`` output into
|
||||
the post-write lint delta filter used by ``write_file`` and ``patch``.
|
||||
|
||||
LSP is **gated on git workspace detection** — if the agent's cwd is
|
||||
inside a git repository, LSP runs against that workspace; otherwise the
|
||||
file_operations layer falls back to its existing in-process syntax
|
||||
checks. This keeps users on user-home cwd's (e.g. Telegram gateway
|
||||
chats) from spawning daemons they don't need.
|
||||
LSP is **gated on git workspace detection**: outside a git repository the
|
||||
file_operations layer falls back to its in-process syntax checks, so users
|
||||
on user-home cwd's (e.g. Telegram gateway chats) never spawn daemons.
|
||||
|
||||
Public API:
|
||||
|
||||
from agent.lsp import get_service
|
||||
|
||||
svc = get_service()
|
||||
if svc and svc.enabled_for(path):
|
||||
await svc.touch_file(path)
|
||||
diags = svc.diagnostics_for(path)
|
||||
|
||||
The bulk of the wiring is internal — most callers only need the layer
|
||||
in :func:`tools.file_operations.FileOperations._check_lint_delta`,
|
||||
which is already wired (see that module).
|
||||
|
||||
Architecture is documented in ``website/docs/user-guide/features/lsp.md``.
|
||||
Public API: ``get_service()`` returns the singleton :class:`LSPService` (or
|
||||
``None`` when disabled); the wiring lives in
|
||||
:func:`tools.file_operations.FileOperations._check_lint_delta`. Architecture
|
||||
docs: ``website/docs/user-guide/features/lsp.md``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -42,46 +29,34 @@ _atexit_registered = False
|
||||
_service_lock = threading.Lock()
|
||||
|
||||
|
||||
def _active(svc: Optional[LSPService]) -> Optional[LSPService]:
|
||||
return svc if (svc is not None and svc.is_active()) else None
|
||||
|
||||
|
||||
def get_service() -> Optional[LSPService]:
|
||||
"""Return the process-wide LSP service singleton, or None when disabled.
|
||||
|
||||
The service is created lazily on first call. ``None`` is returned
|
||||
when LSP is disabled in config, when no workspace can be detected,
|
||||
or when the platform doesn't support subprocess-based LSP servers.
|
||||
|
||||
On first creation, registers an :mod:`atexit` handler that tears
|
||||
down spawned language servers on Python exit so a long-running
|
||||
CLI or gateway session doesn't leak pyright/gopls/etc. processes
|
||||
when it terminates.
|
||||
Created lazily on first call. Also registers an :mod:`atexit` hook so a
|
||||
clean exit tears down spawned language servers: without it every
|
||||
``hermes chat`` exit leaks pyright processes for a few seconds while
|
||||
their stdout buffers drain. (SIGKILL/os._exit skip atexit — fine, the
|
||||
kernel reaps the stateless servers with their parent.)
|
||||
"""
|
||||
global _service, _atexit_registered
|
||||
if _service is not None:
|
||||
return _service if _service.is_active() else None
|
||||
return _active(_service)
|
||||
with _service_lock:
|
||||
if _service is not None:
|
||||
return _service if _service.is_active() else None
|
||||
return _active(_service)
|
||||
_service = LSPService.create_from_config()
|
||||
if not _atexit_registered:
|
||||
# ``atexit`` handlers run in LIFO order on normal Python
|
||||
# exit and on SystemExit, but NOT on os._exit() or
|
||||
# uncaught signals. Language servers are stateless
|
||||
# subprocesses — losing them on SIGKILL is fine; they'll
|
||||
# be reaped by the kernel along with their parent. We
|
||||
# care about clean exits where Python flushes stdio
|
||||
# before terminating; without this hook every
|
||||
# ``hermes chat`` exit would leak pyright processes that
|
||||
# outlive the parent for a few seconds while their
|
||||
# stdout buffers drain.
|
||||
atexit.register(_atexit_shutdown)
|
||||
_atexit_registered = True
|
||||
return _service if (_service is not None and _service.is_active()) else None
|
||||
return _active(_service)
|
||||
|
||||
|
||||
def shutdown_service() -> None:
|
||||
"""Tear down the LSP service if one was started.
|
||||
|
||||
Safe to call multiple times; safe to call when no service was created.
|
||||
"""
|
||||
"""Tear down the LSP service if one was started. Idempotent."""
|
||||
global _service
|
||||
with _service_lock:
|
||||
svc = _service
|
||||
@@ -94,9 +69,7 @@ def shutdown_service() -> None:
|
||||
|
||||
|
||||
def _atexit_shutdown() -> None:
|
||||
"""atexit-registered wrapper. Logs at debug because by the time
|
||||
atexit fires the user has already seen the agent's final output —
|
||||
a noisy shutdown line on top of that is just clutter."""
|
||||
"""atexit wrapper; logs at debug since the user has already seen the final output."""
|
||||
try:
|
||||
shutdown_service()
|
||||
except Exception as e: # noqa: BLE001
|
||||
|
||||
+29
-51
@@ -1,16 +1,6 @@
|
||||
"""``hermes lsp`` CLI subcommand.
|
||||
"""``hermes lsp`` CLI subcommand: status / list / install / install-all / restart / which.
|
||||
|
||||
Subcommands:
|
||||
|
||||
- ``status`` — show service state, configured servers, install status.
|
||||
- ``install <server_id>`` — eagerly install one server's binary.
|
||||
- ``install-all`` — try to install every server with a known recipe.
|
||||
- ``restart`` — tear down running clients so the next edit re-spawns.
|
||||
- ``which <server_id>`` — print the resolved binary path for one server.
|
||||
- ``list`` — print the registry of supported servers.
|
||||
|
||||
The handlers are kept here (rather than in
|
||||
``hermes_cli/main.py``) so the LSP module ships self-contained.
|
||||
Handlers live here (not in ``hermes_cli/main.py``) so the LSP module ships self-contained.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -66,24 +56,25 @@ def register_subparser(subparsers: argparse._SubParsersAction) -> None:
|
||||
parser.set_defaults(func=run_lsp_command)
|
||||
|
||||
|
||||
_COMMANDS = {
|
||||
"status": lambda a: _cmd_status(getattr(a, "json", False)),
|
||||
"list": lambda a: _cmd_list(getattr(a, "installed_only", False)),
|
||||
"install": lambda a: _cmd_install(a.server),
|
||||
"install-all": lambda a: _cmd_install_all(getattr(a, "include_manual", False)),
|
||||
"restart": lambda a: _cmd_restart(),
|
||||
"which": lambda a: _cmd_which(a.server),
|
||||
}
|
||||
|
||||
|
||||
def run_lsp_command(args: argparse.Namespace) -> int:
|
||||
"""Top-level dispatcher for ``hermes lsp <subcommand>``."""
|
||||
sub = getattr(args, "lsp_command", None) or "status"
|
||||
try:
|
||||
if sub == "status":
|
||||
return _cmd_status(getattr(args, "json", False))
|
||||
if sub == "list":
|
||||
return _cmd_list(getattr(args, "installed_only", False))
|
||||
if sub == "install":
|
||||
return _cmd_install(args.server)
|
||||
if sub == "install-all":
|
||||
return _cmd_install_all(getattr(args, "include_manual", False))
|
||||
if sub == "restart":
|
||||
return _cmd_restart()
|
||||
if sub == "which":
|
||||
return _cmd_which(args.server)
|
||||
sys.stderr.write(f"unknown lsp subcommand: {sub}\n")
|
||||
return 2
|
||||
handler = _COMMANDS.get(sub)
|
||||
if handler is None:
|
||||
sys.stderr.write(f"unknown lsp subcommand: {sub}\n")
|
||||
return 2
|
||||
return handler(args)
|
||||
except KeyboardInterrupt:
|
||||
return 130
|
||||
|
||||
@@ -140,9 +131,7 @@ def _cmd_status(emit_json: bool) -> int:
|
||||
if disabled:
|
||||
out.append(f" disabled in cfg: {', '.join(disabled)}")
|
||||
|
||||
# Surface backend-tool gaps that aren't visible in the registry table:
|
||||
# some servers spawn fine but emit no diagnostics without a sidecar
|
||||
# binary (bash-language-server -> shellcheck).
|
||||
# Sidecar gaps the registry table can't show (bash-language-server -> shellcheck).
|
||||
backend_warnings = _backend_warnings()
|
||||
if backend_warnings:
|
||||
out.append("")
|
||||
@@ -259,33 +248,22 @@ def _cmd_which(server_id: str) -> int:
|
||||
return 1
|
||||
|
||||
|
||||
# server_id → install-recipe key, where the two differ.
|
||||
_RECIPE_ALIASES = {
|
||||
"vue-language-server": "@vue/language-server",
|
||||
"astro-language-server": "@astrojs/language-server",
|
||||
"dockerfile-ls": "dockerfile-language-server-nodejs",
|
||||
"typescript": "typescript-language-server",
|
||||
}
|
||||
|
||||
|
||||
def _recipe_pkg_for(server_id: str) -> str:
|
||||
"""Map a registry ``server_id`` to its install-recipe package key."""
|
||||
# The mapping lives here (not in install.py) because it's a CLI
|
||||
# convenience layer. Most server_ids are also their own recipe
|
||||
# key, but a few differ (e.g. ``vue-language-server`` →
|
||||
# ``@vue/language-server``).
|
||||
aliases = {
|
||||
"vue-language-server": "@vue/language-server",
|
||||
"astro-language-server": "@astrojs/language-server",
|
||||
"dockerfile-ls": "dockerfile-language-server-nodejs",
|
||||
"typescript": "typescript-language-server",
|
||||
}
|
||||
return aliases.get(server_id, server_id)
|
||||
return _RECIPE_ALIASES.get(server_id, server_id)
|
||||
|
||||
|
||||
def _backend_warnings() -> list:
|
||||
"""Return human-readable notes about LSP backend tools that are missing
|
||||
in a way that won't surface elsewhere.
|
||||
|
||||
Some language servers ship as thin wrappers around an external CLI for
|
||||
actual diagnostics — they spawn cleanly but never emit any errors when
|
||||
the sidecar binary isn't on PATH. bash-language-server / shellcheck
|
||||
is the load-bearing example.
|
||||
|
||||
Returned strings are short, actionable, and include the install
|
||||
suggestion across common platforms.
|
||||
"""
|
||||
"""Notes about missing sidecar tools that make a server spawn fine but emit nothing (e.g. shellcheck)."""
|
||||
import shutil as _shutil
|
||||
from agent.lsp.install import _existing_binary
|
||||
notes: list = []
|
||||
|
||||
+144
-306
@@ -1,49 +1,25 @@
|
||||
"""Async LSP client over stdin/stdout.
|
||||
|
||||
One :class:`LSPClient` corresponds to one ``(language_server, workspace_root)``
|
||||
pair — exactly what OpenCode keys clients on, and the same shape Claude
|
||||
Code uses. The client owns a child process, drives the JSON-RPC
|
||||
exchange, and exposes:
|
||||
|
||||
- :meth:`open_file` / :meth:`change_file` — text document sync
|
||||
- :meth:`wait_for_diagnostics` — block until the server emits fresh
|
||||
diagnostics for a specific file (or a timeout fires)
|
||||
- :meth:`diagnostics_for` — read the current per-file diagnostic store
|
||||
- :meth:`shutdown` — graceful close + SIGTERM/SIGKILL fallback
|
||||
|
||||
The class is designed for async use from a single asyncio event loop.
|
||||
The :class:`agent.lsp.manager.LSPService` runs an event loop in a
|
||||
background thread so the synchronous file_operations layer can call
|
||||
into it via :func:`agent.lsp.manager.LSPService.touch_file`.
|
||||
One :class:`LSPClient` per ``(language_server, workspace_root)`` pair. It owns
|
||||
the child process, drives JSON-RPC, and exposes :meth:`open_file`,
|
||||
:meth:`wait_for_diagnostics`, :meth:`diagnostics_for` and :meth:`shutdown`.
|
||||
:class:`agent.lsp.manager.LSPService` runs the event loop in a background
|
||||
thread so the synchronous file_operations layer can call in.
|
||||
|
||||
Implementation notes:
|
||||
|
||||
- All per-document state lives in one :class:`_DocState` keyed by
|
||||
absolute path. Freshness is tracked with **document versions**,
|
||||
not timestamps: every didChange bumps ``version``, and each stored
|
||||
push/pull result is tagged with the version it describes. A
|
||||
result is fresh iff its tag >= the version being waited on, so a
|
||||
didChange implicitly invalidates everything older — no clearing,
|
||||
no clock comparisons, no race windows. This is what prevents
|
||||
"ghost diagnostics": a slow server's leftovers from the previous
|
||||
edit can never masquerade as a verdict on the current content.
|
||||
|
||||
- Whole-document sync. Even when the server advertises incremental
|
||||
sync, we send a single ``contentChanges`` entry replacing the
|
||||
entire document. Pretending to be incremental while sending a
|
||||
full replacement is well-tolerated by every major server and saves
|
||||
range bookkeeping. See OpenCode's ``client.ts:584-659`` for the
|
||||
same trick.
|
||||
|
||||
- The "touch-file dance": every ``open_file`` call also fires a
|
||||
``workspace/didChangeWatchedFiles`` notification (CREATED on the
|
||||
first open, CHANGED thereafter). Some servers (clangd, eslint)
|
||||
only re-scan when this notification fires, even though the LSP spec
|
||||
doesn't strictly require it.
|
||||
|
||||
- ``ContentModified`` (-32801) errors get retried with exponential
|
||||
backoff up to 3 times. This matches Claude Code's
|
||||
``LSPServerInstance.sendRequest``.
|
||||
- Freshness is tracked with **document versions**, not timestamps: every
|
||||
didChange bumps ``version`` and each stored push/pull result is tagged with
|
||||
the version it describes. A result is fresh iff its tag >= the version
|
||||
being waited on, so a didChange implicitly invalidates everything older.
|
||||
This is what prevents "ghost diagnostics" — a slow server's leftovers from
|
||||
the previous edit can never masquerade as a verdict on the current content.
|
||||
- Whole-document sync: even when the server advertises incremental sync we
|
||||
send one ``contentChanges`` entry replacing the whole document. Every
|
||||
major server tolerates this and it saves range bookkeeping.
|
||||
- Every ``open_file`` also fires ``workspace/didChangeWatchedFiles`` (CREATED
|
||||
first, CHANGED after) — some servers (clangd, eslint) only re-scan on it.
|
||||
- ``ContentModified`` (-32801) errors are retried with exponential backoff.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -74,7 +50,7 @@ from agent.lsp.protocol import (
|
||||
|
||||
logger = logging.getLogger("agent.lsp.client")
|
||||
|
||||
# Timeouts (seconds) — mirror OpenCode's constants, scaled to seconds.
|
||||
# Timeouts (seconds).
|
||||
INITIALIZE_TIMEOUT = 45.0
|
||||
DIAGNOSTICS_DOCUMENT_WAIT = 5.0
|
||||
DIAGNOSTICS_FULL_WAIT = 10.0
|
||||
@@ -82,21 +58,18 @@ DIAGNOSTICS_REQUEST_TIMEOUT = 3.0
|
||||
PUSH_DEBOUNCE = 0.15
|
||||
SHUTDOWN_GRACE = 1.0 # seconds between SIGTERM and SIGKILL
|
||||
|
||||
# Retry policy for transient ContentModified errors.
|
||||
# Retry policy for transient ContentModified errors: 0.5, 1.0, 2.0s.
|
||||
MAX_CONTENT_MODIFIED_RETRIES = 3
|
||||
RETRY_BASE_DELAY = 0.5 # 0.5, 1.0, 2.0 — exponential
|
||||
RETRY_BASE_DELAY = 0.5
|
||||
|
||||
_WRITE_ERRORS = (BrokenPipeError, ConnectionResetError, OSError)
|
||||
|
||||
|
||||
def file_uri(path: str) -> str:
|
||||
"""Return ``file://`` URI for an absolute filesystem path.
|
||||
|
||||
Mirrors Node's ``pathToFileURL`` — handles spaces, unicode, and
|
||||
Windows drive letters (``C:\\foo`` → ``file:///C:/foo``).
|
||||
"""
|
||||
"""Return a ``file://`` URI for a path (handles spaces, unicode, Windows drive letters)."""
|
||||
abs_path = os.path.abspath(path)
|
||||
if os.name == "nt":
|
||||
# Windows: backslash → forward slash, prepend extra slash so
|
||||
# the drive letter shows up as part of the path component.
|
||||
# ``C:\foo`` → ``file:///C:/foo``: the drive letter must be a path component.
|
||||
abs_path = abs_path.replace("\\", "/")
|
||||
if not abs_path.startswith("/"):
|
||||
abs_path = "/" + abs_path
|
||||
@@ -114,42 +87,27 @@ def uri_to_path(uri: str) -> str:
|
||||
|
||||
|
||||
def _end_position(text: str) -> Dict[str, int]:
|
||||
"""Return the LSP Position at the end of ``text``.
|
||||
|
||||
Used to construct a single-range "replace whole document" change
|
||||
for ``textDocument/didChange`` regardless of the server's declared
|
||||
sync mode.
|
||||
"""
|
||||
"""LSP Position at the end of ``text`` (for a whole-document replace range)."""
|
||||
if not text:
|
||||
return {"line": 0, "character": 0}
|
||||
lines = text.splitlines(keepends=False)
|
||||
last_line = len(lines) - 1
|
||||
last_col = len(lines[-1]) if lines else 0
|
||||
# If the text ends with a trailing newline, ``splitlines`` won't
|
||||
# represent it. The end position is then the start of the next
|
||||
# (empty) line — line index is len(lines), column 0.
|
||||
# A trailing newline isn't represented by splitlines: the end is then
|
||||
# the start of the next (empty) line.
|
||||
if text.endswith(("\n", "\r")):
|
||||
return {"line": last_line + 1, "character": 0}
|
||||
return {"line": last_line, "character": last_col}
|
||||
return {"line": len(lines), "character": 0}
|
||||
return {"line": len(lines) - 1, "character": len(lines[-1])}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _DocState:
|
||||
"""Everything the client tracks for one open document.
|
||||
"""Per-document state.
|
||||
|
||||
``version`` is the LSP document version we last sent (didOpen=0,
|
||||
each didChange +1). It doubles as the freshness token: stored
|
||||
push/pull results are tagged with the version they describe
|
||||
(``push_version`` / ``pull_version``), and a result is *fresh*
|
||||
iff its tag has caught up to ``version``. Bumping the version on
|
||||
didChange therefore invalidates all older results implicitly —
|
||||
no store-clearing, no timestamps.
|
||||
|
||||
``push_version``/``pull_version`` start at -1 = "no data yet".
|
||||
Servers that echo a document version in publishDiagnostics get
|
||||
exact tagging; those that don't are credited with the current
|
||||
version at receipt time (a push observed after we sent the
|
||||
change describes the changed content or newer).
|
||||
``version`` is the LSP document version last sent (didOpen=0, +1 per
|
||||
didChange) and doubles as the freshness token: ``push_version`` /
|
||||
``pull_version`` tag stored results, which are fresh iff tag >= version.
|
||||
Tags start at -1 ("no data yet"). Servers that echo a version in
|
||||
publishDiagnostics get exact tagging; others are credited with the
|
||||
current version at receipt time.
|
||||
"""
|
||||
|
||||
version: int = 0
|
||||
@@ -170,20 +128,10 @@ class _DocState:
|
||||
class LSPClient:
|
||||
"""Async LSP client tied to one server process and one workspace root.
|
||||
|
||||
Lifecycle:
|
||||
|
||||
c = LSPClient(server_id, workspace_root, command, args, init_options)
|
||||
await c.start() # spawn + initialize
|
||||
ver = await c.open_file("/path/to/foo.py")
|
||||
await c.wait_for_diagnostics("/path/to/foo.py", ver)
|
||||
diags = c.diagnostics_for("/path/to/foo.py")
|
||||
await c.shutdown()
|
||||
Lifecycle: ``start()`` → ``open_file()`` → ``wait_for_diagnostics()`` →
|
||||
``diagnostics_for()`` → ``shutdown()``.
|
||||
"""
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# construction + lifecycle
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -203,18 +151,15 @@ class LSPClient:
|
||||
self._init_options = initialization_options or {}
|
||||
self._seed_first_push = seed_diagnostics_on_first_push
|
||||
|
||||
# Process + streams
|
||||
self._proc: Optional[asyncio.subprocess.Process] = None
|
||||
self._stderr_task: Optional[asyncio.Task] = None
|
||||
self._reader_task: Optional[asyncio.Task] = None
|
||||
self._cleanup_lock = asyncio.Lock()
|
||||
|
||||
# Request/response correlation
|
||||
self._next_id: int = 0
|
||||
self._pending: Dict[int, asyncio.Future] = {}
|
||||
|
||||
# Server-side request handlers (server → client requests).
|
||||
# Kept small and explicit; everything else returns method-not-found.
|
||||
# Server → client requests; anything else gets method-not-found.
|
||||
self._request_handlers: Dict[str, Callable[[Any], Awaitable[Any]]] = {
|
||||
"window/workDoneProgress/create": self._handle_work_done_create,
|
||||
"workspace/configuration": self._handle_workspace_configuration,
|
||||
@@ -223,35 +168,24 @@ class LSPClient:
|
||||
"workspace/workspaceFolders": self._handle_workspace_folders,
|
||||
"workspace/diagnostic/refresh": self._handle_diagnostic_refresh,
|
||||
}
|
||||
# Notifications (server → client) we care about.
|
||||
# Server → client notifications; others (showMessage, $/progress) are dropped.
|
||||
self._notification_handlers: Dict[str, Callable[[Any], None]] = {
|
||||
"textDocument/publishDiagnostics": self._handle_publish_diagnostics,
|
||||
# Everything else (window/showMessage, $/progress, etc.)
|
||||
# is silently dropped by default.
|
||||
}
|
||||
|
||||
# Per-document state (version, text, diagnostic stores, and
|
||||
# their freshness tags), keyed by absolute file path (NOT URI).
|
||||
# See _DocState for the version-based freshness model.
|
||||
# Per-document state keyed by absolute path (NOT URI).
|
||||
self._docs: Dict[str, _DocState] = {}
|
||||
# Capability registrations — only diagnostic ones are tracked.
|
||||
# Only diagnostic capability registrations are tracked.
|
||||
self._diagnostic_registrations: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
# State machine
|
||||
self._state: str = "stopped"
|
||||
self._initialize_result: Optional[Dict[str, Any]] = None
|
||||
self._sync_kind: int = 1 # 1=Full, 2=Incremental
|
||||
self._stopping: bool = False
|
||||
|
||||
# Push event for waiters.
|
||||
# Waiters snapshot ``_push_counter`` and treat any increase as "recheck
|
||||
# the predicate" — avoids the asyncio.Event sticky-state trap.
|
||||
self._push_event = asyncio.Event()
|
||||
# Monotonic counter incremented on every publishDiagnostics push.
|
||||
# Waiters snapshot it on entry and treat any increase as
|
||||
# "something happened, recheck the predicate". Avoids the
|
||||
# asyncio.Event sticky-state trap.
|
||||
self._push_counter = 0
|
||||
# Registration change event so wait_for_diagnostics can re-loop
|
||||
# when the server announces a new dynamic provider.
|
||||
self._registration_event = asyncio.Event()
|
||||
|
||||
@property
|
||||
@@ -278,9 +212,7 @@ class LSPClient:
|
||||
async def start(self) -> None:
|
||||
"""Spawn the server and complete the initialize handshake.
|
||||
|
||||
Raises any exception encountered during spawn/init. On failure
|
||||
the process is killed and the client is left in state
|
||||
``"error"`` — re-call ``start()`` to retry.
|
||||
On failure the process is killed and state is ``"error"``; re-call to retry.
|
||||
"""
|
||||
if self._state in {"running", "starting"}:
|
||||
return
|
||||
@@ -299,8 +231,7 @@ class LSPClient:
|
||||
@staticmethod
|
||||
def _win_wrap_cmd(cmd: List[str]) -> List[str]:
|
||||
"""On Windows, wrap .cmd/.bat shims so CreateProcess can run them."""
|
||||
exe = cmd[0]
|
||||
if exe.lower().endswith((".cmd", ".bat")):
|
||||
if cmd[0].lower().endswith((".cmd", ".bat")):
|
||||
return ["cmd.exe", "/c", *cmd]
|
||||
return cmd
|
||||
|
||||
@@ -312,21 +243,13 @@ class LSPClient:
|
||||
cmd = self._command
|
||||
if sys.platform == "win32":
|
||||
cmd = self._win_wrap_cmd(cmd)
|
||||
# Suppress the cmd.exe console window that would otherwise flash
|
||||
# every time we launch a ``.cmd``-wrapped language server
|
||||
# (e.g. pyright-langserver.CMD) from a console-less host such as
|
||||
# a VS Code/Zed extension running the ACP adapter.
|
||||
# windows_hide_flags() is CREATE_NO_WINDOW on Windows, 0 on POSIX.
|
||||
creationflags = windows_hide_flags()
|
||||
|
||||
try:
|
||||
# start_new_session=True detaches the LSP server into its own
|
||||
# process group / session. Without this, the LSP server inherits
|
||||
# the gateway's pgid (= TUI parent PID). When mcp_tool's
|
||||
# _kill_orphaned_mcp_children races with LSP spawn and sweeps the
|
||||
# gateway's child set, it captures the LSP PID, records the
|
||||
# inherited pgid, and killpg() then kills the TUI parent itself.
|
||||
# See tui_gateway_crash.log "killpg → SIGTERM received" stacks.
|
||||
# start_new_session=True gives the server its own process group.
|
||||
# Otherwise it inherits the gateway's pgid and mcp_tool's orphan
|
||||
# sweeper can killpg() the TUI parent along with it.
|
||||
# windows_hide_flags() suppresses the console window a .cmd shim
|
||||
# would flash from a console-less host (CREATE_NO_WINDOW; 0 on POSIX).
|
||||
self._proc = await asyncio.create_subprocess_exec(
|
||||
cmd[0],
|
||||
*cmd[1:],
|
||||
@@ -336,17 +259,15 @@ class LSPClient:
|
||||
env=env,
|
||||
cwd=self._cwd,
|
||||
start_new_session=True,
|
||||
creationflags=creationflags,
|
||||
creationflags=windows_hide_flags(),
|
||||
)
|
||||
except FileNotFoundError as e:
|
||||
raise LSPProtocolError(
|
||||
f"LSP server binary not found: {cmd[0]} ({e})"
|
||||
) from e
|
||||
|
||||
# Drain stderr at debug level — if we don't, the pipe buffer
|
||||
# fills and the server hangs.
|
||||
# stderr must be drained or the pipe buffer fills and the server hangs.
|
||||
self._stderr_task = asyncio.create_task(self._drain_stderr())
|
||||
# Start the reader loop.
|
||||
self._reader_task = asyncio.create_task(self._reader_loop())
|
||||
|
||||
async def _drain_stderr(self) -> None:
|
||||
@@ -389,7 +310,7 @@ class LSPClient:
|
||||
unexpected_close = not self._stopping and self._state in {"starting", "running"}
|
||||
if unexpected_close:
|
||||
self._state = "error"
|
||||
# Wake up any pending requests so they can fail fast.
|
||||
# Fail pending requests fast.
|
||||
for fut in list(self._pending.values()):
|
||||
if not fut.done():
|
||||
fut.set_exception(LSPProtocolError("server connection closed"))
|
||||
@@ -447,14 +368,12 @@ class LSPClient:
|
||||
self._send_request("initialize", params),
|
||||
timeout=INITIALIZE_TIMEOUT,
|
||||
)
|
||||
self._initialize_result = result
|
||||
self._sync_kind = self._extract_sync_kind(result.get("capabilities") or {})
|
||||
|
||||
await self._send_notification("initialized", {})
|
||||
if self._init_options:
|
||||
# Some servers (vtsls, eslint) want config pushed via
|
||||
# didChangeConfiguration even if it was sent in
|
||||
# initializationOptions.
|
||||
# Some servers (vtsls, eslint) only pick config up via
|
||||
# didChangeConfiguration even when it was in initializationOptions.
|
||||
await self._send_notification(
|
||||
"workspace/didChangeConfiguration",
|
||||
{"settings": self._init_options},
|
||||
@@ -472,11 +391,7 @@ class LSPClient:
|
||||
return 1 # default to Full
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
"""Best-effort graceful shutdown.
|
||||
|
||||
Sends ``shutdown`` + ``exit``, then SIGTERMs/SIGKILLs the
|
||||
process if it doesn't exit cleanly. Idempotent.
|
||||
"""
|
||||
"""Best-effort graceful shutdown: ``shutdown`` + ``exit``, then SIGTERM/SIGKILL. Idempotent."""
|
||||
if self._stopping:
|
||||
return
|
||||
self._stopping = True
|
||||
@@ -494,64 +409,61 @@ class LSPClient:
|
||||
self._state = "stopped"
|
||||
await self._cleanup_process()
|
||||
|
||||
@staticmethod
|
||||
async def _cancel_task(task: Optional[asyncio.Task]) -> None:
|
||||
if task is not None and not task.done():
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except (asyncio.CancelledError, Exception): # noqa: BLE001
|
||||
pass
|
||||
|
||||
async def _cleanup_process(self) -> None:
|
||||
async with self._cleanup_lock:
|
||||
current_task = asyncio.current_task()
|
||||
reader_task = self._reader_task
|
||||
self._reader_task = None
|
||||
if (
|
||||
reader_task is not None
|
||||
and reader_task is not current_task
|
||||
and not reader_task.done()
|
||||
):
|
||||
reader_task.cancel()
|
||||
try:
|
||||
await reader_task
|
||||
except (asyncio.CancelledError, Exception): # noqa: BLE001
|
||||
pass
|
||||
if reader_task is not asyncio.current_task():
|
||||
await self._cancel_task(reader_task)
|
||||
stderr_task = self._stderr_task
|
||||
self._stderr_task = None
|
||||
if stderr_task is not None and not stderr_task.done():
|
||||
stderr_task.cancel()
|
||||
try:
|
||||
await stderr_task
|
||||
except (asyncio.CancelledError, Exception): # noqa: BLE001
|
||||
pass
|
||||
await self._cancel_task(stderr_task)
|
||||
proc = self._proc
|
||||
self._proc = None
|
||||
if proc is None:
|
||||
if proc is None or proc.returncode is not None:
|
||||
return
|
||||
if proc.returncode is None:
|
||||
try:
|
||||
proc.terminate()
|
||||
try:
|
||||
proc.terminate()
|
||||
await asyncio.wait_for(proc.wait(), timeout=SHUTDOWN_GRACE)
|
||||
except asyncio.TimeoutError:
|
||||
try:
|
||||
await asyncio.wait_for(proc.wait(), timeout=SHUTDOWN_GRACE)
|
||||
except asyncio.TimeoutError:
|
||||
try:
|
||||
proc.kill()
|
||||
await proc.wait()
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
proc.kill()
|
||||
await proc.wait()
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# request / notification plumbing
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _write(self, msg: dict) -> None:
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
self._proc.stdin.write(encode_message(msg))
|
||||
await self._proc.stdin.drain()
|
||||
|
||||
async def _send_request(self, method: str, params: Any) -> Any:
|
||||
if not self._connection_is_open():
|
||||
raise LSPProtocolError(f"cannot send {method!r}: server connection closed")
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
loop = asyncio.get_running_loop()
|
||||
req_id = self._next_id
|
||||
self._next_id += 1
|
||||
fut: asyncio.Future = loop.create_future()
|
||||
self._pending[req_id] = fut
|
||||
try:
|
||||
self._proc.stdin.write(encode_message(make_request(req_id, method, params)))
|
||||
await self._proc.stdin.drain()
|
||||
except (BrokenPipeError, ConnectionResetError, OSError) as e:
|
||||
await self._write(make_request(req_id, method, params))
|
||||
except _WRITE_ERRORS as e:
|
||||
self._pending.pop(req_id, None)
|
||||
raise LSPProtocolError(f"send failed for {method!r}: {e}") from e
|
||||
try:
|
||||
@@ -560,12 +472,7 @@ class LSPClient:
|
||||
self._pending.pop(req_id, None)
|
||||
|
||||
async def _send_request_with_retry(self, method: str, params: Any, *, timeout: float) -> Any:
|
||||
"""Send a request, retrying on ``ContentModified`` (-32801).
|
||||
|
||||
Other errors propagate. The retry policy matches Claude Code's
|
||||
``LSPServerInstance.sendRequest`` — 3 attempts with delays
|
||||
0.5s, 1.0s, 2.0s.
|
||||
"""
|
||||
"""Send a request, retrying ``ContentModified`` (-32801) with backoff; other errors propagate."""
|
||||
for attempt in range(MAX_CONTENT_MODIFIED_RETRIES + 1):
|
||||
try:
|
||||
return await asyncio.wait_for(self._send_request(method, params), timeout=timeout)
|
||||
@@ -578,29 +485,18 @@ class LSPClient:
|
||||
async def _send_notification(self, method: str, params: Any) -> None:
|
||||
if not self._connection_is_open():
|
||||
raise LSPProtocolError(f"cannot send {method!r}: server connection closed")
|
||||
assert self._proc is not None and self._proc.stdin is not None
|
||||
try:
|
||||
self._proc.stdin.write(encode_message(make_notification(method, params)))
|
||||
await self._proc.stdin.drain()
|
||||
except (BrokenPipeError, ConnectionResetError, OSError) as e:
|
||||
await self._write(make_notification(method, params))
|
||||
except _WRITE_ERRORS as e:
|
||||
logger.debug("[%s] notify %s failed: %s", self.server_id, method, e)
|
||||
|
||||
async def _send_response(self, req_id: Any, result: Any) -> None:
|
||||
async def _send_reply(self, msg: dict) -> None:
|
||||
"""Send a response to a server→client request; silently no-ops when the pipe is gone."""
|
||||
if self._proc is None or self._proc.stdin is None or self._proc.stdin.is_closing():
|
||||
return
|
||||
try:
|
||||
self._proc.stdin.write(encode_message(make_response(req_id, result)))
|
||||
await self._proc.stdin.drain()
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
pass
|
||||
|
||||
async def _send_error_response(self, req_id: Any, code: int, message: str) -> None:
|
||||
if self._proc is None or self._proc.stdin is None or self._proc.stdin.is_closing():
|
||||
return
|
||||
try:
|
||||
self._proc.stdin.write(encode_message(make_error_response(req_id, code, message)))
|
||||
await self._proc.stdin.drain()
|
||||
except (BrokenPipeError, ConnectionResetError, OSError):
|
||||
await self._write(msg)
|
||||
except _WRITE_ERRORS:
|
||||
pass
|
||||
|
||||
def _dispatch_response(self, req_id: int, msg: dict) -> None:
|
||||
@@ -624,15 +520,15 @@ class LSPClient:
|
||||
params = msg.get("params")
|
||||
handler = self._request_handlers.get(method)
|
||||
if handler is None:
|
||||
await self._send_error_response(req_id, ERROR_METHOD_NOT_FOUND, f"method not found: {method}")
|
||||
await self._send_reply(make_error_response(req_id, ERROR_METHOD_NOT_FOUND, f"method not found: {method}"))
|
||||
return
|
||||
try:
|
||||
result = await handler(params)
|
||||
except Exception as e: # noqa: BLE001 — protocol must not blow up
|
||||
logger.warning("[%s] request handler %s failed: %s", self.server_id, method, e)
|
||||
await self._send_error_response(req_id, -32000, f"handler failed: {e}")
|
||||
await self._send_reply(make_error_response(req_id, -32000, f"handler failed: {e}"))
|
||||
return
|
||||
await self._send_response(req_id, result)
|
||||
await self._send_reply(make_response(req_id, result))
|
||||
|
||||
def _dispatch_notification(self, method: str, msg: dict) -> None:
|
||||
handler = self._notification_handlers.get(method)
|
||||
@@ -652,13 +548,11 @@ class LSPClient:
|
||||
return None
|
||||
|
||||
async def _handle_workspace_configuration(self, params: Any) -> Any:
|
||||
# Walk dotted sections through initializationOptions. Mirrors
|
||||
# OpenCode's `client.ts:198-220` — return null when missing.
|
||||
# Walk dotted sections through initializationOptions; null when missing.
|
||||
if not isinstance(params, dict):
|
||||
return [None]
|
||||
items = params.get("items") or []
|
||||
out: List[Any] = []
|
||||
for item in items:
|
||||
for item in params.get("items") or []:
|
||||
if not isinstance(item, dict):
|
||||
out.append(None)
|
||||
continue
|
||||
@@ -682,9 +576,8 @@ class LSPClient:
|
||||
for reg in params.get("registrations") or []:
|
||||
if not isinstance(reg, dict):
|
||||
continue
|
||||
method = reg.get("method")
|
||||
reg_id = reg.get("id")
|
||||
if method == "textDocument/diagnostic" and reg_id:
|
||||
if reg.get("method") == "textDocument/diagnostic" and reg_id:
|
||||
self._diagnostic_registrations[str(reg_id)] = reg
|
||||
self._registration_event.set()
|
||||
return None
|
||||
@@ -707,45 +600,32 @@ class LSPClient:
|
||||
# We don't honour refresh — we re-pull on every touchFile.
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# publishDiagnostics handler
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _handle_publish_diagnostics(self, params: Any) -> None:
|
||||
if not isinstance(params, dict):
|
||||
return
|
||||
uri = params.get("uri")
|
||||
if not isinstance(uri, str):
|
||||
return
|
||||
path = uri_to_path(uri)
|
||||
diagnostics = params.get("diagnostics") or []
|
||||
if not isinstance(diagnostics, list):
|
||||
diagnostics = []
|
||||
version = params.get("version")
|
||||
|
||||
doc = self._docs.setdefault(path, _DocState(version=-1))
|
||||
if self._seed_first_push and not doc.seed_seen:
|
||||
# First push: seed the store WITHOUT a freshness tag. It
|
||||
# arrives before the user-triggered didChange could've
|
||||
# produced fresh diagnostics, so it must never satisfy a
|
||||
# waiter — it's baseline data only.
|
||||
doc.seed_seen = True
|
||||
doc.push = diagnostics
|
||||
return
|
||||
|
||||
doc = self._docs.setdefault(uri_to_path(uri), _DocState(version=-1))
|
||||
is_seed = self._seed_first_push and not doc.seed_seen
|
||||
doc.seed_seen = True
|
||||
doc.push = diagnostics
|
||||
# Tag with the echoed document version when the server provides
|
||||
# one; otherwise credit the current version — a push observed
|
||||
# after we sent the change describes the changed content (or
|
||||
# newer). Note doc.version is -1 for never-opened paths
|
||||
# (e.g. relatedDocuments spillover), keeping them unfresh.
|
||||
if is_seed:
|
||||
# First push is baseline data only: it predates any didChange we
|
||||
# sent, so it's stored WITHOUT a freshness tag and never satisfies a waiter.
|
||||
return
|
||||
# Tag with the echoed version when provided; otherwise credit the
|
||||
# current version (a push observed after our change describes it or
|
||||
# newer). doc.version is -1 for never-opened paths (relatedDocuments
|
||||
# spillover), keeping them unfresh.
|
||||
doc.push_version = version if isinstance(version, int) else doc.version
|
||||
# Bump the monotonic push counter and wake every waiter. We
|
||||
# keep the Event sticky-set so any wait already in progress
|
||||
# resolves; waiters re-check their predicate after waking and
|
||||
# decide whether to keep waiting. ``_push_counter`` is what
|
||||
# they actually compare against to detect a fresh event.
|
||||
# Keep the Event sticky-set so in-progress waits resolve; waiters
|
||||
# compare ``_push_counter`` to detect a genuinely new push.
|
||||
self._push_counter += 1
|
||||
self._push_event.set()
|
||||
|
||||
@@ -754,11 +634,7 @@ class LSPClient:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def open_file(self, path: str, *, language_id: str = "plaintext") -> int:
|
||||
"""Send didOpen (first time) or didChange (subsequent) for ``path``.
|
||||
|
||||
Returns the new document version number that the agent's
|
||||
``wait_for_diagnostics`` should match against.
|
||||
"""
|
||||
"""Send didOpen (first time) or didChange (subsequent); return the new document version."""
|
||||
if not self.is_running:
|
||||
raise LSPProtocolError("client not running")
|
||||
|
||||
@@ -772,20 +648,18 @@ class LSPClient:
|
||||
doc = self._docs.get(abs_path)
|
||||
|
||||
if doc is not None and doc.version >= 0:
|
||||
# Re-open: bump version, fire didChangeWatchedFiles + didChange.
|
||||
await self._send_notification(
|
||||
"workspace/didChangeWatchedFiles",
|
||||
{"changes": [{"uri": uri, "type": 2}]}, # 2 = CHANGED
|
||||
)
|
||||
new_version = doc.version + 1
|
||||
old_text = doc.text
|
||||
content_changes: List[Dict[str, Any]]
|
||||
if self._sync_kind == 2:
|
||||
content_changes = [
|
||||
{
|
||||
"range": {
|
||||
"start": {"line": 0, "character": 0},
|
||||
"end": _end_position(old_text),
|
||||
"end": _end_position(doc.text),
|
||||
},
|
||||
"text": text,
|
||||
}
|
||||
@@ -799,20 +673,17 @@ class LSPClient:
|
||||
"contentChanges": content_changes,
|
||||
},
|
||||
)
|
||||
# Bumping the version is the whole invalidation story:
|
||||
# every stored result tagged with an older version is now
|
||||
# stale by definition (see _DocState).
|
||||
# Bumping the version is the whole invalidation story (see _DocState).
|
||||
doc.version = new_version
|
||||
doc.text = text
|
||||
return new_version
|
||||
|
||||
# First open: didChangeWatchedFiles CREATED + didOpen.
|
||||
await self._send_notification(
|
||||
"workspace/didChangeWatchedFiles",
|
||||
{"changes": [{"uri": uri, "type": 1}]}, # 1 = CREATED
|
||||
)
|
||||
# Fresh doc state — anything stashed under this path by a
|
||||
# pre-open push (relatedDocuments spillover etc.) is discarded.
|
||||
# Fresh state: anything a pre-open push stashed under this path
|
||||
# (relatedDocuments spillover) is discarded.
|
||||
self._docs[abs_path] = _DocState(version=0, text=text)
|
||||
await self._send_notification(
|
||||
"textDocument/didOpen",
|
||||
@@ -831,10 +702,9 @@ class LSPClient:
|
||||
"""Send didSave for ``path``. Some linters re-scan only on save."""
|
||||
if not self.is_running:
|
||||
return
|
||||
abs_path = os.path.abspath(path)
|
||||
await self._send_notification(
|
||||
"textDocument/didSave",
|
||||
{"textDocument": {"uri": file_uri(abs_path)}},
|
||||
{"textDocument": {"uri": file_uri(os.path.abspath(path))}},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -842,25 +712,19 @@ class LSPClient:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _pull_document_diagnostics(self, path: str) -> None:
|
||||
"""Send ``textDocument/diagnostic`` for one file.
|
||||
"""Send ``textDocument/diagnostic`` for one file into the pull store.
|
||||
|
||||
Stores results into the doc's pull store, tagged with the
|
||||
document version captured at request send time. If a didChange
|
||||
races past the in-flight request, the version bump makes the
|
||||
stored result stale automatically — no explicit invalidation.
|
||||
Silently no-ops on errors (server may not support the pull
|
||||
endpoint).
|
||||
Results are tagged with the version captured at send time, so a
|
||||
didChange racing past the request makes them stale automatically.
|
||||
Silently no-ops on errors (server may not support pull).
|
||||
"""
|
||||
abs_path = os.path.abspath(path)
|
||||
doc = self._docs.get(abs_path)
|
||||
sent_version = doc.version if doc else -1
|
||||
try:
|
||||
params: Dict[str, Any] = {
|
||||
"textDocument": {"uri": file_uri(abs_path)}
|
||||
}
|
||||
result = await self._send_request_with_retry(
|
||||
"textDocument/diagnostic",
|
||||
params,
|
||||
{"textDocument": {"uri": file_uri(abs_path)}},
|
||||
timeout=DIAGNOSTICS_REQUEST_TIMEOUT,
|
||||
)
|
||||
except (LSPRequestError, LSPProtocolError, asyncio.TimeoutError) as e:
|
||||
@@ -882,8 +746,7 @@ class LSPClient:
|
||||
if isinstance(sub_items, list):
|
||||
rel = self._docs.setdefault(uri_to_path(uri), _DocState(version=-1))
|
||||
rel.pull = sub_items
|
||||
# Same send-anchored tagging: fresh only if that
|
||||
# doc hasn't changed since the request went out.
|
||||
# Same send-anchored tagging: fresh only if that doc hasn't changed since.
|
||||
rel.pull_version = rel.version
|
||||
|
||||
async def wait_for_diagnostics(
|
||||
@@ -894,22 +757,14 @@ class LSPClient:
|
||||
mode: str = "document",
|
||||
timeout: Optional[float] = None,
|
||||
) -> bool:
|
||||
"""Wait for the server to publish diagnostics for ``path`` at ``version``.
|
||||
"""Wait for fresh diagnostics for ``path`` at ``version``.
|
||||
|
||||
``mode`` is ``"document"`` (5s budget, document pulls) or
|
||||
``"full"`` (10s budget, also workspace pulls). ``timeout``
|
||||
overrides the mode's default budget when provided — this is
|
||||
how the user's ``lsp.wait_timeout`` config reaches the wait
|
||||
loop (slow servers like tsserver on big projects need more
|
||||
than the 5s default).
|
||||
|
||||
Returns ``True`` when *fresh* diagnostics arrived (a push at
|
||||
or after our didChange, or a pull answered after it) and
|
||||
``False`` on timeout. Callers must treat ``False`` as "no
|
||||
data", NOT as "no errors" — the diagnostic stores may still
|
||||
hold stale entries from the previous edit at that point.
|
||||
Best-effort — never throws if the server doesn't support pull
|
||||
diagnostics; we still get the push side.
|
||||
``mode`` is ``"document"`` (5s) or ``"full"`` (10s); ``timeout`` overrides
|
||||
the budget (this is how ``lsp.wait_timeout`` reaches the loop). Returns
|
||||
True when fresh data arrived (push at/after our didChange, or a pull
|
||||
answered after it), False on timeout. Callers must treat False as
|
||||
"no data", NOT "no errors" — the stores may still hold stale entries.
|
||||
Never throws for servers lacking pull support; the push side still works.
|
||||
"""
|
||||
if timeout is not None and timeout > 0:
|
||||
budget = timeout
|
||||
@@ -943,18 +798,10 @@ class LSPClient:
|
||||
except (asyncio.CancelledError, Exception): # noqa: BLE001
|
||||
pass
|
||||
|
||||
# If we got a fresh push for our version, we're done.
|
||||
doc = self._docs.get(abs_path)
|
||||
if doc and doc.fresh_push(version):
|
||||
if doc and (doc.fresh_push(version) or doc.fresh_pull(version)):
|
||||
return True
|
||||
|
||||
# Pull may have answered for the current version — that's
|
||||
# also success.
|
||||
if doc and doc.fresh_pull(version):
|
||||
return True
|
||||
|
||||
# Loop until budget runs out.
|
||||
|
||||
async def _wait_for_fresh_push(self, path: str, version: int, timeout: float) -> None:
|
||||
"""Wait until a fresh publishDiagnostics arrives for ``path`` at ``version``+."""
|
||||
deadline = asyncio.get_event_loop().time() + timeout
|
||||
@@ -962,10 +809,8 @@ class LSPClient:
|
||||
while True:
|
||||
doc = self._docs.get(path)
|
||||
if doc and doc.fresh_push(version):
|
||||
# Debounce — wait a tick in case more diagnostics arrive
|
||||
# immediately after. TS often emits in pairs. We
|
||||
# snapshot the counter so we wake on a *new* push, not
|
||||
# on the one that satisfied us a moment ago.
|
||||
# Debounce: TS often emits in pairs. Snapshot the counter so
|
||||
# we wake on a *new* push, not the one that just satisfied us.
|
||||
debounce_baseline = self._push_counter
|
||||
debounce_deadline = asyncio.get_event_loop().time() + PUSH_DEBOUNCE
|
||||
while self._push_counter == debounce_baseline:
|
||||
@@ -982,8 +827,7 @@ class LSPClient:
|
||||
if remaining <= 0:
|
||||
return
|
||||
if self._push_counter > baseline:
|
||||
# New event arrived but predicate still false — re-check
|
||||
# immediately without waiting again.
|
||||
# New push but predicate still false — re-check without waiting.
|
||||
baseline = self._push_counter
|
||||
continue
|
||||
self._push_event.clear()
|
||||
@@ -993,17 +837,11 @@ class LSPClient:
|
||||
continue
|
||||
|
||||
def diagnostics_for(self, path: str, *, fresh_only: bool = False) -> List[Dict[str, Any]]:
|
||||
"""Return current merged + deduped diagnostics for one file.
|
||||
"""Merged + deduped push/pull diagnostics for one file.
|
||||
|
||||
Diagnostics from push and pull stores are concatenated and
|
||||
deduplicated by ``(severity, code, message, range)`` content
|
||||
key. Empty list if the server hasn't published anything.
|
||||
|
||||
With ``fresh_only=True``, a store only contributes when its
|
||||
version tag has caught up to the document's current version —
|
||||
stale leftovers from the previous edit cycle are excluded.
|
||||
This is what report paths should use: after an edit, "stale
|
||||
errors" and "no errors" must not be conflated.
|
||||
With ``fresh_only=True`` a store only contributes when its version tag
|
||||
has caught up to the document's version — report paths must use this
|
||||
so "stale errors" and "no errors" aren't conflated.
|
||||
"""
|
||||
doc = self._docs.get(os.path.abspath(path))
|
||||
if doc is None:
|
||||
@@ -1032,12 +870,12 @@ def _dedupe(*lists: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
|
||||
|
||||
def _diagnostic_key(d: Dict[str, Any]) -> str:
|
||||
"""Content-equality key for a diagnostic.
|
||||
"""Content-equality key: severity + code + source + message + range.
|
||||
|
||||
Matches the structural-equality used in claude-code's
|
||||
``areDiagnosticsEqual`` — message + severity + source + code +
|
||||
range coords. The range is reduced to a tuple to keep the key
|
||||
stable across dict orderings.
|
||||
Shared with the manager's cross-edit delta filter (as ``_diag_key``) so
|
||||
both layers agree on diagnostic identity. The range is included so an
|
||||
identical error introduced at a second site still surfaces as new; the
|
||||
manager line-shifts its baseline into post-edit coordinates before keying.
|
||||
"""
|
||||
rng = d.get("range") or {}
|
||||
start = rng.get("start") or {}
|
||||
|
||||
+51
-109
@@ -1,39 +1,20 @@
|
||||
"""Structured logging with steady-state silence for the LSP layer.
|
||||
|
||||
The LSP layer fires on every write_file/patch. In a busy session
|
||||
that's hundreds of events. We want users to be able to ``rg`` the
|
||||
log for "did LSP fire on that edit?" without drowning in noise.
|
||||
LSP fires on every write_file/patch, so the level model keeps ``agent.log``
|
||||
greppable (``rg 'lsp\\['``) without noise:
|
||||
|
||||
The level model:
|
||||
- ``DEBUG`` for steady-state events with no novel signal (clean, skipped,
|
||||
repeat "no project root", repeat "server unavailable").
|
||||
- ``INFO`` for once-per-session transitions (``active for <root>`` the first
|
||||
time a client starts, first ``no project root`` per file) and for every
|
||||
diagnostic event (rare and exactly what users grep for).
|
||||
- ``WARNING`` for action-required failures: first ``server unavailable`` per
|
||||
(server_id, binary), ``no server configured`` once per language, and every
|
||||
timeout / unexpected error.
|
||||
|
||||
- ``DEBUG`` for steady-state events that have no novel signal:
|
||||
``clean``, ``feature off``, ``extension not mapped``, ``no project
|
||||
root for already-announced file``, ``server unavailable for
|
||||
already-announced binary``. These never reach ``agent.log`` at the
|
||||
default INFO threshold.
|
||||
|
||||
- ``INFO`` for state transitions worth surfacing exactly once per
|
||||
session: ``active for <root>`` the first time a (server_id,
|
||||
workspace_root) client starts, ``no project root for <path>``
|
||||
the first time we see that file. Plus every diagnostic event
|
||||
(those are inherently rare and per-edit, exactly what users grep
|
||||
for).
|
||||
|
||||
- ``WARNING`` for action-required failures: ``server unavailable``
|
||||
(binary not on PATH) the first time per (server_id, binary),
|
||||
``no server configured`` once per language. Per-call WARNING for
|
||||
timeouts and unexpected bridge exceptions.
|
||||
|
||||
The dedup is in-process module-level sets. Each set grows at most by
|
||||
the number of distinct (server_id, root) and (server_id, binary)
|
||||
pairs touched in one Python process — bytes of memory in even an
|
||||
aggressive monorepo session. Bounded LRU was rejected: evicting an
|
||||
entry would risk re-firing the WARNING/INFO line we explicitly want
|
||||
to suppress.
|
||||
|
||||
Grep recipe::
|
||||
|
||||
tail -f ~/.hermes/logs/agent.log | rg 'lsp\\['
|
||||
Dedup uses module-level sets bounded by the distinct pairs touched in one
|
||||
process. A bounded LRU was rejected: evicting an entry would re-fire the
|
||||
line we explicitly want suppressed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -42,28 +23,19 @@ import os
|
||||
import threading
|
||||
from typing import List, Tuple
|
||||
|
||||
# Dedicated logger name so the documented grep recipe survives a
|
||||
# ``logging.getLogger(__name__)`` rename of any internal module.
|
||||
# Dedicated logger name so the documented grep recipe survives any
|
||||
# ``logging.getLogger(__name__)`` rename of internal modules.
|
||||
event_log = logging.getLogger("hermes.lint.lsp")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Once-per-X dedup sets
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_announce_lock = threading.Lock()
|
||||
_announced_active: set = set() # keys: (server_id, workspace_root)
|
||||
_announced_unavailable: set = set() # keys: (server_id, binary_path_or_name)
|
||||
_announced_no_root: set = set() # keys: (server_id, file_path)
|
||||
_announced_no_server: set = set() # keys: (server_id,)
|
||||
_ALL_BUCKETS = (_announced_active, _announced_unavailable, _announced_no_root)
|
||||
|
||||
|
||||
def _short_path(file_path: str) -> str:
|
||||
"""Render *file_path* relative to the cwd when sensible, else absolute.
|
||||
|
||||
Keeps log lines readable for the common case (the user is inside
|
||||
the project they're editing) without emitting brittle ``../../..``
|
||||
chains for the cross-tree case.
|
||||
"""
|
||||
"""Render *file_path* relative to cwd when it's inside it, else absolute (no ``../..`` chains)."""
|
||||
if not file_path:
|
||||
return file_path
|
||||
try:
|
||||
@@ -80,11 +52,7 @@ def _emit(server_id: str, level: int, message: str) -> None:
|
||||
|
||||
|
||||
def _announce_once(bucket: set, key: Tuple) -> bool:
|
||||
"""Return True if *key* has not been announced for *bucket* yet.
|
||||
|
||||
Atomically marks the key as announced so concurrent callers
|
||||
cannot both win the race and double-log.
|
||||
"""
|
||||
"""Atomically mark *key* announced; True only for the first caller."""
|
||||
with _announce_lock:
|
||||
if key in bucket:
|
||||
return False
|
||||
@@ -92,82 +60,61 @@ def _announce_once(bucket: set, key: Tuple) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _emit_once(bucket: set, key: Tuple, server_id: str, level: int, first: str, repeat: str) -> None:
|
||||
"""Log *first* at *level* the first time *key* is seen, *repeat* at DEBUG thereafter."""
|
||||
if _announce_once(bucket, key):
|
||||
_emit(server_id, level, first)
|
||||
else:
|
||||
_emit(server_id, logging.DEBUG, repeat)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public event helpers — call these from the LSP layer.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def log_clean(server_id: str, file_path: str) -> None:
|
||||
"""No diagnostics emitted for *file_path*. DEBUG (silent at default)."""
|
||||
"""No diagnostics emitted for *file_path*. DEBUG."""
|
||||
_emit(server_id, logging.DEBUG, f"clean ({_short_path(file_path)})")
|
||||
|
||||
|
||||
def log_disabled(server_id: str, file_path: str, reason: str) -> None:
|
||||
"""LSP intentionally skipped for this file (feature off, ext unmapped,
|
||||
backend not local, etc.). DEBUG."""
|
||||
"""LSP intentionally skipped for this file (feature off, ext unmapped, ...). DEBUG."""
|
||||
_emit(server_id, logging.DEBUG, f"skipped: {reason} ({_short_path(file_path)})")
|
||||
|
||||
|
||||
def log_active(server_id: str, workspace_root: str) -> None:
|
||||
"""A new LSP client started for (server_id, workspace_root).
|
||||
|
||||
INFO once per (server_id, workspace_root); DEBUG thereafter.
|
||||
Lets users verify "is LSP actually running?" with a single grep.
|
||||
"""
|
||||
key = (server_id, workspace_root)
|
||||
if _announce_once(_announced_active, key):
|
||||
_emit(server_id, logging.INFO, f"active for {workspace_root}")
|
||||
else:
|
||||
_emit(server_id, logging.DEBUG, f"reused client for {workspace_root}")
|
||||
"""A client started for (server_id, workspace_root). INFO once per pair, DEBUG thereafter."""
|
||||
_emit_once(
|
||||
_announced_active, (server_id, workspace_root), server_id, logging.INFO,
|
||||
f"active for {workspace_root}", f"reused client for {workspace_root}",
|
||||
)
|
||||
|
||||
|
||||
def log_diagnostics(server_id: str, file_path: str, count: int) -> None:
|
||||
"""Diagnostics arrived for a file. INFO every time — these are the
|
||||
failure signals users actually want to grep for, and they are
|
||||
inherently rare per edit."""
|
||||
"""Diagnostics arrived for a file. INFO every time — rare per edit and what users grep for."""
|
||||
_emit(server_id, logging.INFO, f"{count} diags ({_short_path(file_path)})")
|
||||
|
||||
|
||||
def log_no_project_root(server_id: str, file_path: str) -> None:
|
||||
"""File had no recognised project marker. INFO once per file,
|
||||
DEBUG thereafter."""
|
||||
key = (server_id, file_path)
|
||||
if _announce_once(_announced_no_root, key):
|
||||
_emit(server_id, logging.INFO, f"no project root for {_short_path(file_path)}")
|
||||
else:
|
||||
_emit(server_id, logging.DEBUG, f"no project root for {_short_path(file_path)}")
|
||||
"""File had no recognised project marker. INFO once per file, DEBUG thereafter."""
|
||||
msg = f"no project root for {_short_path(file_path)}"
|
||||
_emit_once(_announced_no_root, (server_id, file_path), server_id, logging.INFO, msg, msg)
|
||||
|
||||
|
||||
def log_server_unavailable(server_id: str, binary_or_pkg: str) -> None:
|
||||
"""The server binary couldn't be resolved. WARNING once per
|
||||
(server_id, binary), DEBUG thereafter so a hundred subsequent
|
||||
.py edits don't spam the log."""
|
||||
key = (server_id, binary_or_pkg)
|
||||
if _announce_once(_announced_unavailable, key):
|
||||
_emit(
|
||||
server_id,
|
||||
logging.WARNING,
|
||||
f"server unavailable: {binary_or_pkg} not found "
|
||||
"(install via `hermes lsp install <id>` or set lsp.servers.<id>.command)",
|
||||
)
|
||||
else:
|
||||
_emit(server_id, logging.DEBUG, f"server still unavailable: {binary_or_pkg}")
|
||||
|
||||
|
||||
def log_no_server_configured(server_id: str) -> None:
|
||||
"""No spawn recipe for this language. WARNING once."""
|
||||
if _announce_once(_announced_no_server, (server_id,)):
|
||||
_emit(server_id, logging.WARNING, "no server configured")
|
||||
"""Server binary unresolved. WARNING once per (server_id, binary), DEBUG thereafter."""
|
||||
_emit_once(
|
||||
_announced_unavailable, (server_id, binary_or_pkg), server_id, logging.WARNING,
|
||||
f"server unavailable: {binary_or_pkg} not found "
|
||||
"(install via `hermes lsp install <id>` or set lsp.servers.<id>.command)",
|
||||
f"server still unavailable: {binary_or_pkg}",
|
||||
)
|
||||
|
||||
|
||||
def log_timeout(server_id: str, file_path: str, kind: str = "diagnostics") -> None:
|
||||
"""A request to the server timed out. WARNING every time — these are
|
||||
inherently novel events worth surfacing on each occurrence."""
|
||||
_emit(
|
||||
server_id,
|
||||
logging.WARNING,
|
||||
f"{kind} timed out for {_short_path(file_path)}",
|
||||
)
|
||||
"""A request to the server timed out. WARNING every time."""
|
||||
_emit(server_id, logging.WARNING, f"{kind} timed out for {_short_path(file_path)}")
|
||||
|
||||
|
||||
def log_server_error(server_id: str, file_path: str, exc: BaseException) -> None:
|
||||
@@ -189,12 +136,10 @@ def log_spawn_failed(server_id: str, workspace_root: str, exc: BaseException) ->
|
||||
|
||||
|
||||
def log_reaped(keys: List[Tuple[str, str]], idle_timeout: float) -> None:
|
||||
"""Idle clients were shut down by the reaper. INFO — one line per
|
||||
sweep so users can correlate memory drops with LSP activity.
|
||||
"""Idle clients were reaped. INFO, one line per sweep.
|
||||
|
||||
Also clears the ``log_active`` announce cache for the reaped keys so
|
||||
a later respawn re-announces at INFO instead of logging a misleading
|
||||
DEBUG "reused client".
|
||||
Also forgets the ``log_active`` announcement for those keys so a respawn
|
||||
re-announces at INFO instead of a misleading DEBUG "reused client".
|
||||
"""
|
||||
with _announce_lock:
|
||||
for key in keys:
|
||||
@@ -210,10 +155,8 @@ def log_reaped(keys: List[Tuple[str, str]], idle_timeout: float) -> None:
|
||||
def reset_announce_caches() -> None:
|
||||
"""Test-only: clear the dedup caches. Production code never calls this."""
|
||||
with _announce_lock:
|
||||
_announced_active.clear()
|
||||
_announced_unavailable.clear()
|
||||
_announced_no_root.clear()
|
||||
_announced_no_server.clear()
|
||||
for bucket in _ALL_BUCKETS:
|
||||
bucket.clear()
|
||||
|
||||
|
||||
__all__ = [
|
||||
@@ -224,7 +167,6 @@ __all__ = [
|
||||
"log_diagnostics",
|
||||
"log_no_project_root",
|
||||
"log_server_unavailable",
|
||||
"log_no_server_configured",
|
||||
"log_timeout",
|
||||
"log_server_error",
|
||||
"log_spawn_failed",
|
||||
|
||||
+100
-199
@@ -1,28 +1,14 @@
|
||||
"""Auto-installation of LSP server binaries.
|
||||
|
||||
Tries to install missing servers using whatever package manager is
|
||||
appropriate. All installs go to a Hermes-owned bin staging dir,
|
||||
``<HERMES_HOME>/lsp/bin/``, so we don't pollute the user's global
|
||||
toolchain.
|
||||
Installs go to a Hermes-owned staging dir, ``<HERMES_HOME>/lsp/bin/``, so the
|
||||
user's global toolchain stays untouched. Strategies: ``auto`` (install with
|
||||
the best available package manager), ``manual`` / ``off`` (probe only; a
|
||||
missing binary skips the server and ``hermes lsp status`` reports it).
|
||||
|
||||
Strategies:
|
||||
|
||||
- ``auto`` — attempt to install with the best available package
|
||||
manager. This is the default.
|
||||
- ``manual`` — never install; if a binary is missing, the server is
|
||||
silently skipped and the user is told about it via ``hermes lsp
|
||||
status``.
|
||||
- ``off`` — same as ``manual`` for now (kept distinct so we can
|
||||
evolve behavior later, e.g. logging differently).
|
||||
|
||||
The actual installs happen synchronously the first time a server is
|
||||
needed and concurrent calls to :func:`try_install` for the same
|
||||
package are deduplicated via a per-package lock.
|
||||
|
||||
Failure modes are non-fatal: every install path is wrapped in
|
||||
try/except and returns ``None`` on failure. The tool layer then
|
||||
falls back to its in-process syntax checker, exactly as if the user
|
||||
hadn't enabled LSP at all.
|
||||
Installs run synchronously the first time a server is needed; concurrent
|
||||
:func:`try_install` calls for the same package are serialized per-package.
|
||||
Every failure path returns ``None`` so the tool layer falls back to its
|
||||
in-process syntax checker.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -39,76 +25,35 @@ from hermes_constants import find_node_executable
|
||||
|
||||
logger = logging.getLogger("agent.lsp.install")
|
||||
|
||||
# Package-name → install-strategy hint registry. Each entry is a
|
||||
# tuple of strategy name + package name + executable name. When the
|
||||
# install completes, we look for the executable in
|
||||
# ``<HERMES_HOME>/lsp/bin/`` first, then on PATH.
|
||||
#
|
||||
# Optional fields:
|
||||
# - ``extra_pkgs``: list of sibling packages to install alongside
|
||||
# ``pkg`` in the same node_modules tree. Used when an LSP server
|
||||
# has a runtime peer dependency that npm doesn't auto-pull (e.g.
|
||||
# typescript-language-server needs ``typescript``).
|
||||
|
||||
def _recipe(strategy: str, pkg: str, bin_name: str, **extra: Any) -> Dict[str, Any]:
|
||||
return {"strategy": strategy, "pkg": pkg, "bin": bin_name, **extra}
|
||||
|
||||
|
||||
# Recipe key → {strategy, pkg, bin[, extra_pkgs]}. After install we look for
|
||||
# ``bin`` in ``<HERMES_HOME>/lsp/bin/`` first, then on PATH. ``extra_pkgs``
|
||||
# are sibling npm packages a server needs in the same node_modules tree.
|
||||
INSTALL_RECIPES: Dict[str, Dict[str, Any]] = {
|
||||
# Python
|
||||
"pyright": {"strategy": "npm", "pkg": "pyright", "bin": "pyright-langserver"},
|
||||
# JS/TS family
|
||||
"typescript-language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "typescript-language-server",
|
||||
"bin": "typescript-language-server",
|
||||
# typescript-language-server requires the `typescript` SDK
|
||||
# (tsserver) to be importable from the same node_modules tree;
|
||||
# otherwise initialize() fails with "Could not find a valid
|
||||
# TypeScript installation". Install them together.
|
||||
"extra_pkgs": ["typescript"],
|
||||
},
|
||||
"@vue/language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "@vue/language-server",
|
||||
"bin": "vue-language-server",
|
||||
},
|
||||
"svelte-language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "svelte-language-server",
|
||||
"bin": "svelteserver",
|
||||
},
|
||||
"@astrojs/language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "@astrojs/language-server",
|
||||
"bin": "astro-ls",
|
||||
},
|
||||
"yaml-language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "yaml-language-server",
|
||||
"bin": "yaml-language-server",
|
||||
},
|
||||
"bash-language-server": {
|
||||
"strategy": "npm",
|
||||
"pkg": "bash-language-server",
|
||||
"bin": "bash-language-server",
|
||||
},
|
||||
"intelephense": {"strategy": "npm", "pkg": "intelephense", "bin": "intelephense"},
|
||||
"dockerfile-language-server-nodejs": {
|
||||
"strategy": "npm",
|
||||
"pkg": "dockerfile-language-server-nodejs",
|
||||
"bin": "docker-langserver",
|
||||
},
|
||||
# Go
|
||||
"gopls": {"strategy": "go", "pkg": "golang.org/x/tools/gopls@latest", "bin": "gopls"},
|
||||
# Rust — too heavy (hundreds of MB to bootstrap). We do NOT
|
||||
# auto-install rust-analyzer; users install via rustup.
|
||||
"rust-analyzer": {"strategy": "manual", "pkg": "", "bin": "rust-analyzer"},
|
||||
# C/C++ — manual (clangd ships with LLVM, very heavy)
|
||||
"clangd": {"strategy": "manual", "pkg": "", "bin": "clangd"},
|
||||
# Lua — manual (LuaLS is platform-specific binaries from GitHub
|
||||
# releases; complex enough that we punt to the user)
|
||||
"lua-language-server": {"strategy": "manual", "pkg": "", "bin": "lua-language-server"},
|
||||
# PowerShell — PowerShellEditorServices ships as a GitHub release
|
||||
# zip driven by a pwsh bootstrap script, not a single binary. We
|
||||
# require a manual bundle install and probe for the pwsh host so
|
||||
# `hermes lsp status` reports the host's presence.
|
||||
"powershell": {"strategy": "manual", "pkg": "", "bin": "pwsh"},
|
||||
"pyright": _recipe("npm", "pyright", "pyright-langserver"),
|
||||
# tsserver must be importable from the same node_modules tree or
|
||||
# initialize() fails with "Could not find a valid TypeScript installation".
|
||||
"typescript-language-server": _recipe("npm", "typescript-language-server", "typescript-language-server", extra_pkgs=["typescript"]),
|
||||
"@vue/language-server": _recipe("npm", "@vue/language-server", "vue-language-server"),
|
||||
"svelte-language-server": _recipe("npm", "svelte-language-server", "svelteserver"),
|
||||
"@astrojs/language-server": _recipe("npm", "@astrojs/language-server", "astro-ls"),
|
||||
"yaml-language-server": _recipe("npm", "yaml-language-server", "yaml-language-server"),
|
||||
"bash-language-server": _recipe("npm", "bash-language-server", "bash-language-server"),
|
||||
"intelephense": _recipe("npm", "intelephense", "intelephense"),
|
||||
"dockerfile-language-server-nodejs": _recipe("npm", "dockerfile-language-server-nodejs", "docker-langserver"),
|
||||
"gopls": _recipe("go", "golang.org/x/tools/gopls@latest", "gopls"),
|
||||
# Manual: rust-analyzer (via rustup) and clangd (ships with LLVM) are far too
|
||||
# heavy to bootstrap; LuaLS is platform-specific GitHub release binaries.
|
||||
"rust-analyzer": _recipe("manual", "", "rust-analyzer"),
|
||||
"clangd": _recipe("manual", "", "clangd"),
|
||||
"lua-language-server": _recipe("manual", "", "lua-language-server"),
|
||||
# PowerShellEditorServices is a release-zip bundle driven by pwsh; we probe
|
||||
# the host so `hermes lsp status` reports its presence.
|
||||
"powershell": _recipe("manual", "", "pwsh"),
|
||||
}
|
||||
|
||||
|
||||
@@ -171,29 +116,20 @@ def _get_lock(pkg: str) -> threading.Lock:
|
||||
|
||||
|
||||
def try_install(pkg: str, strategy: str = "auto") -> Optional[str]:
|
||||
"""Try to install ``pkg`` and return the binary path if successful.
|
||||
"""Try to install ``pkg``; return the binary path or ``None``.
|
||||
|
||||
``strategy`` is ``"auto"``, ``"manual"``, or ``"off"``. In
|
||||
``manual``/``off`` mode, this function only probes for an
|
||||
existing binary and returns ``None`` if not found.
|
||||
|
||||
The install is cached per-package — a second call returns the
|
||||
same path (or ``None``) without reinstalling. Concurrent calls
|
||||
Only ``"auto"`` installs; ``"manual"``/``"off"`` just probe for an
|
||||
existing binary. Results are cached per package and concurrent calls
|
||||
are serialized.
|
||||
"""
|
||||
if strategy not in {"auto",}:
|
||||
# Only ``auto`` triggers an actual install. In manual/off,
|
||||
# we still check whether the binary already exists.
|
||||
recipe = INSTALL_RECIPES.get(pkg, {})
|
||||
bin_name = recipe.get("bin", pkg)
|
||||
return _existing_binary(bin_name)
|
||||
return _existing_binary(recipe.get("bin", pkg))
|
||||
|
||||
if pkg in _install_results:
|
||||
return _install_results[pkg]
|
||||
|
||||
lock = _get_lock(pkg)
|
||||
with lock:
|
||||
# Double-check after acquiring lock.
|
||||
with _get_lock(pkg):
|
||||
if pkg in _install_results:
|
||||
return _install_results[pkg]
|
||||
result = _do_install(pkg)
|
||||
@@ -210,7 +146,6 @@ def _do_install(pkg: str) -> Optional[str]:
|
||||
strategy = recipe.get("strategy", "manual")
|
||||
bin_name = recipe.get("bin", pkg)
|
||||
|
||||
# Check if already present (shutil.which or staging dir)
|
||||
existing = _existing_binary(bin_name)
|
||||
if existing:
|
||||
return existing
|
||||
@@ -234,71 +169,75 @@ def _do_install(pkg: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
|
||||
def _run_installer(tool: str, pkg: str, cmd: list, *, timeout: int, env: Optional[dict] = None) -> bool:
|
||||
"""Run one install subprocess; log and return False on non-zero exit or error."""
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
cmd,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True, encoding="utf-8", errors="replace",
|
||||
timeout=timeout,
|
||||
env=env,
|
||||
stdin=subprocess.DEVNULL,
|
||||
creationflags=windows_hide_flags(),
|
||||
)
|
||||
if proc.returncode != 0:
|
||||
logger.warning(
|
||||
"[install] %s install failed for %s: %s", tool, pkg, proc.stderr.strip()[:500]
|
||||
)
|
||||
return False
|
||||
except (subprocess.TimeoutExpired, OSError) as e:
|
||||
logger.warning("[install] %s install errored for %s: %s", tool, pkg, e)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _link_into_bin(target: Path) -> str:
|
||||
"""Symlink (or copy, where symlinks fail) ``target`` into ``lsp/bin/`` and return the path to use."""
|
||||
link = hermes_lsp_bin_dir() / target.name
|
||||
if not link.exists():
|
||||
try:
|
||||
link.symlink_to(target)
|
||||
except (OSError, NotImplementedError):
|
||||
# Symlinks fail on some Windows setups — copy instead.
|
||||
try:
|
||||
shutil.copy2(target, link)
|
||||
except OSError:
|
||||
return str(target)
|
||||
return str(link if link.exists() else target)
|
||||
|
||||
|
||||
def _install_npm(
|
||||
pkg: str,
|
||||
bin_name: str,
|
||||
extra_pkgs: Optional[list] = None,
|
||||
) -> Optional[str]:
|
||||
"""Install an npm package into our staging dir.
|
||||
|
||||
Uses ``npm install --prefix`` so the binaries land in
|
||||
``<staging>/node_modules/.bin/<bin_name>`` and we symlink them up
|
||||
one level for direct PATH-style access.
|
||||
|
||||
``extra_pkgs`` is a list of sibling packages to install in the
|
||||
same ``node_modules`` tree. Used for LSP servers with runtime
|
||||
peer deps that npm doesn't auto-pull (typescript-language-server
|
||||
needs ``typescript`` next to it; intelephense ships standalone).
|
||||
"""
|
||||
# Managed npm first: $HERMES_HOME/node is not on an arbitrary process's
|
||||
# PATH, so a bare which() misses the Node that Hermes installed and
|
||||
# reports "npm not on PATH" on a machine that has a perfectly good one.
|
||||
"""``npm install --prefix <staging>`` then link ``node_modules/.bin/<bin_name>`` into ``lsp/bin/``."""
|
||||
# Managed npm first: $HERMES_HOME/node isn't on an arbitrary process's
|
||||
# PATH, so a bare which() would miss the Node that Hermes installed.
|
||||
npm = find_node_executable("npm")
|
||||
if npm is None:
|
||||
logger.info("[install] cannot install %s: no usable npm found", pkg)
|
||||
return None
|
||||
staging = hermes_lsp_bin_dir().parent # <HERMES_HOME>/lsp/
|
||||
install_targets = [pkg] + list(extra_pkgs or [])
|
||||
try:
|
||||
logger.info(
|
||||
"[install] npm install --prefix %s %s",
|
||||
staging,
|
||||
" ".join(install_targets),
|
||||
)
|
||||
proc = subprocess.run(
|
||||
[npm, "install", "--prefix", str(staging), "--silent", "--no-fund", "--no-audit", *install_targets],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True, encoding="utf-8", errors="replace",
|
||||
timeout=300,
|
||||
stdin=subprocess.DEVNULL,
|
||||
creationflags=windows_hide_flags(),
|
||||
)
|
||||
if proc.returncode != 0:
|
||||
logger.warning(
|
||||
"[install] npm install failed for %s: %s", pkg, proc.stderr.strip()[:500]
|
||||
)
|
||||
return None
|
||||
except (subprocess.TimeoutExpired, OSError) as e:
|
||||
logger.warning("[install] npm install errored for %s: %s", pkg, e)
|
||||
logger.info(
|
||||
"[install] npm install --prefix %s %s",
|
||||
staging,
|
||||
" ".join(install_targets),
|
||||
)
|
||||
if not _run_installer(
|
||||
"npm", pkg,
|
||||
[npm, "install", "--prefix", str(staging), "--silent", "--no-fund", "--no-audit", *install_targets],
|
||||
timeout=300,
|
||||
):
|
||||
return None
|
||||
|
||||
# Find the bin
|
||||
nm_bin = staging / "node_modules" / ".bin" / bin_name
|
||||
for c in _native_binary_candidates(nm_bin):
|
||||
if c.exists():
|
||||
# Symlink into our `lsp/bin/` for stable PATH access.
|
||||
link = hermes_lsp_bin_dir() / c.name
|
||||
if not link.exists():
|
||||
try:
|
||||
link.symlink_to(c)
|
||||
except (OSError, NotImplementedError):
|
||||
# Symlinks fail on some Windows setups — copy instead.
|
||||
try:
|
||||
shutil.copy2(c, link)
|
||||
except OSError:
|
||||
return str(c)
|
||||
return str(link if link.exists() else c)
|
||||
return _link_into_bin(c)
|
||||
logger.warning("[install] npm install for %s succeeded but bin %s not found", pkg, bin_name)
|
||||
return None
|
||||
|
||||
@@ -312,25 +251,8 @@ def _install_go(pkg: str, bin_name: str) -> Optional[str]:
|
||||
staging = hermes_lsp_bin_dir()
|
||||
env = dict(os.environ)
|
||||
env["GOBIN"] = str(staging)
|
||||
try:
|
||||
logger.info("[install] go install %s (GOBIN=%s)", pkg, staging)
|
||||
proc = subprocess.run(
|
||||
[go, "install", pkg],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True, encoding="utf-8", errors="replace",
|
||||
timeout=600,
|
||||
env=env,
|
||||
stdin=subprocess.DEVNULL,
|
||||
creationflags=windows_hide_flags(),
|
||||
)
|
||||
if proc.returncode != 0:
|
||||
logger.warning(
|
||||
"[install] go install failed for %s: %s", pkg, proc.stderr.strip()[:500]
|
||||
)
|
||||
return None
|
||||
except (subprocess.TimeoutExpired, OSError) as e:
|
||||
logger.warning("[install] go install errored for %s: %s", pkg, e)
|
||||
logger.info("[install] go install %s (GOBIN=%s)", pkg, staging)
|
||||
if not _run_installer("go", pkg, [go, "install", pkg], timeout=600, env=env):
|
||||
return None
|
||||
bin_path = staging / bin_name
|
||||
if _is_windows():
|
||||
@@ -342,14 +264,7 @@ def _install_go(pkg: str, bin_name: str) -> Optional[str]:
|
||||
|
||||
|
||||
def _install_pip(pkg: str, bin_name: str) -> Optional[str]:
|
||||
"""Install a Python package into a hermes-owned target dir.
|
||||
|
||||
We avoid polluting the user's site-packages by using
|
||||
``pip install --target``. Bins go into
|
||||
``<staging>/python-packages/bin/`` which we symlink into
|
||||
``<staging>/bin``. Note: this only works for packages that ship a
|
||||
console script.
|
||||
"""
|
||||
"""``pip install --target <staging>/python-packages`` then link the console script into ``lsp/bin/``."""
|
||||
pip_target = hermes_lsp_bin_dir().parent / "python-packages"
|
||||
pip_target.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
@@ -368,33 +283,19 @@ def _install_pip(pkg: str, bin_name: str) -> Optional[str]:
|
||||
except (subprocess.TimeoutExpired, OSError) as e:
|
||||
logger.warning("[install] pip install errored for %s: %s", pkg, e)
|
||||
return None
|
||||
# Look for the console script. POSIX wheels generally write to bin/,
|
||||
# while native Windows installs use Scripts/.
|
||||
# POSIX wheels write console scripts to bin/, native Windows to Scripts/.
|
||||
script_dirs = [pip_target / "bin"]
|
||||
if _is_windows():
|
||||
script_dirs.append(pip_target / "Scripts")
|
||||
for script_dir in script_dirs:
|
||||
for bin_path in _native_binary_candidates(script_dir / bin_name):
|
||||
if bin_path.exists():
|
||||
link = hermes_lsp_bin_dir() / bin_path.name
|
||||
if not link.exists():
|
||||
try:
|
||||
link.symlink_to(bin_path)
|
||||
except (OSError, NotImplementedError):
|
||||
try:
|
||||
shutil.copy2(bin_path, link)
|
||||
except OSError:
|
||||
return str(bin_path)
|
||||
return str(link if link.exists() else bin_path)
|
||||
return _link_into_bin(bin_path)
|
||||
return None
|
||||
|
||||
|
||||
def detect_status(pkg: str) -> str:
|
||||
"""Return ``installed``, ``missing``, or ``manual-only`` for a package.
|
||||
|
||||
Used by the ``hermes lsp status`` CLI to give users a quick
|
||||
overview of what's available without spawning anything.
|
||||
"""
|
||||
"""Return ``installed``, ``missing``, or ``manual-only`` (for ``hermes lsp status``; spawns nothing)."""
|
||||
recipe = INSTALL_RECIPES.get(pkg)
|
||||
bin_name = recipe.get("bin", pkg) if recipe else pkg
|
||||
if _existing_binary(bin_name):
|
||||
|
||||
+98
-239
@@ -1,36 +1,19 @@
|
||||
"""Service-level orchestration for LSP clients.
|
||||
|
||||
The :class:`LSPService` is the bridge between the synchronous
|
||||
file_operations layer and the async :class:`agent.lsp.client.LSPClient`.
|
||||
:class:`LSPService` bridges the synchronous file_operations layer and the
|
||||
async :class:`agent.lsp.client.LSPClient`:
|
||||
|
||||
Design choices:
|
||||
- One asyncio loop in a background thread; :meth:`get_diagnostics_sync`
|
||||
opens + waits + drains in one blocking call.
|
||||
- One lazily spawned client per ``(server_id, workspace_root)``.
|
||||
- A **broken-set** of pairs that failed to spawn/initialize — never retried
|
||||
for the life of the service.
|
||||
- A **delta baseline** per file: ``snapshot_baseline()`` runs BEFORE a write,
|
||||
and the next ``get_diagnostics_sync()`` returns only diagnostics not in it.
|
||||
|
||||
- A **single asyncio event loop** runs in a background thread. All
|
||||
client work happens on that loop. Synchronous callers from
|
||||
``tools/file_operations.py`` use :meth:`get_diagnostics_sync` to
|
||||
open + wait + drain in one blocking call.
|
||||
|
||||
- One client per ``(server_id, workspace_root)`` key. Lazy spawn:
|
||||
the first request for a key spawns the client; subsequent requests
|
||||
re-use it.
|
||||
|
||||
- A **broken-set** records ``(server_id, workspace_root)`` pairs that
|
||||
failed to spawn or initialize. These are never retried for the
|
||||
life of the service. Mirrors OpenCode's design.
|
||||
|
||||
- A **delta baseline** map keeps "diagnostics-as-of-the-last-snapshot"
|
||||
per file. ``snapshot_baseline()`` is called BEFORE a write; the
|
||||
next ``get_diagnostics_sync()`` returns only diagnostics that
|
||||
weren't in the baseline. This is the lift from Claude Code's
|
||||
``beforeFileEdited`` / ``getNewDiagnostics`` pattern, except wired
|
||||
to the local LSP layer instead of MCP IDE RPC.
|
||||
|
||||
The service is **off by default** — call :meth:`is_active` to check
|
||||
whether it's actually doing anything. When LSP is disabled in
|
||||
config, when no git workspace can be detected, when all configured
|
||||
servers are missing binaries and auto-install is off, ``is_active``
|
||||
returns False and the file_operations layer falls through to the
|
||||
in-process syntax check.
|
||||
The service is off unless config enables it; :meth:`is_active` says whether
|
||||
it does anything, and file_operations falls back to the in-process syntax
|
||||
check otherwise.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -45,9 +28,11 @@ from agent.lsp import eventlog
|
||||
from agent.lsp.client import (
|
||||
DIAGNOSTICS_DOCUMENT_WAIT,
|
||||
LSPClient,
|
||||
_diagnostic_key as _diag_key,
|
||||
)
|
||||
from agent.lsp.servers import (
|
||||
ServerContext,
|
||||
ServerDef,
|
||||
find_server_for_file,
|
||||
language_id_for,
|
||||
)
|
||||
@@ -63,11 +48,7 @@ MIN_IDLE_TIMEOUT = 30 # floor for config values; must exceed any per-op wait bu
|
||||
|
||||
|
||||
class _BackgroundLoop:
|
||||
"""A daemon thread that owns one asyncio event loop.
|
||||
|
||||
Provides :meth:`run` for synchronous callers — submits a coroutine
|
||||
to the loop and blocks until it finishes (or a timeout fires).
|
||||
"""
|
||||
"""A daemon thread owning one asyncio loop; :meth:`run` blocks on a coroutine."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
@@ -99,10 +80,7 @@ class _BackgroundLoop:
|
||||
pass
|
||||
|
||||
def run(self, coro, *, timeout: Optional[float] = None) -> Any:
|
||||
"""Submit a coroutine to the loop and block until done.
|
||||
|
||||
Returns the coroutine's result, or raises its exception.
|
||||
"""
|
||||
"""Submit a coroutine to the loop and block for its result (or raise)."""
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
if self._loop is None:
|
||||
if asyncio.iscoroutine(coro):
|
||||
@@ -132,17 +110,7 @@ class _BackgroundLoop:
|
||||
|
||||
|
||||
class LSPService:
|
||||
"""The process-wide LSP service.
|
||||
|
||||
Created once via :meth:`create_from_config`; the
|
||||
:func:`agent.lsp.get_service` accessor manages the singleton.
|
||||
Most callers should use that accessor rather than constructing
|
||||
:class:`LSPService` directly.
|
||||
"""
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# construction + factory
|
||||
# ------------------------------------------------------------------
|
||||
"""The process-wide LSP service; use :func:`agent.lsp.get_service` rather than constructing directly."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -179,10 +147,7 @@ class LSPService:
|
||||
self._state_lock = threading.Lock()
|
||||
self._idle_reaper_task: Optional[asyncio.Task] = None
|
||||
|
||||
# Delta baseline: file path → snapshot of diagnostics taken
|
||||
# immediately before a write. ``get_diagnostics_sync`` filters
|
||||
# out anything in the baseline so the agent only sees errors
|
||||
# introduced by the current edit.
|
||||
# file path → diagnostics snapshot taken immediately before a write.
|
||||
self._delta_baseline: Dict[str, List[Dict[str, Any]]] = {}
|
||||
|
||||
if self._enabled and self._idle_timeout > 0:
|
||||
@@ -190,11 +155,7 @@ class LSPService:
|
||||
|
||||
@classmethod
|
||||
def create_from_config(cls) -> Optional["LSPService"]:
|
||||
"""Build a service from ``hermes_cli.config`` settings.
|
||||
|
||||
Returns ``None`` if the config can't be loaded. The service
|
||||
itself returns ``is_active()`` False when LSP is disabled.
|
||||
"""
|
||||
"""Build a service from ``hermes_cli.config``; ``None`` if config can't load."""
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
cfg = load_config_readonly()
|
||||
@@ -215,10 +176,9 @@ class LSPService:
|
||||
except (TypeError, ValueError):
|
||||
idle_timeout = DEFAULT_IDLE_TIMEOUT
|
||||
if 0 < idle_timeout < MIN_IDLE_TIMEOUT:
|
||||
# A timeout below the per-operation wait budget could reap a
|
||||
# client mid-flight; the resulting outer timeout would then
|
||||
# mark the (server, workspace) pair broken for the process
|
||||
# lifetime. Clamp to a safe floor (0 still disables).
|
||||
# Below the per-op wait budget the reaper could kill a client
|
||||
# mid-flight and the outer timeout would then mark the pair broken
|
||||
# for the process lifetime. Clamp (0 still disables).
|
||||
idle_timeout = MIN_IDLE_TIMEOUT
|
||||
servers_cfg = lsp_cfg.get("servers") or {}
|
||||
disabled = []
|
||||
@@ -261,55 +221,48 @@ class LSPService:
|
||||
"""Return True iff this service should be consulted at all."""
|
||||
return self._enabled
|
||||
|
||||
def _broken_key(self, srv: ServerDef, file_path: str) -> Optional[Tuple[str, str]]:
|
||||
"""``(server_id, per-server root)`` broken-set key, or ``None`` when the file isn't gated in.
|
||||
|
||||
Falls back to the workspace root when the per-server resolver fails —
|
||||
the same key ``_get_or_spawn`` would have used when it failed.
|
||||
"""
|
||||
ws_root, gated = resolve_workspace_for_file(file_path)
|
||||
if not (ws_root and gated):
|
||||
return None
|
||||
try:
|
||||
per_server_root = srv.resolve_root(file_path, ws_root) or ws_root
|
||||
except Exception: # noqa: BLE001
|
||||
per_server_root = ws_root
|
||||
return (srv.server_id, per_server_root)
|
||||
|
||||
def enabled_for(self, file_path: str) -> bool:
|
||||
"""Return True iff LSP should run for this specific file.
|
||||
"""Return True iff LSP should run for this file.
|
||||
|
||||
Gates on workspace detection (file or cwd inside a git worktree),
|
||||
on whether any registered server matches the extension, and
|
||||
on whether the (server_id, workspace_root) pair is in the
|
||||
broken-set from a previous spawn failure.
|
||||
|
||||
Files in already-broken pairs return False so the file_operations
|
||||
layer skips the LSP path entirely — no spawn attempts, no
|
||||
timeout cost — until the service is restarted (``hermes lsp
|
||||
restart``) or the process exits.
|
||||
Gates on a registered, non-disabled server for the extension, on
|
||||
git-workspace detection, and on the pair not being in the broken-set
|
||||
(so a failed server costs no spawn attempts or timeouts until
|
||||
``hermes lsp restart`` or process exit).
|
||||
"""
|
||||
if not self._enabled:
|
||||
return False
|
||||
srv = find_server_for_file(file_path)
|
||||
if srv is None or srv.server_id in self._disabled_servers:
|
||||
return False
|
||||
ws_root, gated_in = resolve_workspace_for_file(file_path)
|
||||
if not (ws_root and gated_in):
|
||||
return False
|
||||
# Broken-set short-circuit. Use the per-server root if we can
|
||||
# compute one cheaply; otherwise fall back to the workspace
|
||||
# root as the broken key (which is what _get_or_spawn would
|
||||
# have used anyway when it failed).
|
||||
try:
|
||||
per_server_root = srv.resolve_root(file_path, ws_root) or ws_root
|
||||
except Exception: # noqa: BLE001
|
||||
per_server_root = ws_root
|
||||
if (srv.server_id, per_server_root) in self._broken:
|
||||
return False
|
||||
return True
|
||||
key = self._broken_key(srv, file_path)
|
||||
return key is not None and key not in self._broken
|
||||
|
||||
def snapshot_baseline(self, file_path: str) -> None:
|
||||
"""Snapshot current diagnostics for ``file_path`` as the delta baseline.
|
||||
"""Snapshot current diagnostics for ``file_path`` as the delta baseline (call BEFORE a write).
|
||||
|
||||
Called BEFORE a write so the next ``get_diagnostics_sync()``
|
||||
can filter out pre-existing errors. Best-effort — failures
|
||||
are silently swallowed so a flaky server can't break a write.
|
||||
|
||||
Outer timeouts (e.g. server hangs during initialize) mark the
|
||||
(server_id, workspace_root) pair as broken so subsequent edits
|
||||
skip it instantly instead of re-paying the timeout cost.
|
||||
Best-effort: failures are swallowed so a flaky server can't break a
|
||||
write, but outer timeouts mark the pair broken so later edits skip it.
|
||||
"""
|
||||
if not self.enabled_for(file_path):
|
||||
return
|
||||
try:
|
||||
# Outer join budget must exceed the inner wait budget or a
|
||||
# slow-but-alive server gets falsely marked broken.
|
||||
# Outer budget must exceed the inner wait or a slow-but-alive
|
||||
# server gets falsely marked broken.
|
||||
t = max(8.0, self._wait_timeout + 3.0)
|
||||
diags = self._loop.run(self._snapshot_async(file_path), timeout=t)
|
||||
self._delta_baseline[os.path.abspath(file_path)] = diags or []
|
||||
@@ -326,35 +279,21 @@ class LSPService:
|
||||
timeout: Optional[float] = None,
|
||||
line_shift: Optional[Callable[[int], Optional[int]]] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Synchronously open ``file_path`` in the right server, wait for
|
||||
diagnostics, return them.
|
||||
"""Synchronously open ``file_path``, wait for diagnostics, return them.
|
||||
|
||||
If ``delta`` is True (default), the result is filtered against
|
||||
any baseline previously captured via :meth:`snapshot_baseline`.
|
||||
Diagnostics present in the baseline are removed so the caller
|
||||
only sees errors introduced by the current edit.
|
||||
With ``delta`` (default) the result excludes anything in the baseline
|
||||
from :meth:`snapshot_baseline`. ``line_shift`` (built by
|
||||
:func:`agent.lsp.range_shift.build_line_shift`) remaps the baseline
|
||||
into post-edit coordinates first, so pre-existing diagnostics that
|
||||
merely moved don't look introduced by this edit.
|
||||
|
||||
When ``line_shift`` is provided, baseline diagnostics are
|
||||
remapped through it before the set-difference. This handles
|
||||
the case where the edit deleted or inserted lines, causing
|
||||
pre-existing diagnostics below the edit point to surface at
|
||||
different line numbers in the post-edit snapshot — without
|
||||
the shift, they'd all look "introduced by this edit". Pass
|
||||
a callable built by
|
||||
:func:`agent.lsp.range_shift.build_line_shift` (pre_text,
|
||||
post_text). Omit when pre/post content isn't available;
|
||||
the unshifted comparison still catches diagnostics that
|
||||
didn't move.
|
||||
|
||||
Returns an empty list when LSP is disabled, when no workspace
|
||||
can be detected, when no server matches, or when the server
|
||||
can't be spawned. Never raises.
|
||||
Returns ``[]`` when LSP is disabled, no workspace/server matches, or
|
||||
the server can't be spawned. Never raises.
|
||||
"""
|
||||
if not self.enabled_for(file_path):
|
||||
return []
|
||||
|
||||
# Resolve server_id eagerly so we can emit structured logs even
|
||||
# when the request errors out below.
|
||||
# Resolve server_id eagerly for structured logs on the error paths.
|
||||
srv = find_server_for_file(file_path)
|
||||
server_id = srv.server_id if srv else "?"
|
||||
|
||||
@@ -373,13 +312,10 @@ class LSPService:
|
||||
return []
|
||||
|
||||
if diags is None:
|
||||
# The server is alive but never produced diagnostics for the
|
||||
# post-edit content within the wait budget (common for
|
||||
# tsserver on large projects). Report "no data" rather than
|
||||
# whatever stale state is in the stores — surfacing the
|
||||
# previous edit's errors as if they were current is the
|
||||
# ghost-diagnostics bug. The server is NOT marked broken:
|
||||
# slow is not dead, and the next edit may well succeed.
|
||||
# Server alive but no verdict on the post-edit content in budget
|
||||
# (common for tsserver on big projects). Report "no data" rather
|
||||
# than stale stores — that would be the ghost-diagnostics bug.
|
||||
# Not marked broken: slow is not dead.
|
||||
eventlog.log_timeout(server_id, file_path, kind="fresh diagnostics")
|
||||
return []
|
||||
|
||||
@@ -388,18 +324,12 @@ class LSPService:
|
||||
baseline = self._delta_baseline.get(abs_path) or []
|
||||
if baseline:
|
||||
if line_shift is not None:
|
||||
# Remap baseline diagnostics into post-edit
|
||||
# coordinates so shifted-but-otherwise-identical
|
||||
# entries hash equal under _diag_key. Entries
|
||||
# that mapped into a deleted region drop out
|
||||
# silently — they no longer apply.
|
||||
# Entries that map into a deleted region drop out — they no longer apply.
|
||||
from agent.lsp.range_shift import shift_baseline
|
||||
baseline = shift_baseline(baseline, line_shift)
|
||||
seen = {_diag_key(d) for d in baseline}
|
||||
diags = [d for d in diags if _diag_key(d) not in seen]
|
||||
# Roll baseline forward — next call returns deltas relative
|
||||
# to the just-emitted state, mirroring claude-code's
|
||||
# diagnosticTracking.
|
||||
# Roll the baseline forward so the next call is a delta against this state.
|
||||
try:
|
||||
fresh = self._loop.run(self._current_diags_async(file_path), timeout=2.0) or []
|
||||
except Exception: # noqa: BLE001
|
||||
@@ -414,53 +344,35 @@ class LSPService:
|
||||
return diags
|
||||
|
||||
def _mark_broken_for_file(self, file_path: str, exc: BaseException) -> None:
|
||||
"""Mark the (server_id, workspace_root) pair as broken so subsequent
|
||||
edits skip it instantly instead of re-paying timeout cost.
|
||||
"""Mark the file's ``(server_id, root)`` pair broken after an outer timeout/error.
|
||||
|
||||
Called when the outer ``_loop.run`` timeout cancels an in-flight
|
||||
spawn/initialize that the inner ``_get_or_spawn`` task was still
|
||||
holding open. Without this, every subsequent write would re-enter
|
||||
the spawn path and re-pay the full ``snapshot_baseline``
|
||||
timeout (8s) until the binary is fixed.
|
||||
|
||||
Also kills any orphan client process that survived the cancelled
|
||||
future, and emits a single eventlog WARNING so the user knows
|
||||
which server gave up.
|
||||
|
||||
``exc`` is whatever exception the outer wrapper caught — used
|
||||
only for logging, never re-raised.
|
||||
The outer ``_loop.run`` timeout cancels the in-flight spawn before
|
||||
``_get_or_spawn`` could record the failure, so without this every
|
||||
later write would re-pay the full timeout. Also kills any
|
||||
half-initialized client left in ``_clients`` and logs the failure once.
|
||||
``exc`` is used only for logging.
|
||||
"""
|
||||
srv = find_server_for_file(file_path)
|
||||
if srv is None:
|
||||
return
|
||||
ws_root, gated = resolve_workspace_for_file(file_path)
|
||||
if not (ws_root and gated):
|
||||
key = self._broken_key(srv, file_path)
|
||||
if key is None:
|
||||
return
|
||||
try:
|
||||
per_server_root = srv.resolve_root(file_path, ws_root) or ws_root
|
||||
except Exception: # noqa: BLE001
|
||||
per_server_root = ws_root
|
||||
key = (srv.server_id, per_server_root)
|
||||
already_broken = key in self._broken
|
||||
self._broken.add(key)
|
||||
|
||||
# Kill any client we managed to spawn before the timeout. The
|
||||
# cancelled future never reached the broken-set add inside
|
||||
# ``_get_or_spawn`` so the client may still be hanging in
|
||||
# ``_clients`` with a half-initialized state.
|
||||
with self._state_lock:
|
||||
client = self._clients.pop(key, None)
|
||||
self._last_used.pop(key, None)
|
||||
if client is not None:
|
||||
try:
|
||||
# Fire-and-forget shutdown — give it a second to cleanup,
|
||||
# but don't block. We're already on a slow path.
|
||||
# Fire-and-forget shutdown — we're already on a slow path.
|
||||
self._loop.run(client.shutdown(), timeout=1.0)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
if not already_broken:
|
||||
eventlog.log_spawn_failed(srv.server_id, per_server_root, exc)
|
||||
eventlog.log_spawn_failed(srv.server_id, key[1], exc)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Tear down all clients and stop the background loop."""
|
||||
@@ -478,43 +390,35 @@ class LSPService:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _snapshot_async(self, file_path: str) -> List[Dict[str, Any]]:
|
||||
client = await self._get_or_spawn(file_path)
|
||||
if client is None:
|
||||
return []
|
||||
try:
|
||||
version = await client.open_file(file_path, language_id=language_id_for(file_path))
|
||||
fresh = await client.wait_for_diagnostics(file_path, version, mode=self._wait_mode)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("snapshot open/wait failed: %s", e)
|
||||
return []
|
||||
self._touch(client)
|
||||
if not fresh:
|
||||
# No fresh data for the pre-edit content — an empty baseline
|
||||
# is safe: worst case the delta filter removes less, never
|
||||
# more. Never seed the baseline from stale stores.
|
||||
return []
|
||||
return list(client.diagnostics_for(file_path, fresh_only=True))
|
||||
# No fresh data for the pre-edit content → empty baseline. Safe: the
|
||||
# delta filter then removes less, never more. Never seed from stale stores.
|
||||
return await self._open_and_wait_async(file_path, snapshot=True) or []
|
||||
|
||||
async def _open_and_wait_async(self, file_path: str) -> Optional[List[Dict[str, Any]]]:
|
||||
async def _open_and_wait_async(self, file_path: str, *, snapshot: bool = False) -> Optional[List[Dict[str, Any]]]:
|
||||
"""Open + wait for FRESH diagnostics.
|
||||
|
||||
Returns the fresh diagnostic list, or ``None`` when the server
|
||||
never produced post-change data within the wait budget. The
|
||||
distinction matters: ``[]`` means "server checked the new
|
||||
content, it's clean", ``None`` means "no verdict" — the caller
|
||||
must not substitute stale data for either.
|
||||
Returns the fresh list, or ``None`` when the server produced no
|
||||
post-change data in budget. ``[]`` means "checked, clean"; ``None``
|
||||
means "no verdict" — callers must not substitute stale data for either.
|
||||
``snapshot`` mode (pre-write baseline) skips didSave and uses the
|
||||
default wait budget.
|
||||
"""
|
||||
client = await self._get_or_spawn(file_path)
|
||||
if client is None:
|
||||
return None
|
||||
try:
|
||||
version = await client.open_file(file_path, language_id=language_id_for(file_path))
|
||||
await client.save_file(file_path)
|
||||
if not snapshot:
|
||||
await client.save_file(file_path)
|
||||
fresh = await client.wait_for_diagnostics(
|
||||
file_path, version, mode=self._wait_mode, timeout=self._wait_timeout
|
||||
file_path, version, mode=self._wait_mode,
|
||||
timeout=None if snapshot else self._wait_timeout,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("open/wait failed for %s: %s", file_path, e)
|
||||
if snapshot:
|
||||
logger.debug("snapshot open/wait failed: %s", e)
|
||||
else:
|
||||
logger.debug("open/wait failed for %s: %s", file_path, e)
|
||||
return None
|
||||
self._touch(client)
|
||||
if not fresh:
|
||||
@@ -548,7 +452,7 @@ class LSPService:
|
||||
eventlog.log_disabled(
|
||||
srv.server_id, file_path, "exclude marker hit (server gated off)"
|
||||
)
|
||||
return None # exclude marker hit, server gated off
|
||||
return None
|
||||
|
||||
key = (srv.server_id, per_server_root)
|
||||
if key in self._broken:
|
||||
@@ -566,7 +470,6 @@ class LSPService:
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
|
||||
# Begin spawn
|
||||
loop = asyncio.get_running_loop()
|
||||
spawn_future: asyncio.Future = loop.create_future()
|
||||
with self._state_lock:
|
||||
@@ -581,10 +484,8 @@ class LSPService:
|
||||
)
|
||||
spec = srv.build_spawn(per_server_root, ctx)
|
||||
if spec is None:
|
||||
# ``build_spawn`` returns None when the binary can't be
|
||||
# located (auto-install disabled, manual-only server,
|
||||
# or install attempt failed). Surface this once via
|
||||
# the structured logger so the user can act on it.
|
||||
# Binary not locatable (auto-install off, manual-only, or
|
||||
# install failed) — surface once via the structured logger.
|
||||
eventlog.log_server_unavailable(srv.server_id, srv.server_id)
|
||||
self._broken.add(key)
|
||||
spawn_future.set_result(None)
|
||||
@@ -619,13 +520,7 @@ class LSPService:
|
||||
self._idle_reaper_task = asyncio.create_task(self._idle_reaper_loop())
|
||||
|
||||
def _touch(self, client: LSPClient) -> None:
|
||||
"""Refresh the last-used timestamp for a client we just used.
|
||||
|
||||
Guarded on membership so a reaped-mid-operation client can't
|
||||
resurrect an orphan ``_last_used`` entry after the reaper popped
|
||||
the key. All writers and the reaper run on the background loop
|
||||
thread; the lock keeps this consistent with the reader anyway.
|
||||
"""
|
||||
"""Refresh last-used; guarded on membership so a client reaped mid-operation can't resurrect its entry."""
|
||||
key = (client.server_id, client.workspace_root)
|
||||
with self._state_lock:
|
||||
if key in self._clients:
|
||||
@@ -640,9 +535,8 @@ class LSPService:
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e: # noqa: BLE001
|
||||
# A transient sweep error must not kill the reaper —
|
||||
# otherwise one bad shutdown permanently re-opens the
|
||||
# unbounded-accumulation leak this loop exists to fix.
|
||||
# A transient sweep error must not kill the reaper, or the
|
||||
# unbounded-accumulation leak it exists to fix comes back.
|
||||
logger.debug("LSP idle reaper sweep error: %s", e)
|
||||
|
||||
async def _reap_idle_once(self) -> None:
|
||||
@@ -682,12 +576,8 @@ class LSPService:
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# status / introspection (used by ``hermes lsp status``)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_status(self) -> Dict[str, Any]:
|
||||
"""Return a snapshot of the service for the CLI status command."""
|
||||
"""Return a snapshot of the service for ``hermes lsp status``."""
|
||||
with self._state_lock:
|
||||
clients = [
|
||||
{
|
||||
@@ -710,35 +600,4 @@ class LSPService:
|
||||
}
|
||||
|
||||
|
||||
def _diag_key(d: Dict[str, Any]) -> str:
|
||||
"""Content equality key used for cross-edit delta filtering.
|
||||
|
||||
Includes the diagnostic's position range — when used together
|
||||
with :func:`agent.lsp.range_shift.shift_baseline`, the baseline
|
||||
is line-shifted into post-edit coordinates BEFORE this key is
|
||||
computed, so identical-but-shifted diagnostics hash equal. Two
|
||||
genuinely distinct diagnostics at different lines (e.g. the same
|
||||
error class introduced at a second site) hash differently and
|
||||
are surfaced as new.
|
||||
|
||||
Mirrors :func:`agent.lsp.client._diagnostic_key`; intentionally
|
||||
identical so the two layers agree on diagnostic identity.
|
||||
"""
|
||||
rng = d.get("range") or {}
|
||||
start = rng.get("start") or {}
|
||||
end = rng.get("end") or {}
|
||||
code = d.get("code")
|
||||
if code is not None and not isinstance(code, str):
|
||||
code = str(code)
|
||||
return "\x00".join(
|
||||
[
|
||||
str(d.get("severity") or 1),
|
||||
str(code or ""),
|
||||
str(d.get("source") or ""),
|
||||
str(d.get("message") or "").strip(),
|
||||
f"{start.get('line', 0)}:{start.get('character', 0)}-{end.get('line', 0)}:{end.get('character', 0)}",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["LSPService"]
|
||||
|
||||
+27
-56
@@ -1,17 +1,9 @@
|
||||
"""Minimal LSP JSON-RPC 2.0 framer over async streams.
|
||||
|
||||
LSP wire format:
|
||||
|
||||
Content-Length: <bytes>\\r\\n
|
||||
\\r\\n
|
||||
<utf-8 JSON body>
|
||||
|
||||
The body is a JSON-RPC 2.0 envelope: request, response, or notification.
|
||||
|
||||
This module replaces what ``vscode-jsonrpc/node`` would do in a
|
||||
TypeScript implementation. We keep it deliberately small — just the
|
||||
framer + envelope helpers — so :class:`agent.lsp.client.LSPClient` can
|
||||
focus on protocol semantics.
|
||||
Wire format: ``Content-Length: <bytes>\\r\\n\\r\\n<utf-8 JSON body>`` where the
|
||||
body is a JSON-RPC 2.0 request, response, or notification. Just the framer
|
||||
plus envelope helpers, so :class:`agent.lsp.client.LSPClient` can focus on
|
||||
protocol semantics.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -22,27 +14,21 @@ from typing import Any, Optional, Tuple
|
||||
|
||||
logger = logging.getLogger("agent.lsp.protocol")
|
||||
|
||||
# LSP error codes we care about. Full list in
|
||||
# https://microsoft.github.io/language-server-protocol/specifications/lsp/3.17/specification/#errorCodes
|
||||
# LSP error codes we care about (spec 3.17 #errorCodes).
|
||||
ERROR_CONTENT_MODIFIED = -32801
|
||||
ERROR_REQUEST_CANCELLED = -32800
|
||||
ERROR_METHOD_NOT_FOUND = -32601
|
||||
|
||||
_MAX_HEADER_BYTES = 8192 # a well-behaved server fits in well under 200 bytes
|
||||
_MAX_BODY_BYTES = 64 * 1024 * 1024
|
||||
|
||||
|
||||
class LSPProtocolError(Exception):
|
||||
"""Raised when the wire protocol is violated.
|
||||
|
||||
Distinct from :class:`LSPRequestError` which represents a server
|
||||
returning a JSON-RPC error response — that's protocol-conformant.
|
||||
This exception means the framing or envelope itself is broken.
|
||||
"""
|
||||
"""The framing or envelope itself is broken (vs. :class:`LSPRequestError`, a conformant error response)."""
|
||||
|
||||
|
||||
class LSPRequestError(Exception):
|
||||
"""Raised when an LSP request returns an error response.
|
||||
|
||||
Carries the JSON-RPC ``code``, ``message``, and optional ``data``.
|
||||
"""
|
||||
"""An LSP request returned a JSON-RPC error response; carries ``code``, ``message``, ``data``."""
|
||||
|
||||
def __init__(self, code: int, message: str, data: Any = None) -> None:
|
||||
super().__init__(f"LSP error {code}: {message}")
|
||||
@@ -52,25 +38,17 @@ class LSPRequestError(Exception):
|
||||
|
||||
|
||||
def encode_message(obj: dict) -> bytes:
|
||||
"""Encode a JSON-RPC envelope as a Content-Length framed byte string.
|
||||
|
||||
The body is encoded as compact UTF-8 JSON (no spaces between
|
||||
separators) — matches what ``vscode-jsonrpc`` emits and keeps the
|
||||
Content-Length count exact.
|
||||
"""
|
||||
"""Encode an envelope as compact UTF-8 JSON with an exact Content-Length header."""
|
||||
body = json.dumps(obj, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
|
||||
header = f"Content-Length: {len(body)}\r\n\r\n".encode("ascii")
|
||||
return header + body
|
||||
|
||||
|
||||
async def read_message(reader: asyncio.StreamReader) -> Optional[dict]:
|
||||
"""Read one Content-Length framed JSON-RPC message from the stream.
|
||||
"""Read one framed message.
|
||||
|
||||
Returns ``None`` on clean EOF (server closed stdout cleanly between
|
||||
messages — typical shutdown). Raises :class:`LSPProtocolError` on
|
||||
malformed framing.
|
||||
|
||||
The reader is advanced to just past the JSON body on success.
|
||||
Returns ``None`` on clean EOF between messages (typical shutdown);
|
||||
raises :class:`LSPProtocolError` on malformed framing.
|
||||
"""
|
||||
headers: dict = {}
|
||||
header_bytes = 0
|
||||
@@ -78,18 +56,15 @@ async def read_message(reader: asyncio.StreamReader) -> Optional[dict]:
|
||||
try:
|
||||
line = await reader.readuntil(b"\r\n")
|
||||
except asyncio.IncompleteReadError as e:
|
||||
# EOF while reading headers. If we hadn't started a header
|
||||
# block, treat as clean EOF; otherwise the framing is bad.
|
||||
# EOF before any header started is a clean close; mid-block is bad framing.
|
||||
if not e.partial and not headers:
|
||||
return None
|
||||
raise LSPProtocolError(
|
||||
f"unexpected EOF while reading LSP headers (partial={e.partial!r})"
|
||||
) from e
|
||||
# Defensive cap against a server streaming headers without ever
|
||||
# emitting CRLF-CRLF. Caps total header bytes at 8 KiB — a
|
||||
# well-behaved server fits in well under 200 bytes.
|
||||
# Cap against a server streaming headers without ever emitting CRLF-CRLF.
|
||||
header_bytes += len(line)
|
||||
if header_bytes > 8192:
|
||||
if header_bytes > _MAX_HEADER_BYTES:
|
||||
raise LSPProtocolError(
|
||||
"LSP header block exceeded 8 KiB without terminator"
|
||||
)
|
||||
@@ -111,7 +86,7 @@ async def read_message(reader: asyncio.StreamReader) -> Optional[dict]:
|
||||
n = int(cl)
|
||||
except ValueError as e:
|
||||
raise LSPProtocolError(f"non-integer Content-Length: {cl!r}") from e
|
||||
if n < 0 or n > 64 * 1024 * 1024: # 64 MiB sanity cap
|
||||
if n < 0 or n > _MAX_BODY_BYTES:
|
||||
raise LSPProtocolError(f"unreasonable Content-Length: {n}")
|
||||
|
||||
try:
|
||||
@@ -129,17 +104,17 @@ async def read_message(reader: asyncio.StreamReader) -> Optional[dict]:
|
||||
raise LSPProtocolError(f"non-UTF-8 LSP body: {e}") from e
|
||||
|
||||
|
||||
def make_request(req_id: int, method: str, params: Any) -> dict:
|
||||
"""Build a JSON-RPC 2.0 request envelope."""
|
||||
msg: dict = {"jsonrpc": "2.0", "id": req_id, "method": method}
|
||||
def make_notification(method: str, params: Any) -> dict:
|
||||
"""Build a JSON-RPC 2.0 notification envelope (no ``id``)."""
|
||||
msg: dict = {"jsonrpc": "2.0", "method": method}
|
||||
if params is not None:
|
||||
msg["params"] = params
|
||||
return msg
|
||||
|
||||
|
||||
def make_notification(method: str, params: Any) -> dict:
|
||||
"""Build a JSON-RPC 2.0 notification envelope (no ``id``)."""
|
||||
msg: dict = {"jsonrpc": "2.0", "method": method}
|
||||
def make_request(req_id: int, method: str, params: Any) -> dict:
|
||||
"""Build a JSON-RPC 2.0 request envelope."""
|
||||
msg: dict = {"jsonrpc": "2.0", "id": req_id, "method": method}
|
||||
if params is not None:
|
||||
msg["params"] = params
|
||||
return msg
|
||||
@@ -159,15 +134,11 @@ def make_error_response(req_id: Any, code: int, message: str, data: Any = None)
|
||||
|
||||
|
||||
def classify_message(msg: dict) -> Tuple[str, Any]:
|
||||
"""Return ``(kind, key)`` where kind is one of ``request``,
|
||||
``response``, ``notification``, ``invalid``.
|
||||
"""Return ``(kind, key)``: kind ∈ request/response/notification/invalid.
|
||||
|
||||
The key is the request id for request/response, the method name
|
||||
for notifications, and ``None`` for invalid messages.
|
||||
Key is the id for request/response, the method for notifications, ``None`` for invalid.
|
||||
"""
|
||||
if not isinstance(msg, dict):
|
||||
return "invalid", None
|
||||
if msg.get("jsonrpc") != "2.0":
|
||||
if not isinstance(msg, dict) or msg.get("jsonrpc") != "2.0":
|
||||
return "invalid", None
|
||||
has_id = "id" in msg
|
||||
has_method = "method" in msg
|
||||
|
||||
+23
-84
@@ -1,28 +1,14 @@
|
||||
"""Diff-aware line-shift map for cross-edit LSP delta filtering.
|
||||
|
||||
When an edit deletes or inserts lines in the middle of a file, every
|
||||
diagnostic below the edit point shifts to a new line number. The
|
||||
LSPService delta filter subtracts the pre-edit baseline from the
|
||||
post-edit diagnostics keyed on ``(severity, code, source, message,
|
||||
range)`` — without an adjustment, the shifted-but-otherwise-identical
|
||||
diagnostics look brand-new and the agent gets flooded with noise.
|
||||
When an edit inserts or deletes lines, every diagnostic below the edit point
|
||||
moves. The delta filter keys on ``(severity, code, source, message, range)``,
|
||||
so without adjustment the shifted-but-identical diagnostics look brand-new.
|
||||
We build a pre→post line map from ``difflib.SequenceMatcher.get_opcodes()``
|
||||
and apply it to the baseline before the set-difference; diagnostics in a
|
||||
deleted region map to ``None`` and drop out (they genuinely no longer apply).
|
||||
|
||||
The fix used here is the same trick git's blame and unified diff use:
|
||||
build a piecewise-linear map from pre-edit line numbers to post-edit
|
||||
line numbers, then apply that map to baseline diagnostics before the
|
||||
set-difference. Diagnostics whose pre-edit line is in a region the
|
||||
edit deleted return ``None`` and are dropped from the baseline (they
|
||||
genuinely no longer apply).
|
||||
|
||||
Trade-off vs. dropping range from the key entirely (the previous
|
||||
fix): preserves the "new instance of an identical error at a
|
||||
different line" signal — if the model introduces a second instance
|
||||
of the same error class at a different location, that one will be
|
||||
surfaced as new instead of swallowed by content-only dedup.
|
||||
|
||||
The map is derived from ``difflib.SequenceMatcher.get_opcodes()`` and
|
||||
exposed as a single callable so callers don't have to reason about
|
||||
diff regions.
|
||||
Keeping range in the key (rather than content-only dedup) preserves the
|
||||
"new instance of an identical error at a different line" signal.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -31,58 +17,29 @@ from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
|
||||
def build_line_shift(pre_text: str, post_text: str) -> Callable[[int], Optional[int]]:
|
||||
"""Build a function mapping pre-edit line numbers to post-edit line numbers.
|
||||
"""Return ``shift(pre_line) -> post_line | None`` over 0-indexed lines (LSP convention).
|
||||
|
||||
Lines are 0-indexed to match the LSP wire format
|
||||
(``range.start.line`` is 0-indexed).
|
||||
|
||||
The returned callable takes a pre-edit 0-indexed line number and
|
||||
returns the corresponding post-edit 0-indexed line number, or
|
||||
``None`` if that line was deleted by the edit (no post-edit
|
||||
counterpart exists).
|
||||
|
||||
Cost: one ``SequenceMatcher.get_opcodes()`` call up front; the
|
||||
returned closure is O(log n) per call (binary search over opcode
|
||||
regions). Cheap enough to call once per write/patch and apply to
|
||||
every baseline diagnostic.
|
||||
``None`` means the line was deleted. One ``get_opcodes()`` call up front;
|
||||
the closure scans the (small) opcode list per lookup.
|
||||
"""
|
||||
pre_lines = pre_text.splitlines() if pre_text else []
|
||||
post_lines = post_text.splitlines() if post_text else []
|
||||
|
||||
# Trivial case: identical content or no content — identity map.
|
||||
if pre_lines == post_lines:
|
||||
return lambda line: line
|
||||
|
||||
# SequenceMatcher.get_opcodes() returns a list of
|
||||
# (tag, i1, i2, j1, j2) where tag is 'equal', 'replace', 'delete',
|
||||
# or 'insert'. i1:i2 is the range in pre, j1:j2 is the range in
|
||||
# post. We build a list of (i1, i2, j1, j2, tag) tuples and
|
||||
# binary-search by i for each lookup.
|
||||
sm = difflib.SequenceMatcher(a=pre_lines, b=post_lines, autojunk=False)
|
||||
opcodes = sm.get_opcodes()
|
||||
# Opcodes are (tag, i1, i2, j1, j2): i-range in pre, j-range in post.
|
||||
opcodes = difflib.SequenceMatcher(a=pre_lines, b=post_lines, autojunk=False).get_opcodes()
|
||||
|
||||
def shift(line: int) -> Optional[int]:
|
||||
# Find the opcode region whose i1 <= line < i2.
|
||||
# Linear scan is fine — typical opcode count is small (single
|
||||
# digits for a typical patch-tool edit).
|
||||
for tag, i1, i2, j1, j2 in opcodes:
|
||||
if i1 <= line < i2:
|
||||
if tag == "equal":
|
||||
# Pre-line N → post-line (N - i1 + j1).
|
||||
return line - i1 + j1
|
||||
if tag == "delete":
|
||||
# Pre-line is in a deleted region — no post counterpart.
|
||||
return None
|
||||
if tag == "replace":
|
||||
# Replace == delete + insert; the pre-line has no
|
||||
# post counterpart in any meaningful sense. Drop.
|
||||
return None
|
||||
# 'insert' has i1 == i2 so line < i2 can't be hit.
|
||||
# 'equal' maps by offset; 'delete'/'replace' lines have no
|
||||
# post counterpart. 'insert' has i1 == i2 and can't match.
|
||||
return line - i1 + j1 if tag == "equal" else None
|
||||
if line < i1:
|
||||
# Past the relevant region — handled in earlier iteration.
|
||||
break
|
||||
# Past the last opcode region (line >= len(pre_lines)).
|
||||
# Anchor at end of post.
|
||||
# Past the last pre line: anchor at end of post.
|
||||
return max(0, len(post_lines) - 1) if post_lines else None
|
||||
|
||||
return shift
|
||||
@@ -90,45 +47,27 @@ def build_line_shift(pre_text: str, post_text: str) -> Callable[[int], Optional[
|
||||
|
||||
def shift_diagnostic_range(diag: Dict[str, Any],
|
||||
shift: Callable[[int], Optional[int]]) -> Optional[Dict[str, Any]]:
|
||||
"""Return a copy of ``diag`` with its line range remapped through ``shift``.
|
||||
"""Copy of ``diag`` with its line range remapped; ``None`` if the start line was deleted.
|
||||
|
||||
Returns ``None`` if the diagnostic's start line maps to ``None``
|
||||
(the line was deleted by the edit) — caller drops it from the
|
||||
baseline since the diagnostic no longer applies.
|
||||
|
||||
Both ``start.line`` and ``end.line`` are remapped independently;
|
||||
when only the end maps to ``None`` (rare, multi-line diagnostic
|
||||
straddling the edit boundary) we collapse to a single-line range
|
||||
at the shifted start to keep the diagnostic in the baseline.
|
||||
|
||||
The original ``diag`` is not mutated.
|
||||
A multi-line diagnostic whose end straddles the deletion collapses to a
|
||||
single-line range at the shifted start so it stays in the baseline.
|
||||
"""
|
||||
rng = diag.get("range") or {}
|
||||
start = rng.get("start") or {}
|
||||
end = rng.get("end") or {}
|
||||
|
||||
pre_start_line = int(start.get("line", 0))
|
||||
pre_end_line = int(end.get("line", pre_start_line))
|
||||
|
||||
new_start_line = shift(pre_start_line)
|
||||
if new_start_line is None:
|
||||
return None
|
||||
|
||||
new_end_line = shift(pre_end_line)
|
||||
new_end_line = shift(int(end.get("line", pre_start_line)))
|
||||
if new_end_line is None:
|
||||
# Diagnostic straddled the deletion — collapse to start.
|
||||
new_end_line = new_start_line
|
||||
|
||||
shifted = dict(diag)
|
||||
shifted["range"] = {
|
||||
"start": {
|
||||
"line": new_start_line,
|
||||
"character": int(start.get("character", 0)),
|
||||
},
|
||||
"end": {
|
||||
"line": new_end_line,
|
||||
"character": int(end.get("character", 0)),
|
||||
},
|
||||
"start": {"line": new_start_line, "character": int(start.get("character", 0))},
|
||||
"end": {"line": new_end_line, "character": int(end.get("character", 0))},
|
||||
}
|
||||
return shifted
|
||||
|
||||
|
||||
+18
-51
@@ -1,27 +1,23 @@
|
||||
"""Format LSP diagnostics for inclusion in tool output.
|
||||
|
||||
The model sees a compact, severity-filtered, line-bounded summary of
|
||||
diagnostics introduced by the latest edit. Format matches what
|
||||
OpenCode's ``lsp/diagnostic.ts`` and Claude Code's
|
||||
``formatDiagnosticsSummary`` produce — ``<diagnostics>`` blocks with
|
||||
1-indexed line/column, capped at ``MAX_PER_FILE`` errors.
|
||||
The model sees a compact, severity-filtered, line-bounded ``<diagnostics>``
|
||||
block (1-indexed line/column, capped at ``MAX_PER_FILE``) for diagnostics
|
||||
introduced by the latest edit.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import html
|
||||
from typing import Any, Dict, List
|
||||
|
||||
# Severity-1 only by default — warnings/info/hints would flood the
|
||||
# agent. Lift this in config under ``lsp.severities`` if needed.
|
||||
# ERROR only by default — warnings/info/hints would flood the agent.
|
||||
SEVERITY_NAMES = {1: "ERROR", 2: "WARN", 3: "INFO", 4: "HINT"}
|
||||
DEFAULT_SEVERITIES = frozenset({1}) # ERROR only
|
||||
DEFAULT_SEVERITIES = frozenset({1})
|
||||
|
||||
MAX_PER_FILE = 20
|
||||
MAX_TOTAL_CHARS = 4000
|
||||
|
||||
# Per-field caps for diagnostic content sourced from the language server.
|
||||
# These bound the length of any single attacker-controlled identifier that
|
||||
# can ride into the model's tool output via an LSP diagnostic message.
|
||||
# Per-field caps bound any single attacker-controlled identifier that can
|
||||
# ride into the model's tool output via an LSP diagnostic message.
|
||||
MAX_MESSAGE_CHARS = 300
|
||||
MAX_CODE_CHARS = 80
|
||||
MAX_SOURCE_CHARS = 80
|
||||
@@ -30,45 +26,22 @@ MAX_SOURCE_CHARS = 80
|
||||
def _sanitize_field(value: Any, *, limit: int) -> str:
|
||||
"""Make a language-server field safe to embed in a tool-result block.
|
||||
|
||||
Diagnostic ``message``, ``code``, and ``source`` originate from a
|
||||
language server that has just parsed user-controlled source code, so
|
||||
they're untrusted from the agent's point of view. A hostile repo can
|
||||
place instruction-shaped text inside identifier names, type aliases,
|
||||
or import paths so the resulting diagnostic echoes that text back
|
||||
into the ``<diagnostics>`` block the model reads.
|
||||
|
||||
This helper:
|
||||
|
||||
* Collapses CR/LF so a raw newline can't synthesize a new line in the
|
||||
formatted block.
|
||||
* Drops non-printable ASCII control characters that have no business
|
||||
in a single-line summary.
|
||||
* Caps length per-field so a long identifier can't push past the
|
||||
block boundary.
|
||||
* HTML-escapes ``< > &`` so the result can't close ``<diagnostics>``
|
||||
early or open a new tag.
|
||||
|
||||
Returns ``""`` for ``None`` / empty so the surrounding format string
|
||||
naturally omits the part (mirrors the prior ``if code not in {None,
|
||||
""}`` check at call sites).
|
||||
``message``/``code``/``source`` come from a server that just parsed
|
||||
user-controlled code, so a hostile repo can smuggle instruction-shaped
|
||||
text through identifier names. We collapse CR/LF, drop control chars,
|
||||
cap the length, and HTML-escape ``< > &`` so the text can't close
|
||||
``<diagnostics>`` early. ``None``/empty → ``""`` so callers can omit the part.
|
||||
"""
|
||||
if value is None:
|
||||
return ""
|
||||
raw = str(value)
|
||||
# Collapse newlines so identifier text with raw \n can't fake new lines.
|
||||
raw = raw.replace("\r", " ").replace("\n", " ")
|
||||
# Drop ASCII control chars; keep regular spaces.
|
||||
raw = str(value).replace("\r", " ").replace("\n", " ")
|
||||
raw = "".join(ch for ch in raw if ch == " " or ch.isprintable())
|
||||
raw = raw.strip()[:limit]
|
||||
return html.escape(raw, quote=False)
|
||||
|
||||
|
||||
def format_diagnostic(d: Dict[str, Any]) -> str:
|
||||
"""One-line representation of a single diagnostic.
|
||||
|
||||
``message``, ``code``, and ``source`` are sanitized before
|
||||
interpolation — see ``_sanitize_field``.
|
||||
"""
|
||||
"""One-line representation of a single diagnostic (fields sanitized)."""
|
||||
sev = SEVERITY_NAMES.get(d.get("severity") or 1, "ERROR")
|
||||
rng = d.get("range") or {}
|
||||
start = rng.get("start") or {}
|
||||
@@ -89,11 +62,7 @@ def report_for_file(
|
||||
severities: frozenset = DEFAULT_SEVERITIES,
|
||||
max_per_file: int = MAX_PER_FILE,
|
||||
) -> str:
|
||||
"""Build a ``<diagnostics file=...>`` block for one file.
|
||||
|
||||
Returns an empty string when no diagnostics pass the severity
|
||||
filter, so callers can do ``if block:`` to skip empty cases.
|
||||
"""
|
||||
"""Build a ``<diagnostics file=...>`` block; ``""`` when nothing passes the severity filter."""
|
||||
if not diagnostics:
|
||||
return ""
|
||||
filtered = [d for d in diagnostics if (d.get("severity") or 1) in severities]
|
||||
@@ -101,13 +70,11 @@ def report_for_file(
|
||||
return ""
|
||||
limited = filtered[:max_per_file]
|
||||
extra = len(filtered) - len(limited)
|
||||
lines = [format_diagnostic(d) for d in limited]
|
||||
body = "\n".join(lines)
|
||||
body = "\n".join(format_diagnostic(d) for d in limited)
|
||||
if extra > 0:
|
||||
body += f"\n... and {extra} more"
|
||||
# quote=True escapes both ``"`` and ``&`` so a crafted file name like
|
||||
# ``foo"><script`` can't break out of the ``file="..."`` attribute and
|
||||
# synthesize new tags inside the tool output.
|
||||
# quote=True also escapes ``"`` so a crafted file name can't break out of
|
||||
# the ``file="..."`` attribute and synthesize new tags.
|
||||
safe_path = html.escape(file_path, quote=True)
|
||||
return f"<diagnostics file=\"{safe_path}\">\n{body}\n</diagnostics>"
|
||||
|
||||
|
||||
+176
-883
File diff suppressed because it is too large
Load Diff
+61
-107
@@ -1,107 +1,97 @@
|
||||
"""Workspace and project-root resolution for LSP.
|
||||
|
||||
Two concerns live here:
|
||||
|
||||
1. **Workspace gate** — the upper-level "is this directory a project?"
|
||||
check. Hermes only runs LSP when the cwd (or the file being edited)
|
||||
sits inside a git worktree. Files outside any git root never
|
||||
trigger LSP, even if a server is configured. This keeps Telegram
|
||||
gateway users on user-home cwd's from spawning daemons.
|
||||
|
||||
2. **NearestRoot** — the per-server project-root walk. Each language
|
||||
server cares about a different marker (``pyproject.toml`` for
|
||||
Python, ``Cargo.toml`` for Rust, ``go.mod`` for Go, etc.) and
|
||||
wants the directory containing that marker. ``nearest_root()``
|
||||
walks up from a starting path looking for any of a list of marker
|
||||
files, optionally bailing if an exclude marker shows up first.
|
||||
1. **Workspace gate** — LSP only runs when the cwd (or the edited file) sits
|
||||
inside a git worktree. Files outside any git root never trigger LSP,
|
||||
which keeps gateway users on user-home cwd's from spawning daemons.
|
||||
2. **nearest_root** — the per-server project-root walk: up from a start path
|
||||
looking for marker files (``pyproject.toml``, ``Cargo.toml``, ...),
|
||||
optionally bailing if an exclude marker shows up first.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Iterable, Optional, Tuple
|
||||
from typing import Iterable, Iterator, Optional, Tuple
|
||||
|
||||
logger = logging.getLogger("agent.lsp.workspace")
|
||||
|
||||
# Cache: cwd → (worktree_root, is_git) so repeated calls don't re-stat.
|
||||
# Cleared on shutdown. Keyed by absolute resolved path so symlink
|
||||
# folds collapse to one entry.
|
||||
# Cache: start dir → (worktree_root, is_git) so repeated calls don't re-stat.
|
||||
# Cleared on shutdown.
|
||||
_workspace_cache: dict = {}
|
||||
|
||||
# Walk cap: the deepest reasonable monorepo is well under 64 levels; bounds a
|
||||
# pathological cwd or symlink cycle even though parent-equality normally stops us.
|
||||
_MAX_WALK = 64
|
||||
|
||||
|
||||
def normalize_path(path: str) -> str:
|
||||
"""Normalize a path for use as a stable map key.
|
||||
"""Expand ``~``, make absolute, collapse ``.``/``..``.
|
||||
|
||||
Resolves ``~``, makes absolute, and collapses ``.``/``..``. We do
|
||||
NOT resolve symlinks here — symlink stability matters for some
|
||||
LSP servers (rust-analyzer cares about Cargo workspace identity)
|
||||
and we want the canonical path the user typed when possible.
|
||||
Symlinks are deliberately NOT resolved — some servers (rust-analyzer's
|
||||
Cargo workspace identity) care, and we want the path the user typed.
|
||||
"""
|
||||
return os.path.abspath(os.path.expanduser(path))
|
||||
|
||||
|
||||
def find_git_worktree(start: str) -> Optional[str]:
|
||||
"""Walk up from ``start`` looking for a ``.git`` entry (file or dir).
|
||||
|
||||
Returns the directory containing ``.git``, or ``None`` if no git
|
||||
root is found before hitting the filesystem root.
|
||||
|
||||
A ``.git`` *file* (not directory) means we're inside a git
|
||||
worktree set up via ``git worktree add`` — both forms count.
|
||||
"""
|
||||
def _start_dir(start: str) -> Optional[Path]:
|
||||
"""Normalized start directory (a file's parent), or ``None`` on pathological input."""
|
||||
try:
|
||||
start_path = Path(normalize_path(start))
|
||||
if start_path.is_file():
|
||||
start_path = start_path.parent
|
||||
except (OSError, RuntimeError, ValueError):
|
||||
# Pathological input (loop in symlinks, encoding error, etc.) —
|
||||
# bail out rather than crash the lint hook.
|
||||
# Symlink loop, encoding error, etc. — bail rather than crash the lint hook.
|
||||
return None
|
||||
return start_path
|
||||
|
||||
|
||||
def _walk_up(start: Path) -> Iterator[Path]:
|
||||
"""Yield ``start`` and its ancestors up to the filesystem root, bounded by ``_MAX_WALK``."""
|
||||
cur = start
|
||||
for _ in range(_MAX_WALK):
|
||||
yield cur
|
||||
parent = cur.parent
|
||||
if parent == cur:
|
||||
return
|
||||
cur = parent
|
||||
|
||||
|
||||
def find_git_worktree(start: str) -> Optional[str]:
|
||||
"""Return the nearest ancestor dir containing ``.git`` (file or dir — worktrees count), else ``None``."""
|
||||
start_path = _start_dir(start)
|
||||
if start_path is None:
|
||||
return None
|
||||
|
||||
# Cache check
|
||||
cached = _workspace_cache.get(str(start_path))
|
||||
if cached is not None:
|
||||
root, _is_git = cached
|
||||
return root
|
||||
return cached[0]
|
||||
|
||||
cur = start_path
|
||||
# Defensive cap: the deepest reasonable monorepo is well under 64
|
||||
# levels. Caps the walk so a pathological cwd or a symlink cycle
|
||||
# we somehow traverse can't keep us looping.
|
||||
for _ in range(64):
|
||||
git_marker = cur / ".git"
|
||||
for cur in _walk_up(start_path):
|
||||
try:
|
||||
if git_marker.exists():
|
||||
if (cur / ".git").exists():
|
||||
resolved = str(cur)
|
||||
_workspace_cache[str(start_path)] = (resolved, True)
|
||||
return resolved
|
||||
except OSError:
|
||||
# Permission error on a parent dir — bail out cleanly.
|
||||
break
|
||||
parent = cur.parent
|
||||
if parent == cur:
|
||||
break
|
||||
cur = parent
|
||||
|
||||
_workspace_cache[str(start_path)] = (None, False)
|
||||
return None
|
||||
|
||||
|
||||
def is_inside_workspace(path: str, workspace_root: str) -> bool:
|
||||
"""Return True iff ``path`` is inside (or equal to) ``workspace_root``.
|
||||
"""True iff ``path`` is inside (or equal to) ``workspace_root``.
|
||||
|
||||
Uses absolute paths but does not resolve symlinks — a file accessed
|
||||
via a symlink that points outside the workspace still counts as
|
||||
outside. This is the conservative interpretation; matches LSP
|
||||
behaviour where servers reject didOpen for unrelated files.
|
||||
Symlinks are not resolved: a symlink pointing outside still counts as
|
||||
outside, matching servers that reject didOpen for unrelated files.
|
||||
"""
|
||||
p = normalize_path(path)
|
||||
root = normalize_path(workspace_root)
|
||||
if p == root:
|
||||
return True
|
||||
# Use os.path.commonpath to handle case-insensitive filesystems
|
||||
# correctly on macOS/Windows.
|
||||
# commonpath handles case-insensitive filesystems on macOS/Windows.
|
||||
try:
|
||||
common = os.path.commonpath([p, root])
|
||||
except ValueError:
|
||||
@@ -117,56 +107,37 @@ def nearest_root(
|
||||
excludes: Optional[Iterable[str]] = None,
|
||||
ceiling: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
"""Walk up from ``start`` looking for any of the given marker files.
|
||||
"""Walk up from ``start`` for the directory containing the first matched marker.
|
||||
|
||||
Returns the **directory containing** the first matched marker, or
|
||||
``None`` if no marker is found before hitting ``ceiling`` (or the
|
||||
filesystem root if no ceiling).
|
||||
|
||||
If ``excludes`` is provided and an exclude marker matches *first*
|
||||
in the upward walk, returns ``None`` — the server is gated off
|
||||
for that file. Mirrors OpenCode's NearestRoot exclude semantics
|
||||
(e.g. typescript skips deno projects when ``deno.json`` is found
|
||||
before ``package.json``).
|
||||
Returns ``None`` past ``ceiling`` (or the filesystem root), or when an
|
||||
exclude marker is found first — the server is gated off for that file
|
||||
(e.g. typescript skips deno projects when ``deno.json`` precedes
|
||||
``package.json``). Marker names are exact filenames — no globs.
|
||||
"""
|
||||
start_path = Path(normalize_path(start))
|
||||
try:
|
||||
if start_path.is_file():
|
||||
start_path = start_path.parent
|
||||
except (OSError, RuntimeError, ValueError):
|
||||
start_path = _start_dir(start)
|
||||
if start_path is None:
|
||||
return None
|
||||
ceiling_path = Path(normalize_path(ceiling)) if ceiling else None
|
||||
|
||||
markers_list = list(markers)
|
||||
excludes_list = list(excludes) if excludes else []
|
||||
|
||||
cur = start_path
|
||||
# Defensive cap matching ``find_git_worktree``. Bounded walk
|
||||
# protects against pathological inputs even though the
|
||||
# parent-equality stop normally terminates within ~10 steps.
|
||||
for _ in range(64):
|
||||
# Check excludes first — if an exclude is found at this level,
|
||||
# the server is gated off for this file.
|
||||
for cur in _walk_up(start_path):
|
||||
# Excludes are checked before markers at each level.
|
||||
for exc in excludes_list:
|
||||
try:
|
||||
if (cur / exc).exists():
|
||||
return None
|
||||
except OSError:
|
||||
continue
|
||||
# Then check markers.
|
||||
for marker in markers_list:
|
||||
try:
|
||||
if (cur / marker).exists():
|
||||
return str(cur)
|
||||
except OSError:
|
||||
continue
|
||||
# Stop conditions.
|
||||
if ceiling_path is not None and cur == ceiling_path:
|
||||
return None
|
||||
parent = cur.parent
|
||||
if parent == cur:
|
||||
return None
|
||||
cur = parent
|
||||
return None
|
||||
|
||||
|
||||
@@ -175,29 +146,16 @@ def resolve_workspace_for_file(
|
||||
*,
|
||||
cwd: Optional[str] = None,
|
||||
) -> Tuple[Optional[str], bool]:
|
||||
"""Resolve the workspace root for a file.
|
||||
"""Return ``(workspace_root, gated_in)`` for a file.
|
||||
|
||||
Returns ``(workspace_root, gated_in)`` where ``gated_in`` is True
|
||||
iff LSP should run for this file at all. Currently the gate is
|
||||
"file is inside a git worktree found by walking up from cwd OR
|
||||
from the file itself".
|
||||
|
||||
The cwd path takes precedence — if the agent was launched in a
|
||||
git project, that worktree is the workspace, and any edit inside
|
||||
it (regardless of where the file lives) is in-scope. If the cwd
|
||||
isn't in a git worktree, we try the file's own location as a
|
||||
fallback.
|
||||
|
||||
Returns ``(None, False)`` when neither path is in a git worktree.
|
||||
The cwd's worktree wins when the file is inside it; otherwise the file's
|
||||
own worktree is the fallback anchor (monorepos / unrelated checkouts).
|
||||
``(None, False)`` when neither is in a git worktree.
|
||||
"""
|
||||
cwd = cwd or os.getcwd()
|
||||
cwd_root = find_git_worktree(cwd)
|
||||
if cwd_root is not None:
|
||||
if is_inside_workspace(file_path, cwd_root):
|
||||
return cwd_root, True
|
||||
# File is outside the cwd's worktree — try the file's own
|
||||
# location as a secondary anchor. Useful for monorepos where
|
||||
# the user opens an unrelated checkout.
|
||||
if cwd_root is not None and is_inside_workspace(file_path, cwd_root):
|
||||
return cwd_root, True
|
||||
file_root = find_git_worktree(file_path)
|
||||
if file_root is not None:
|
||||
return file_root, True
|
||||
@@ -205,11 +163,7 @@ def resolve_workspace_for_file(
|
||||
|
||||
|
||||
def clear_cache() -> None:
|
||||
"""Clear the workspace-resolution cache.
|
||||
|
||||
Called on service shutdown so a subsequent re-init doesn't pick
|
||||
up stale results from a previous session.
|
||||
"""
|
||||
"""Clear the workspace-resolution cache (on service shutdown, so re-init doesn't see stale results)."""
|
||||
_workspace_cache.clear()
|
||||
|
||||
|
||||
|
||||
+372
-609
File diff suppressed because it is too large
Load Diff
+84
-278
@@ -1,34 +1,11 @@
|
||||
"""Abstract base class for pluggable memory providers.
|
||||
|
||||
Memory providers give the agent persistent recall across sessions.
|
||||
The MemoryManager enforces a one-external-provider limit to prevent
|
||||
tool schema bloat and conflicting memory backends.
|
||||
|
||||
External providers (Honcho, Hindsight, Mem0, etc.) are registered
|
||||
and managed via MemoryManager. Only one external provider runs at a
|
||||
time.
|
||||
|
||||
Registration:
|
||||
Plugins ship in plugins/memory/<name>/ and are activated via
|
||||
the memory.provider config key.
|
||||
|
||||
Lifecycle (called by MemoryManager, wired in run_agent.py):
|
||||
initialize() — connect, create resources, warm up
|
||||
system_prompt_block() — static text for the system prompt
|
||||
prefetch(query) — background recall before each turn
|
||||
sync_turn(user, asst) — async write after each turn
|
||||
get_tool_schemas() — tool schemas to expose to the model
|
||||
handle_tool_call() — dispatch a tool call
|
||||
shutdown() — clean exit
|
||||
|
||||
Optional hooks (override to opt in):
|
||||
on_turn_start(turn, message, **kwargs) — per-turn tick with runtime context
|
||||
on_session_end(messages) — end-of-session extraction
|
||||
on_session_switch(new_session_id, **kwargs) — mid-process session_id rotation
|
||||
on_pre_compress(messages) -> str — extract before context compression
|
||||
on_memory_write(action, target, content, metadata=None) — mirror built-in memory writes
|
||||
on_delegation(task, result, **kwargs) — parent-side observation of subagent work
|
||||
backup_paths() -> list[str] — extra on-disk paths to include in `hermes backup`
|
||||
Memory providers give the agent persistent recall across sessions. Plugins
|
||||
ship in ``plugins/memory/<name>/`` and are activated via ``memory.provider``;
|
||||
MemoryManager allows only ONE external provider at a time (tool-schema bloat,
|
||||
conflicting backends). Lifecycle is driven by MemoryManager: initialize ->
|
||||
system_prompt_block / prefetch / sync_turn per turn -> tool dispatch ->
|
||||
shutdown, plus the optional ``on_*`` hooks below.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -41,43 +18,30 @@ from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Version 1 is the historical, implicit contract every provider is already
|
||||
# on: best-effort on_pre_compress() with the raw message list. Version 2 is
|
||||
# the opt-in fail-closed checkpoint contract (normalized evidence handoff +
|
||||
# strict-mode failure propagation).
|
||||
# v1 = historical implicit contract (best-effort on_pre_compress() with the raw
|
||||
# message list); v2 = opt-in fail-closed checkpoint (normalized evidence handoff
|
||||
# + strict-mode failure propagation).
|
||||
PRE_COMPRESS_CHECKPOINT_API_VERSION = 2
|
||||
|
||||
# Default glyph for the deterministic memory indicators. Providers override
|
||||
# per-status with their own brand mark (e.g. Hindsight uses "👁️").
|
||||
# Default glyph for recall indicators; providers may use their own brand mark.
|
||||
INDICATOR_GLYPH = "🧠"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RecallStatus:
|
||||
"""Summary of what a provider's most recent prefetch injected this turn.
|
||||
|
||||
Returned by :meth:`MemoryProvider.recall_status` so the agent can emit a
|
||||
deterministic, model-independent "memory was used" indicator (see
|
||||
``MemoryManager.describe_recall``). ``count`` is the number of discrete
|
||||
memories injected; ``0`` means content was injected but has no discrete
|
||||
count (e.g. a synthesized reflect answer), which the indicator renders
|
||||
generically rather than as "0 memories". ``glyph`` is the brand mark the
|
||||
indicator leads with.
|
||||
"""
|
||||
"""What the last prefetch injected, for the deterministic recall indicator
|
||||
(``MemoryManager.describe_recall``). ``count == 0`` means content without a
|
||||
discrete count (e.g. a synthesized reflect answer) and renders generically."""
|
||||
|
||||
provider_label: str
|
||||
count: int
|
||||
glyph: str = INDICATOR_GLYPH
|
||||
|
||||
|
||||
# Prompts that carry no semantic signal — trivial acknowledgements, greetings,
|
||||
# slash commands, empty input. Single source of truth shared by the core
|
||||
# per-turn prefetch gate (agent/turn_context.py, run_agent.py) and provider-
|
||||
# side classifiers (plugins/memory/honcho) so the two can never drift apart.
|
||||
# The alternation is anchored and may only be followed by whitespace or
|
||||
# punctuation, so words that merely START with a trivial word ("k8s", "yolo",
|
||||
# "note", "hindsight") do NOT match, while trailing-punctuation variants
|
||||
# ("hi!", "hey.", "thanks :)", "done???") do.
|
||||
# Prompts with no semantic signal. Single source of truth for the core prefetch
|
||||
# gate (turn_context.py, run_agent.py) and provider-side classifiers (honcho).
|
||||
# Anchored and followed only by whitespace/punctuation, so "k8s"/"yolo"/"note"
|
||||
# do NOT match while "hi!"/"thanks :)"/"done???" do.
|
||||
TRIVIAL_PROMPT_RE = re.compile(
|
||||
r'^(yes|no|ok|okay|sure|thanks|thank you|y|n|yep|nope|yeah|nah|'
|
||||
r'hi|hey|hello|yo|sup|'
|
||||
@@ -88,21 +52,13 @@ TRIVIAL_PROMPT_RE = re.compile(
|
||||
|
||||
|
||||
def is_trivial_prompt(text: Optional[str]) -> bool:
|
||||
"""Return True if a user prompt is too trivial to warrant memory recall.
|
||||
"""True for empty input, slash commands and bare greetings/acknowledgements.
|
||||
|
||||
Empty/whitespace-only input, slash commands, and bare greetings or
|
||||
acknowledgements (with optional trailing punctuation) all count as
|
||||
trivial. Callers use this to skip memory-provider prefetch/injection
|
||||
on turns that carry no semantic signal — saving a blocking network
|
||||
round-trip and preventing stale user-model context from derailing
|
||||
one-word replies.
|
||||
Skipping recall on these saves a blocking network round-trip and keeps
|
||||
stale user-model context from derailing one-word replies.
|
||||
"""
|
||||
if not text:
|
||||
return True
|
||||
stripped = text.strip()
|
||||
if not stripped:
|
||||
return True
|
||||
if stripped.startswith("/"):
|
||||
stripped = (text or "").strip()
|
||||
if not stripped or stripped.startswith("/"):
|
||||
return True
|
||||
return bool(TRIVIAL_PROMPT_RE.match(stripped))
|
||||
|
||||
@@ -110,10 +66,8 @@ def is_trivial_prompt(text: Optional[str]) -> bool:
|
||||
class MemoryProvider(ABC):
|
||||
"""Abstract base class for memory providers."""
|
||||
|
||||
# Providers that durably checkpoint every successful on_pre_compress()
|
||||
# call may opt into that host contract by setting the current version
|
||||
# (PRE_COMPRESS_CHECKPOINT_API_VERSION). Version 1 is the implicit
|
||||
# historical contract: best-effort semantics, raw message list.
|
||||
# Providers that durably checkpoint every successful on_pre_compress() opt
|
||||
# in by setting PRE_COMPRESS_CHECKPOINT_API_VERSION; 1 = best-effort legacy.
|
||||
pre_compress_checkpoint_api_version = 1
|
||||
|
||||
@property
|
||||
@@ -125,89 +79,47 @@ class MemoryProvider(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def is_available(self) -> bool:
|
||||
"""Return True if this provider is configured, has credentials, and is ready.
|
||||
|
||||
Called during agent init to decide whether to activate the provider.
|
||||
Should not make network calls — just check config and installed deps.
|
||||
"""
|
||||
"""Configured, credentialed and ready? Gates activation at agent init;
|
||||
check config/deps only — no network calls."""
|
||||
|
||||
@abstractmethod
|
||||
def initialize(self, session_id: str, **kwargs) -> None:
|
||||
"""Initialize for a session.
|
||||
"""Initialize once at agent startup (connections, resources, threads).
|
||||
|
||||
Called once at agent startup. May create resources (banks, tables),
|
||||
establish connections, start background threads, etc.
|
||||
|
||||
kwargs always include:
|
||||
- hermes_home (str): The active HERMES_HOME directory path. Use this
|
||||
for profile-scoped storage instead of hardcoding ``~/.hermes``.
|
||||
- platform (str): "cli", "telegram", "discord", "cron", etc.
|
||||
|
||||
kwargs may also include:
|
||||
- agent_context (str): "primary", "subagent", "cron", or "flush".
|
||||
Providers should skip writes for non-primary contexts (cron system
|
||||
prompts would corrupt user representations).
|
||||
- agent_identity (str): Profile name (e.g. "coder"). Use for
|
||||
per-profile provider identity scoping.
|
||||
- agent_workspace (str): Shared workspace name (e.g. "hermes").
|
||||
- parent_session_id (str): For subagents, the parent's session_id.
|
||||
- user_id (str): Platform user identifier (gateway sessions).
|
||||
- user_id_alt (str): Optional alternate stable platform user identifier.
|
||||
kwargs always include ``hermes_home`` (use it for profile-scoped storage,
|
||||
never hardcode ``~/.hermes``) and ``platform``. May include
|
||||
``agent_context`` ("primary" | "subagent" | "cron" | "flush" — skip
|
||||
writes for non-primary contexts, cron prompts would corrupt user
|
||||
representations), ``agent_identity`` (profile name), ``agent_workspace``,
|
||||
``parent_session_id``, ``user_id``, ``user_id_alt``.
|
||||
"""
|
||||
|
||||
def unavailable_reason(self) -> str:
|
||||
"""Actionable reason this provider reports unavailable, for the caller.
|
||||
|
||||
``is_available()`` gates initialization, so a provider that reports
|
||||
unavailable is never initialized — any diagnostic it would log from
|
||||
``initialize()`` is unreachable. Return a short, user-facing hint here
|
||||
(e.g. which package to install) so the caller's "provider unavailable"
|
||||
warning can surface it. Empty string (the default) adds nothing.
|
||||
"""
|
||||
"""Short user-facing hint for the "provider unavailable" warning (e.g.
|
||||
which package to install) — ``initialize()`` never runs when unavailable,
|
||||
so this is the only place such a diagnostic can surface."""
|
||||
return ""
|
||||
|
||||
def system_prompt_block(self) -> str:
|
||||
"""Return text to include in the system prompt.
|
||||
|
||||
Called during system prompt assembly. Return empty string to skip.
|
||||
This is for STATIC provider info (instructions, status). Prefetched
|
||||
recall context is injected separately via prefetch().
|
||||
"""
|
||||
"""STATIC system-prompt text (instructions, status); "" to skip.
|
||||
Recalled context goes through prefetch(), not here."""
|
||||
return ""
|
||||
|
||||
def prefetch(self, query: str, *, session_id: str = "") -> str:
|
||||
"""Recall relevant context for the upcoming turn.
|
||||
"""Formatted recall context for the upcoming turn ("" if none).
|
||||
|
||||
Called before each API call. Return formatted text to inject as
|
||||
context, or empty string if nothing relevant. Implementations
|
||||
should be fast — use background threads for the actual recall
|
||||
and return cached results here.
|
||||
|
||||
session_id is provided for providers serving concurrent sessions
|
||||
(gateway group chats, cached agents). Providers that don't need
|
||||
per-session scoping can ignore it.
|
||||
Must be fast — do the recall in the background and return cached
|
||||
results. ``session_id`` scopes concurrent sessions (gateway, cached agents).
|
||||
"""
|
||||
return ""
|
||||
|
||||
def queue_prefetch(self, query: str, *, session_id: str = "") -> None:
|
||||
"""Queue a background recall for the NEXT turn.
|
||||
|
||||
Called after each turn completes. The result will be consumed
|
||||
by prefetch() on the next turn. Default is no-op — providers
|
||||
that do background prefetching should override this.
|
||||
"""
|
||||
"""Queue a background recall after each turn; prefetch() consumes it next turn."""
|
||||
|
||||
def recall_status(self) -> Optional[RecallStatus]:
|
||||
"""Describe what the most recent :meth:`prefetch` injected, for the UI.
|
||||
|
||||
Called by the agent right after prefetch, on the same (single) turn
|
||||
thread, so it can surface a deterministic "👁️ recalled N memories"
|
||||
status line that does not depend on the model choosing to mention it.
|
||||
|
||||
Return ``None`` (the default) when this provider injected nothing this
|
||||
turn or does not want a visible indicator. Providers that override it
|
||||
must reflect only the LAST prefetch — never a stale prior count.
|
||||
"""
|
||||
"""What the most recent :meth:`prefetch` injected, for a deterministic
|
||||
"recalled N memories" indicator. ``None`` = nothing / no indicator.
|
||||
Must reflect only the LAST prefetch, never a stale prior count."""
|
||||
return None
|
||||
|
||||
def sync_turn(
|
||||
@@ -218,32 +130,16 @@ class MemoryProvider(ABC):
|
||||
session_id: str = "",
|
||||
messages: Optional[List[Dict[str, Any]]] = None,
|
||||
) -> None:
|
||||
"""Persist a completed turn to the backend.
|
||||
|
||||
Called after each turn. Should be non-blocking — queue for
|
||||
background processing if the backend has latency.
|
||||
|
||||
``messages`` is the OpenAI-style conversation message list as of the
|
||||
completed turn, including any assistant tool calls and tool results.
|
||||
Providers that do not need raw turn context can ignore it.
|
||||
"""
|
||||
"""Persist a completed turn; should be non-blocking. ``messages`` is the
|
||||
OpenAI-style list as of this turn, including tool calls/results."""
|
||||
|
||||
@abstractmethod
|
||||
def get_tool_schemas(self) -> List[Dict[str, Any]]:
|
||||
"""Return tool schemas this provider exposes.
|
||||
|
||||
Each schema follows the OpenAI function calling format:
|
||||
{"name": "...", "description": "...", "parameters": {...}}
|
||||
|
||||
Return empty list if this provider has no tools (context-only).
|
||||
"""
|
||||
"""OpenAI function-calling schemas ({"name", "description", "parameters"});
|
||||
[] for context-only providers."""
|
||||
|
||||
def handle_tool_call(self, tool_name: str, args: Dict[str, Any], **kwargs) -> str:
|
||||
"""Handle a tool call for one of this provider's tools.
|
||||
|
||||
Must return a JSON string (the tool result).
|
||||
Only called for tool names returned by get_tool_schemas().
|
||||
"""
|
||||
"""Handle one of this provider's tools; must return a JSON string."""
|
||||
raise NotImplementedError(f"Provider {self.name} does not handle tool {tool_name}")
|
||||
|
||||
def shutdown(self) -> None:
|
||||
@@ -252,23 +148,12 @@ class MemoryProvider(ABC):
|
||||
# -- Optional hooks (override to opt in) ---------------------------------
|
||||
|
||||
def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None:
|
||||
"""Called at the start of each turn with the user message.
|
||||
|
||||
Use for turn-counting, scope management, periodic maintenance.
|
||||
|
||||
kwargs may include: remaining_tokens, model, platform, tool_count.
|
||||
Providers use what they need; extras are ignored.
|
||||
"""
|
||||
"""Per-turn tick (turn-counting, scope management, maintenance).
|
||||
kwargs may include remaining_tokens, model, platform, tool_count."""
|
||||
|
||||
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
|
||||
"""Called when a session ends (explicit exit or timeout).
|
||||
|
||||
Use for end-of-session fact extraction, summarization, etc.
|
||||
messages is the full conversation history.
|
||||
|
||||
NOT called after every turn — only at actual session boundaries
|
||||
(CLI exit, /reset, gateway session expiry).
|
||||
"""
|
||||
"""End-of-session extraction over the full history. Fires only at real
|
||||
session boundaries (CLI exit, /reset, gateway expiry), never per-turn."""
|
||||
|
||||
def on_session_switch(
|
||||
self,
|
||||
@@ -279,104 +164,43 @@ class MemoryProvider(ABC):
|
||||
rewound: bool = False,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""Called when the agent switches session_id mid-process.
|
||||
"""session_id reassigned mid-process (/resume, /branch, /reset, /new,
|
||||
gateway equivalents, context compression) without a provider teardown.
|
||||
|
||||
Fires on ``/resume``, ``/branch``, ``/reset``, ``/new`` (CLI), the
|
||||
gateway equivalents, and context compression — any path that
|
||||
reassigns ``AIAgent.session_id`` without tearing the provider down.
|
||||
|
||||
Providers that cache per-session state in ``initialize()``
|
||||
(``_session_id``, ``_document_id``, accumulated turn buffers,
|
||||
counters) should update or reset that state here so subsequent
|
||||
writes land in the correct session's record.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
new_session_id:
|
||||
The session_id the agent just switched to.
|
||||
parent_session_id:
|
||||
The previous session_id, if meaningful — set for ``/branch``
|
||||
(fork lineage), context compression (continuation lineage),
|
||||
and ``/resume`` (the session we're leaving). Empty string
|
||||
when no lineage applies.
|
||||
reset:
|
||||
``True`` when this is a genuinely new conversation, not a
|
||||
resumption of an existing one. Fired by ``/reset`` / ``/new``.
|
||||
Providers should flush accumulated per-session buffers
|
||||
(``_session_turns``, ``_turn_counter``, etc.) when this is
|
||||
set. ``False`` for ``/resume`` / ``/branch`` / compression
|
||||
where the logical conversation continues under the new id.
|
||||
rewound:
|
||||
``True`` if session_id is unchanged but the transcript was
|
||||
truncated; providers caching per-turn document state should
|
||||
invalidate.
|
||||
|
||||
Default is no-op for backward compatibility.
|
||||
Update or reset any per-session state cached in ``initialize()`` so
|
||||
later writes land in the right record. ``parent_session_id`` carries
|
||||
lineage for /branch, compression and /resume ("" when none). ``reset``
|
||||
is True only for a genuinely new conversation (/reset, /new) — flush
|
||||
per-session buffers; False when the logical conversation continues
|
||||
under a new id. ``rewound``: same id but the transcript was truncated,
|
||||
so invalidate per-turn document state.
|
||||
"""
|
||||
|
||||
def on_pre_compress(self, messages: List[Dict[str, Any]]) -> str:
|
||||
"""Called before context compression discards old messages.
|
||||
|
||||
Use to extract insights from messages about to be compressed.
|
||||
messages is the list that will be summarized/discarded.
|
||||
|
||||
Return text to include in the compression summary prompt so the
|
||||
compressor preserves provider-extracted insights. Return empty
|
||||
string for no contribution (backwards-compatible default).
|
||||
"""
|
||||
"""Extract insights from ``messages`` about to be compressed; the returned
|
||||
text is fed into the compression summary prompt ("" = nothing)."""
|
||||
return ""
|
||||
|
||||
def on_delegation(self, task: str, result: str, *,
|
||||
child_session_id: str = "", **kwargs) -> None:
|
||||
"""Called on the PARENT agent when a subagent completes.
|
||||
|
||||
The parent's memory provider gets the task+result pair as an
|
||||
observation of what was delegated and what came back. The subagent
|
||||
itself has no provider session (skip_memory=True).
|
||||
|
||||
task: the delegation prompt
|
||||
result: the subagent's final response
|
||||
child_session_id: the subagent's session_id
|
||||
"""
|
||||
"""PARENT-side observation of a completed delegation (task prompt + final
|
||||
result); the subagent itself has no provider session (skip_memory=True)."""
|
||||
|
||||
def get_config_schema(self) -> List[Dict[str, Any]]:
|
||||
"""Return config fields this provider needs for setup.
|
||||
"""Setup fields for ``hermes memory setup`` ([] if none).
|
||||
|
||||
Used by 'hermes memory setup' to walk the user through configuration.
|
||||
Each field is a dict with:
|
||||
key: config key name (e.g. 'api_key', 'mode')
|
||||
description: human-readable description
|
||||
secret: True if this should go to .env (default: False)
|
||||
required: True if required (default: False)
|
||||
default: default value (optional)
|
||||
choices: list of valid values (optional)
|
||||
type: text, integer, number, or boolean (optional)
|
||||
minimum: numeric lower bound for integer/number fields (optional)
|
||||
maximum: numeric upper bound for integer/number fields (optional)
|
||||
step: numeric input step for Dashboard rendering (optional)
|
||||
url: URL where user can get this credential (optional)
|
||||
env_var: explicit env var name for secrets (default: auto-generated)
|
||||
|
||||
Return empty list if no config needed (e.g. local-only providers).
|
||||
Each field: ``key``, ``description``, optional ``secret`` (goes to .env),
|
||||
``required``, ``default``, ``choices``, ``type`` (text | integer |
|
||||
number | boolean), ``minimum`` / ``maximum`` / ``step`` (numeric,
|
||||
Dashboard rendering), ``url`` (where to get the credential), ``env_var``
|
||||
(explicit secret env var; default auto-generated).
|
||||
"""
|
||||
return []
|
||||
|
||||
def save_config(self, values: Dict[str, Any], hermes_home: str) -> None:
|
||||
"""Write non-secret config to the provider's native location.
|
||||
|
||||
Called by 'hermes memory setup' after collecting user inputs.
|
||||
``values`` contains only non-secret fields (secrets go to .env).
|
||||
``hermes_home`` is the active HERMES_HOME directory path.
|
||||
|
||||
Providers with native config files (JSON, YAML) should override
|
||||
this to write to their expected location. Providers that use only
|
||||
env vars can leave the default (no-op).
|
||||
|
||||
All new memory provider plugins MUST implement either:
|
||||
- save_config() for native config file formats, OR
|
||||
- use only env vars (in which case get_config_schema() fields
|
||||
should all have ``env_var`` set and this method stays no-op).
|
||||
"""
|
||||
"""Write non-secret setup ``values`` (secrets go to .env) to the provider's
|
||||
native config location. Plugins MUST either override this or use only
|
||||
env vars (every schema field carrying ``env_var``) and keep the no-op."""
|
||||
|
||||
def on_memory_write(
|
||||
self,
|
||||
@@ -385,32 +209,14 @@ class MemoryProvider(ABC):
|
||||
content: str,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""Called when the built-in memory tool writes an entry.
|
||||
|
||||
action: 'add', 'replace', or 'remove'
|
||||
target: 'memory' or 'user'
|
||||
content: the entry content
|
||||
metadata: structured provenance for the write, when available. Common
|
||||
keys include ``write_origin``, ``execution_context``, ``session_id``,
|
||||
``parent_session_id``, ``platform``, and ``tool_name``.
|
||||
|
||||
Use to mirror built-in memory writes to your backend.
|
||||
"""
|
||||
"""Mirror a built-in memory-tool write. ``action`` is add | replace |
|
||||
remove, ``target`` is memory | user; ``metadata`` (when available) has
|
||||
provenance such as write_origin, execution_context, session_id,
|
||||
parent_session_id, platform, tool_name."""
|
||||
|
||||
def backup_paths(self) -> List[str]:
|
||||
"""Return extra on-disk paths this provider stores OUTSIDE HERMES_HOME.
|
||||
|
||||
``hermes backup`` only walks HERMES_HOME, so any provider state kept
|
||||
under ``~/.honcho``, ``~/.hindsight``, ``~/.openviking``, etc. is lost
|
||||
across a backup/import cycle unless it's declared here.
|
||||
|
||||
Return a list of absolute path strings (files or directories). The
|
||||
backup command resolves each, captures the ones that exist and live
|
||||
under the user's home directory into a reserved ``_external/`` subtree
|
||||
of the archive, and ``hermes import`` restores them to their original
|
||||
locations. Paths outside the home directory are skipped for safety.
|
||||
|
||||
MUST be callable without ``initialize()`` and without network — resolve
|
||||
from config/env only. Default returns an empty list (nothing external).
|
||||
"""
|
||||
"""Absolute paths of provider state OUTSIDE HERMES_HOME (e.g. ``~/.honcho``)
|
||||
so ``hermes backup`` can capture them under ``_external/`` and
|
||||
``hermes import`` restore them; paths outside the home dir are skipped.
|
||||
MUST work without ``initialize()`` or network — resolve from config/env."""
|
||||
return []
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Shared base classes for the pluggable-backend provider ABCs.
|
||||
|
||||
Every tool-provider ABC (browser, TTS, image/video gen, transcription, web
|
||||
search, terminal env) shares the same identity + ``hermes tools`` picker
|
||||
surface; the previous per-ABC copies of these defaults were byte-identical.
|
||||
Concrete ABCs subclass :class:`ProviderBase` (or :class:`CatalogProviderBase`
|
||||
when the backend also exposes a model catalog and is available by default) and
|
||||
add only their domain methods. Plugins keep subclassing the concrete ABC, so
|
||||
``isinstance`` checks and abstract-method sets are unchanged.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
class ProviderBase(abc.ABC):
|
||||
"""Identity + picker metadata common to every provider ABC."""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used as the provider's config-key value.
|
||||
|
||||
Lowercase, no spaces (hyphens allowed where they preserve an existing
|
||||
user-visible name). Registries key providers by this string.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name``."""
|
||||
return self.name
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Provider row for the ``hermes tools`` picker.
|
||||
|
||||
Shape: ``{"name", "badge", "tag", "env_vars": [{"key", "prompt", "url"}, ...]}``
|
||||
(browser providers may add ``"post_setup"``). Default: a minimal entry
|
||||
derived from ``display_name`` with no env vars — override to expose API
|
||||
key prompts and badges.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
|
||||
class CatalogProviderBase(ProviderBase):
|
||||
"""Provider with a model catalog; available by default, ``display_name`` is title-cased."""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name.title()``."""
|
||||
return self.name.title()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""True when this provider can service calls (API key present, SDK importable).
|
||||
|
||||
Default True. Must NOT raise and must NOT make network calls — the picker
|
||||
and ``hermes setup`` call it on every paint.
|
||||
"""
|
||||
return True
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
"""Model catalog entries (``{"id": ..., "display": ...}`` plus optional
|
||||
provider-specific keys). Default: empty (no user-selectable models)."""
|
||||
return []
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
"""Id of the first catalog entry, or None when the catalog is empty."""
|
||||
models = self.list_models()
|
||||
if models:
|
||||
return models[0].get("id")
|
||||
return None
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Shared ``$HERMES_HOME/cache/<kind>/`` materialisation helpers for the
|
||||
image/video generation provider ABCs.
|
||||
|
||||
Several backends (xAI, OpenAI, DeepInfra, FAL) return *ephemeral* delivery URLs
|
||||
that expire before a downstream consumer (Telegram ``send_photo``, browser
|
||||
fetch) can resolve them, so providers materialise the bytes locally at
|
||||
tool-completion time. Filenames are ``<prefix>_<YYYYMMDD_HHMMSS>_<uuid8>.<ext>``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import datetime
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Dict, Tuple
|
||||
|
||||
|
||||
def cache_dir(kind: str) -> Path:
|
||||
"""Return ``$HERMES_HOME/cache/<kind>/``, creating parents as needed."""
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
path = get_hermes_home() / "cache" / kind
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def cache_path(kind: str, prefix: str, extension: str) -> Path:
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
return cache_dir(kind) / f"{prefix}_{ts}_{short}.{extension}"
|
||||
|
||||
|
||||
def save_bytes(kind: str, raw: bytes, *, prefix: str, extension: str) -> Path:
|
||||
"""Write raw bytes to the cache and return the absolute path."""
|
||||
path = cache_path(kind, prefix, extension)
|
||||
path.write_bytes(raw)
|
||||
return path
|
||||
|
||||
|
||||
def save_b64(kind: str, b64_data: str, *, prefix: str, extension: str) -> Path:
|
||||
"""Decode base64 data into the cache and return the absolute path."""
|
||||
return save_bytes(kind, base64.b64decode(b64_data), prefix=prefix, extension=extension)
|
||||
|
||||
|
||||
def save_url(
|
||||
kind: str,
|
||||
url: str,
|
||||
*,
|
||||
prefix: str,
|
||||
timeout: float,
|
||||
max_bytes: int,
|
||||
chunk_size: int,
|
||||
content_types: Dict[str, str],
|
||||
url_extensions: Tuple[str, ...],
|
||||
default_extension: str,
|
||||
label: str,
|
||||
empty_error: str,
|
||||
) -> Path:
|
||||
"""Stream-download *url* into the cache with a size cap.
|
||||
|
||||
The extension comes from the response ``Content-Type`` (a small explicit
|
||||
table — never inherit a type that points at HTML/JSON from a degenerate
|
||||
response), then the URL suffix (some CDNs return
|
||||
``application/octet-stream``), then *default_extension*. Raises on any
|
||||
network / HTTP / oversize / empty error so callers can fall back to the bare
|
||||
URL; a partial file is never left behind.
|
||||
"""
|
||||
import requests
|
||||
|
||||
response = requests.get(url, timeout=timeout, stream=True)
|
||||
response.raise_for_status()
|
||||
|
||||
content_type = (response.headers.get("Content-Type") or "").split(";", 1)[0].strip().lower()
|
||||
extension = content_types.get(content_type)
|
||||
if extension is None:
|
||||
url_path = url.split("?", 1)[0].lower()
|
||||
for ext in url_extensions:
|
||||
if url_path.endswith(f".{ext}"):
|
||||
extension = "jpg" if ext == "jpeg" else ext
|
||||
break
|
||||
if extension is None:
|
||||
extension = default_extension
|
||||
|
||||
path = cache_path(kind, prefix, extension)
|
||||
|
||||
bytes_written = 0
|
||||
with path.open("wb") as fh:
|
||||
for chunk in response.iter_content(chunk_size=chunk_size):
|
||||
if not chunk:
|
||||
continue
|
||||
bytes_written += len(chunk)
|
||||
if bytes_written > max_bytes:
|
||||
fh.close()
|
||||
_unlink_quiet(path)
|
||||
raise ValueError(
|
||||
f"{label} at {url} exceeds {max_bytes // (1024 * 1024)}MB cap; refusing to cache."
|
||||
)
|
||||
fh.write(chunk)
|
||||
|
||||
if bytes_written == 0:
|
||||
_unlink_quiet(path)
|
||||
raise ValueError(empty_error.format(url=url))
|
||||
|
||||
return path
|
||||
|
||||
|
||||
def _unlink_quiet(path: Path) -> None:
|
||||
try:
|
||||
path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
@@ -0,0 +1,244 @@
|
||||
"""Shared engine behind the ``agent.*_registry`` provider registries.
|
||||
|
||||
Every pluggable-backend registry (browser, TTS, image/video gen, transcription,
|
||||
web search, terminal env) has the same shape: a global name->provider map plus
|
||||
per-profile *scoped* maps (multiplexed gateways), a lock, registration with
|
||||
re-registration logging, and the snapshot/restore pair that
|
||||
:mod:`hermes_cli.plugins` uses to unwind a plugin's registrations. Each
|
||||
``*_registry`` module instantiates one :class:`ProviderRegistry` and re-exports
|
||||
its bound methods under the historical module-level names via
|
||||
:meth:`ProviderRegistry.export`, so call sites, ``patch("agent.x_registry.get_provider")``
|
||||
targets, and the ``_providers`` / ``_scoped_providers`` / ``_lock`` test hooks are unchanged.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Any, Callable, Dict, FrozenSet, Generic, List, Optional, TypeVar
|
||||
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
P = TypeVar("P")
|
||||
|
||||
|
||||
def strip_key(name: str) -> str:
|
||||
return name.strip()
|
||||
|
||||
|
||||
def lower_key(name: str) -> str:
|
||||
return name.strip().lower()
|
||||
|
||||
|
||||
class ProviderRegistry(Generic[P]):
|
||||
"""Global + per-scope provider map with plugin snapshot/restore support.
|
||||
|
||||
Args:
|
||||
label: Human label used in log/error strings (``"Browser"``, ``"TTS"``).
|
||||
provider_cls: ABC every registered instance must satisfy (TypeError otherwise).
|
||||
logger: The owning module's logger, so record names stay per-registry.
|
||||
normalize: Key normalizer — ``strip_key`` or ``lower_key`` (case-insensitive
|
||||
registries mirror how their dispatcher normalizes the configured name).
|
||||
builtin_names: Reserved names owned by in-tree implementations; a collision
|
||||
calls ``on_builtin_collision(key)`` and, if that returns, skips registration.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
label: str,
|
||||
provider_cls: type,
|
||||
logger: logging.Logger,
|
||||
normalize: Callable[[str], str] = strip_key,
|
||||
builtin_names: FrozenSet[str] = frozenset(),
|
||||
on_builtin_collision: Optional[Callable[[str], None]] = None,
|
||||
) -> None:
|
||||
self.label = label
|
||||
self.provider_cls = provider_cls
|
||||
self.logger = logger
|
||||
self.normalize = normalize
|
||||
self.builtin_names = builtin_names
|
||||
self._on_builtin_collision = on_builtin_collision
|
||||
self._providers: Dict[str, P] = {}
|
||||
self._scoped_providers: Dict[str, Dict[str, P]] = {}
|
||||
self._generation = 0
|
||||
self._scoped_generations: Dict[str, int] = {}
|
||||
self._lock = threading.Lock()
|
||||
# "TTS provider" but "Registered browser provider": acronyms keep their case.
|
||||
self._log_label = label if label.isupper() else label[0].lower() + label[1:]
|
||||
|
||||
# -- internal helpers (caller holds the lock) ---------------------------
|
||||
|
||||
def _target(self, scope: Optional[str], *, create: bool) -> Dict[str, P]:
|
||||
if scope is None:
|
||||
return self._providers
|
||||
if create:
|
||||
return self._scoped_providers.setdefault(scope, {})
|
||||
return self._scoped_providers.get(scope, {})
|
||||
|
||||
def _bump(self, scope: Optional[str]) -> None:
|
||||
if scope is None:
|
||||
self._generation += 1
|
||||
else:
|
||||
self._scoped_generations[scope] = self._scoped_generations.get(scope, 0) + 1
|
||||
|
||||
# -- registration -------------------------------------------------------
|
||||
|
||||
def register(self, provider: P, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a provider; same-name re-registration overwrites (hot reload)."""
|
||||
if not isinstance(provider, self.provider_cls):
|
||||
article = "an" if self.provider_cls.__name__[0] in "AEIOU" else "a"
|
||||
raise TypeError(
|
||||
f"register_provider() expects {article} {self.provider_cls.__name__} "
|
||||
f"instance, got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = getattr(provider, "name")
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError(f"{self.label} provider .name must be a non-empty string")
|
||||
key = self.normalize(raw_name)
|
||||
if key in self.builtin_names:
|
||||
if self._on_builtin_collision is not None:
|
||||
self._on_builtin_collision(key)
|
||||
return
|
||||
with self._lock:
|
||||
target = self._target(scope, create=True)
|
||||
existing = target.get(key)
|
||||
target[key] = provider
|
||||
self._bump(scope)
|
||||
if existing is not None:
|
||||
self.logger.debug(
|
||||
f"{self.label} provider '%s' re-registered (was %r)",
|
||||
key, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
self.logger.debug(
|
||||
f"Registered {self._log_label} provider '%s' (%s)",
|
||||
key, type(provider).__name__,
|
||||
)
|
||||
|
||||
# -- lookup ---------------------------------------------------------------
|
||||
|
||||
def merged(self, scope: Optional[str] = None) -> Dict[str, P]:
|
||||
"""Global map overlaid with the active profile's scoped map (a copy)."""
|
||||
with self._lock:
|
||||
merged = dict(self._providers)
|
||||
merged.update(self._scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
return merged
|
||||
|
||||
def list_providers(self, *, scope: Optional[str] = None) -> List[P]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
return sorted(self.merged(scope).values(), key=lambda p: p.name)
|
||||
|
||||
def get_provider(self, name: str, *, scope: Optional[str] = None) -> Optional[P]:
|
||||
"""Return the provider registered under *name* (scoped first), or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
key = self.normalize(name)
|
||||
with self._lock:
|
||||
return (
|
||||
self._scoped_providers.get(scope or hermes_home_key(), {}).get(key)
|
||||
or self._providers.get(key)
|
||||
)
|
||||
|
||||
def registry_generation(self, *, scope: Optional[str] = None) -> tuple:
|
||||
"""Cache fingerprint ``(global_generation, scoped_generation)``."""
|
||||
active_scope = scope or hermes_home_key()
|
||||
with self._lock:
|
||||
return self._generation, self._scoped_generations.get(active_scope, 0)
|
||||
|
||||
# -- plugin unload support (hermes_cli.plugins) -----------------------------
|
||||
|
||||
def snapshot_registration(self, name: str, *, scope: Optional[str] = None) -> Optional[P]:
|
||||
"""Exact-slot lookup (no global fallback) used to detect plugin ownership."""
|
||||
with self._lock:
|
||||
return self._target(scope, create=False).get(self.normalize(name))
|
||||
|
||||
def restore_registration(
|
||||
self, name: str, current: P, previous: Optional[P], *, scope: Optional[str] = None
|
||||
) -> bool:
|
||||
"""Restore *previous* only when *current* is still installed under *name*."""
|
||||
key = self.normalize(name)
|
||||
with self._lock:
|
||||
target = self._target(scope, create=True)
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
self._bump(scope)
|
||||
if scope is not None and not target:
|
||||
self._scoped_providers.pop(scope, None)
|
||||
return True
|
||||
|
||||
def reset_for_tests(self) -> None:
|
||||
"""Clear every registration. **Test-only.**"""
|
||||
with self._lock:
|
||||
self._providers.clear()
|
||||
self._scoped_providers.clear()
|
||||
self._scoped_generations.clear()
|
||||
self._generation += 1
|
||||
|
||||
def export(self, namespace: Dict[str, Any]) -> None:
|
||||
"""Bind the historical module-level API into a ``*_registry`` module.
|
||||
|
||||
Installs ``register_provider``/``list_providers``/``get_provider``/
|
||||
``snapshot_registration``/``restore_registration``/``registry_generation``/
|
||||
``_reset_for_tests`` plus the ``_providers``/``_scoped_providers``/``_lock``
|
||||
test hooks, so ``patch("agent.x_registry.get_provider")`` and direct
|
||||
``_providers`` manipulation in tests keep working unchanged.
|
||||
"""
|
||||
namespace.update(
|
||||
_providers=self._providers,
|
||||
_scoped_providers=self._scoped_providers,
|
||||
_lock=self._lock,
|
||||
register_provider=self.register,
|
||||
list_providers=self.list_providers,
|
||||
get_provider=self.get_provider,
|
||||
snapshot_registration=self.snapshot_registration,
|
||||
restore_registration=self.restore_registration,
|
||||
registry_generation=self.registry_generation,
|
||||
_reset_for_tests=self.reset_for_tests,
|
||||
)
|
||||
|
||||
|
||||
def is_available_safe(
|
||||
provider: Any,
|
||||
logger: logging.Logger,
|
||||
fmt: str,
|
||||
*,
|
||||
level: int = logging.DEBUG,
|
||||
exc_info: bool = False,
|
||||
) -> bool:
|
||||
"""``bool(provider.is_available())`` that treats a raising provider as unavailable."""
|
||||
try:
|
||||
return bool(provider.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.log(level, fmt, provider.name, exc, exc_info=exc_info)
|
||||
return False
|
||||
|
||||
|
||||
def configured_provider_name(section: str, logger: logging.Logger) -> Optional[str]:
|
||||
"""Read ``<section>.provider`` from config.yaml, mapping the managed Nous
|
||||
selection to ``fal`` (the FAL plugin services it via the managed gateway)."""
|
||||
configured: Optional[str] = None
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
|
||||
cfg = load_config_readonly()
|
||||
block = cfg.get(section) if isinstance(cfg, dict) else None
|
||||
if isinstance(block, dict):
|
||||
raw = block.get("provider")
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
configured = raw.strip()
|
||||
except Exception as exc:
|
||||
logger.debug("Could not read %s.provider from config: %s", section, exc)
|
||||
if configured:
|
||||
try:
|
||||
from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER
|
||||
|
||||
if configured.lower() == NOUS_MANAGED_PROVIDER:
|
||||
configured = "fal"
|
||||
except Exception: # pragma: no cover — helpers are in-repo
|
||||
pass
|
||||
return configured
|
||||
+60
-132
@@ -2,65 +2,42 @@
|
||||
Terminal Environment Provider ABC
|
||||
=================================
|
||||
|
||||
Defines the pluggable-backend interface for terminal execution environments
|
||||
(cloud sandboxes, remote runners). Providers register instances via
|
||||
:meth:`PluginContext.register_terminal_environment_provider`; the dispatch
|
||||
ladder in :func:`tools.terminal_tool._create_environment` consults the
|
||||
registry for any ``TERMINAL_ENV`` / ``terminal.backend`` value that is not a
|
||||
built-in backend (local, docker, singularity, modal, daytona,
|
||||
vercel_sandbox, ssh).
|
||||
|
||||
Providers live in ``~/.hermes/plugins/<name>/`` (user, opt-in via
|
||||
``plugins.enabled``) or ship as standalone plugin repos. Built-in backends
|
||||
stay in-tree under ``tools/environments/`` — this extension point exists so
|
||||
third-party sandbox vendors do NOT have to live in core (see AGENTS.md:
|
||||
"Third-party products ... ship them as a standalone plugin repo").
|
||||
|
||||
This ABC mirrors :class:`agent.browser_provider.BrowserProvider` — same
|
||||
registration flow, same scope semantics, same plugin-context gating.
|
||||
Pluggable-backend interface for terminal execution environments (cloud
|
||||
sandboxes, remote runners). Providers register via
|
||||
:meth:`PluginContext.register_terminal_environment_provider`;
|
||||
:func:`tools.terminal_tool._create_environment` consults the registry for any
|
||||
``TERMINAL_ENV`` / ``terminal.backend`` value that is not a built-in backend.
|
||||
Built-ins stay in-tree under ``tools/environments/``; this extension point
|
||||
exists so third-party sandbox vendors do NOT have to live in core.
|
||||
|
||||
Classification contract
|
||||
-----------------------
|
||||
A backend participates in core policy decisions that were historically
|
||||
frozensets of built-in names. Each is a declarative attribute so a new backend
|
||||
cannot silently miss a classification site:
|
||||
|
||||
Beyond creating environments, a terminal backend participates in several
|
||||
core policy decisions that were historically frozensets of built-in names.
|
||||
Each is expressed as a declarative attribute so the class of "new backend
|
||||
missed classification site N" bugs (see PR #30112's seven-site sweep) cannot
|
||||
recur for plugin backends:
|
||||
|
||||
* ``is_remote`` — commands run somewhere other than the host machine.
|
||||
Suppresses host OS/home/cwd hints in the system prompt, the host Python
|
||||
env probe, and remote-aware skill env handling.
|
||||
* ``is_container`` — the backend behaves like a container/sandbox with its
|
||||
own filesystem rooted away from the host: container resource config is
|
||||
passed through, host-looking cwds are sanitized, and file tools use
|
||||
container path resolution.
|
||||
* ``skip_container_guards`` — the sandbox is isolated enough that
|
||||
dangerous-command approval prompts are skipped (a wiped filesystem is
|
||||
disposable). Defaults to ``is_container``. Backends that can mount host
|
||||
paths should override to ``False``.
|
||||
* ``cache_path_base`` — where auto-synced ``~/.hermes/cache`` files land
|
||||
inside the backend (e.g. ``"~/.hermes"`` for home-synced backends,
|
||||
``"/root/.hermes"`` for root-homed containers), or ``None`` when host
|
||||
paths remain correct (nothing is translated).
|
||||
* ``strip_env_keys`` — credential env var names owned by this backend
|
||||
(API tokens for the sandbox vendor). Stripped from every subprocess the
|
||||
agent spawns so a model-authored command can never read them.
|
||||
* ``is_remote`` — commands run off-host: suppresses host OS/home/cwd hints in
|
||||
the system prompt, the host Python probe, and remote-aware skill env handling.
|
||||
* ``is_container`` — own filesystem rooted away from the host: container
|
||||
resource config is passed through, host-looking cwds are sanitized, file
|
||||
tools use container path resolution.
|
||||
* ``skip_container_guards`` — sandbox is disposable enough to skip
|
||||
dangerous-command approval prompts. Defaults to ``is_container``; backends
|
||||
that can mount host paths should override to ``False``.
|
||||
* ``cache_path_base`` — where auto-synced ``~/.hermes/cache`` files land inside
|
||||
the backend (``"~/.hermes"``, ``"/root/.hermes"``), or ``None`` when host
|
||||
paths remain correct.
|
||||
* ``strip_env_keys`` — vendor credential env vars, stripped from every
|
||||
subprocess the agent spawns so a model-authored command can never read them.
|
||||
* ``session_isolated_when_nonpersistent`` — non-persistent mode gives each
|
||||
session its own sandbox identity instead of sharing one (the #82731
|
||||
contract; opt in when a shared name would let two ephemeral runs attach
|
||||
and destroy each other's sandbox).
|
||||
session its own sandbox identity; opt in when a shared name would let two
|
||||
ephemeral runs attach to and destroy each other's sandbox.
|
||||
|
||||
Environment object contract
|
||||
---------------------------
|
||||
|
||||
:meth:`create_environment` returns an object satisfying the same duck-typed
|
||||
interface as :class:`tools.environments.base.BaseEnvironment` (``execute()``,
|
||||
``cleanup()`` …). Subclassing ``BaseEnvironment`` is recommended but not
|
||||
required — the registry does not isinstance-check the returned environment.
|
||||
The factory stamps ``_hermes_backend_name`` on the returned object so
|
||||
file-path resolution can identify plugin backends without class-name
|
||||
sniffing.
|
||||
:meth:`create_environment` returns any object satisfying the
|
||||
:class:`tools.environments.base.BaseEnvironment` duck-typed interface
|
||||
(``execute()``, ``cleanup()`` …); the registry does not isinstance-check it.
|
||||
The factory stamps ``_hermes_backend_name`` on the result so file-path
|
||||
resolution can identify plugin backends without class-name sniffing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -68,38 +45,22 @@ from __future__ import annotations
|
||||
import abc
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from agent.provider_base import ProviderBase
|
||||
|
||||
class TerminalEnvironmentProvider(abc.ABC):
|
||||
"""Abstract base class for a pluggable terminal execution backend."""
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Identity
|
||||
# ------------------------------------------------------------------
|
||||
class TerminalEnvironmentProvider(ProviderBase):
|
||||
"""Abstract base class for a pluggable terminal execution backend.
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used as the ``terminal.backend`` /
|
||||
``TERMINAL_ENV`` value.
|
||||
|
||||
Lowercase, ``[a-z0-9_]``. Must not collide with a built-in backend
|
||||
name (local, docker, singularity, modal, managed_modal, daytona,
|
||||
vercel_sandbox, ssh) — the registry rejects such registrations.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label for pickers. Defaults to ``name``."""
|
||||
return self.name
|
||||
:attr:`name` is the ``terminal.backend`` / ``TERMINAL_ENV`` value
|
||||
(``[a-z0-9_]``); the registry rejects built-in backend names.
|
||||
"""
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
"""One-line description shown in backend pickers."""
|
||||
return f"Run commands in a {self.display_name} environment."
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Classification flags (see module docstring)
|
||||
# ------------------------------------------------------------------
|
||||
# -- Classification flags (see module docstring) -----------------------
|
||||
|
||||
is_remote: bool = True
|
||||
is_container: bool = True
|
||||
@@ -122,63 +83,42 @@ class TerminalEnvironmentProvider(abc.ABC):
|
||||
|
||||
@property
|
||||
def env_description(self) -> str:
|
||||
"""Prompt-builder fallback description of where commands run.
|
||||
|
||||
Used when the live backend probe fails at system-prompt build time,
|
||||
e.g. ``"a Daytona workspace (Linux)"``.
|
||||
"""
|
||||
"""Prompt-builder fallback for where commands run when the live backend
|
||||
probe fails at system-prompt build time (e.g. ``"a Daytona workspace (Linux)"``)."""
|
||||
return f"a {self.display_name} environment (likely Linux)"
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Availability / setup UX
|
||||
# ------------------------------------------------------------------
|
||||
# -- Availability / setup UX -------------------------------------------
|
||||
|
||||
@abc.abstractmethod
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this backend can service commands.
|
||||
|
||||
Cheap check only (env var present, SDK importable). Must NOT make
|
||||
network calls — this runs during requirement checks and UI paints.
|
||||
"""
|
||||
"""True when this backend can service commands. Cheap check only — must
|
||||
NOT make network calls; runs during requirement checks and UI paints."""
|
||||
|
||||
def check_requirements(self, config: Dict[str, Any]) -> bool:
|
||||
"""Full requirements check for :func:`check_terminal_requirements`.
|
||||
|
||||
``config`` is the merged terminal env config dict. Default defers to
|
||||
:meth:`is_available`. Log actionable errors before returning False.
|
||||
"""
|
||||
"""Full requirements check for :func:`check_terminal_requirements` with the
|
||||
merged terminal env config. Default defers to :meth:`is_available`; log
|
||||
actionable errors before returning False."""
|
||||
return self.is_available()
|
||||
|
||||
def probe(self) -> Tuple[str, str]:
|
||||
"""Dashboard picker health probe: ``(status, detail)``.
|
||||
|
||||
``status`` is ``"ready"`` / ``"needs_setup"`` / ``"unavailable"``;
|
||||
``detail`` carries setup guidance for non-ready rows. Must never
|
||||
raise and must stay fast (<~2s).
|
||||
"""
|
||||
"""Dashboard picker health probe ``(status, detail)`` with status in
|
||||
``ready`` / ``needs_setup`` / ``unavailable``. Must never raise; stay fast (<~2s)."""
|
||||
if self.is_available():
|
||||
return ("ready", "")
|
||||
return ("needs_setup", f"{self.display_name} is not configured.")
|
||||
|
||||
def setup_instructions(self) -> List[str]:
|
||||
"""Lines printed by ``hermes setup`` after this backend is selected.
|
||||
|
||||
Use for token acquisition hints, SDK install commands, etc. The
|
||||
wizard persists ``terminal.backend`` itself; providers that need an
|
||||
interactive flow can run it in :meth:`post_setup`.
|
||||
"""
|
||||
"""Lines printed by ``hermes setup`` after this backend is selected. The
|
||||
wizard persists ``terminal.backend`` itself; interactive flows go in
|
||||
:meth:`post_setup`."""
|
||||
return []
|
||||
|
||||
def post_setup(self) -> None:
|
||||
"""Optional interactive setup hook run by ``hermes setup`` after the
|
||||
backend is selected (prompt for tokens, install SDKs). Default no-op.
|
||||
"""
|
||||
"""Optional interactive hook run by ``hermes setup`` after selection
|
||||
(prompt for tokens, install SDKs). Default no-op."""
|
||||
|
||||
def doctor_checks(self) -> List[Tuple[bool, str, str]]:
|
||||
"""``hermes doctor`` rows: ``(ok, label, detail)`` triples.
|
||||
|
||||
Default: a single row reflecting :meth:`is_available`.
|
||||
"""
|
||||
"""``hermes doctor`` rows ``(ok, label, detail)``; default reflects :meth:`is_available`."""
|
||||
ok = False
|
||||
try:
|
||||
ok = bool(self.is_available())
|
||||
@@ -187,9 +127,7 @@ class TerminalEnvironmentProvider(abc.ABC):
|
||||
detail = "(configured)" if ok else "(not configured — see setup instructions)"
|
||||
return [(ok, f"{self.display_name} backend", detail)]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# The factory
|
||||
# ------------------------------------------------------------------
|
||||
# -- The factory -------------------------------------------------------
|
||||
|
||||
@abc.abstractmethod
|
||||
def create_environment(
|
||||
@@ -202,21 +140,11 @@ class TerminalEnvironmentProvider(abc.ABC):
|
||||
container_config: Optional[Dict[str, Any]] = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""Create and return an execution environment instance.
|
||||
"""Create and return an execution environment (``BaseEnvironment`` duck type).
|
||||
|
||||
MUST accept ``**kwargs`` and ignore unknown keys — the forward-compat
|
||||
contract that lets the factory signature evolve without breaking
|
||||
older plugins.
|
||||
|
||||
Args:
|
||||
cwd: Working directory inside the backend.
|
||||
timeout: Default per-command timeout in seconds.
|
||||
task_id: Task identifier for environment reuse/persistence keying.
|
||||
image: Configured container image name (may be irrelevant).
|
||||
container_config: Resource config dict (``container_cpu``,
|
||||
``container_memory``, ``container_disk``,
|
||||
``container_persistent``) when :attr:`is_container` is True.
|
||||
|
||||
Returns:
|
||||
An object satisfying the ``BaseEnvironment`` duck-typed contract.
|
||||
MUST accept ``**kwargs`` and ignore unknown keys so the factory signature
|
||||
can evolve without breaking older plugins. ``task_id`` keys environment
|
||||
reuse/persistence; ``container_config`` carries ``container_cpu`` /
|
||||
``container_memory`` / ``container_disk`` / ``container_persistent`` when
|
||||
:attr:`is_container` is True.
|
||||
"""
|
||||
|
||||
+21
-133
@@ -26,11 +26,10 @@ into a per-profile scope (multiplexed gateways) or the global base map.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import List, Optional
|
||||
|
||||
from agent.provider_registry import ProviderRegistry, lower_key
|
||||
from agent.terminal_env_provider import TerminalEnvironmentProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -43,86 +42,27 @@ BUILTIN_BACKEND_NAMES = frozenset({
|
||||
})
|
||||
|
||||
|
||||
_providers: Dict[str, TerminalEnvironmentProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, TerminalEnvironmentProvider]] = {}
|
||||
_generation = 0
|
||||
_scoped_generations: Dict[str, int] = {}
|
||||
_lock = threading.Lock()
|
||||
def _reject_builtin_collision(name: str) -> None:
|
||||
raise ValueError(
|
||||
f"Terminal backend name '{name}' is reserved for the built-in "
|
||||
f"{name} backend and cannot be registered by a plugin"
|
||||
)
|
||||
|
||||
|
||||
def register_provider(
|
||||
provider: TerminalEnvironmentProvider, *, scope: Optional[str] = None
|
||||
) -> None:
|
||||
"""Register a terminal environment provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry — makes
|
||||
hot-reload scenarios (tests, dev loops) behave predictably.
|
||||
|
||||
Raises:
|
||||
TypeError: not a TerminalEnvironmentProvider instance.
|
||||
ValueError: empty name or collision with a built-in backend name.
|
||||
"""
|
||||
if not isinstance(provider, TerminalEnvironmentProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a TerminalEnvironmentProvider "
|
||||
f"instance, got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Terminal environment provider .name must be a non-empty string")
|
||||
name = raw_name.strip().lower()
|
||||
if name in BUILTIN_BACKEND_NAMES:
|
||||
raise ValueError(
|
||||
f"Terminal backend name '{name}' is reserved for the built-in "
|
||||
f"{name} backend and cannot be registered by a plugin"
|
||||
)
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"Terminal environment provider '%s' re-registered (was %r)",
|
||||
name, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered terminal environment provider '%s' (%s)",
|
||||
name, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[TerminalEnvironmentProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[TerminalEnvironmentProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
return (
|
||||
_scoped_providers.get(scope or hermes_home_key(), {}).get(key)
|
||||
or _providers.get(key)
|
||||
)
|
||||
_registry: ProviderRegistry[TerminalEnvironmentProvider] = ProviderRegistry(
|
||||
label="Terminal environment",
|
||||
provider_cls=TerminalEnvironmentProvider,
|
||||
logger=logger,
|
||||
normalize=lower_key,
|
||||
builtin_names=BUILTIN_BACKEND_NAMES,
|
||||
on_builtin_collision=_reject_builtin_collision,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
def plugin_backend_names(*, scope: Optional[str] = None) -> List[str]:
|
||||
"""Names of all registered plugin backends (sorted)."""
|
||||
return [p.name.strip().lower() for p in list_providers(scope=scope)]
|
||||
return [p.name.strip().lower() for p in _registry.list_providers(scope=scope)]
|
||||
|
||||
|
||||
def provider_flag(name: str, attr: str, default=False):
|
||||
@@ -132,7 +72,7 @@ def provider_flag(name: str, attr: str, default=False):
|
||||
misbehaving plugin degrades to built-in-equivalent behavior instead of
|
||||
taking the terminal tool down.
|
||||
"""
|
||||
provider = get_provider(name)
|
||||
provider = _registry.get_provider(name)
|
||||
if provider is None:
|
||||
return default
|
||||
try:
|
||||
@@ -154,9 +94,9 @@ def plugin_strip_env_keys() -> frozenset:
|
||||
the static tier-1 set unconditionally).
|
||||
"""
|
||||
keys: set = set()
|
||||
with _lock:
|
||||
all_providers = list(_providers.values())
|
||||
for scoped in _scoped_providers.values():
|
||||
with _registry._lock:
|
||||
all_providers = list(_registry._providers.values())
|
||||
for scoped in _registry._scoped_providers.values():
|
||||
all_providers.extend(scoped.values())
|
||||
for provider in all_providers:
|
||||
try:
|
||||
@@ -167,55 +107,3 @@ def plugin_strip_env_keys() -> frozenset:
|
||||
exc_info=True,
|
||||
)
|
||||
return frozenset(keys)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[TerminalEnvironmentProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip().lower())
|
||||
|
||||
|
||||
def registry_generation(*, scope: Optional[str] = None) -> tuple:
|
||||
"""Return a cache fingerprint for the global base and one profile."""
|
||||
active_scope = scope or hermes_home_key()
|
||||
with _lock:
|
||||
return _generation, _scoped_generations.get(active_scope, 0)
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: TerminalEnvironmentProvider,
|
||||
previous: Optional[TerminalEnvironmentProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip().lower()
|
||||
global _generation
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is None:
|
||||
_generation += 1
|
||||
else:
|
||||
_scoped_generations[scope] = _scoped_generations.get(scope, 0) + 1
|
||||
if not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear all registrations. Test hook — mirrors sibling registries."""
|
||||
global _generation
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
_scoped_generations.clear()
|
||||
_generation = 0
|
||||
|
||||
+23
-163
@@ -2,41 +2,16 @@
|
||||
Transcription Provider ABC
|
||||
==========================
|
||||
|
||||
Defines the pluggable-backend interface for speech-to-text. Providers
|
||||
register instances via
|
||||
:meth:`PluginContext.register_transcription_provider`; the active one
|
||||
(selected via ``stt.provider`` in ``config.yaml``) services every
|
||||
:func:`tools.transcription_tools.transcribe_audio` call **when the
|
||||
configured name is neither a built-in (``local``, ``local_command``,
|
||||
``groq``, ``openai``, ``mistral``, ``xai``) nor disabled**.
|
||||
Pluggable-backend interface for speech-to-text. Providers register via
|
||||
:meth:`PluginContext.register_transcription_provider`; the one named by
|
||||
``stt.provider`` services :func:`tools.transcription_tools.transcribe_audio`
|
||||
**when that name is not a built-in**. Built-ins (``BUILTIN_STT_PROVIDERS`` in
|
||||
:mod:`tools.transcription_tools`) always win: the registry rejects colliding
|
||||
names at registration and the dispatcher re-checks at dispatch time. The
|
||||
``HERMES_LOCAL_STT_COMMAND`` shell escape hatch stays on the built-in
|
||||
``local_command`` path.
|
||||
|
||||
Two coexisting STT extension surfaces — in resolution order:
|
||||
|
||||
1. **Built-in providers** (``BUILTIN_STT_PROVIDERS`` in
|
||||
:mod:`tools.transcription_tools`) — native Python implementations
|
||||
for the 6 backends shipped today (faster-whisper, local_command,
|
||||
Groq, OpenAI, Mistral, xAI). **Always win** — plugins cannot
|
||||
shadow them. The single-env-var shell escape hatch
|
||||
``HERMES_LOCAL_STT_COMMAND`` is preserved via the built-in
|
||||
``local_command`` path.
|
||||
2. **Plugin-registered providers** (this ABC). For new STT backends —
|
||||
OpenRouter, SenseAudio, Gemini-STT, custom proprietary engines —
|
||||
that need a Python implementation without modifying
|
||||
``tools/transcription_tools.py``.
|
||||
|
||||
Built-ins-always-win is enforced at registration time
|
||||
(:func:`agent.transcription_registry.register_provider` rejects names
|
||||
in ``BUILTIN_STT_PROVIDERS`` with a warning) AND at dispatch time
|
||||
(:func:`tools.transcription_tools._dispatch_to_plugin_provider`
|
||||
re-checks defensively).
|
||||
|
||||
Providers live in ``<repo>/plugins/transcription/<name>/`` (built-in
|
||||
plugins, none shipped today) or
|
||||
``~/.hermes/plugins/transcription/<name>/`` (user-installed).
|
||||
|
||||
Response contract
|
||||
-----------------
|
||||
:meth:`TranscriptionProvider.transcribe` returns a dict with keys::
|
||||
Response contract for :meth:`TranscriptionProvider.transcribe`::
|
||||
|
||||
success bool
|
||||
transcript str transcribed text (empty when success=False)
|
||||
@@ -48,106 +23,20 @@ from __future__ import annotations
|
||||
|
||||
import abc
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from agent.provider_base import CatalogProviderBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TranscriptionProvider(abc.ABC):
|
||||
class TranscriptionProvider(CatalogProviderBase):
|
||||
"""Abstract base class for a speech-to-text backend.
|
||||
|
||||
Subclasses must implement :attr:`name` and :meth:`transcribe`.
|
||||
Everything else has sane defaults — override only what your provider
|
||||
needs.
|
||||
Subclasses must implement :attr:`name` (rejected at registration if it
|
||||
collides with a built-in STT name) and :meth:`transcribe`.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in ``stt.provider`` config.
|
||||
|
||||
Lowercase, no spaces. Examples: ``openrouter``, ``sensaudio``,
|
||||
``gemini``, ``deepgram``. Names that collide with a built-in STT
|
||||
provider (``local``, ``local_command``, ``groq``, ``openai``,
|
||||
``mistral``, ``xai``) are rejected at registration time.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``.
|
||||
|
||||
Defaults to ``name.title()``.
|
||||
"""
|
||||
return self.name.title()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
|
||||
Typically checks for a required API key + that the SDK is
|
||||
importable. Default: True (providers with no external
|
||||
dependencies are always available).
|
||||
|
||||
Must NOT raise — used by the picker and ``hermes setup`` for
|
||||
availability displays and should fail gracefully.
|
||||
"""
|
||||
return True
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
"""Return model catalog entries.
|
||||
|
||||
Each entry::
|
||||
|
||||
{
|
||||
"id": "whisper-large-v3-turbo", # required
|
||||
"display": "Whisper Large v3 Turbo", # optional
|
||||
"languages": ["en", "es", "fr"], # optional
|
||||
"max_audio_seconds": 1500, # optional
|
||||
}
|
||||
|
||||
Default: empty list (provider has a single fixed model or
|
||||
doesn't expose model selection).
|
||||
"""
|
||||
return []
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
"""Return the default model id, or None if not applicable."""
|
||||
models = self.list_models()
|
||||
if models:
|
||||
return models[0].get("id")
|
||||
return None
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by ``tools_config.py`` to inject this provider as a row in
|
||||
the Speech-to-Text provider list. Shape::
|
||||
|
||||
{
|
||||
"name": "OpenRouter STT", # picker label
|
||||
"badge": "paid", # optional short tag
|
||||
"tag": "Whisper via OpenRouter API", # optional subtitle
|
||||
"env_vars": [ # keys to prompt for
|
||||
{"key": "OPENROUTER_API_KEY",
|
||||
"prompt": "OpenRouter API key",
|
||||
"url": "https://openrouter.ai/keys"},
|
||||
],
|
||||
}
|
||||
|
||||
Default: minimal entry derived from ``display_name`` with no
|
||||
env vars. Override to expose API key prompts and custom badges.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
@abc.abstractmethod
|
||||
def transcribe(
|
||||
self,
|
||||
@@ -157,42 +46,13 @@ class TranscriptionProvider(abc.ABC):
|
||||
language: Optional[str] = None,
|
||||
**extra: Any,
|
||||
) -> Dict[str, Any]:
|
||||
"""Transcribe the audio file at ``file_path``.
|
||||
"""Transcribe the audio file at ``file_path`` into the module-docstring envelope.
|
||||
|
||||
Returns a dict with the standard envelope::
|
||||
|
||||
{
|
||||
"success": True,
|
||||
"transcript": "the transcribed text",
|
||||
"provider": "<this provider's name>",
|
||||
}
|
||||
|
||||
or on failure::
|
||||
|
||||
{
|
||||
"success": False,
|
||||
"transcript": "",
|
||||
"error": "human-readable error message",
|
||||
"provider": "<this provider's name>",
|
||||
}
|
||||
|
||||
Implementations should NOT raise — convert exceptions to the
|
||||
error envelope so the dispatcher can deliver a consistent shape
|
||||
to the gateway/CLI caller.
|
||||
|
||||
Args:
|
||||
file_path: Absolute path to the audio file. The dispatcher
|
||||
has already validated existence + size before calling.
|
||||
model: Model identifier from :meth:`list_models`, or None
|
||||
to use :meth:`default_model`.
|
||||
language: Optional BCP-47 language hint (e.g. ``"en"``,
|
||||
``"ja"``) — providers without language hints should
|
||||
ignore this argument.
|
||||
**extra: Forward-compat parameters future schema versions
|
||||
may expose. Implementations should ignore unknown keys.
|
||||
The dispatcher currently forwards ``prompt`` here when a
|
||||
transcription prompt is set (via the ``stt.prompt`` config
|
||||
key or a ``pre_transcription`` hook) — prompt-capable
|
||||
providers may use it as a vocabulary/context hint; others
|
||||
should ignore it.
|
||||
Implementations should NOT raise — convert exceptions to the error
|
||||
envelope so the gateway/CLI caller always gets a consistent shape. The
|
||||
dispatcher has already validated existence + size. ``model`` None →
|
||||
:meth:`default_model`; ``language`` is an optional BCP-47 hint. The
|
||||
dispatcher forwards ``prompt`` in ``extra`` when ``stt.prompt`` or a
|
||||
``pre_transcription`` hook sets one — prompt-capable providers may use
|
||||
it as a vocabulary hint; unknown keys must be ignored.
|
||||
"""
|
||||
|
||||
+31
-133
@@ -2,42 +2,31 @@
|
||||
Transcription Provider Registry
|
||||
================================
|
||||
|
||||
Central map of registered STT providers. Populated by plugins at
|
||||
import-time via :meth:`PluginContext.register_transcription_provider`;
|
||||
consumed by :mod:`tools.transcription_tools` to dispatch
|
||||
:func:`transcribe_audio` calls to the active plugin backend **when**
|
||||
the configured ``stt.provider`` name is not a built-in.
|
||||
Central map of registered STT providers. Populated by plugins at import-time
|
||||
via :meth:`PluginContext.register_transcription_provider`; consumed by
|
||||
:mod:`tools.transcription_tools` to dispatch :func:`transcribe_audio` calls
|
||||
to the active plugin backend **when** the configured ``stt.provider`` name is
|
||||
not a built-in.
|
||||
|
||||
Built-ins-always-win
|
||||
--------------------
|
||||
Plugin names that collide with a built-in STT provider (``local``,
|
||||
``local_command``, ``groq``, ``openai``, ``mistral``, ``xai``) are
|
||||
rejected at registration with a warning. This invariant is also
|
||||
re-checked at dispatch time in
|
||||
:func:`tools.transcription_tools._dispatch_to_plugin_provider`.
|
||||
Built-ins-always-win: a plugin name colliding with a built-in STT provider is
|
||||
rejected at registration with a warning (re-checked at dispatch time in
|
||||
:func:`tools.transcription_tools._dispatch_to_plugin_provider`).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from agent.provider_registry import ProviderRegistry, lower_key
|
||||
from agent.transcription_provider import TranscriptionProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Names reserved for native built-in STT handlers. Plugins cannot
|
||||
# register a name in this set — the registration call is rejected with
|
||||
# a warning. **Kept in sync with ``BUILTIN_STT_PROVIDERS`` in
|
||||
# :mod:`tools.transcription_tools`** — a regression test in
|
||||
# ``tests/agent/test_transcription_registry.py::TestBuiltinSync``
|
||||
# fails if the two lists drift. Importing from
|
||||
# ``tools.transcription_tools`` directly would create a circular
|
||||
# dependency (``tools.transcription_tools`` imports
|
||||
# ``agent.transcription_registry`` for dispatch).
|
||||
# Names reserved for native built-in STT handlers. **Kept in sync with
|
||||
# ``BUILTIN_STT_PROVIDERS`` in :mod:`tools.transcription_tools`** (a regression
|
||||
# test in ``tests/agent/test_transcription_registry.py::TestBuiltinSync`` fails
|
||||
# on drift); importing it directly would be a circular import.
|
||||
_BUILTIN_NAMES = frozenset({
|
||||
"local",
|
||||
"local_command",
|
||||
@@ -50,114 +39,23 @@ _BUILTIN_NAMES = frozenset({
|
||||
})
|
||||
|
||||
|
||||
_providers: Dict[str, TranscriptionProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, TranscriptionProvider]] = {}
|
||||
_lock = threading.Lock()
|
||||
def _warn_builtin_collision(key: str) -> None:
|
||||
logger.warning(
|
||||
"Transcription provider '%s' shadows a built-in name; registration "
|
||||
"ignored. Built-in STT providers (%s) always win — pick a different "
|
||||
"name.",
|
||||
key, ", ".join(sorted(_BUILTIN_NAMES)),
|
||||
)
|
||||
|
||||
|
||||
def register_provider(provider: TranscriptionProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a transcription provider.
|
||||
|
||||
Rejects:
|
||||
|
||||
- Non-:class:`TranscriptionProvider` instances (raises :class:`TypeError`).
|
||||
- Empty/whitespace ``.name`` (raises :class:`ValueError`).
|
||||
- Names colliding with a built-in (logs a warning, silently
|
||||
ignores — built-ins-always-win invariant).
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and
|
||||
logs a debug message — makes hot-reload scenarios (tests, dev
|
||||
loops) behave predictably.
|
||||
"""
|
||||
if not isinstance(provider, TranscriptionProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a TranscriptionProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
name = provider.name
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("Transcription provider .name must be a non-empty string")
|
||||
key = name.strip().lower()
|
||||
if key in _BUILTIN_NAMES:
|
||||
logger.warning(
|
||||
"Transcription provider '%s' shadows a built-in name; registration "
|
||||
"ignored. Built-in STT providers (%s) always win — pick a different "
|
||||
"name.",
|
||||
key, ", ".join(sorted(_BUILTIN_NAMES)),
|
||||
)
|
||||
return
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(key)
|
||||
target[key] = provider
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"Transcription provider '%s' re-registered (was %r)",
|
||||
key, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered transcription provider '%s' (%s)",
|
||||
key, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[TranscriptionProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[TranscriptionProvider]:
|
||||
"""Return the provider registered under *name*, or None.
|
||||
|
||||
Name matching is case-insensitive and whitespace-tolerant — mirrors
|
||||
how ``tools.transcription_tools._get_provider`` normalizes the
|
||||
configured ``stt.provider`` value.
|
||||
"""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[TranscriptionProvider]:
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(key)
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: TranscriptionProvider,
|
||||
previous: Optional[TranscriptionProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is not None and not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
# Case-insensitive, whitespace-tolerant keys mirror how
|
||||
# ``tools.transcription_tools`` normalizes the configured ``stt.provider``.
|
||||
_registry: ProviderRegistry[TranscriptionProvider] = ProviderRegistry(
|
||||
label="Transcription",
|
||||
provider_cls=TranscriptionProvider,
|
||||
logger=logger,
|
||||
normalize=lower_key,
|
||||
builtin_names=_BUILTIN_NAMES,
|
||||
on_builtin_collision=_warn_builtin_collision,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
+45
-203
@@ -2,45 +2,22 @@
|
||||
Text-to-Speech Provider ABC
|
||||
============================
|
||||
|
||||
Defines the pluggable-backend interface for text-to-speech synthesis.
|
||||
Providers register instances via
|
||||
``PluginContext.register_tts_provider()``; the active one (selected via
|
||||
``tts.provider`` in ``config.yaml``) services every ``text_to_speech``
|
||||
tool call **only when the configured name is neither a built-in nor a
|
||||
command-type provider declared under ``tts.providers.<name>``**.
|
||||
Pluggable-backend interface for TTS synthesis. Providers register via
|
||||
``PluginContext.register_tts_provider()``; the one named by ``tts.provider``
|
||||
services ``text_to_speech`` **only when that name is neither a built-in nor a
|
||||
``tts.providers.<name>: type: command`` entry**. Resolution order:
|
||||
|
||||
Three coexisting TTS extension surfaces — in resolution order:
|
||||
1. Built-in providers (``BUILTIN_TTS_PROVIDERS`` in :mod:`tools.tts_tool`) —
|
||||
always win; :func:`agent.tts_registry.register_provider` rejects colliding
|
||||
names and the dispatcher re-checks at dispatch time.
|
||||
2. Command-type providers from ``config.yaml`` — win over a same-name plugin
|
||||
because config is more local than a plugin install.
|
||||
3. Plugin providers (this ABC) — for backends needing a Python SDK, streaming
|
||||
bytes, OAuth refresh, or voice-listing APIs the shell template can't express.
|
||||
|
||||
1. **Built-in providers** (``BUILTIN_TTS_PROVIDERS`` in
|
||||
:mod:`tools.tts_tool`) — native Python implementations (edge, openai,
|
||||
elevenlabs, …). **Always win** — plugins cannot shadow them.
|
||||
2. **Command-type providers** declared under ``tts.providers.<name>:
|
||||
type: command`` (PR #17843, commit ``2facea7f7``). Wire any local
|
||||
CLI into Hermes with shell-template placeholders. **Wins over a
|
||||
same-name plugin** — config is more local than plugin install.
|
||||
3. **Plugin-registered providers** (this ABC). For backends that need a
|
||||
Python SDK, streaming bytes, OAuth refresh, or voice-listing APIs
|
||||
the shell-template grammar can't reasonably express.
|
||||
|
||||
Built-ins-always-win is enforced at registration time
|
||||
(:func:`agent.tts_registry.register_provider` rejects names in
|
||||
``BUILTIN_TTS_PROVIDERS`` with a warning) AND at dispatch time
|
||||
(:func:`tools.tts_tool._dispatch_to_plugin_provider` re-checks
|
||||
defensively). The dispatcher also rejects plugin dispatch when a same-
|
||||
name command provider is configured.
|
||||
|
||||
Providers live in ``<repo>/plugins/tts/<name>/`` (built-in plugins, no
|
||||
shipped today) or ``~/.hermes/plugins/tts/<name>/`` (user-installed).
|
||||
None ship in-tree as of issue #30398 — the hook is additive
|
||||
infrastructure waiting for a real consumer (Cartesia, Fish Audio, …).
|
||||
|
||||
Response contract
|
||||
-----------------
|
||||
:meth:`TTSProvider.synthesize` writes the audio bytes to ``output_path``
|
||||
and returns the path as a string. Implementations should raise on
|
||||
failure — the dispatcher converts exceptions into the standard
|
||||
``{success: False, error: …}`` JSON envelope the rest of Hermes
|
||||
expects.
|
||||
:meth:`TTSProvider.synthesize` writes audio to ``output_path`` and returns the
|
||||
path; it should raise on failure — the dispatcher converts exceptions into the
|
||||
standard ``{success: False, error: …}`` envelope.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -49,6 +26,8 @@ import abc
|
||||
import logging
|
||||
from typing import Any, Dict, Iterator, List, Optional
|
||||
|
||||
from agent.provider_base import CatalogProviderBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -56,122 +35,20 @@ DEFAULT_OUTPUT_FORMAT = "mp3"
|
||||
VALID_OUTPUT_FORMATS = frozenset({"mp3", "wav", "ogg", "opus", "flac"})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TTSProvider(abc.ABC):
|
||||
class TTSProvider(CatalogProviderBase):
|
||||
"""Abstract base class for a text-to-speech backend.
|
||||
|
||||
Subclasses must implement :attr:`name` and :meth:`synthesize`.
|
||||
Everything else has sane defaults — override only what your provider
|
||||
needs.
|
||||
Subclasses must implement :attr:`name` (rejected at registration if it
|
||||
collides with a built-in TTS provider name) and :meth:`synthesize`.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in ``tts.provider`` config.
|
||||
|
||||
Lowercase, no spaces. Examples: ``cartesia``, ``fishaudio``,
|
||||
``deepgram``. Names that collide with a built-in TTS provider
|
||||
(``edge``, ``openai``, ``elevenlabs``, ``minimax``, ``gemini``,
|
||||
``mistral``, ``xai``, ``piper``, ``kittentts``, ``neutts``) are
|
||||
rejected at registration time.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``.
|
||||
|
||||
Defaults to ``name.title()`` (e.g. ``Cartesia`` for ``cartesia``).
|
||||
"""
|
||||
return self.name.title()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
|
||||
Typically checks for a required API key + that the SDK is
|
||||
importable. Default: True (providers with no external
|
||||
dependencies are always available).
|
||||
|
||||
Must NOT raise — used by the picker and ``hermes setup`` for
|
||||
availability displays and should fail gracefully.
|
||||
"""
|
||||
return True
|
||||
|
||||
def list_voices(self) -> List[Dict[str, Any]]:
|
||||
"""Return voice catalog entries.
|
||||
|
||||
Each entry::
|
||||
|
||||
{
|
||||
"id": "voice-abc-123", # required
|
||||
"display": "Aria — neutral female", # optional; defaults to id
|
||||
"language": "en-US", # optional
|
||||
"gender": "female", # optional
|
||||
"preview_url": "https://...mp3", # optional
|
||||
}
|
||||
|
||||
Default: empty list (provider has no enumerable voices or
|
||||
doesn't surface them via API).
|
||||
"""
|
||||
"""Voice catalog entries: ``{"id"}`` required; ``display`` / ``language``
|
||||
/ ``gender`` / ``preview_url`` optional. Default: empty."""
|
||||
return []
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
"""Return model catalog entries.
|
||||
|
||||
Each entry::
|
||||
|
||||
{
|
||||
"id": "sonic-2", # required
|
||||
"display": "Sonic 2", # optional
|
||||
"languages": ["en", "es", "fr"], # optional
|
||||
"max_text_length": 5000, # optional
|
||||
}
|
||||
|
||||
Default: empty list (provider has a single fixed model or
|
||||
doesn't expose model selection).
|
||||
"""
|
||||
return []
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by ``tools_config.py`` to inject this provider as a row in
|
||||
the Text-to-Speech provider list. Shape::
|
||||
|
||||
{
|
||||
"name": "Cartesia", # picker label
|
||||
"badge": "paid", # optional short tag
|
||||
"tag": "Ultra-low-latency streaming", # optional subtitle
|
||||
"env_vars": [ # keys to prompt for
|
||||
{"key": "CARTESIA_API_KEY",
|
||||
"prompt": "Cartesia API key",
|
||||
"url": "https://play.cartesia.ai/console"},
|
||||
],
|
||||
}
|
||||
|
||||
Default: minimal entry derived from ``display_name`` with no
|
||||
env vars. Override to expose API key prompts and custom badges.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
"""Return the default model id, or None if not applicable."""
|
||||
models = self.list_models()
|
||||
if models:
|
||||
return models[0].get("id")
|
||||
return None
|
||||
|
||||
def default_voice(self) -> Optional[str]:
|
||||
"""Return the default voice id, or None if not applicable."""
|
||||
"""Id of the first voice entry, or None if not applicable."""
|
||||
voices = self.list_voices()
|
||||
if voices:
|
||||
return voices[0].get("id")
|
||||
@@ -189,31 +66,14 @@ class TTSProvider(abc.ABC):
|
||||
format: str = DEFAULT_OUTPUT_FORMAT,
|
||||
**extra: Any,
|
||||
) -> str:
|
||||
"""Synthesize ``text`` and write audio bytes to ``output_path``.
|
||||
"""Synthesize ``text`` into ``output_path`` and return the written path.
|
||||
|
||||
Returns the absolute path to the written file as a string
|
||||
(typically just echoes ``output_path``). Raises on failure —
|
||||
the dispatcher converts exceptions to the standard
|
||||
``{success: False, error: ...}`` JSON envelope.
|
||||
|
||||
Args:
|
||||
text: The text to synthesize. Already truncated to the
|
||||
provider's max length by the dispatcher.
|
||||
output_path: Absolute path where the audio file should be
|
||||
written. Parent directory is guaranteed to exist.
|
||||
voice: Voice identifier from :meth:`list_voices`, or None
|
||||
to use :meth:`default_voice`.
|
||||
model: Model identifier from :meth:`list_models`, or None
|
||||
to use :meth:`default_model`.
|
||||
speed: Optional speech-rate multiplier (1.0 = normal).
|
||||
Providers that don't support speed control should
|
||||
ignore this argument.
|
||||
format: Output audio format. Implementations should match
|
||||
the requested format when possible; if unsupported,
|
||||
pick the closest equivalent and ensure ``output_path``
|
||||
ends with the correct extension.
|
||||
**extra: Forward-compat parameters future schema versions
|
||||
may expose. Implementations should ignore unknown keys.
|
||||
``text`` is already truncated to the provider's max length and the
|
||||
parent directory exists. ``voice`` / ``model`` fall back to
|
||||
:meth:`default_voice` / :meth:`default_model` when None; ``speed`` is a
|
||||
rate multiplier providers may ignore. If ``format`` is unsupported, pick
|
||||
the closest equivalent and make ``output_path`` carry the right
|
||||
extension. Unknown ``extra`` keys must be ignored. Raise on failure.
|
||||
"""
|
||||
|
||||
def stream(
|
||||
@@ -225,15 +85,12 @@ class TTSProvider(abc.ABC):
|
||||
format: str = "opus",
|
||||
**extra: Any,
|
||||
) -> Iterator[bytes]:
|
||||
"""Stream synthesized audio bytes.
|
||||
"""Stream synthesized audio bytes (optional).
|
||||
|
||||
Optional. Providers that don't support streaming raise
|
||||
:class:`NotImplementedError` (the default) and the dispatcher
|
||||
falls back to :meth:`synthesize` + read-whole-file.
|
||||
|
||||
Args mirror :meth:`synthesize`. Default ``format`` is ``opus``
|
||||
because the primary streaming use case is voice-bubble
|
||||
delivery (Telegram et al.) which requires Opus.
|
||||
Default raises :class:`NotImplementedError`; the dispatcher then falls
|
||||
back to :meth:`synthesize` + read-whole-file. ``format`` defaults to
|
||||
``opus`` because the primary streaming consumer is voice-bubble
|
||||
delivery (Telegram et al.), which requires Opus.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
f"TTS provider {self.name!r} does not implement streaming "
|
||||
@@ -244,44 +101,29 @@ class TTSProvider(abc.ABC):
|
||||
def warm(self) -> None:
|
||||
"""Speech output was just turned on; pre-load so the first reply is hot.
|
||||
|
||||
Optional. Called from the TTS lease path (Desktop read-aloud / voice
|
||||
conversation, ``/voice tts``) when this provider is the configured
|
||||
``tts.provider`` — e.g. ask a local model server to load its model.
|
||||
Best-effort: exceptions are logged at debug and ignored. Default: no-op.
|
||||
Called from the TTS lease path (Desktop read-aloud / voice conversation)
|
||||
when this is the configured provider. Best-effort; default no-op.
|
||||
"""
|
||||
|
||||
def release(self) -> None:
|
||||
"""The last speech-output lease was released; free resident resources.
|
||||
|
||||
Optional counterpart of :meth:`warm` — e.g. tell a local model server
|
||||
to unload. Best-effort; default: no-op.
|
||||
"""
|
||||
"""Last speech-output lease released; free resident resources (counterpart
|
||||
of :meth:`warm`). Best-effort; default no-op."""
|
||||
|
||||
@property
|
||||
def voice_compatible(self) -> bool:
|
||||
"""Whether output is suitable for voice-bubble delivery.
|
||||
"""Whether output suits voice-bubble delivery (mirrors
|
||||
``tts.providers.<name>.voice_compatible``).
|
||||
|
||||
Mirrors the ``tts.providers.<name>.voice_compatible`` field
|
||||
from PR #17843. When True, the gateway's voice-message
|
||||
delivery pipeline runs ffmpeg conversion to Opus if needed.
|
||||
When False, output is delivered as a regular audio attachment.
|
||||
|
||||
Default: False (safe — providers opt in explicitly).
|
||||
True → the gateway converts to Opus via ffmpeg if needed; False →
|
||||
delivered as a regular audio attachment. Default False (opt in).
|
||||
"""
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def resolve_output_format(value: Optional[str]) -> str:
|
||||
"""Clamp an output_format value to the valid set.
|
||||
|
||||
Invalid values are coerced to :data:`DEFAULT_OUTPUT_FORMAT` rather
|
||||
than rejected so the tool surface is forgiving of agent mistakes.
|
||||
"""
|
||||
"""Clamp an output_format to :data:`VALID_OUTPUT_FORMATS`; invalid values
|
||||
coerce to :data:`DEFAULT_OUTPUT_FORMAT` so the tool surface forgives agent
|
||||
mistakes instead of rejecting them."""
|
||||
if not isinstance(value, str):
|
||||
return DEFAULT_OUTPUT_FORMAT
|
||||
v = value.strip().lower()
|
||||
|
||||
+34
-139
@@ -2,50 +2,36 @@
|
||||
TTS Provider Registry
|
||||
=====================
|
||||
|
||||
Central map of registered TTS providers. Populated by plugins at
|
||||
import-time via :meth:`PluginContext.register_tts_provider`; consumed
|
||||
by :mod:`tools.tts_tool` to dispatch ``text_to_speech`` tool calls to
|
||||
the active plugin backend **when** the configured ``tts.provider``
|
||||
name is neither a built-in nor a command-type provider.
|
||||
Central map of registered TTS providers. Populated by plugins at import-time
|
||||
via :meth:`PluginContext.register_tts_provider`; consumed by
|
||||
:mod:`tools.tts_tool` to dispatch ``text_to_speech`` calls to the active
|
||||
plugin backend **when** the configured ``tts.provider`` name is neither a
|
||||
built-in nor a command-type provider.
|
||||
|
||||
Built-ins-always-win
|
||||
--------------------
|
||||
Plugin names that collide with a built-in TTS provider (``edge``,
|
||||
``openai``, ``elevenlabs``, ``minimax``, ``gemini``, ``mistral``,
|
||||
``xai``, ``piper``, ``kittentts``, ``neutts``) are rejected at
|
||||
registration with a warning. This invariant is also re-checked at
|
||||
dispatch time in :func:`tools.tts_tool._dispatch_to_plugin_provider`.
|
||||
Built-ins-always-win: a plugin name colliding with a built-in TTS provider is
|
||||
rejected at registration with a warning (re-checked at dispatch time in
|
||||
:func:`tools.tts_tool._dispatch_to_plugin_provider`).
|
||||
|
||||
Command-providers-win-over-plugins
|
||||
----------------------------------
|
||||
This registry doesn't enforce the command-vs-plugin precedence — that
|
||||
lives in the dispatcher, which checks for a same-name
|
||||
``tts.providers.<name>: type: command`` entry before consulting the
|
||||
registry. The rationale is locality: a name declared in the user's
|
||||
``config.yaml`` is more specific to their setup than a plugin that
|
||||
happens to be installed.
|
||||
Command-providers-win-over-plugins is enforced by the dispatcher, not here:
|
||||
it checks for a same-name ``tts.providers.<name>: type: command`` entry before
|
||||
consulting the registry (a name declared in the user's config.yaml is more
|
||||
specific to their setup than an installed plugin).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from agent.provider_registry import ProviderRegistry, lower_key
|
||||
from agent.tts_provider import TTSProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Names reserved for native built-in TTS handlers. Plugins cannot
|
||||
# register a name in this set — the registration call is rejected with
|
||||
# a warning. **Kept in sync with ``BUILTIN_TTS_PROVIDERS`` in
|
||||
# :mod:`tools.tts_tool`** — a regression test in
|
||||
# ``tests/agent/test_tts_registry.py::TestBuiltinSync`` fails if the
|
||||
# two lists drift. Importing from ``tools.tts_tool`` directly would
|
||||
# create a circular dependency (``tools.tts_tool`` imports
|
||||
# ``agent.tts_registry`` for dispatch).
|
||||
# Names reserved for native built-in TTS handlers. **Kept in sync with
|
||||
# ``BUILTIN_TTS_PROVIDERS`` in :mod:`tools.tts_tool`** (a regression test in
|
||||
# ``tests/agent/test_tts_registry.py::TestBuiltinSync`` fails on drift);
|
||||
# importing it directly would be a circular import.
|
||||
_BUILTIN_NAMES = frozenset({
|
||||
"edge",
|
||||
"elevenlabs",
|
||||
@@ -61,113 +47,22 @@ _BUILTIN_NAMES = frozenset({
|
||||
})
|
||||
|
||||
|
||||
_providers: Dict[str, TTSProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, TTSProvider]] = {}
|
||||
_lock = threading.Lock()
|
||||
def _warn_builtin_collision(key: str) -> None:
|
||||
logger.warning(
|
||||
"TTS provider '%s' shadows a built-in name; registration ignored. "
|
||||
"Built-in TTS providers (%s) always win — pick a different name.",
|
||||
key, ", ".join(sorted(_BUILTIN_NAMES)),
|
||||
)
|
||||
|
||||
|
||||
def register_provider(provider: TTSProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a TTS provider.
|
||||
|
||||
Rejects:
|
||||
|
||||
- Non-:class:`TTSProvider` instances (raises :class:`TypeError`).
|
||||
- Empty/whitespace ``.name`` (raises :class:`ValueError`).
|
||||
- Names colliding with a built-in (logs a warning, silently
|
||||
ignores — built-ins-always-win invariant).
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and
|
||||
logs a debug message — makes hot-reload scenarios (tests, dev
|
||||
loops) behave predictably.
|
||||
"""
|
||||
if not isinstance(provider, TTSProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a TTSProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
name = provider.name
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("TTS provider .name must be a non-empty string")
|
||||
key = name.strip().lower()
|
||||
if key in _BUILTIN_NAMES:
|
||||
logger.warning(
|
||||
"TTS provider '%s' shadows a built-in name; registration ignored. "
|
||||
"Built-in TTS providers (%s) always win — pick a different name.",
|
||||
key, ", ".join(sorted(_BUILTIN_NAMES)),
|
||||
)
|
||||
return
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(key)
|
||||
target[key] = provider
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"TTS provider '%s' re-registered (was %r)",
|
||||
key, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered TTS provider '%s' (%s)",
|
||||
key, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[TTSProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[TTSProvider]:
|
||||
"""Return the provider registered under *name*, or None.
|
||||
|
||||
Name matching is case-insensitive and whitespace-tolerant — mirrors
|
||||
how ``tools.tts_tool._get_provider`` normalizes the configured
|
||||
``tts.provider`` value.
|
||||
"""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[TTSProvider]:
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(key)
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: TTSProvider,
|
||||
previous: Optional[TTSProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip().lower()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is not None and not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
# Case-insensitive, whitespace-tolerant keys mirror how
|
||||
# ``tools.tts_tool._get_provider`` normalizes the configured ``tts.provider``.
|
||||
_registry: ProviderRegistry[TTSProvider] = ProviderRegistry(
|
||||
label="TTS",
|
||||
provider_cls=TTSProvider,
|
||||
logger=logger,
|
||||
normalize=lower_key,
|
||||
builtin_names=_BUILTIN_NAMES,
|
||||
on_builtin_collision=_warn_builtin_collision,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
+60
-253
@@ -2,35 +2,20 @@
|
||||
Video Generation Provider ABC
|
||||
=============================
|
||||
|
||||
Defines the pluggable-backend interface for video generation. Providers register
|
||||
instances via ``PluginContext.register_video_gen_provider()``; the active one
|
||||
(selected via ``video_gen.provider`` in ``config.yaml``) services every
|
||||
``video_generate`` tool call.
|
||||
Pluggable-backend interface for video generation. Providers register via
|
||||
``PluginContext.register_video_gen_provider()``; the one selected by
|
||||
``video_gen.provider`` services every ``video_generate`` call. Providers live in
|
||||
``<repo>/plugins/video_gen/<name>/`` (built-in) or
|
||||
``~/.hermes/plugins/video_gen/<name>/`` (user, opt-in). Mirrors
|
||||
``agent/image_gen_provider.py`` so the two surfaces stay learnable together.
|
||||
|
||||
Providers live in ``<repo>/plugins/video_gen/<name>/`` (built-in, auto-loaded
|
||||
as ``kind: backend``) or ``~/.hermes/plugins/video_gen/<name>/`` (user, opt-in
|
||||
via ``plugins.enabled``).
|
||||
One tool covers text-to-video and image-to-video: ``image_url`` present routes
|
||||
to the provider's image-to-video endpoint, absent routes to text-to-video. Users
|
||||
pick one model family; the provider picks the FAL/xAI endpoint. Video edit and
|
||||
extend are deliberately NOT exposed — backends are too inconsistent for one
|
||||
unified tool.
|
||||
|
||||
Mirrors the ``image_gen`` provider design (``agent/image_gen_provider.py``) so
|
||||
the two surfaces stay learnable together.
|
||||
|
||||
Unified surface
|
||||
---------------
|
||||
One tool — ``video_generate`` — covers **text-to-video** and **image-to-video**.
|
||||
The router is the presence of ``image_url``: if it's set, the provider routes
|
||||
to its image-to-video endpoint; if it's omitted, the provider routes to
|
||||
text-to-video. Users pick one **model family** (e.g. Pixverse v6, Veo 3.1,
|
||||
Kling O3 Standard); the provider handles which underlying FAL/xAI endpoint
|
||||
to hit.
|
||||
|
||||
Video edit and video extend are intentionally NOT exposed in this surface —
|
||||
the inconsistency across backends is too large for one unified tool. If
|
||||
those use cases warrant attention later they can ship as separate tools.
|
||||
|
||||
Response shape
|
||||
--------------
|
||||
All providers return a dict built by :func:`success_response` /
|
||||
:func:`error_response`. Keys:
|
||||
Response shape (built by :func:`success_response` / :func:`error_response`)::
|
||||
|
||||
success bool
|
||||
video str | None URL or absolute file path
|
||||
@@ -47,19 +32,18 @@ All providers return a dict built by :func:`success_response` /
|
||||
from __future__ import annotations
|
||||
|
||||
import abc
|
||||
import base64
|
||||
import datetime
|
||||
import logging
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from agent import provider_media
|
||||
from agent.provider_base import CatalogProviderBase
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Common aspect ratios across providers (Veo / Kling / xAI / Pixverse). The
|
||||
# tool schema advertises this set as an enum hint, but providers may accept
|
||||
# a narrower or wider set — they are responsible for clamping.
|
||||
# Advertised as an enum hint in the tool schema; providers may accept a narrower
|
||||
# or wider set and are responsible for clamping.
|
||||
COMMON_ASPECT_RATIOS: Tuple[str, ...] = ("16:9", "9:16", "1:1", "4:3", "3:4", "3:2", "2:3")
|
||||
DEFAULT_ASPECT_RATIO = "16:9"
|
||||
|
||||
@@ -67,98 +51,23 @@ COMMON_RESOLUTIONS: Tuple[str, ...] = ("480p", "540p", "720p", "1080p")
|
||||
DEFAULT_RESOLUTION = "720p"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class VideoGenProvider(abc.ABC):
|
||||
class VideoGenProvider(CatalogProviderBase):
|
||||
"""Abstract base class for a video generation backend.
|
||||
|
||||
Subclasses must implement :meth:`generate`. Everything else has sane
|
||||
defaults — override only what your provider needs.
|
||||
Subclasses must implement :attr:`name` and :meth:`generate`; everything else
|
||||
has defaults. ``list_models`` entries are **model families** and may add
|
||||
``speed`` / ``strengths`` / ``price`` / advisory ``modalities``.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in ``video_gen.provider`` config.
|
||||
|
||||
Lowercase, no spaces. Examples: ``xai``, ``fal``, ``google``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name.title()``."""
|
||||
return self.name.title()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
|
||||
Typically checks for a required API key and optional-dependency
|
||||
import. Default: True.
|
||||
"""
|
||||
return True
|
||||
|
||||
def list_models(self) -> List[Dict[str, Any]]:
|
||||
"""Return catalog entries for ``hermes tools`` model picker.
|
||||
|
||||
Each entry represents a **model family** that supports text-to-video
|
||||
and/or image-to-video routing internally::
|
||||
|
||||
{
|
||||
"id": "veo-3.1", # required
|
||||
"display": "Veo 3.1", # optional; defaults to id
|
||||
"speed": "~60s", # optional
|
||||
"strengths": "...", # optional
|
||||
"price": "$0.20/s", # optional
|
||||
"modalities": ["text", "image"], # optional, advisory
|
||||
}
|
||||
|
||||
Default: empty list (provider has no user-selectable models).
|
||||
"""
|
||||
return []
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker."""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
"""Return the default model id, or None if not applicable."""
|
||||
models = self.list_models()
|
||||
if models:
|
||||
return models[0].get("id")
|
||||
return None
|
||||
|
||||
def capabilities(self) -> Dict[str, Any]:
|
||||
"""Return what this provider supports.
|
||||
"""What this provider supports (all keys optional): ``modalities``,
|
||||
``aspect_ratios``, ``resolutions``, ``max_duration`` / ``min_duration``,
|
||||
``supports_audio`` / ``supports_negative_prompt`` / ``supports_seed`` /
|
||||
``supports_upscale``, ``max_reference_images``.
|
||||
|
||||
Returned dict (all keys optional)::
|
||||
|
||||
{
|
||||
"modalities": ["text", "image"], # which inputs the backend accepts
|
||||
"aspect_ratios": ["16:9", "9:16", ...],
|
||||
"resolutions": ["720p", "1080p"],
|
||||
"max_duration": 15, # seconds
|
||||
"min_duration": 1,
|
||||
"supports_audio": True,
|
||||
"supports_negative_prompt": True,
|
||||
"supports_seed": True,
|
||||
"supports_upscale": True,
|
||||
"max_reference_images": 7,
|
||||
}
|
||||
|
||||
Used by the tool layer for soft validation, for capability-gated
|
||||
param rendering in the dynamic ``video_generate`` schema (args a
|
||||
backend can't honor are not advertised), and by ``hermes tools``
|
||||
for the picker. Default fails closed: text-only, no optional
|
||||
features — a provider that doesn't declare a capability doesn't
|
||||
advertise it.
|
||||
Used for soft validation, capability-gated params in the dynamic
|
||||
``video_generate`` schema (args a backend can't honor aren't advertised),
|
||||
and the picker. Default fails closed: text-only, no optional features.
|
||||
"""
|
||||
return {
|
||||
"modalities": ["text"],
|
||||
@@ -189,24 +98,12 @@ class VideoGenProvider(abc.ABC):
|
||||
seed: Optional[int] = None,
|
||||
**kwargs: Any,
|
||||
) -> Dict[str, Any]:
|
||||
"""Generate a video from a prompt (text-to-video) or animate an image
|
||||
(image-to-video).
|
||||
"""Generate a video from a prompt, or animate ``image_url`` when given.
|
||||
|
||||
Routing: if ``image_url`` is provided, the provider should route to
|
||||
its image-to-video endpoint; otherwise text-to-video. The plugin
|
||||
is responsible for picking the right underlying endpoint within
|
||||
the user's chosen model family.
|
||||
|
||||
Implementations should return the dict from :func:`success_response`
|
||||
or :func:`error_response`. ``kwargs`` may contain forward-compat
|
||||
parameters future versions of the schema will expose —
|
||||
implementations MUST ignore unknown keys (no TypeError).
|
||||
|
||||
Known optional kwarg: ``upscale`` (bool) — when true, the caller
|
||||
requests a post-generation high-resolution pass through the
|
||||
backend's video upscaler. Providers without an upscaler simply
|
||||
ignore it; providers that honor it should report ``upscaled: True``
|
||||
in the response ``extra``.
|
||||
Return :func:`success_response` / :func:`error_response`. Unknown
|
||||
``kwargs`` MUST be ignored (forward compat). Known optional kwarg:
|
||||
``upscale`` (bool) — a post-generation high-res pass; providers that
|
||||
honor it report ``upscaled: True`` in ``extra``.
|
||||
"""
|
||||
|
||||
|
||||
@@ -215,33 +112,14 @@ class VideoGenProvider(abc.ABC):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _videos_cache_dir() -> Path:
|
||||
"""Return ``$HERMES_HOME/cache/videos/``, creating parents as needed."""
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
path = get_hermes_home() / "cache" / "videos"
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def save_b64_video(
|
||||
b64_data: str,
|
||||
*,
|
||||
prefix: str = "video",
|
||||
extension: str = "mp4",
|
||||
) -> Path:
|
||||
"""Decode base64 video data and write under ``$HERMES_HOME/cache/videos/``.
|
||||
|
||||
Returns the absolute :class:`Path` to the saved file.
|
||||
|
||||
Filename format: ``<prefix>_<YYYYMMDD_HHMMSS>_<short-uuid>.<ext>``.
|
||||
"""
|
||||
raw = base64.b64decode(b64_data)
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
path = _videos_cache_dir() / f"{prefix}_{ts}_{short}.{extension}"
|
||||
path.write_bytes(raw)
|
||||
return path
|
||||
"""Decode base64 video data into ``$HERMES_HOME/cache/videos/``; return the path."""
|
||||
return provider_media.save_b64("videos", b64_data, prefix=prefix, extension=extension)
|
||||
|
||||
|
||||
def save_bytes_video(
|
||||
@@ -251,11 +129,7 @@ def save_bytes_video(
|
||||
extension: str = "mp4",
|
||||
) -> Path:
|
||||
"""Write raw video bytes (e.g. an HTTP download body) to the cache."""
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
path = _videos_cache_dir() / f"{prefix}_{ts}_{short}.{extension}"
|
||||
path.write_bytes(raw)
|
||||
return path
|
||||
return provider_media.save_bytes("videos", raw, prefix=prefix, extension=extension)
|
||||
|
||||
|
||||
_URL_VIDEO_CONTENT_TYPES = {
|
||||
@@ -273,61 +147,17 @@ def save_url_video(
|
||||
timeout: float = 180.0,
|
||||
max_bytes: int = 200 * 1024 * 1024,
|
||||
) -> Path:
|
||||
"""Download a video URL and write it under ``$HERMES_HOME/cache/videos/``.
|
||||
"""Download an (often ephemeral) video URL into ``$HERMES_HOME/cache/videos/``.
|
||||
|
||||
The video twin of :func:`agent.image_gen_provider.save_url_image`: several
|
||||
backends (DeepInfra, FAL) return an *ephemeral* delivery URL that expires
|
||||
before a downstream consumer can fetch it, so we materialise the bytes
|
||||
locally at tool-completion time. Streams with a size cap.
|
||||
|
||||
Raises on any network / HTTP / oversize error so callers can fall back to
|
||||
returning the bare URL.
|
||||
Raises on network / HTTP / oversize / empty errors so callers can fall back
|
||||
to returning the bare URL. See :mod:`agent.provider_media`.
|
||||
"""
|
||||
import requests
|
||||
|
||||
response = requests.get(url, timeout=timeout, stream=True)
|
||||
response.raise_for_status()
|
||||
|
||||
content_type = (response.headers.get("Content-Type") or "").split(";", 1)[0].strip().lower()
|
||||
extension = _URL_VIDEO_CONTENT_TYPES.get(content_type)
|
||||
if extension is None:
|
||||
url_path = url.split("?", 1)[0].lower()
|
||||
for ext in ("mp4", "webm", "mov", "mkv"):
|
||||
if url_path.endswith(f".{ext}"):
|
||||
extension = ext
|
||||
break
|
||||
if extension is None:
|
||||
extension = "mp4"
|
||||
|
||||
ts = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
short = uuid.uuid4().hex[:8]
|
||||
path = _videos_cache_dir() / f"{prefix}_{ts}_{short}.{extension}"
|
||||
|
||||
bytes_written = 0
|
||||
with path.open("wb") as fh:
|
||||
for chunk in response.iter_content(chunk_size=256 * 1024):
|
||||
if not chunk:
|
||||
continue
|
||||
bytes_written += len(chunk)
|
||||
if bytes_written > max_bytes:
|
||||
fh.close()
|
||||
try:
|
||||
path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
raise ValueError(
|
||||
f"Video at {url} exceeds {max_bytes // (1024 * 1024)}MB cap; refusing to cache."
|
||||
)
|
||||
fh.write(chunk)
|
||||
|
||||
if bytes_written == 0:
|
||||
try:
|
||||
path.unlink()
|
||||
except OSError:
|
||||
pass
|
||||
raise ValueError(f"Video at {url} was empty (0 bytes).")
|
||||
|
||||
return path
|
||||
return provider_media.save_url(
|
||||
"videos", url, prefix=prefix, timeout=timeout, max_bytes=max_bytes,
|
||||
chunk_size=256 * 1024, content_types=_URL_VIDEO_CONTENT_TYPES,
|
||||
url_extensions=("mp4", "webm", "mov", "mkv"), default_extension="mp4",
|
||||
label="Video", empty_error="Video at {url} was empty (0 bytes).",
|
||||
)
|
||||
|
||||
|
||||
def success_response(
|
||||
@@ -341,12 +171,7 @@ def success_response(
|
||||
provider: str,
|
||||
extra: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Build a uniform success response dict.
|
||||
|
||||
``video`` may be an HTTP URL or an absolute filesystem path.
|
||||
``modality`` is ``"text"`` (text-to-video) or ``"image"`` (image-to-video) —
|
||||
indicates which endpoint was actually hit, useful for diagnostics.
|
||||
"""
|
||||
"""Uniform success dict; ``extra`` keys are added without overriding standard ones."""
|
||||
payload: Dict[str, Any] = {
|
||||
"success": True,
|
||||
"video": video,
|
||||
@@ -393,10 +218,9 @@ def error_response(
|
||||
class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
"""Generic text/image-to-video over the OpenAI ``client.videos`` API.
|
||||
|
||||
DeepInfra, OpenAI/Sora, and OpenRouter all expose the same
|
||||
``POST /videos`` async-job shape (``create`` → poll → ``download_content``),
|
||||
so the SDK call lives here once. A concrete backend only needs to declare
|
||||
its identity and credentials::
|
||||
DeepInfra, OpenAI/Sora, and OpenRouter share the ``POST /videos`` async-job
|
||||
shape (``create`` → poll → ``download_content``), so the SDK call lives here
|
||||
once; a concrete backend declares identity and credentials::
|
||||
|
||||
class FooVideoGenProvider(OpenAICompatibleVideoGenProvider):
|
||||
name = "foo"
|
||||
@@ -413,12 +237,9 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
_env_key: str = "OPENAI_API_KEY"
|
||||
_default_base_url: str = "https://api.openai.com/v1"
|
||||
|
||||
# Polling cadence for the async video job. The OpenAI SDK's
|
||||
# ``create_and_poll`` defaults to ~1 poll/second and loops forever on a
|
||||
# non-terminal status, so a multi-minute job issues hundreds of sequential
|
||||
# requests and a stuck job pins its tool-executor worker thread with no way
|
||||
# out. We hand-roll a bounded poll instead: a coarse interval plus a hard
|
||||
# wall-clock deadline that surfaces a timeout error.
|
||||
# The SDK's ``create_and_poll`` polls ~1/s forever on a non-terminal status,
|
||||
# pinning the tool-executor thread on a stuck job; we poll coarsely with a
|
||||
# hard wall-clock deadline instead.
|
||||
_poll_interval_s: float = 5.0
|
||||
_poll_deadline_s: float = 900.0
|
||||
|
||||
@@ -431,13 +252,8 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
return bool(self._api_key())
|
||||
|
||||
def _create_and_poll(self, client: Any, call_kwargs: Dict[str, Any]) -> Any:
|
||||
"""Create the video job and poll to completion with a hard deadline.
|
||||
|
||||
Replaces ``client.videos.create_and_poll`` (unbounded 1/s loop) with a
|
||||
coarse interval and a wall-clock cap. Returns the terminal video object
|
||||
(any status); raises :class:`TimeoutError` if the deadline passes
|
||||
first.
|
||||
"""
|
||||
"""Create the job and poll to a terminal status (any); raise
|
||||
:class:`TimeoutError` when ``_poll_deadline_s`` passes first."""
|
||||
import time
|
||||
|
||||
video = client.videos.create(**call_kwargs)
|
||||
@@ -502,8 +318,7 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
provider=self.name,
|
||||
)
|
||||
|
||||
# Provider-specific fields the OpenAI ``videos.create`` signature does
|
||||
# not name natively — pass them through ``extra_body``.
|
||||
# Fields ``videos.create`` doesn't name natively ride in ``extra_body``.
|
||||
extra_body = {
|
||||
k: v
|
||||
for k, v in {
|
||||
@@ -537,13 +352,10 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
aspect_ratio=aspect_ratio,
|
||||
)
|
||||
|
||||
# Terminal success status differs across backends: DeepInfra reports
|
||||
# "succeeded", OpenAI/Sora reports "completed". Accept both.
|
||||
# DeepInfra reports "succeeded", OpenAI/Sora "completed" — accept both.
|
||||
status = getattr(video, "status", None)
|
||||
if status not in ("completed", "succeeded"):
|
||||
# ``video.error`` is a structured SDK object (pydantic
|
||||
# VideoCreateError), not a string — str() it so the response
|
||||
# dict stays JSON-serializable for the tool layer.
|
||||
# ``video.error`` is a pydantic object — str() keeps the dict JSON-serializable.
|
||||
job_error = getattr(video, "error", None)
|
||||
return error_response(
|
||||
error=str(job_error) if job_error else f"video job ended with status={status!r}",
|
||||
@@ -554,11 +366,9 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
aspect_ratio=aspect_ratio,
|
||||
)
|
||||
|
||||
# Resolve the output. Providers expose it either as a delivery URL in
|
||||
# the job's ``data`` list (DeepInfra, FAL-style) or only via the SDK
|
||||
# download endpoint (OpenAI/Sora). Download the bytes and save locally
|
||||
# so the caller gets a durable file — DeepInfra's delivery URLs in
|
||||
# particular are short-lived. Matches plugins/image_gen/deepinfra.
|
||||
# Output is a delivery URL in ``data`` (DeepInfra/FAL) or only reachable
|
||||
# via the SDK download endpoint (OpenAI/Sora). Save locally either way —
|
||||
# DeepInfra's delivery URLs are short-lived.
|
||||
url = None
|
||||
for item in getattr(video, "data", None) or []:
|
||||
candidate = item.get("url") if isinstance(item, dict) else getattr(item, "url", None)
|
||||
@@ -568,15 +378,12 @@ class OpenAICompatibleVideoGenProvider(VideoGenProvider):
|
||||
|
||||
try:
|
||||
if url:
|
||||
# Materialise the (often short-lived) delivery URL locally.
|
||||
video_ref = str(save_url_video(url, prefix=self.name))
|
||||
else:
|
||||
# OpenAI/Sora style: no public URL — pull bytes via the SDK.
|
||||
raw = client.videos.download_content(video.id).read()
|
||||
video_ref = str(save_bytes_video(raw, prefix=self.name))
|
||||
except Exception as exc: # noqa: BLE001
|
||||
if url:
|
||||
# Best-effort: hand back the URL rather than fail outright.
|
||||
logger.debug("%s: saving video locally failed (%s); returning URL", self.name, exc)
|
||||
video_ref = url
|
||||
else:
|
||||
|
||||
+18
-138
@@ -15,138 +15,34 @@ If unset, :func:`get_active_provider` applies fallback logic:
|
||||
2. Otherwise return ``None`` (the tool surfaces a helpful error pointing
|
||||
the user at ``hermes tools``).
|
||||
|
||||
Mirrors ``agent/image_gen_registry.py`` so the two surfaces behave the
|
||||
same: the unconfigured fallback is filtered by ``is_available()`` so a box
|
||||
that has credentials for only one backend (e.g. DeepInfra, while the
|
||||
``fal``/``xai`` plugins also register unconditionally) auto-selects it
|
||||
instead of returning ``None``.
|
||||
Mirrors ``agent/image_gen_registry.py``: the unconfigured fallback is
|
||||
filtered by ``is_available()`` so a box with credentials for only one backend
|
||||
(e.g. DeepInfra, while ``fal``/``xai`` register unconditionally) auto-selects
|
||||
it instead of returning ``None``. Unlike image gen there is no legacy ``fal``
|
||||
preference, and a configured-but-unregistered name fails closed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from agent.provider_registry import ProviderRegistry, configured_provider_name, is_available_safe
|
||||
from agent.video_gen_provider import VideoGenProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_providers: Dict[str, VideoGenProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, VideoGenProvider]] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def register_provider(provider: VideoGenProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a video generation provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and logs
|
||||
a debug message — this makes hot-reload scenarios (tests, dev loops)
|
||||
behave predictably.
|
||||
"""
|
||||
if not isinstance(provider, VideoGenProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a VideoGenProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Video gen provider .name must be a non-empty string")
|
||||
name = raw_name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if existing is not None:
|
||||
logger.debug("Video gen provider '%s' re-registered (was %r)", name, type(existing).__name__)
|
||||
else:
|
||||
logger.debug("Registered video gen provider '%s' (%s)", name, type(provider).__name__)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[VideoGenProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[VideoGenProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
with _lock:
|
||||
key = name.strip()
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[VideoGenProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip())
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: VideoGenProvider,
|
||||
previous: Optional[VideoGenProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is not None and not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
_registry: ProviderRegistry[VideoGenProvider] = ProviderRegistry(
|
||||
label="Video gen", provider_cls=VideoGenProvider, logger=logger,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
def get_active_provider() -> Optional[VideoGenProvider]:
|
||||
"""Resolve the currently-active provider.
|
||||
|
||||
Reads ``video_gen.provider`` from config.yaml; falls back per the
|
||||
module docstring.
|
||||
"""
|
||||
configured: Optional[str] = None
|
||||
try:
|
||||
from hermes_cli.config import load_config_readonly
|
||||
|
||||
cfg = load_config_readonly()
|
||||
section = cfg.get("video_gen") if isinstance(cfg, dict) else None
|
||||
if isinstance(section, dict):
|
||||
raw = section.get("provider")
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
configured = raw.strip()
|
||||
except Exception as exc:
|
||||
logger.debug("Could not read video_gen.provider from config: %s", exc)
|
||||
|
||||
# The managed "Nous Subscription" selection is serviced by the FAL
|
||||
# plugin through the managed fal-queue gateway (the plugin's resolver
|
||||
# routes managed when the stored selection is "nous").
|
||||
if configured:
|
||||
try:
|
||||
from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER
|
||||
|
||||
if configured.lower() == NOUS_MANAGED_PROVIDER:
|
||||
configured = "fal"
|
||||
except Exception: # pragma: no cover — helpers are in-repo
|
||||
pass
|
||||
|
||||
with _lock:
|
||||
snapshot = dict(_providers)
|
||||
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
|
||||
"""Resolve the currently-active provider (see module docstring)."""
|
||||
configured = configured_provider_name("video_gen", logger)
|
||||
snapshot = _registry.merged()
|
||||
|
||||
if configured:
|
||||
provider = snapshot.get(configured)
|
||||
@@ -158,27 +54,11 @@ def get_active_provider() -> Optional[VideoGenProvider]:
|
||||
)
|
||||
return None
|
||||
|
||||
def _is_available_safe(p: VideoGenProvider) -> bool:
|
||||
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
|
||||
try:
|
||||
return bool(p.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("video_gen provider %s.is_available() raised %s", p.name, exc)
|
||||
return False
|
||||
|
||||
# Fallback: single *available* provider — filter by is_available() so a
|
||||
# box with credentials for only one backend auto-selects it even when
|
||||
# other providers (fal/xai) register unconditionally without keys.
|
||||
# Mirrors agent/image_gen_registry.get_active_provider().
|
||||
available = [p for p in snapshot.values() if _is_available_safe(p)]
|
||||
available = [
|
||||
p for p in snapshot.values()
|
||||
if is_available_safe(p, logger, "video_gen provider %s.is_available() raised %s")
|
||||
]
|
||||
if len(available) == 1:
|
||||
return available[0]
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
|
||||
+45
-163
@@ -2,51 +2,25 @@
|
||||
Web Search Provider ABC
|
||||
=======================
|
||||
|
||||
Defines the pluggable-backend interface for web search and content extraction.
|
||||
Providers register instances via ``PluginContext.register_web_search_provider()``;
|
||||
the active one (selected via ``web.search_backend`` / ``web.extract_backend`` /
|
||||
``web.backend`` in ``config.yaml``) services every ``web_search`` /
|
||||
``web_extract`` tool call.
|
||||
Pluggable-backend interface for web search and content extraction — the SINGLE
|
||||
plugin-facing surface every in-tree web provider (brave-free, ddgs, searxng,
|
||||
exa, parallel, tavily, keenable, firecrawl) implements. Providers register via
|
||||
``PluginContext.register_web_search_provider()``; the active one (selected by
|
||||
``web.search_backend`` / ``web.extract_backend`` / ``web.backend``) services
|
||||
every ``web_search`` / ``web_extract`` call.
|
||||
|
||||
Providers live in ``<repo>/plugins/web/<name>/`` (built-in, auto-loaded as
|
||||
``kind: backend``) or ``~/.hermes/plugins/web/<name>/`` (user, opt-in via
|
||||
``plugins.enabled``).
|
||||
Response shape (preserved from the legacy contract so the tool wrapper does not
|
||||
translate). Search::
|
||||
|
||||
This ABC is the SINGLE plugin-facing surface for web providers — every
|
||||
provider in the tree (brave-free, ddgs, searxng, exa, parallel, tavily,
|
||||
keenable, firecrawl) implements it. The legacy in-tree ``tools.web_providers.base``
|
||||
ABCs were deleted in PR #25182 along with the per-vendor inline helpers
|
||||
in ``tools/web_tools.py``; the response-shape contract documented below
|
||||
is preserved bit-for-bit so the tool wrapper does not have to translate.
|
||||
{"success": True, "data": {"web": [
|
||||
{"title": str, "url": str, "description": str, "position": int}, ...]}}
|
||||
|
||||
Response shape (preserved from the legacy contract):
|
||||
Extract::
|
||||
|
||||
Search results::
|
||||
{"success": True, "data": [
|
||||
{"url": str, "title": str, "content": str, "raw_content": str, "metadata": dict}, ...]}
|
||||
|
||||
{
|
||||
"success": True,
|
||||
"data": {
|
||||
"web": [
|
||||
{"title": str, "url": str, "description": str, "position": int},
|
||||
...
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
Extract results::
|
||||
|
||||
{
|
||||
"success": True,
|
||||
"data": [
|
||||
{"url": str, "title": str, "content": str,
|
||||
"raw_content": str, "metadata": dict},
|
||||
...
|
||||
]
|
||||
}
|
||||
|
||||
On failure (either capability)::
|
||||
|
||||
{"success": False, "error": str}
|
||||
On failure (either capability): ``{"success": False, "error": str}``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -55,19 +29,16 @@ import abc
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from agent.provider_base import ProviderBase
|
||||
|
||||
|
||||
def get_provider_env(name: str) -> str:
|
||||
"""Config-aware env lookup for web providers.
|
||||
"""Config-aware env lookup: ``os.environ`` first, then ``~/.hermes/.env``.
|
||||
|
||||
Resolves *name* via :func:`hermes_cli.config.get_env_value` (checks
|
||||
``os.environ`` first, then ``~/.hermes/.env``) so credentials set
|
||||
through Hermes' config layer are visible even when they were never
|
||||
exported into the process environment — gateway sessions, delegate
|
||||
children, and subprocess agent runs (issue #40190). Falls back to a
|
||||
bare ``os.getenv`` when the config module is unavailable (stripped
|
||||
installs, early import contexts).
|
||||
|
||||
Returns the stripped value, or ``""`` when unset.
|
||||
Credentials set through Hermes' config layer must be visible even when never
|
||||
exported into the process environment (gateway sessions, delegate children,
|
||||
subprocess agent runs). Falls back to bare ``os.getenv`` when the config
|
||||
module is unavailable. Returns the stripped value, or ``""`` when unset.
|
||||
"""
|
||||
val: Optional[str] = None
|
||||
try:
|
||||
@@ -81,147 +52,58 @@ def get_provider_env(name: str) -> str:
|
||||
return (val or "").strip()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ABC
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class WebSearchProvider(abc.ABC):
|
||||
class WebSearchProvider(ProviderBase):
|
||||
"""Abstract base class for a web search/extract backend.
|
||||
|
||||
Subclasses must implement :meth:`is_available` and at least one of
|
||||
:meth:`search` / :meth:`extract`. The :meth:`supports_search` /
|
||||
:meth:`supports_extract` capability flags let the registry route each
|
||||
tool call to the right provider, and let multi-capability providers
|
||||
(Firecrawl, Tavily, Exa, …) advertise multiple capabilities from a
|
||||
single class.
|
||||
Subclasses implement :meth:`is_available` and at least one of :meth:`search`
|
||||
/ :meth:`extract`; the :meth:`supports_search` / :meth:`supports_extract`
|
||||
flags let the registry route each capability, so one class can serve both.
|
||||
"""
|
||||
|
||||
@property
|
||||
@abc.abstractmethod
|
||||
def name(self) -> str:
|
||||
"""Stable short identifier used in ``web.search_backend`` /
|
||||
``web.extract_backend`` / ``web.backend`` config keys.
|
||||
|
||||
Lowercase, no spaces; hyphens permitted to preserve existing
|
||||
user-visible names. Examples: ``brave-free``, ``ddgs``,
|
||||
``searxng``, ``firecrawl``.
|
||||
"""
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Human-readable label shown in ``hermes tools``. Defaults to ``name``."""
|
||||
return self.name
|
||||
|
||||
@abc.abstractmethod
|
||||
def is_available(self) -> bool:
|
||||
"""Return True when this provider can service calls.
|
||||
"""True when this provider can service calls.
|
||||
|
||||
Typically a cheap check (env var present, optional Python dep
|
||||
importable, instance URL set). Must NOT make network calls — this
|
||||
runs at tool-registration time and on every ``hermes tools`` paint.
|
||||
Cheap check only (env var present, dep importable, instance URL set) —
|
||||
must NOT make network calls; runs at tool-registration time and on every
|
||||
``hermes tools`` paint.
|
||||
"""
|
||||
|
||||
def supports_search(self) -> bool:
|
||||
"""Return True if this provider implements :meth:`search`."""
|
||||
"""True if this provider implements :meth:`search`."""
|
||||
return True
|
||||
|
||||
def is_keyless_available(self) -> bool:
|
||||
"""Return True when this provider can serve calls WITHOUT credentials.
|
||||
"""True when this provider can serve calls WITHOUT credentials.
|
||||
|
||||
A separate, weaker tier than :meth:`is_available`: providers with a
|
||||
public anonymous free tier (Exa / Parallel MCP endpoints) return
|
||||
True here so the registry can fall back to them when NO provider is
|
||||
configured or keyed — and only then. Keyless availability must never
|
||||
make :meth:`is_available` return True, or the legacy preference walk
|
||||
would route users with real credentials for a lower-priority backend
|
||||
onto the free tier of a higher-priority one.
|
||||
|
||||
Like :meth:`is_available`, this must be cheap and must NOT make
|
||||
network calls. Default: False.
|
||||
A weaker tier than :meth:`is_available`, used only when NO provider is
|
||||
configured or keyed (public anonymous free tiers such as Exa / Parallel
|
||||
MCP). It must never make :meth:`is_available` True, or the legacy
|
||||
preference walk would route users holding real credentials for a
|
||||
lower-priority backend onto a higher-priority backend's free tier.
|
||||
Cheap, no network. Default False.
|
||||
"""
|
||||
return False
|
||||
|
||||
def supports_extract(self) -> bool:
|
||||
"""Return True if this provider implements :meth:`extract`.
|
||||
|
||||
Both sync and async :meth:`extract` implementations are valid — the
|
||||
dispatcher detects coroutine functions via
|
||||
:func:`inspect.iscoroutinefunction` and awaits as needed. Sync
|
||||
implementations that perform blocking I/O (HTTP, SDK calls) should
|
||||
ideally wrap in :func:`asyncio.to_thread` at the call site; small
|
||||
providers can keep their sync shape and let the dispatcher handle
|
||||
threading.
|
||||
"""
|
||||
"""True if this provider implements :meth:`extract` (sync or ``async def`` —
|
||||
the dispatcher awaits coroutine functions)."""
|
||||
return False
|
||||
|
||||
def search(self, query: str, limit: int = 5) -> Dict[str, Any]:
|
||||
"""Execute a web search.
|
||||
|
||||
Override when :meth:`supports_search` returns True. The default
|
||||
raises NotImplementedError; callers should gate on
|
||||
:meth:`supports_search` before calling.
|
||||
"""
|
||||
"""Execute a web search. Callers gate on :meth:`supports_search`."""
|
||||
raise NotImplementedError(
|
||||
f"{self.name} does not support search (override supports_search)"
|
||||
)
|
||||
|
||||
def extract(self, urls: List[str], **kwargs: Any) -> Any:
|
||||
"""Extract content from one or more URLs.
|
||||
"""Extract content from URLs. Callers gate on :meth:`supports_extract`.
|
||||
|
||||
Override when :meth:`supports_extract` returns True. The default
|
||||
raises NotImplementedError; callers should gate on
|
||||
:meth:`supports_extract` before calling.
|
||||
|
||||
Return shape: a list of result dicts matching what the legacy
|
||||
:func:`tools.web_tools.web_extract_tool` post-processing pipeline
|
||||
expects::
|
||||
|
||||
[
|
||||
{
|
||||
"url": str,
|
||||
"title": str,
|
||||
"content": str,
|
||||
"raw_content": str,
|
||||
"metadata": dict, # optional
|
||||
"error": str, # optional, only on per-URL failure
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
Implementations MAY be ``async def`` — the dispatcher detects
|
||||
coroutines via :func:`inspect.iscoroutinefunction` and awaits.
|
||||
|
||||
``kwargs`` may carry forward-compat fields (``format``, ``include_raw``,
|
||||
``max_chars``) — implementations should ignore unknown keys.
|
||||
Returns a list of ``{"url", "title", "content", "raw_content",
|
||||
"metadata"?, "error"?}`` dicts (``error`` only on per-URL failure).
|
||||
May be ``async def``. ``kwargs`` may carry forward-compat fields
|
||||
(``format``, ``include_raw``, ``max_chars``) — ignore unknown keys.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
f"{self.name} does not support extract (override supports_extract)"
|
||||
)
|
||||
|
||||
def get_setup_schema(self) -> Dict[str, Any]:
|
||||
"""Return provider metadata for the ``hermes tools`` picker.
|
||||
|
||||
Used by ``hermes_cli/tools_config.py`` to inject this provider as a
|
||||
row in the Web Search / Web Extract picker. Shape::
|
||||
|
||||
{
|
||||
"name": "Brave Search (Free)",
|
||||
"badge": "free",
|
||||
"tag": "No paid tier needed — uses Brave's free API.",
|
||||
"env_vars": [
|
||||
{"key": "BRAVE_SEARCH_API_KEY",
|
||||
"prompt": "Brave Search API key",
|
||||
"url": "https://brave.com/search/api/"},
|
||||
],
|
||||
}
|
||||
|
||||
Default: minimal entry derived from ``display_name``. Override to
|
||||
expose API key prompts, badges, and instance URL fields.
|
||||
"""
|
||||
return {
|
||||
"name": self.display_name,
|
||||
"badge": "",
|
||||
"tag": "",
|
||||
"env_vars": [],
|
||||
}
|
||||
|
||||
+41
-192
@@ -33,98 +33,18 @@ extract-capable backend.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from agent.provider_registry import ProviderRegistry, is_available_safe
|
||||
from agent.web_search_provider import WebSearchProvider
|
||||
from hermes_constants import hermes_home_key
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_providers: Dict[str, WebSearchProvider] = {}
|
||||
_scoped_providers: Dict[str, Dict[str, WebSearchProvider]] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
|
||||
def register_provider(provider: WebSearchProvider, *, scope: Optional[str] = None) -> None:
|
||||
"""Register a web search/extract provider.
|
||||
|
||||
Re-registration (same ``name``) overwrites the previous entry and logs
|
||||
a debug message — makes hot-reload scenarios (tests, dev loops) behave
|
||||
predictably.
|
||||
"""
|
||||
if not isinstance(provider, WebSearchProvider):
|
||||
raise TypeError(
|
||||
f"register_provider() expects a WebSearchProvider instance, "
|
||||
f"got {type(provider).__name__}"
|
||||
)
|
||||
raw_name = provider.name
|
||||
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||
raise ValueError("Web provider .name must be a non-empty string")
|
||||
name = raw_name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
existing = target.get(name)
|
||||
target[name] = provider
|
||||
if existing is not None:
|
||||
logger.debug(
|
||||
"Web provider '%s' re-registered (was %r)",
|
||||
name, type(existing).__name__,
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Registered web provider '%s' (%s)",
|
||||
name, type(provider).__name__,
|
||||
)
|
||||
|
||||
|
||||
def list_providers(*, scope: Optional[str] = None) -> List[WebSearchProvider]:
|
||||
"""Return all registered providers, sorted by name."""
|
||||
with _lock:
|
||||
merged = dict(_providers)
|
||||
merged.update(_scoped_providers.get(scope or hermes_home_key(), {}))
|
||||
items = list(merged.values())
|
||||
return sorted(items, key=lambda p: p.name)
|
||||
|
||||
|
||||
def get_provider(name: str, *, scope: Optional[str] = None) -> Optional[WebSearchProvider]:
|
||||
"""Return the provider registered under *name*, or None."""
|
||||
if not isinstance(name, str):
|
||||
return None
|
||||
with _lock:
|
||||
key = name.strip()
|
||||
return _scoped_providers.get(scope or hermes_home_key(), {}).get(key) or _providers.get(key)
|
||||
|
||||
|
||||
def snapshot_registration(
|
||||
name: str, *, scope: Optional[str] = None
|
||||
) -> Optional[WebSearchProvider]:
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.get(scope, {})
|
||||
return target.get(name.strip())
|
||||
|
||||
|
||||
def restore_registration(
|
||||
name: str,
|
||||
current: WebSearchProvider,
|
||||
previous: Optional[WebSearchProvider],
|
||||
*,
|
||||
scope: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""Restore a plugin registration only when *current* is still installed."""
|
||||
key = name.strip()
|
||||
with _lock:
|
||||
target = _providers if scope is None else _scoped_providers.setdefault(scope, {})
|
||||
if target.get(key) is not current:
|
||||
return False
|
||||
if previous is None:
|
||||
target.pop(key, None)
|
||||
else:
|
||||
target[key] = previous
|
||||
if scope is not None and not target:
|
||||
_scoped_providers.pop(scope, None)
|
||||
return True
|
||||
_registry: ProviderRegistry[WebSearchProvider] = ProviderRegistry(
|
||||
label="Web", provider_cls=WebSearchProvider, logger=logger,
|
||||
)
|
||||
_registry.export(globals())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -150,6 +70,11 @@ def _read_config_key(*path: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
|
||||
def _configured_backend(capability: str) -> Optional[str]:
|
||||
"""``web.<capability>_backend`` (preferred) or ``web.backend`` (shared fallback)."""
|
||||
return _read_config_key("web", f"{capability}_backend") or _read_config_key("web", "backend")
|
||||
|
||||
|
||||
# Legacy preference order — preserves behaviour for users who set no
|
||||
# ``web.backend`` / ``web.<capability>_backend`` config key at all. Matches
|
||||
# the historic candidate order in :func:`tools.web_tools._get_backend`
|
||||
@@ -207,36 +132,13 @@ def _keyless_preference() -> tuple:
|
||||
def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearchProvider]:
|
||||
"""Resolve the active provider for a capability ("search" | "extract").
|
||||
|
||||
Resolution rules (in order):
|
||||
|
||||
1. **Explicit config wins, ignoring availability.** If
|
||||
``web.{capability}_backend`` or ``web.backend`` names a registered
|
||||
provider that supports *capability*, return it even if its
|
||||
:meth:`is_available` returns False — the dispatcher will surface a
|
||||
precise "X_API_KEY is not set" error to the user instead of silently
|
||||
routing somewhere else. Matches legacy
|
||||
:func:`tools.web_tools._get_backend` behavior for configured names.
|
||||
|
||||
2. **Single-provider shortcut.** When only one registered provider
|
||||
supports *capability* AND ``is_available()`` reports True, return it.
|
||||
|
||||
3. **Legacy preference walk, filtered by availability.** Walk the
|
||||
:data:`_LEGACY_PREFERENCE` order (firecrawl → parallel → tavily →
|
||||
exa → searxng → brave-free → ddgs) looking for a provider whose
|
||||
``supports_<capability>()`` is True AND whose ``is_available()`` is
|
||||
True. Matches the historic ``tools.web_tools._get_backend()``
|
||||
candidate order so users with credentials but no explicit config
|
||||
key keep landing on the same provider as pre-migration. This is
|
||||
the path that fires when no config key is set — pick the
|
||||
highest-priority backend the user actually has credentials for.
|
||||
|
||||
Returns None when no provider is configured AND no available provider
|
||||
matches the legacy preference; the dispatcher then returns a "set up a
|
||||
provider" error to the user.
|
||||
Rules, in order (see module docstring): explicit config wins even when
|
||||
``is_available()`` is False (the dispatcher surfaces a precise
|
||||
"X_API_KEY is not set" error instead of a silent switch); then the single
|
||||
available capable provider; then the availability-filtered legacy walk;
|
||||
then the keyless free-tier walk; else None.
|
||||
"""
|
||||
with _lock:
|
||||
snapshot = dict(_providers)
|
||||
snapshot.update(_scoped_providers.get(hermes_home_key(), {}))
|
||||
snapshot = _registry.merged()
|
||||
|
||||
def _capable(p: WebSearchProvider) -> bool:
|
||||
if capability == "search":
|
||||
@@ -245,17 +147,9 @@ def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearc
|
||||
return bool(p.supports_extract())
|
||||
return False
|
||||
|
||||
def _is_available_safe(p: WebSearchProvider) -> bool:
|
||||
"""Wrap ``is_available()`` so a buggy provider doesn't kill resolution."""
|
||||
try:
|
||||
return bool(p.is_available())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("provider %s.is_available() raised %s", p.name, exc)
|
||||
return False
|
||||
def _available(p: WebSearchProvider) -> bool:
|
||||
return is_available_safe(p, logger, "provider %s.is_available() raised %s")
|
||||
|
||||
# 1. Explicit config wins — return regardless of is_available() so the
|
||||
# user gets a precise downstream error message rather than a silent
|
||||
# backend switch. Matches _get_backend() in web_tools.py.
|
||||
if configured:
|
||||
provider = snapshot.get(configured)
|
||||
if provider is not None and _capable(provider):
|
||||
@@ -271,31 +165,20 @@ def _resolve(configured: Optional[str], *, capability: str) -> Optional[WebSearc
|
||||
configured, capability,
|
||||
)
|
||||
|
||||
# 2. + 3. Fallback path — filter by availability so we don't surface
|
||||
# a provider the user has no credentials for. Without this filter,
|
||||
# a registered-but-unconfigured provider could end up "active" on
|
||||
# a fresh install with no API keys at all.
|
||||
eligible = [
|
||||
p for p in snapshot.values()
|
||||
if _capable(p) and _is_available_safe(p)
|
||||
]
|
||||
# Fallbacks are availability-filtered so a registered-but-keyless provider
|
||||
# never becomes "active" on a fresh install.
|
||||
eligible = [p for p in snapshot.values() if _capable(p) and _available(p)]
|
||||
if len(eligible) == 1:
|
||||
return eligible[0]
|
||||
|
||||
for legacy in _LEGACY_PREFERENCE:
|
||||
provider = snapshot.get(legacy)
|
||||
if (
|
||||
provider is not None
|
||||
and _capable(provider)
|
||||
and _is_available_safe(provider)
|
||||
):
|
||||
if provider is not None and _capable(provider) and _available(provider):
|
||||
return provider
|
||||
|
||||
# 4. Keyless free-tier walk — the user has NO credentialed/importable
|
||||
# backend at all. Fall back to providers that can serve anonymously
|
||||
# (public MCP free tiers), unless disabled via
|
||||
# ``web.keyless_fallback: false``. This tier never pre-empts a keyed
|
||||
# setup: it is only reachable when the legacy walk found nothing.
|
||||
# Keyless free tier (anonymous public MCP tiers) is last-resort only: it is
|
||||
# reachable solely when the legacy walk found nothing, never pre-empting a
|
||||
# keyed setup. Disabled via ``web.keyless_fallback: false``.
|
||||
if _keyless_tier_enabled():
|
||||
for name in _keyless_preference():
|
||||
provider = snapshot.get(name)
|
||||
@@ -325,41 +208,23 @@ def _keyless_tier_enabled() -> bool:
|
||||
|
||||
|
||||
def _disabled_web_plugin_for(configured: Optional[str] = None, *, capability: Optional[str] = None) -> Optional[str]:
|
||||
"""Return the plugin key of a *disabled* bundled web plugin that would
|
||||
have provided the configured backend, or None.
|
||||
"""Plugin key of a *disabled* bundled web plugin that would have provided
|
||||
the configured backend (``web.<capability>_backend`` → ``web.backend``),
|
||||
or None.
|
||||
|
||||
When a user sets ``web.extract_backend: firecrawl`` (or the search
|
||||
equivalent) but also lists ``web-firecrawl`` in ``plugins.disabled``,
|
||||
the provider never registers and the dispatcher would otherwise emit a
|
||||
misleading "No web extract provider configured. Set web.extract_backend
|
||||
to ..." error — even though the backend IS configured correctly. The
|
||||
real fix is to re-enable the plugin. This helper detects that case so
|
||||
the dispatcher can point the user at the actual cause (issue #40190
|
||||
follow-up: pi314's disabled-plugin symptom).
|
||||
|
||||
Pass ``capability`` ("search" | "extract") to resolve the configured
|
||||
name straight from ``config.yaml`` (``web.<capability>_backend`` →
|
||||
``web.backend``). This is more reliable than the resolved backend the
|
||||
dispatcher fell back to, since a disabled provider fails the
|
||||
``_is_backend_available`` gate and the dispatcher silently drops to
|
||||
the shared default. An explicit ``configured`` name still wins when
|
||||
given.
|
||||
|
||||
Matching is by convention: bundled web plugins live under the
|
||||
``web/<vendor>`` key with the provider ``name`` differing only in
|
||||
hyphen/underscore (``brave-free`` provider ⇄ ``web/brave_free`` key,
|
||||
``firecrawl`` ⇄ ``web/firecrawl``). We normalize both sides before
|
||||
comparing so every bundled provider is covered without hardcoding a
|
||||
per-vendor table.
|
||||
Lets the dispatcher say "re-enable web-firecrawl" instead of a misleading
|
||||
"No web extract provider configured" when the backend IS configured but
|
||||
listed in ``plugins.disabled``. Resolving from config.yaml (rather than
|
||||
the resolved backend) matters because a disabled provider fails the
|
||||
availability gate and the dispatcher silently drops to the default.
|
||||
Bundled web plugins live under ``web/<vendor>`` with the provider name
|
||||
differing only by hyphen/underscore, so both sides are normalized.
|
||||
"""
|
||||
def _norm(s: str) -> str:
|
||||
return s.strip().lower().replace("-", "_")
|
||||
|
||||
if not configured and capability in ("search", "extract"):
|
||||
configured = (
|
||||
_read_config_key("web", f"{capability}_backend")
|
||||
or _read_config_key("web", "backend")
|
||||
)
|
||||
configured = _configured_backend(capability)
|
||||
if not configured:
|
||||
return None
|
||||
|
||||
@@ -384,27 +249,11 @@ def _disabled_web_plugin_for(configured: Optional[str] = None, *, capability: Op
|
||||
|
||||
|
||||
def get_active_search_provider() -> Optional[WebSearchProvider]:
|
||||
"""Resolve the currently-active web search provider.
|
||||
|
||||
Reads ``web.search_backend`` (preferred) or ``web.backend`` (shared
|
||||
fallback) from config.yaml; falls back per the module docstring.
|
||||
"""
|
||||
explicit = _read_config_key("web", "search_backend") or _read_config_key("web", "backend")
|
||||
return _resolve(explicit, capability="search")
|
||||
"""Resolve the currently-active web search provider."""
|
||||
return _resolve(_configured_backend("search"), capability="search")
|
||||
|
||||
|
||||
def get_active_extract_provider() -> Optional[WebSearchProvider]:
|
||||
"""Resolve the currently-active web extract provider.
|
||||
"""Resolve the currently-active web extract provider."""
|
||||
return _resolve(_configured_backend("extract"), capability="extract")
|
||||
|
||||
Reads ``web.extract_backend`` (preferred) or ``web.backend`` (shared
|
||||
fallback) from config.yaml; falls back per the module docstring.
|
||||
"""
|
||||
explicit = _read_config_key("web", "extract_backend") or _read_config_key("web", "backend")
|
||||
return _resolve(explicit, capability="extract")
|
||||
|
||||
|
||||
def _reset_for_tests() -> None:
|
||||
"""Clear the registry. **Test-only.**"""
|
||||
with _lock:
|
||||
_providers.clear()
|
||||
_scoped_providers.clear()
|
||||
|
||||
Reference in New Issue
Block a user