Merge branch 'simp/hclia-providers' into simp/integration

This commit is contained in:
Teknium
2026-09-02 14:19:39 -07:00
32 changed files with 2509 additions and 6110 deletions
+30 -121
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+24 -51
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+61 -107
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+84 -278
View File
@@ -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 []
+77
View File
@@ -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
+112
View File
@@ -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
+244
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()