diff --git a/agent/browser_provider.py b/agent/browser_provider.py index d7c6dae61f..1ea59a40a3 100644 --- a/agent/browser_provider.py +++ b/agent/browser_provider.py @@ -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 ``/plugins/browser//`` +(built-in) or ``~/.hermes/plugins/browser//`` (user, opt-in). -Providers live in ``/plugins/browser//`` (built-in, auto-loaded as -``kind: backend``) or ``~/.hermes/plugins/browser//`` (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`.""" diff --git a/agent/browser_registry.py b/agent/browser_registry.py index 4348237af1..93a88dd011 100644 --- a/agent/browser_registry.py +++ b/agent/browser_registry.py @@ -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//`` 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 diff --git a/agent/context_engine.py b/agent/context_engine.py index b772125c0c..ff153fa2f9 100644 --- a/agent/context_engine.py +++ b/agent/context_engine.py @@ -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//`` 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//``). 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 ``. - 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 `` (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", {}), diff --git a/agent/image_gen_provider.py b/agent/image_gen_provider.py index b76bc8c01a..86857eed51 100644 --- a/agent/image_gen_provider.py +++ b/agent/image_gen_provider.py @@ -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 +``/plugins/image_gen//`` (built-in) or +``~/.hermes/plugins/image_gen//`` (user, opt-in). -Providers live in ``/plugins/image_gen//`` (built-in, auto-loaded -as ``kind: backend``) or ``~/.hermes/plugins/image_gen//`` (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: ``__.``. - """ - 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, diff --git a/agent/image_gen_registry.py b/agent/image_gen_registry.py index 6239bbe891..087144f8be 100644 --- a/agent/image_gen_registry.py +++ b/agent/image_gen_registry.py @@ -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() diff --git a/agent/image_routing.py b/agent/image_routing.py index a861cd29bf..02c21b9ef6 100644 --- a/agent/image_routing.py +++ b/agent/image_routing.py @@ -3,38 +3,25 @@ Two modes: native — attach images as OpenAI-style ``image_url`` content parts on the - user turn. Provider adapters (Anthropic, Gemini, Bedrock, Codex, - OpenAI chat.completions) already translate these into their - vendor-specific multimodal formats. - + user turn; provider adapters translate these to vendor formats. text — run ``vision_analyze`` on each image up-front and prepend the - description to the user's text. The model never sees the pixels; - it only sees a lossy text summary. This is the pre-existing - behaviour and still the right choice for non-vision models. + (lossy) description to the user's text. Still the right choice + for non-vision models. -The decision is made once per message turn by :func:`decide_image_input_mode`. -It reads ``agent.image_input_mode`` from config.yaml (``auto`` | ``native`` -| ``text``, default ``auto``) and the active model's capability metadata. +:func:`decide_image_input_mode` decides once per message turn from +``agent.image_input_mode`` (``auto`` | ``native`` | ``text``, default ``auto``) +plus the active model's capability metadata. In ``auto`` mode: -In ``auto`` mode: - - If the user has explicitly configured ``auxiliary.vision`` - (provider/model/base_url not ``auto``/empty), images route through - that backend — the DE-FACTO choice: a user who named a dedicated - vision model wants it used, even when the main model has native - vision (maintainer decision 2026-08-28, reversing #29135's - fallback-only posture). - - Otherwise, if the active model reports ``supports_vision=True`` (via - config override or models.dev metadata), we attach natively. - - Otherwise (non-vision model, no aux backend), text via the default - vision_analyze flow. - ``agent.image_input_mode: native`` remains the absolute override for - users who want native attach despite a configured aux backend. + - An explicitly configured ``auxiliary.vision`` backend is the DE-FACTO route + (``text``): a user who named a dedicated vision model wants it used even + when the main model has native vision. ``image_input_mode: native`` is the + absolute override. + - Otherwise, ``supports_vision=True`` (config override or catalog) → native. + - Otherwise text via the default vision_analyze flow. -This keeps ``vision_analyze`` surfaced as a tool in every session — skills -and agent flows that chain it (browser screenshots, deeper inspection of -URL-referenced images, style-gating loops) keep working. The routing only -affects *how user-attached images on the current turn* are presented to the -main model. +``vision_analyze`` stays surfaced as a tool in every session so skills that +chain it keep working; routing only affects how *user-attached images on the +current turn* are presented to the main model. """ from __future__ import annotations @@ -45,7 +32,7 @@ import mimetypes import os import re from pathlib import Path -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Callable, Dict, Iterable, List, Optional, Tuple logger = logging.getLogger(__name__) @@ -53,27 +40,24 @@ logger = logging.getLogger(__name__) _VALID_MODES = frozenset({"auto", "native", "text"}) -# Image extensions used by extract_image_refs(). Kept tight on purpose — we -# only auto-attach things the model can actually see. Documents/archives are -# excluded because the gateway's broader extract_local_files() also routes -# them differently (send_document), and we don't want to attach a PDF as a -# vision part. +# Extensions extract_image_refs() auto-attaches. Kept tight: documents/archives +# are excluded because the gateway routes them via send_document, and we never +# want a PDF attached as a vision part. _IMAGE_EXTS = ( ".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".tiff", ".tif", ".heic", ) _IMAGE_EXT_PATTERN = "|".join(e.lstrip(".") for e in _IMAGE_EXTS) -# Absolute / home-relative local image path. Matches the same shape gateway's -# extract_local_files() uses: anchors to ``~/`` or ``/``, ignores matches inside -# URLs (the ``(?\"']+?\.(?:" + _IMAGE_EXT_PATTERN + r")(?:\?[^\s<>\"']*)?", re.IGNORECASE, @@ -83,80 +67,55 @@ _IMAGE_URL_RE = re.compile( def extract_image_refs(text: str) -> Tuple[List[str], List[str]]: """Scan free-form text for image references the model should see. - Returns ``(local_paths, urls)``: - - * ``local_paths`` — absolute (``/``) or home-relative (``~/``) paths - whose suffix is an image extension AND whose expanded form exists - on disk as a file. Order-preserving, deduplicated. - * ``urls`` — ``http(s)://…`` URLs whose path ends in an image - extension (a ``?query`` is allowed after the extension). - Order-preserving, deduplicated. - - Matches inside fenced code blocks (``` ``` ```) and inline backticks - (`` `…` ``) are skipped so that snippets pasted into a task body for - reference aren't mistaken for live attachments. This mirrors the - behaviour of ``gateway.platforms.base.BaseAdapter.extract_local_files``. - - Local paths are validated against the filesystem; URLs are not - (the provider fetches them at request time). + Returns ``(local_paths, urls)``, each order-preserving and deduplicated. + Local paths (``/`` or ``~/``) must exist on disk as files; URLs are not + validated (the provider fetches them). Matches inside fenced code blocks + and inline backticks are skipped so pasted example snippets aren't treated + as live attachments (mirrors ``BaseAdapter.extract_local_files``). """ if not isinstance(text, str) or not text: return [], [] - # Build spans covered by fenced code blocks and inline code so we can - # ignore references the author embedded purely as example text. - code_spans: list[tuple[int, int]] = [] - for m in re.finditer(r"```[^\n]*\n.*?```", text, re.DOTALL): - code_spans.append((m.start(), m.end())) - for m in re.finditer(r"`[^`\n]+`", text): - code_spans.append((m.start(), m.end())) + code_spans: list[tuple[int, int]] = [ + (m.start(), m.end()) + for pattern, flags in ((r"```[^\n]*\n.*?```", re.DOTALL), (r"`[^`\n]+`", 0)) + for m in re.finditer(pattern, text, flags) + ] def _in_code(pos: int) -> bool: return any(s <= pos < e for s, e in code_spans) local_paths: list[str] = [] - seen_paths: set[str] = set() for match in _LOCAL_IMAGE_PATH_RE.finditer(text): if _in_code(match.start()): continue - raw = match.group(0) - expanded = os.path.expanduser(raw) + expanded = os.path.expanduser(match.group(0)) try: if not os.path.isfile(expanded): continue except OSError: # ENAMETOOLONG / EINVAL on pathological inputs — skip rather than crash. continue - if expanded in seen_paths: - continue - seen_paths.add(expanded) - local_paths.append(expanded) + if expanded not in local_paths: + local_paths.append(expanded) urls: list[str] = [] - seen_urls: set[str] = set() for match in _IMAGE_URL_RE.finditer(text): if _in_code(match.start()): continue - url = match.group(0) - # Strip trailing punctuation that's almost certainly prose, not part - # of the URL (e.g. "see https://x.com/a.png." or "/a.png)"). - url = url.rstrip(".,;:!?)]>") - if url in seen_urls: - continue - seen_urls.add(url) - urls.append(url) + # Trailing punctuation is almost certainly prose ("see https://x/a.png."). + url = match.group(0).rstrip(".,;:!?)]>") + if url not in urls: + urls.append(url) return local_paths, urls -# Strict YAML/JSON boolean coercion for capability overrides. -# -# ``bool("false")`` is True in Python because non-empty strings are truthy, so -# a user writing ``supports_vision: "false"`` (quoted — a common YAML mistake) -# would silently enable native vision routing on a model that can't actually -# handle it. Accept only the values YAML 1.1 / 1.2 treat as booleans, plus -# real ``bool`` and integer 0/1. Anything else returns None so the caller -# falls through to models.dev rather than honouring garbage. +# Strict YAML/JSON boolean coercion for capability overrides. ``bool("false")`` +# is True, so a quoted ``supports_vision: "false"`` would silently enable native +# routing on a model that can't handle it. Accept only YAML boolean tokens, +# real bools and 0/1; anything else is None so the caller falls through to +# models.dev rather than honouring garbage. _TRUE_TOKENS = frozenset({"true", "yes", "on", "1"}) _FALSE_TOKENS = frozenset({"false", "no", "off", "0"}) @@ -166,9 +125,7 @@ def _coerce_capability_bool(raw: Any) -> Optional[bool]: if isinstance(raw, bool): return raw if isinstance(raw, int): - if raw in (0, 1): - return bool(raw) - return None + return bool(raw) if raw in (0, 1) else None if isinstance(raw, str): s = raw.strip().lower() if s in _TRUE_TOKENS: @@ -178,6 +135,47 @@ def _coerce_capability_bool(raw: Any) -> Optional[bool]: return None +def _dict_or_empty(raw: Any) -> Dict[str, Any]: + return raw if isinstance(raw, dict) else {} + + +def _clean_str(raw: Any) -> str: + return str(raw or "").strip() + + +def _runtime_main(key: str) -> str: + """Stripped context-local main-runtime value, or "" when unavailable.""" + try: + from agent.auxiliary_client import _runtime_main_value + + return _clean_str(_runtime_main_value(key)) + except Exception: + return "" + + +def _model_supports_vision_override(models_cfg: Any, model: str) -> Optional[bool]: + """Per-model ``supports_vision`` (or ``vision`` alias) from a ``models`` mapping.""" + per_model = _dict_or_empty(_dict_or_empty(models_cfg).get(model)) + return _coerce_capability_bool(per_model.get("supports_vision", per_model.get("vision"))) + + +def _custom_provider_entries(cfg: Dict[str, Any], names: Iterable[str]) -> Iterable[Dict[str, Any]]: + """Yield legacy ``custom_providers`` list entries whose ``name`` matches ``names``. + + Iterates ``names`` in the given priority order (outer loop) so list order + cannot let a persisted default shadow the live route. + """ + custom_providers = cfg.get("custom_providers") + if not isinstance(custom_providers, list): + return + entries = [e for e in custom_providers if isinstance(e, dict)] + for name in names: + wanted = name.strip().lower() + for entry in entries: + if _clean_str(entry.get("name")).lower() == wanted: + yield entry + + def _supports_vision_override( cfg: Optional[Dict[str, Any]], provider: str, @@ -185,123 +183,72 @@ def _supports_vision_override( *, requested_provider: str = "", ) -> Optional[bool]: - """Resolve user-declared vision capability from config.yaml. + """Resolve user-declared vision capability from config.yaml; None when unset. - Resolution order, first hit wins: - 1. ``model.supports_vision`` (top-level shortcut for the active model) - 2. ``providers..models..supports_vision`` - (named custom providers — ``provider`` may be the runtime-resolved - value ``"custom"``, the runtime's originally requested provider, - and/or the user-declared name under ``model.provider``; all are - tried. For ``custom:`` syntax, the stripped ```` is also - tried as a provider key.) - 2b. ``custom_providers`` (legacy list form) ``.models.`` - - Under (2) and (2b), the per-model capability key may be written as - either ``supports_vision`` or the shorter ``vision`` alias; both work. - - Returns None when no override is set, so the caller falls through to - models.dev. Returns False explicitly only when the user wrote a - recognised boolean false token. + First hit wins: ``model.supports_vision`` → ``providers.

.models.`` + → legacy ``custom_providers[].models.``. Named custom providers are + rewritten to ``provider="custom"`` at runtime while config keeps the user's + name under ``model.provider``, so the requested, runtime and config + identities are all tried, plus the bare ```` of any ``custom:``. """ if not isinstance(cfg, dict): return None - # 1. Top-level shortcut - model_cfg_raw = cfg.get("model") - model_cfg: Dict[str, Any] = model_cfg_raw if isinstance(model_cfg_raw, dict) else {} + model_cfg = _dict_or_empty(cfg.get("model")) top = _coerce_capability_bool(model_cfg.get("supports_vision")) if top is not None: return top - # 2. Per-provider, per-model. Named custom providers (e.g. "my-vllm") - # get rewritten to provider="custom" at runtime - # (hermes_cli/runtime_provider.py:_resolve_named_custom_runtime), so the - # config still holds the user-declared name under model.provider. Try - # both as candidate provider keys. Either identity may use the - # "custom:" form while providers: is keyed by bare . - config_provider = str(model_cfg.get("provider") or "").strip() provider_candidates: List[str] = [] - for candidate in (requested_provider, provider, config_provider): - if not candidate: - continue - provider_candidates.append(candidate) - if candidate.startswith("custom:"): - stripped_candidate = candidate[len("custom:"):] - if stripped_candidate: - provider_candidates.append(stripped_candidate) - providers_raw = cfg.get("providers") - providers_cfg: Dict[str, Any] = providers_raw if isinstance(providers_raw, dict) else {} - for p in dict.fromkeys(provider_candidates): - entry_raw = providers_cfg.get(p) - entry: Dict[str, Any] = entry_raw if isinstance(entry_raw, dict) else {} - models_raw = entry.get("models") - models_cfg: Dict[str, Any] = models_raw if isinstance(models_raw, dict) else {} - per_model_raw = models_cfg.get(model) - per_model: Dict[str, Any] = per_model_raw if isinstance(per_model_raw, dict) else {} - coerced = _coerce_capability_bool( - per_model.get("supports_vision", per_model.get("vision")) - ) + for candidate in (requested_provider, provider, _clean_str(model_cfg.get("provider"))): + if candidate: + provider_candidates.append(candidate) + if candidate.startswith("custom:") and candidate[len("custom:"):]: + provider_candidates.append(candidate[len("custom:"):]) + provider_candidates = list(dict.fromkeys(provider_candidates)) + + providers_cfg = _dict_or_empty(cfg.get("providers")) + for p in provider_candidates: + coerced = _model_supports_vision_override(_dict_or_empty(providers_cfg.get(p)).get("models"), model) if coerced is not None: return coerced - # 2b. Legacy list-style custom_providers. Entries are dicts with a - # "name" key and a nested "models" dict. Match by provider name (which - # may appear as the raw name or "custom:" at runtime). - custom_providers = cfg.get("custom_providers") - if isinstance(custom_providers, list): - # Candidate priority matters when the CLI-selected provider differs - # from model.provider. Walk identities first, then config entries, so - # list order cannot let the persisted default shadow the live route. - for candidate in dict.fromkeys(provider_candidates): - candidate_name = candidate.strip().lower() - for entry_raw in custom_providers: - if not isinstance(entry_raw, dict): - continue - entry_name = str(entry_raw.get("name") or "").strip().lower() - if entry_name != candidate_name: - continue - models_raw = entry_raw.get("models") - models_cfg = models_raw if isinstance(models_raw, dict) else {} - per_model_raw = models_cfg.get(model) - per_model = per_model_raw if isinstance(per_model_raw, dict) else {} - coerced = _coerce_capability_bool( - per_model.get("supports_vision", per_model.get("vision")) - ) - if coerced is not None: - return coerced + for entry in _custom_provider_entries(cfg, provider_candidates): + coerced = _model_supports_vision_override(entry.get("models"), model) + if coerced is not None: + return coerced return None -def _resolve_inference_base_url( +def _resolve_inference_value( cfg: Optional[Dict[str, Any]], provider: str, + key: str, + *, + runtime_ok: Callable[[str], bool], ) -> str: - """Best-effort base URL for the active inference provider.""" - try: - from agent.auxiliary_client import _runtime_main_value + """Shared resolution for ``base_url`` / ``api_key`` of the active inference provider. - runtime = str(_runtime_main_value("base_url") or "").strip() - runtime_provider = str(_runtime_main_value("provider") or "").strip().lower() - requested_provider = str(provider or "").strip().lower() - if runtime and (not requested_provider or requested_provider == runtime_provider): - return runtime - except Exception: - pass + Order: context-local runtime value (when ``runtime_ok`` accepts it) → + ``model.`` → ``providers..`` → ``custom_providers[].``, + where ```` covers the provider and ``model.provider`` in both bare + and ``custom:``-prefixed forms. + """ + runtime = _runtime_main(key) + if runtime and runtime_ok(runtime): + return runtime if not isinstance(cfg, dict): return "" - model_cfg_raw = cfg.get("model") - model_cfg: Dict[str, Any] = model_cfg_raw if isinstance(model_cfg_raw, dict) else {} - base_url = str(model_cfg.get("base_url") or "").strip() - if base_url: - return base_url + model_cfg = _dict_or_empty(cfg.get("model")) + value = _clean_str(model_cfg.get(key)) + if value: + return value - config_provider = str(model_cfg.get("provider") or "").strip() candidate_names: set[str] = set() - for p in filter(None, (provider, config_provider)): + for p in filter(None, (provider, _clean_str(model_cfg.get("provider")))): candidate_names.add(p) if p.lower().startswith("custom:"): candidate_names.add(p.split(":", 1)[1]) @@ -313,9 +260,9 @@ def _resolve_inference_base_url( for name in candidate_names: entry = providers_cfg.get(name) if isinstance(entry, dict): - bu = str(entry.get("base_url") or "").strip() - if bu: - return bu + value = _clean_str(entry.get(key)) + if value: + return value custom_providers = cfg.get("custom_providers") if isinstance(custom_providers, list): @@ -323,79 +270,44 @@ def _resolve_inference_base_url( for entry_raw in custom_providers: if not isinstance(entry_raw, dict): continue - entry_name = str(entry_raw.get("name") or "").strip() + entry_name = _clean_str(entry_raw.get("name")) if entry_name not in candidate_names and entry_name.lower() not in lowered: continue - bu = str(entry_raw.get("base_url") or "").strip() - if bu: - return bu + value = _clean_str(entry_raw.get(key)) + if value: + return value return "" +def _resolve_inference_base_url( + cfg: Optional[Dict[str, Any]], + provider: str, +) -> str: + """Best-effort base URL for the active inference provider. + + The runtime base_url is only trusted when it belongs to the requested + provider (or no provider was requested). + """ + requested_provider = _clean_str(provider).lower() + + def _runtime_ok(_: str) -> bool: + return not requested_provider or requested_provider == _runtime_main("provider").lower() + + return _resolve_inference_value(cfg, provider, "base_url", runtime_ok=_runtime_ok) + + def _resolve_inference_api_key( cfg: Optional[Dict[str, Any]], provider: str, ) -> str: """Best-effort API key for the active inference provider. - Mirrors :func:`_resolve_inference_base_url`'s resolution order (runtime - value, then ``model.api_key``, then the providers blocks) so the key - matches the base URL actually being probed. Without this, the local - server-type probe fires at a remote API-keyed endpoint without an - Authorization header — 5×401 per image-bearing turn on a keyed - sglang/vLLM deployment (#89863). + Mirrors :func:`_resolve_inference_base_url` so the key matches the base URL + actually probed; otherwise the local server-type probe hits a keyed remote + endpoint without Authorization and sprays 401s on every image turn. """ - try: - from agent.auxiliary_client import _runtime_main_value - - runtime_key = str(_runtime_main_value("api_key") or "").strip() - if runtime_key: - return runtime_key - except Exception: - pass - - if not isinstance(cfg, dict): - return "" - - model_cfg_raw = cfg.get("model") - model_cfg: Dict[str, Any] = model_cfg_raw if isinstance(model_cfg_raw, dict) else {} - key = str(model_cfg.get("api_key") or "").strip() - if key: - return key - - config_provider = str(model_cfg.get("provider") or "").strip() - candidate_names: set[str] = set() - for p in filter(None, (provider, config_provider)): - candidate_names.add(p) - if p.lower().startswith("custom:"): - candidate_names.add(p.split(":", 1)[1]) - else: - candidate_names.add(f"custom:{p}") - - providers_cfg = cfg.get("providers") - if isinstance(providers_cfg, dict): - for name in candidate_names: - entry = providers_cfg.get(name) - if isinstance(entry, dict): - k = str(entry.get("api_key") or "").strip() - if k: - return k - - custom_providers = cfg.get("custom_providers") - if isinstance(custom_providers, list): - lowered = {n.lower() for n in candidate_names} - for entry_raw in custom_providers: - if not isinstance(entry_raw, dict): - continue - entry_name = str(entry_raw.get("name") or "").strip() - if entry_name not in candidate_names and entry_name.lower() not in lowered: - continue - k = str(entry_raw.get("api_key") or "").strip() - if k: - return k - - return "" + return _resolve_inference_value(cfg, provider, "api_key", runtime_ok=lambda _: True) def _should_probe_ollama_vision( @@ -403,57 +315,38 @@ def _should_probe_ollama_vision( ) -> bool: """True when the active provider likely fronts a local Ollama server. - Server-fingerprint probing is only meaningful for *local* endpoints — - remote OpenAI-compatible APIs (sglang, vLLM, etc.) should never be probed, - and probing them without an api_key sprays 401s at the inference backend - (issue #89863). + Server-fingerprint probing is only valid for LOCAL endpoints: remote + OpenAI-compatible APIs (sglang, vLLM) expose Ollama-compat routes that can + misidentify, and probing them without an api_key returns 401 on every leg. """ - p = (provider or "").strip().lower() - if p == "ollama": + if (provider or "").strip().lower() == "ollama": return True if not base_url: return False - # Remote endpoints must never be fingerprinted: the probe waterfall is - # only valid for local/LM-Studio/Ollama boxes. Non-Ollama remotes (sglang, - # vLLM, OpenAI-compat) expose Ollama-compat endpoints that can misidentify - # and, without an api_key, return 401 on every leg (issue #89863). - if p != "ollama": - try: - from agent.model_metadata import is_local_endpoint - - if not is_local_endpoint(base_url): - return False - except Exception: - return False try: - from agent.model_metadata import detect_local_server_type + from agent.model_metadata import detect_local_server_type, is_local_endpoint - # Forward the API key: a remote API-keyed endpoint answers the - # probe waterfall with 401s without it, and an unauthorized probe - # can never produce a positive verdict (#89863). + if not is_local_endpoint(base_url): + return False + # Forward the key: an unauthorized probe can never produce a positive verdict. return detect_local_server_type(base_url, api_key=api_key) == "ollama" except Exception: return False def _coerce_mode(raw: Any) -> str: - """Normalize a config value into one of the valid modes.""" - if not isinstance(raw, str): - return "auto" - val = raw.strip().lower() - if val in _VALID_MODES: - return val + """Normalize a config value into one of the valid modes (default ``auto``).""" + if isinstance(raw, str) and raw.strip().lower() in _VALID_MODES: + return raw.strip().lower() return "auto" def _explicit_aux_vision_override(cfg: Optional[Dict[str, Any]]) -> bool: - """True when the user configured a specific auxiliary vision backend. + """True when the user configured a specific ``auxiliary.vision`` backend. - An explicit backend is the DE-FACTO image route in ``auto`` mode — - the user named a dedicated vision model, so images go through it even - when the main model could take them natively (maintainer decision, - reversing #29135). ``agent.image_input_mode: native`` still forces - native; unset/auto aux config leaves native as the default. + An explicit backend is the DE-FACTO image route in ``auto`` mode even when + the main model could take images natively. ``auto``/empty provider with no + model and no base_url is not explicit. """ if not isinstance(cfg, dict): return False @@ -464,14 +357,12 @@ def _explicit_aux_vision_override(cfg: Optional[Dict[str, Any]]) -> bool: if not isinstance(vision, dict): return False - provider = str(vision.get("provider") or "").strip().lower() - model = str(vision.get("model") or "").strip() - base_url = str(vision.get("base_url") or "").strip() - - # "auto" / "" / blank = not explicit - if provider in {"", "auto"} and not model and not base_url: - return False - return True + provider = _clean_str(vision.get("provider")).lower() + return not ( + provider in {"", "auto"} + and not _clean_str(vision.get("model")) + and not _clean_str(vision.get("base_url")) + ) def _lookup_supports_vision( @@ -481,33 +372,21 @@ def _lookup_supports_vision( *, requested_provider: str = "", ) -> Optional[bool]: - """Return True/False if we can resolve caps, None if unknown. + """Return True/False if vision capability can be resolved, None if unknown. - Consults the user's ``supports_vision`` override in config.yaml first - (so custom/local models declared as vision-capable don't fall through to - text routing in ``auto`` mode), then falls back to models.dev. + Order: config ``supports_vision`` override → managed local runtime → + models.dev catalog → Ollama probe for local endpoints. """ - # Named custom providers are canonicalized to ``provider="custom"`` by - # runtime resolution. The original CLI/config name is carried in the - # context-local main runtime so capability lookup can still select the - # exact custom_providers entry. Require an exact provider+model match: - # background/auxiliary lookups must never borrow another turn's identity. + # Named custom providers are canonicalized to ``provider="custom"``; the + # original name lives in the context-local main runtime. Borrow it only on an + # exact provider+model match so background/auxiliary lookups never take + # another turn's identity. if not requested_provider: - try: - from agent.auxiliary_client import _runtime_main_value - - runtime_provider = str( - _runtime_main_value("provider") or "" - ).strip().lower() - runtime_model = str(_runtime_main_value("model") or "").strip() - lookup_provider = str(provider or "").strip().lower() - lookup_model = str(model or "").strip() - if runtime_provider == lookup_provider and runtime_model == lookup_model: - requested_provider = str( - _runtime_main_value("requested_provider") or "" - ).strip() - except Exception: - pass + if ( + _runtime_main("provider").lower() == _clean_str(provider).lower() + and _runtime_main("model") == _clean_str(model) + ): + requested_provider = _runtime_main("requested_provider") override = _supports_vision_override( cfg, @@ -520,13 +399,10 @@ def _lookup_supports_vision( if not provider or not model: return None - # Managed local runtime: the server that would receive the image is - # the authority on whether it can see (its /props reports modalities - # when a vision projector is loaded; the catalog covers staged-but- - # unloaded models). Cloud catalogs have never heard of a local GGUF, - # so without this answer every local model reads as text-only and - # images detour to a cloud auxiliary — wrong twice for a local-first - # user (broken feature, and a screenshot leaving the machine). + # Managed local runtime: the server receiving the image is the authority on + # whether it can see (its /props reports modalities). Cloud catalogs have + # never heard of a local GGUF, so without this every local model reads as + # text-only and screenshots detour to a cloud auxiliary. try: from hermes_cli.local_runtime.capabilities import ( is_managed_provider, @@ -544,13 +420,10 @@ def _lookup_supports_vision( caps = None try: from agent.models_dev import get_model_capabilities - # allow_network=True on purpose: vision-capability lookup runs when - # an image actually needs routing (not per turn), and the #31179 - # text-only-main guard depends on catalog data — a cold cache - # returning "unknown" would fall back to attempting the call and - # reintroduce the bug. This preserves the historical - # network-on-cold-cache behavior for this one path; the fetch is - # cached (4h TTL) and backoff-limited after failures. + # allow_network=True on purpose: this runs only when an image needs + # routing, and the text-only-main guard depends on catalog data — a cold + # cache returning "unknown" would reintroduce attempting the call. The + # fetch is cached (4h TTL) and backoff-limited. caps = get_model_capabilities(provider, model, allow_network=True) except Exception as exc: # pragma: no cover - defensive logger.debug("image_routing: caps lookup failed for %s:%s — %s", provider, model, exc) @@ -561,8 +434,6 @@ def _lookup_supports_vision( if not base_url and (provider or "").strip().lower() == "ollama": base_url = "http://localhost:11434/v1" - # Resolve the provider's API key so probe requests at keyed endpoints - # carry Authorization and don't spray 401s (issue #89863). resolved_api_key = _resolve_inference_api_key(cfg, provider) if _should_probe_ollama_vision(provider, base_url, api_key=resolved_api_key): @@ -594,9 +465,9 @@ def decide_image_input_mode( """Return ``"native"`` or ``"text"`` for the given turn. Args: - provider: active inference provider ID (e.g. ``"anthropic"``, ``"openrouter"``). - model: active model slug as it would be sent to the provider. - cfg: loaded config.yaml dict, or None. When None, behaves as auto. + provider: active inference provider ID (e.g. ``"anthropic"``). + model: active model slug as sent to the provider. + cfg: loaded config.yaml dict, or None (behaves as auto). requested_provider: provider identity before runtime canonicalization. """ mode_cfg = "auto" @@ -605,19 +476,11 @@ def decide_image_input_mode( if isinstance(agent_cfg, dict): mode_cfg = _coerce_mode(agent_cfg.get("image_input_mode")) - if mode_cfg == "native": - return "native" - if mode_cfg == "text": - return "text" + if mode_cfg != "auto": + return mode_cfg - # auto: an explicitly configured auxiliary.vision backend is the - # DE-FACTO choice — the user named a dedicated vision model, so that's - # what they want images to go through, even when the main model has - # native vision (maintainer decision, 2026-08-28, reversing #29135's - # fallback-only posture: config that only takes effect when the main - # model gets worse is a trap, not a setting). Native vision remains - # the default for unconfigured installs, and the fallback when the - # aux backend is unset. + # auto: an explicit auxiliary.vision backend wins (see module docstring); + # native remains the default for unconfigured installs. if _explicit_aux_vision_override(cfg): return "text" if requested_provider: @@ -628,107 +491,80 @@ def decide_image_input_mode( requested_provider=requested_provider, ) else: - # Keep the long-standing three-argument call contract for callers and - # tests that replace the capability lookup hook. + # Keep the three-argument call contract for callers/tests that replace + # the capability lookup hook. supports = _lookup_supports_vision(provider, model, cfg) - if supports is True: - return "native" - return "text" + return "native" if supports is True else "text" -# Image size handling is REACTIVE rather than proactive: we attempt native -# attachment at full size regardless of provider, and rely on -# ``run_agent._try_shrink_image_parts_in_messages`` to shrink + retry if -# the provider rejects the request (e.g. Anthropic's hard 5 MB per-image -# ceiling returned as HTTP 400 "image exceeds 5 MB maximum"). -# -# Why reactive: our knowledge of provider ceilings is partial and evolving -# (OpenAI accepts 49 MB+, Anthropic 5 MB, Gemini 100 MB, others unknown). -# A proactive per-provider table would be stale the moment a provider raises -# or lowers its limit, and silently degrading quality for users on providers -# that would have accepted the full image is the worse failure mode. -# The shrink-on-reject path loses 1 API call + maybe 1s of Pillow work when -# it fires, which is cheaper than permanent quality loss. +# Image size handling is REACTIVE: attach at full size regardless of provider +# and let ``run_agent._try_shrink_image_parts_in_messages`` shrink + retry on +# rejection (e.g. Anthropic's 5 MB ceiling as HTTP 400). Provider ceilings are +# partial and evolving (OpenAI 49 MB+, Anthropic 5 MB, Gemini 100 MB); a +# proactive table would go stale and silently degrade quality for providers +# that would have accepted the full image — worse than one extra API call. + + +# Magic-byte signatures, checked in order. Filename-based detection is +# unreliable when platforms lie about content-type (Discord serves PNG as +# ``image/webp`` for proxied stickers); Anthropic rejects a media_type that +# does not match the bytes with HTTP 400, so we sniff. +_HEIC_BRANDS = frozenset({ + b"heic", b"heix", b"hevc", b"hevx", b"mif1", b"msf1", b"heim", b"heis", +}) +_MAGIC_PREFIXES: Tuple[Tuple[bytes, str], ...] = ( + (b"\x89PNG\r\n\x1a\n", "image/png"), + (b"\xff\xd8\xff", "image/jpeg"), + (b"GIF87a", "image/gif"), + (b"GIF89a", "image/gif"), +) def _sniff_mime_from_bytes(raw: bytes) -> Optional[str]: - """Detect image MIME from magic bytes. Returns None if unrecognised. - - Filename-based detection (``mimetypes.guess_type``) is unreliable when - upstream platforms lie about content-type. Discord, for example, can - serve a PNG with ``content_type=image/webp`` for proxied/animated - stickers, custom emoji previews, or images uploaded via certain bots. - Anthropic strictly validates that declared media_type matches the - actual bytes and returns HTTP 400 on mismatch, so we sniff to be safe. - """ + """Detect image MIME from magic bytes; None if unrecognised.""" if not raw: return None - # PNG: 89 50 4E 47 0D 0A 1A 0A - if raw.startswith(b"\x89PNG\r\n\x1a\n"): - return "image/png" - # JPEG: FF D8 FF - if raw.startswith(b"\xff\xd8\xff"): - return "image/jpeg" - # GIF87a / GIF89a - if raw[:6] in {b"GIF87a", b"GIF89a"}: - return "image/gif" - # WEBP: "RIFF" .... "WEBP" + for prefix, mime in _MAGIC_PREFIXES: + if raw.startswith(prefix): + return mime if len(raw) >= 12 and raw[:4] == b"RIFF" and raw[8:12] == b"WEBP": return "image/webp" - # BMP: "BM" if raw.startswith(b"BM"): return "image/bmp" - # ISO-BMFF family (HEIC/HEIF/AVIF): bytes 4..8 == 'ftyp', major brand at 8..12 + # ISO-BMFF family (HEIC/HEIF/AVIF): 'ftyp' at 4..8, major brand at 8..12. if len(raw) >= 12 and raw[4:8] == b"ftyp": brand = raw[8:12] if brand in {b"avif", b"avis"}: return "image/avif" - if brand in { - b"heic", b"heix", b"hevc", b"hevx", - b"mif1", b"msf1", b"heim", b"heis", - }: + if brand in _HEIC_BRANDS: return "image/heic" - # TIFF: II*\0 (little-endian) or MM\0* (big-endian) if raw[:4] in {b"II*\x00", b"MM\x00*"}: return "image/tiff" - # ICO: 00 00 01 00 (reserved=0, type=1=icon) if raw[:4] == b"\x00\x00\x01\x00": return "image/x-icon" - # SVG: text-based, look for an Optional[bytes]: - """Decode arbitrary image bytes with Pillow and re-encode as PNG. + """Decode image bytes with Pillow and re-encode as PNG; None when impossible. - Returns None if Pillow isn't installed or can't decode the input - (rare formats, corrupted bytes, missing optional decoder plugin for - HEIC/AVIF, or vector formats like SVG). Caller falls back to skipping - the image so the rest of the turn still works. - - HEIC/HEIF and AVIF need optional Pillow plugins; we try to register - them on demand and swallow ImportError so a missing plugin just - looks like 'Pillow can't decode this' rather than crashing. + HEIC/HEIF and AVIF need optional Pillow plugins, registered on demand; a + missing plugin just looks like "Pillow can't decode this" so the caller + skips the image and the rest of the turn proceeds. """ try: from PIL import Image @@ -739,8 +575,6 @@ def _transcode_to_png(raw: bytes) -> Optional[bytes]: "(and `pillow-heif` / `pillow-avif-plugin` for those formats)." ) return None - # Optional plugin registration. Silent on failure: an unsupported - # format will just fall through to Image.open raising below. try: import pillow_heif # type: ignore @@ -755,9 +589,8 @@ def _transcode_to_png(raw: bytes) -> Optional[bytes]: from io import BytesIO with Image.open(BytesIO(raw)) as im: - # Pick an output mode PNG can serialise. Anything other than - # the standard set gets normalised to RGBA so transparency is - # preserved where the source had it. + # Normalise exotic modes to RGBA so PNG can serialise and any + # source transparency survives. if im.mode not in {"RGB", "RGBA", "L", "LA", "P"}: im = im.convert("RGBA") buf = BytesIO() @@ -770,12 +603,18 @@ def _transcode_to_png(raw: bytes) -> Optional[bytes]: return None -def _guess_mime(path: Path, raw: Optional[bytes] = None) -> str: - """Return image MIME type for *path*. +_SUFFIX_MIMES = { + ".jpg": "image/jpeg", + ".jpeg": "image/jpeg", + ".png": "image/png", + ".gif": "image/gif", + ".webp": "image/webp", + ".bmp": "image/bmp", +} - If *raw* bytes are provided, magic-byte sniffing wins (authoritative). - Otherwise we fall back to ``mimetypes`` then suffix-based defaults. - """ + +def _guess_mime(path: Path, raw: Optional[bytes] = None) -> str: + """Image MIME for *path*: magic bytes (authoritative) → ``mimetypes`` → suffix → jpeg.""" if raw is not None: sniffed = _sniff_mime_from_bytes(raw) if sniffed: @@ -783,40 +622,18 @@ def _guess_mime(path: Path, raw: Optional[bytes] = None) -> str: mime, _ = mimetypes.guess_type(str(path)) if mime and mime.startswith("image/"): return mime - # mimetypes on some Linux distros mis-maps .jpg; default to jpeg when - # the suffix looks imagey. - suffix = path.suffix.lower() - return { - ".jpg": "image/jpeg", - ".jpeg": "image/jpeg", - ".png": "image/png", - ".gif": "image/gif", - ".webp": "image/webp", - ".bmp": "image/bmp", - }.get(suffix, "image/jpeg") + # mimetypes on some Linux distros mis-maps .jpg; default to jpeg. + return _SUFFIX_MIMES.get(path.suffix.lower(), "image/jpeg") def _file_to_data_url(path: Path) -> Optional[str]: """Encode a local image as a base64 data URL at its native size. - Size limits are NOT enforced here — the agent retry loop - (``run_agent._try_shrink_image_parts_in_messages``) shrinks on the - provider's first rejection. Keeping this simple means providers that - accept large images (OpenAI 49 MB+, Gemini 100 MB) don't pay a silent - quality tax just because one other provider is stricter. - - Format compatibility IS handled here: if the sniffed MIME isn't one - of ``_UNIVERSALLY_SUPPORTED_MIMES`` (i.e. it's something like AVIF, - HEIC, BMP, TIFF, or ICO that some providers reject outright), we - transcode to PNG with Pillow before declaring media_type. This fixes - the user-visible "Could not process image" HTTP 400 from Anthropic on - Discord-attached AVIF/HEIC/BMP files. - - Returns None if the file can't be read OR if the format isn't - universally supported AND Pillow can't transcode it (Pillow missing, - HEIC/AVIF plugin missing, vector format like SVG, corrupt bytes). The - caller reports those paths in ``skipped`` and the rest of the turn - proceeds. + Size is NOT limited here (the agent retry loop shrinks on the provider's + first rejection, so lenient providers pay no silent quality tax). Format + compatibility IS handled: MIMEs outside the accepted set are transcoded to + PNG. Returns None when the file can't be read, is blocked by the read + guard, or can't be transcoded — the caller reports it in ``skipped``. """ try: from agent.file_safety import raise_if_read_blocked @@ -836,11 +653,9 @@ def _file_to_data_url(path: Path) -> Optional[str]: return None mime = _guess_mime(path, raw=raw) accepted = _UNIVERSALLY_SUPPORTED_MIMES - # The managed local server decodes fewer formats than cloud providers - # (no WebP — and a WebP part fails SILENTLY: the model never sees an - # image and confabulates a description). When the active main model is - # served by the managed runtime, narrow the accepted set so those - # formats transcode to PNG here instead of vanishing server-side. + # The managed local server decodes fewer formats (no WebP — and a WebP part + # fails SILENTLY: the model confabulates a description). Narrow the accepted + # set so those formats transcode here instead of vanishing server-side. try: from agent.auxiliary_client import _runtime_main_value from hermes_cli.local_runtime.capabilities import ( @@ -881,86 +696,45 @@ def build_native_content_parts( ) -> Tuple[List[Dict[str, Any]], List[str]]: """Build an OpenAI-style ``content`` list for a user turn. - Shape: - [{"type": "text", "text": "...\\n\\n[Image attached at: /local/path]"}, - {"type": "image_url", "image_url": {"url": "data:image/png;base64,..."}}, - {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}, - ...] + Local paths are embedded as base64 ``data:`` URLs; remote URLs pass through + verbatim. When at least one image attaches, a single text part combines the + caption (or a neutral default) with one hint per image — + ``[Image attached at: ]`` / ``[Image attached: ]`` — giving the + model a string handle for tools that take an image path/URL, mirroring the + text-mode hint from ``Runner._enrich_message_with_vision``. - Local paths are read from disk and embedded as base64 ``data:`` URLs. - Remote URLs (``http(s)://``) are passed through verbatim — the provider - fetches them server-side. The model still sees the pixels either way. - - For each successfully attached image, a hint is appended to the text - part: - - * local path → ``[Image attached at: ]`` - * URL → ``[Image attached: ]`` - - The hint gives the model a string handle so MCP/skill tools that take - an image path or URL argument can be invoked on the same image without - an extra round-trip. This parallels the text-mode hint produced by - ``Runner._enrich_message_with_vision`` (``vision_analyze using image_url: - ``) so behaviour is consistent across both image input modes. - - Images are attached at their native size. If a provider rejects the - request because an image is too large (e.g. Anthropic's 5 MB per-image - ceiling), the agent's retry loop transparently shrinks and retries - once — see ``run_agent._try_shrink_image_parts_in_messages``. - - Returns (content_parts, skipped). Skipped entries are local paths - that couldn't be read from disk; URLs are never skipped (they're - not validated here). + Returns ``(content_parts, skipped)``; ``skipped`` holds local paths that + could not be read. URLs are never skipped (not validated here). """ skipped: List[str] = [] image_parts: List[Dict[str, Any]] = [] - attached_paths: List[str] = [] - attached_urls: List[str] = [] + hint_lines: List[str] = [] for raw_path in image_paths: p = Path(raw_path) - if not p.exists() or not p.is_file(): - skipped.append(str(raw_path)) - continue - data_url = _file_to_data_url(p) + data_url = _file_to_data_url(p) if p.exists() and p.is_file() else None if not data_url: skipped.append(str(raw_path)) continue - image_parts.append({ - "type": "image_url", - "image_url": {"url": data_url}, - }) - attached_paths.append(str(raw_path)) + image_parts.append({"type": "image_url", "image_url": {"url": data_url}}) + hint_lines.append(f"[Image attached at: {raw_path}]") for url in image_urls or []: url = (url or "").strip() if not url: continue - image_parts.append({ - "type": "image_url", - "image_url": {"url": url}, - }) - attached_urls.append(url) + image_parts.append({"type": "image_url", "image_url": {"url": url}}) + hint_lines.append(f"[Image attached: {url}]") text = (user_text or "").strip() - # If at least one image attached, build a single text part that combines - # the user's caption (or a neutral default) with one hint per image. - if attached_paths or attached_urls: + if image_parts: base_text = text or "What do you see in this image?" - hint_lines: List[str] = [] - hint_lines.extend(f"[Image attached at: {p}]" for p in attached_paths) - hint_lines.extend(f"[Image attached: {u}]" for u in attached_urls) combined_text = f"{base_text}\n\n" + "\n".join(hint_lines) - parts: List[Dict[str, Any]] = [{"type": "text", "text": combined_text}] - parts.extend(image_parts) - return parts, skipped + return [{"type": "text", "text": combined_text}, *image_parts], skipped - # No images successfully attached — fall back to plain text-only behaviour. - parts = [] - if text: - parts.append({"type": "text", "text": text}) - return parts, skipped + # No images attached — plain text-only behaviour. + return ([{"type": "text", "text": text}] if text else []), skipped __all__ = [ diff --git a/agent/lsp/__init__.py b/agent/lsp/__init__.py index 7819162dd4..21d072b8ec 100644 --- a/agent/lsp/__init__.py +++ b/agent/lsp/__init__.py @@ -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 diff --git a/agent/lsp/cli.py b/agent/lsp/cli.py index 607c156d10..70e9b13a15 100644 --- a/agent/lsp/cli.py +++ b/agent/lsp/cli.py @@ -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 `` — 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 `` — 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 ``.""" 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 = [] diff --git a/agent/lsp/client.py b/agent/lsp/client.py index 36c2721027..4740652a7c 100644 --- a/agent/lsp/client.py +++ b/agent/lsp/client.py @@ -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 {} diff --git a/agent/lsp/eventlog.py b/agent/lsp/eventlog.py index f118ccf0ac..60d452fb3e 100644 --- a/agent/lsp/eventlog.py +++ b/agent/lsp/eventlog.py @@ -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 `` 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 `` the first time a (server_id, - workspace_root) client starts, ``no project root for `` - 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 ` or set lsp.servers..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 ` or set lsp.servers..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", diff --git a/agent/lsp/install.py b/agent/lsp/install.py index fc9bea5930..01040b9ff7 100644 --- a/agent/lsp/install.py +++ b/agent/lsp/install.py @@ -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, -``/lsp/bin/``, so we don't pollute the user's global -toolchain. +Installs go to a Hermes-owned staging dir, ``/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 -# ``/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 ``/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 - ``/node_modules/.bin/`` 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 `` then link ``node_modules/.bin/`` 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 # /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 - ``/python-packages/bin/`` which we symlink into - ``/bin``. Note: this only works for packages that ship a - console script. - """ + """``pip install --target /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): diff --git a/agent/lsp/manager.py b/agent/lsp/manager.py index 7ba1b914f7..275d8721ad 100644 --- a/agent/lsp/manager.py +++ b/agent/lsp/manager.py @@ -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"] diff --git a/agent/lsp/protocol.py b/agent/lsp/protocol.py index 2b35b741f5..51616eeec0 100644 --- a/agent/lsp/protocol.py +++ b/agent/lsp/protocol.py @@ -1,17 +1,9 @@ """Minimal LSP JSON-RPC 2.0 framer over async streams. -LSP wire format: - - Content-Length: \\r\\n - \\r\\n - - -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: \\r\\n\\r\\n`` 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 diff --git a/agent/lsp/range_shift.py b/agent/lsp/range_shift.py index 8efdfc3098..fa3aa85e89 100644 --- a/agent/lsp/range_shift.py +++ b/agent/lsp/range_shift.py @@ -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 diff --git a/agent/lsp/reporter.py b/agent/lsp/reporter.py index 2be1779cce..468ae579c8 100644 --- a/agent/lsp/reporter.py +++ b/agent/lsp/reporter.py @@ -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 — ```` blocks with -1-indexed line/column, capped at ``MAX_PER_FILE`` errors. +The model sees a compact, severity-filtered, line-bounded ```` +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 ```` 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 ```` - 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 + ```` 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 ```` 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 ```` 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">\n{body}\n" diff --git a/agent/lsp/servers.py b/agent/lsp/servers.py index fc2a0b2616..97341b0936 100644 --- a/agent/lsp/servers.py +++ b/agent/lsp/servers.py @@ -1,21 +1,10 @@ """Server registry — per-language LSP server definitions. -Each :class:`ServerDef` knows how to: - -- match a file by extension (or basename for extensionless files like - ``Dockerfile``), -- resolve a project root from a file path (often via - :func:`agent.lsp.workspace.nearest_root`), -- assemble the spawn command (binary, args, env, cwd), -- compute LSP ``initializationOptions``. - -Auto-installation is a separate concern handled by -:mod:`agent.lsp.install`. This module describes WHAT to spawn; the -install module makes the binary appear on PATH if it isn't there. - -The full set of servers ships with the package, but most are only -*invoked* when the user actually edits a file in that language. This -keeps cold-start fast — we don't probe binaries until needed. +Each :class:`ServerDef` matches files (by extension or basename for +extensionless files like ``Dockerfile``), resolves a project root, and +assembles the spawn command. Auto-installation lives in +:mod:`agent.lsp.install`; nothing here probes binaries until a file in +that language is actually edited. """ from __future__ import annotations @@ -29,9 +18,8 @@ from agent.lsp.workspace import nearest_root logger = logging.getLogger("agent.lsp.servers") -# Language IDs per LSP spec. Used for ``textDocument/didOpen.languageId``. -# Most servers don't care exactly, but a few (typescript-language-server, -# vue-language-server) refuse files with the wrong ID. +# LSP languageId for ``textDocument/didOpen``. A few servers +# (typescript-language-server, vue-language-server) refuse wrong IDs. LANGUAGE_BY_EXT: Dict[str, str] = { ".py": "python", ".pyi": "python", @@ -110,13 +98,7 @@ LANGUAGE_BY_EXT: Dict[str, str] = { @dataclass class SpawnSpec: - """The result of resolving a server for a file. - - Returned by :meth:`ServerDef.resolve` when a server is applicable - to a file. ``None`` is returned instead when the server should - be skipped (binary missing and auto-install disabled, project - marker not found, exclude marker hit, etc.). - """ + """Result of resolving a server for a file (``None`` means skip).""" command: List[str] workspace_root: str @@ -130,14 +112,9 @@ class SpawnSpec: class ServerDef: """Definition of one language server. - The :func:`resolve_root` callable receives the absolute file path - plus the workspace root (git worktree) and returns either the - project-specific root for this server (e.g. the directory - containing ``pyproject.toml``) or ``None`` to skip. - - The :func:`build_spawn` callable receives the resolved root and - returns a :class:`SpawnSpec` (or ``None`` if the binary can't be - found and auto-install isn't configured). + ``resolve_root(file_path, workspace_root)`` returns the per-server + project root or ``None`` to skip; ``build_spawn(root, ctx)`` returns + a :class:`SpawnSpec` or ``None`` when the binary can't be found. """ server_id: str @@ -149,18 +126,12 @@ class ServerDef: def matches(self, file_path: str) -> bool: """Return True iff this server handles ``file_path``.""" - ext = _file_ext_or_basename(file_path) - return ext in self.extensions + return _file_ext_or_basename(file_path) in self.extensions @dataclass class ServerContext: - """Context passed into :meth:`ServerDef.build_spawn`. - - Carries the user's auto-install policy, any user-overridden - binary paths, and helpers the spawn builder needs. All fields - are optional; defaults yield "auto-install allowed, no overrides". - """ + """User policy passed into :meth:`ServerDef.build_spawn` (install strategy, overrides).""" workspace_root: str install_strategy: str = "auto" # "auto" | "manual" | "off" @@ -175,17 +146,10 @@ class ServerContext: def _file_ext_or_basename(path: str) -> str: - """Return the lower-cased extension OR full basename for extensionless files. - - Mirrors OpenCode's ``path.parse(file).ext || file`` — files like - ``Dockerfile`` or ``Makefile`` match by basename, while normal - files match by extension (``.py``, ``.ts``). - """ + """Lower-cased extension, or the full basename for extensionless files (``Dockerfile``).""" base = os.path.basename(path) _root, ext = os.path.splitext(base) - if ext: - return ext.lower() - return base + return ext.lower() if ext else base def _which(*names: str) -> Optional[str]: @@ -198,67 +162,91 @@ def _which(*names: str) -> Optional[str]: def _root_or_workspace(file_path: str, workspace: str, markers: Sequence[str], excludes: Sequence[str] = ()) -> Optional[str]: - """Common pattern: try ``nearest_root``, fall back to workspace root. - - Returns ``None`` if an exclude marker matches first (server gated off). - """ - found = nearest_root( - file_path, - markers, - excludes=excludes, - ceiling=os.path.dirname(workspace) if workspace else None, - ) + """``nearest_root`` with workspace fallback; ``None`` iff an exclude marker hit.""" + ceiling = os.path.dirname(workspace) if workspace else None + found = nearest_root(file_path, markers, excludes=excludes, ceiling=ceiling) if found is None and excludes: - # Distinguish "no marker found" from "exclude hit": when - # excludes are configured, None means gated off. - # Re-check without excludes — if still None, we fall back to - # workspace; if found, the exclude hit and we return None. - recheck = nearest_root( - file_path, - markers, - ceiling=os.path.dirname(workspace) if workspace else None, - ) - if recheck is not None: - return None # exclude triggered + # None is ambiguous with excludes configured: re-check without them — + # a hit now means the exclude fired (gated off), else fall back. + if nearest_root(file_path, markers, ceiling=ceiling) is not None: + return None return workspace return found or workspace +def _resolve_override(ctx: ServerContext, server_id: str) -> Optional[str]: + """User can pin a binary path in config.""" + override = ctx.binary_overrides.get(server_id) + if override and override[0] and os.path.exists(override[0]): + return override[0] + return None + + +def _find_binary(ctx: ServerContext, server_id: str, which: Sequence[str], install_pkg: Optional[str]) -> Optional[str]: + """Override → PATH → (optional) auto-install; ``None`` when nothing resolves.""" + bin_path = _resolve_override(ctx, server_id) or _which(*which) + if bin_path is None and install_pkg is not None: + from agent.lsp.install import try_install + bin_path = try_install(install_pkg, ctx.install_strategy) + return bin_path + + +def _make_spec(root: str, ctx: ServerContext, server_id: str, command: List[str], + base_init: Optional[Dict[str, Any]] = None, seed: bool = False) -> SpawnSpec: + if base_init is None: + init = ctx.init_overrides.get(server_id, {}) + else: + init = dict(base_init) + init.update(ctx.init_overrides.get(server_id, {})) + return SpawnSpec( + command=command, + workspace_root=root, + cwd=root, + env=ctx.env_overrides.get(server_id, {}), + initialization_options=init, + seed_diagnostics_on_first_push=seed, + ) + + +def _simple_spawn(server_id: str, which: Sequence[str], args: Sequence[str] = (), + install_pkg: Optional[str] = None, base_init: Optional[Dict[str, Any]] = None, + seed: bool = False) -> Callable[[str, ServerContext], Optional[SpawnSpec]]: + """Build a spawn function for the common single-binary server shape.""" + def build(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: + bin_path = _find_binary(ctx, server_id, which, install_pkg) + if bin_path is None: + return None + return _make_spec(root, ctx, server_id, [bin_path, *args], base_init, seed) + return build + + +def _markers_root(markers: Optional[Sequence[str]], excludes: Sequence[str] = ()) -> Callable[[str, str], Optional[str]]: + """Root resolver over marker files; ``None`` markers means "always the workspace root".""" + if markers is None: + return lambda fp, ws: ws + return lambda fp, ws: _root_or_workspace(fp, ws, markers, excludes=excludes) + + # --------------------------------------------------------------------------- -# per-server spawn builders +# bespoke spawn builders # --------------------------------------------------------------------------- def _spawn_pyright(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "pyright") or _which( - "pyright-langserver", "pyright" - ) + bin_path = _find_binary(ctx, "pyright", ("pyright-langserver", "pyright"), "pyright") if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("pyright", ctx.install_strategy) - if bin_path is None: - return None + return None # If we got the cli ``pyright``, the langserver is its sibling. - base = os.path.basename(bin_path) - if base in {"pyright", "pyright.exe"}: + if os.path.basename(bin_path) in {"pyright", "pyright.exe"}: sibling = os.path.join(os.path.dirname(bin_path), "pyright-langserver") if os.path.exists(sibling): bin_path = sibling init: Dict[str, Any] = {} - # Pick the project's venv interpreter if there is one — otherwise - # pyright defaults to "python on PATH" which is rarely the venv. + # Point pyright at the project venv; its default "python on PATH" rarely is. py = _detect_python(root) if py: init["python"] = {"pythonPath": py} - if "pyright" in ctx.init_overrides: - init.update(ctx.init_overrides["pyright"]) - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("pyright", {}), - initialization_options=init, - ) + return _make_spec(root, ctx, "pyright", [bin_path, "--stdio"], init) def _detect_python(root: str) -> Optional[str]: @@ -274,85 +262,15 @@ def _detect_python(root: str) -> Optional[str]: return None -def _spawn_typescript(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "typescript") or _which("typescript-language-server") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("typescript-language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("typescript", {}), - initialization_options=ctx.init_overrides.get("typescript", {}), - seed_diagnostics_on_first_push=True, - ) - - -def _spawn_gopls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "gopls") or _which("gopls") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("gopls", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("gopls", {}), - initialization_options=ctx.init_overrides.get("gopls", {}), - ) - - -def _spawn_rust_analyzer(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "rust-analyzer") or _which("rust-analyzer") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("rust-analyzer", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("rust-analyzer", {}), - initialization_options=ctx.init_overrides.get("rust-analyzer", {}), - ) - - -def _spawn_clangd(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "clangd") or _which("clangd") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("clangd", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--background-index", "--clang-tidy"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("clangd", {}), - initialization_options=ctx.init_overrides.get("clangd", {}), - ) - - _BASH_SHELLCHECK_WARNED = False def _spawn_bash_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "bash-language-server") or _which("bash-language-server") + bin_path = _find_binary(ctx, "bash-language-server", ("bash-language-server",), "bash-language-server") if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("bash-language-server", ctx.install_strategy) - if bin_path is None: - return None - # bash-language-server delegates diagnostics to ``shellcheck``. Without - # it on PATH the server starts and accepts requests but never reports - # any problems — to the user it looks like a working integration that - # never finds bugs. Warn once so the gap is visible. + return None + # bash-language-server delegates diagnostics to shellcheck; without it the + # server runs but never reports anything. Warn once so the gap is visible. global _BASH_SHELLCHECK_WARNED if not _BASH_SHELLCHECK_WARNED and _which("shellcheck") is None: _BASH_SHELLCHECK_WARNED = True @@ -361,344 +279,18 @@ def _spawn_bash_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: "diagnostics will be empty until shellcheck is installed " "(apt: shellcheck, brew: shellcheck, scoop: shellcheck)." ) - return SpawnSpec( - command=[bin_path, "start"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("bash-language-server", {}), - initialization_options=ctx.init_overrides.get("bash-language-server", {}), - ) - - -def _spawn_yaml_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "yaml-language-server") or _which("yaml-language-server") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("yaml-language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("yaml-language-server", {}), - initialization_options=ctx.init_overrides.get("yaml-language-server", {}), - ) - - -def _spawn_lua_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "lua-language-server") or _which("lua-language-server") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("lua-language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("lua-language-server", {}), - initialization_options=ctx.init_overrides.get("lua-language-server", {}), - ) - - -def _spawn_intelephense(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "intelephense") or _which("intelephense") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("intelephense", ctx.install_strategy) - if bin_path is None: - return None - init = {"telemetry": {"enabled": False}} - init.update(ctx.init_overrides.get("intelephense", {})) - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("intelephense", {}), - initialization_options=init, - ) - - -def _spawn_ocamllsp(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "ocaml-lsp") or _which("ocamllsp") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("ocaml-lsp", {}), - initialization_options=ctx.init_overrides.get("ocaml-lsp", {}), - ) - - -def _spawn_dockerfile_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "dockerfile-ls") or _which("docker-langserver") - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("dockerfile-language-server-nodejs", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("dockerfile-ls", {}), - initialization_options=ctx.init_overrides.get("dockerfile-ls", {}), - ) - - -def _spawn_terraform_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "terraform-ls") or _which("terraform-ls") - if bin_path is None: - return None # terraform-ls is heavy to auto-install; require user - init = { - "experimentalFeatures": { - "prefillRequiredFields": True, - "validateOnSave": True, - } - } - init.update(ctx.init_overrides.get("terraform-ls", {})) - return SpawnSpec( - command=[bin_path, "serve"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("terraform-ls", {}), - initialization_options=init, - ) - - -def _spawn_dart(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "dart") or _which("dart") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "language-server", "--lsp"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("dart", {}), - initialization_options=ctx.init_overrides.get("dart", {}), - ) - - -def _spawn_haskell_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "haskell-language-server") or _which( - "haskell-language-server-wrapper", "haskell-language-server" - ) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--lsp"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("haskell-language-server", {}), - initialization_options=ctx.init_overrides.get("haskell-language-server", {}), - ) - - -def _spawn_julia(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "julia") or _which("julia") - if bin_path is None: - return None - return SpawnSpec( - command=[ - bin_path, - "--startup-file=no", - "--history-file=no", - "-e", - "using LanguageServer; runserver()", - ], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("julia", {}), - initialization_options=ctx.init_overrides.get("julia", {}), - ) - - -def _spawn_clojure_lsp(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "clojure-lsp") or _which("clojure-lsp") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "listen"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("clojure-lsp", {}), - initialization_options=ctx.init_overrides.get("clojure-lsp", {}), - ) - - -def _spawn_nixd(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "nixd") or _which("nixd") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("nixd", {}), - initialization_options=ctx.init_overrides.get("nixd", {}), - ) - - -def _spawn_zls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "zls") or _which("zls") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("zls", {}), - initialization_options=ctx.init_overrides.get("zls", {}), - ) - - -def _spawn_gleam(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "gleam") or _which("gleam") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "lsp"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("gleam", {}), - initialization_options=ctx.init_overrides.get("gleam", {}), - ) - - -def _spawn_elixir_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "elixir-ls") or _which("elixir-ls", "language_server.sh") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("elixir-ls", {}), - initialization_options=ctx.init_overrides.get("elixir-ls", {}), - ) - - -def _spawn_prisma(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "prisma") or _which("prisma") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "language-server"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("prisma", {}), - initialization_options=ctx.init_overrides.get("prisma", {}), - ) - - -def _spawn_kotlin_ls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "kotlin-language-server") or _which( - "kotlin-language-server" - ) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("kotlin-language-server", {}), - initialization_options=ctx.init_overrides.get("kotlin-language-server", {}), - ) - - -def _spawn_jdtls(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - # jdtls has a complex install flow. We require a manual install - # for now and look for the wrapper script that the jdtls install - # produces. - bin_path = _resolve_override(ctx, "jdtls") or _which("jdtls") - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("jdtls", {}), - initialization_options=ctx.init_overrides.get("jdtls", {}), - ) - - -def _spawn_vue(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "vue-language-server") or _which( - "vue-language-server" - ) - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("@vue/language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("vue-language-server", {}), - initialization_options=ctx.init_overrides.get("vue-language-server", {}), - ) - - -def _spawn_svelte(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "svelte-language-server") or _which( - "svelteserver", "svelte-language-server" - ) - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("svelte-language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("svelte-language-server", {}), - initialization_options=ctx.init_overrides.get("svelte-language-server", {}), - ) - - -def _spawn_astro(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - bin_path = _resolve_override(ctx, "astro-language-server") or _which( - "astro-ls", "astro-language-server" - ) - if bin_path is None: - from agent.lsp.install import try_install - bin_path = try_install("@astrojs/language-server", ctx.install_strategy) - if bin_path is None: - return None - return SpawnSpec( - command=[bin_path, "--stdio"], - workspace_root=root, - cwd=root, - env=ctx.env_overrides.get("astro-language-server", {}), - initialization_options=ctx.init_overrides.get("astro-language-server", {}), - ) + return _make_spec(root, ctx, "bash-language-server", [bin_path, "start"]) _PSES_BUNDLE_WARNED = False def _find_pses_bundle(ctx: ServerContext) -> Optional[str]: - """Locate the PowerShellEditorServices module bundle directory. + """Locate the PowerShellEditorServices bundle dir (release zip, manual install). - PSES ships as a GitHub release zip (not an npm/go/pip package), so - there's no auto-install recipe — the user downloads it and points us - at the extracted bundle. Resolution order: - - 1. ``command`` override in config (``lsp.servers.powershell.command``) — - the FIRST element is treated as the bundle path when it's a - directory. This is the documented config knob. - 2. ``init_overrides["powershell"]["bundlePath"]``. - 3. ``PSES_BUNDLE_PATH`` env var. - 4. ``/lsp/PowerShellEditorServices`` staging dir (where a - user-run unzip would naturally land). - - Returns the bundle directory containing ``PowerShellEditorServices/``, - or ``None`` when it can't be found. + Resolution order: ``lsp.servers.powershell.command[0]`` when a directory, + ``init_overrides["powershell"]["bundlePath"]``, ``PSES_BUNDLE_PATH`` env, + then ``/lsp/PowerShellEditorServices``. """ candidates: List[str] = [] override = ctx.binary_overrides.get("powershell") @@ -712,32 +304,21 @@ def _find_pses_bundle(ctx: ServerContext) -> Optional[str]: candidates.append(env_path) from hermes_constants import get_hermes_home - home = str(get_hermes_home()) - candidates.append(os.path.join(home, "lsp", "PowerShellEditorServices")) + candidates.append(os.path.join(str(get_hermes_home()), "lsp", "PowerShellEditorServices")) for cand in candidates: if not cand: continue # Accept either the bundle root or the inner module dir. - start_script = os.path.join( - cand, "PowerShellEditorServices", "Start-EditorServices.ps1" - ) - if os.path.isfile(start_script): + if os.path.isfile(os.path.join(cand, "PowerShellEditorServices", "Start-EditorServices.ps1")): return cand - inner = os.path.join(cand, "Start-EditorServices.ps1") - if os.path.isfile(inner): + if os.path.isfile(os.path.join(cand, "Start-EditorServices.ps1")): return os.path.dirname(cand) return None def _spawn_powershell_es(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: - """Spawn PowerShellEditorServices over stdio. - - Unlike the single-binary servers, PSES is a PowerShell module driven - by a bootstrap script. We need both a PowerShell host (``pwsh`` for - PowerShell 7+, or Windows ``powershell``) and the PSES module bundle. - The bundle is manual-install (release zip) — see ``_find_pses_bundle``. - """ + """Spawn PowerShellEditorServices: needs a ``pwsh``/``powershell`` host plus the module bundle.""" pwsh = _which("pwsh", "powershell") if pwsh is None: return None @@ -755,13 +336,9 @@ def _spawn_powershell_es(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: "/lsp/PowerShellEditorServices." ) return None - start_script = os.path.join( - bundle, "PowerShellEditorServices", "Start-EditorServices.ps1" - ) - # Session details file: PSES writes connection info here on startup. - session_path = os.path.join( - hermes_lsp_session_dir(), f"pses-session-{os.getpid()}.json" - ) + start_script = os.path.join(bundle, "PowerShellEditorServices", "Start-EditorServices.ps1") + # PSES writes connection info to the session details file on startup. + session_path = os.path.join(hermes_lsp_session_dir(), f"pses-session-{os.getpid()}.json") log_path = os.path.join(hermes_lsp_session_dir(), "pses.log") inner = ( f"& '{start_script}' " @@ -773,16 +350,7 @@ def _spawn_powershell_es(root: str, ctx: ServerContext) -> Optional[SpawnSpec]: f"-Stdio -LogLevel Normal" ) return SpawnSpec( - command=[ - pwsh, - "-NoLogo", - "-NoProfile", - "-NonInteractive", - "-ExecutionPolicy", - "Bypass", - "-Command", - inner, - ], + command=[pwsh, "-NoLogo", "-NoProfile", "-NonInteractive", "-ExecutionPolicy", "Bypass", "-Command", inner], workspace_root=root, cwd=root, env=ctx.env_overrides.get("powershell", {}), @@ -798,367 +366,93 @@ def hermes_lsp_session_dir() -> str: """Return (and create) the dir for PSES session/log scratch files.""" from hermes_constants import get_hermes_home - home = str(get_hermes_home()) - d = os.path.join(home, "lsp", "pses") + d = os.path.join(str(get_hermes_home()), "lsp", "pses") os.makedirs(d, exist_ok=True) return d -def _resolve_override(ctx: ServerContext, server_id: str) -> Optional[str]: - """User can pin a binary path in config.""" - override = ctx.binary_overrides.get(server_id) - if override and override[0] and os.path.exists(override[0]): - return override[0] - return None - - -# --------------------------------------------------------------------------- -# root resolvers -# --------------------------------------------------------------------------- - - -def _root_python(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - ["pyproject.toml", "setup.py", "setup.cfg", "requirements.txt", "Pipfile", "pyrightconfig.json"], - ) - - -def _root_typescript(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - [ - "package-lock.json", - "bun.lockb", - "bun.lock", - "pnpm-lock.yaml", - "yarn.lock", - "package.json", - "tsconfig.json", - ], - excludes=["deno.json", "deno.jsonc"], - ) - - -def _root_go(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - ["go.work", "go.mod", "go.sum"], - ) - - -def _root_rust(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["Cargo.toml", "Cargo.lock"]) - - -def _root_ruby(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["Gemfile"]) - - -def _root_clangd(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - ["compile_commands.json", "compile_flags.txt", ".clangd"], - ) - - -def _root_bash(file_path: str, workspace: str) -> str: - return workspace - - -def _root_yaml(file_path: str, workspace: str) -> str: - return workspace - - -def _root_lua(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - [".luarc.json", ".luarc.jsonc", ".luacheckrc", ".stylua.toml", "stylua.toml", "selene.toml", "selene.yml"], - ) - - -def _root_php(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["composer.json", "composer.lock", ".php-version"]) - - -def _root_ocaml(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["dune-project", "dune-workspace", ".merlin", "opam"]) - - -def _root_docker(file_path: str, workspace: str) -> str: - return workspace - - -def _root_terraform(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, [".terraform.lock.hcl", "terraform.tfstate"]) - - -def _root_dart(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["pubspec.yaml", "analysis_options.yaml"]) - - -def _root_haskell(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["stack.yaml", "cabal.project", "hie.yaml"]) - - -def _root_julia(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["Project.toml", "Manifest.toml"]) - - -def _root_clojure(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, workspace, ["deps.edn", "project.clj", "shadow-cljs.edn", "bb.edn", "build.boot"] - ) - - -def _root_nix(file_path: str, workspace: str) -> str: - found = nearest_root(file_path, ["flake.nix"]) - return found or workspace - - -def _root_zig(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["build.zig"]) - - -def _root_elixir(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace(file_path, workspace, ["mix.exs", "mix.lock"]) - - -def _root_prisma(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, workspace, ["schema.prisma", "prisma/schema.prisma"] - ) - - -def _root_kotlin(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - ["settings.gradle", "settings.gradle.kts", "build.gradle", "build.gradle.kts", "pom.xml"], - ) - - -def _root_java(file_path: str, workspace: str) -> Optional[str]: - return _root_or_workspace( - file_path, - workspace, - ["pom.xml", "build.gradle", "build.gradle.kts", ".project", ".classpath", "settings.gradle"], - ) - - -def _root_powershell(file_path: str, workspace: str) -> Optional[str]: - # PowerShell projects rarely have a universal root marker. Use the - # PSScriptAnalyzer settings file when present, otherwise fall back to - # the git workspace root (nearest_root does exact-name matching only, - # so no globs here). - return _root_or_workspace( - file_path, - workspace, - ["PSScriptAnalyzerSettings.psd1"], - ) - - # --------------------------------------------------------------------------- # the registry # --------------------------------------------------------------------------- +_JS_MARKERS = ["package-lock.json", "bun.lockb", "bun.lock", "pnpm-lock.yaml", "yarn.lock", "package.json", "tsconfig.json"] +_DENO_EXCLUDES = ["deno.json", "deno.jsonc"] +_root_typescript = _markers_root(_JS_MARKERS, _DENO_EXCLUDES) + + +def _server(server_id: str, extensions: Tuple[str, ...], description: str, *, + markers: Optional[Sequence[str]] = None, excludes: Sequence[str] = (), + resolve_root: Optional[Callable[[str, str], Optional[str]]] = None, + build_spawn: Optional[Callable[[str, ServerContext], Optional[SpawnSpec]]] = None, + which: Sequence[str] = (), args: Sequence[str] = (), install_pkg: Optional[str] = None, + base_init: Optional[Dict[str, Any]] = None, seed: bool = False) -> ServerDef: + return ServerDef( + server_id=server_id, + extensions=extensions, + resolve_root=resolve_root or _markers_root(markers, excludes), + build_spawn=build_spawn or _simple_spawn(server_id, which or (server_id,), args, install_pkg, base_init, seed), + seed_first_push=seed, + description=description, + ) + SERVERS: List[ServerDef] = [ - ServerDef( - server_id="pyright", - extensions=(".py", ".pyi"), - resolve_root=_root_python, - build_spawn=_spawn_pyright, - description="Python — Microsoft pyright", - ), - ServerDef( - server_id="typescript", - extensions=(".ts", ".tsx", ".js", ".jsx", ".mjs", ".cjs", ".mts", ".cts"), - resolve_root=_root_typescript, - build_spawn=_spawn_typescript, - seed_first_push=True, - description="JavaScript/TypeScript — typescript-language-server", - ), - ServerDef( - server_id="vue-language-server", - extensions=(".vue",), - resolve_root=_root_typescript, - build_spawn=_spawn_vue, - description="Vue.js — @vue/language-server", - ), - ServerDef( - server_id="svelte-language-server", - extensions=(".svelte",), - resolve_root=_root_typescript, - build_spawn=_spawn_svelte, - description="Svelte — svelte-language-server", - ), - ServerDef( - server_id="astro-language-server", - extensions=(".astro",), - resolve_root=_root_typescript, - build_spawn=_spawn_astro, - description="Astro — @astrojs/language-server", - ), - ServerDef( - server_id="gopls", - extensions=(".go",), - resolve_root=_root_go, - build_spawn=_spawn_gopls, - description="Go — gopls", - ), - ServerDef( - server_id="rust-analyzer", - extensions=(".rs",), - resolve_root=_root_rust, - build_spawn=_spawn_rust_analyzer, - description="Rust — rust-analyzer", - ), - ServerDef( - server_id="clangd", - extensions=(".c", ".cpp", ".cc", ".cxx", ".h", ".hh", ".hpp", ".hxx"), - resolve_root=_root_clangd, - build_spawn=_spawn_clangd, - description="C/C++ — clangd", - ), - ServerDef( - server_id="bash-language-server", - extensions=(".sh", ".bash", ".zsh", ".ksh"), - resolve_root=_root_bash, - build_spawn=_spawn_bash_ls, - description="Bash — bash-language-server", - ), - ServerDef( - server_id="yaml-language-server", - extensions=(".yaml", ".yml"), - resolve_root=_root_yaml, - build_spawn=_spawn_yaml_ls, - description="YAML — yaml-language-server", - ), - ServerDef( - server_id="lua-language-server", - extensions=(".lua",), - resolve_root=_root_lua, - build_spawn=_spawn_lua_ls, - description="Lua — lua-language-server", - ), - ServerDef( - server_id="intelephense", - extensions=(".php",), - resolve_root=_root_php, - build_spawn=_spawn_intelephense, - description="PHP — intelephense", - ), - ServerDef( - server_id="ocaml-lsp", - extensions=(".ml", ".mli"), - resolve_root=_root_ocaml, - build_spawn=_spawn_ocamllsp, - description="OCaml — ocaml-lsp", - ), - ServerDef( - server_id="dockerfile-ls", - extensions=(".dockerfile", "Dockerfile"), - resolve_root=_root_docker, - build_spawn=_spawn_dockerfile_ls, - description="Dockerfile — dockerfile-language-server-nodejs", - ), - ServerDef( - server_id="terraform-ls", - extensions=(".tf", ".tfvars"), - resolve_root=_root_terraform, - build_spawn=_spawn_terraform_ls, - description="Terraform — terraform-ls", - ), - ServerDef( - server_id="dart", - extensions=(".dart",), - resolve_root=_root_dart, - build_spawn=_spawn_dart, - description="Dart — built-in language server", - ), - ServerDef( - server_id="haskell-language-server", - extensions=(".hs", ".lhs"), - resolve_root=_root_haskell, - build_spawn=_spawn_haskell_ls, - description="Haskell — haskell-language-server", - ), - ServerDef( - server_id="julia", - extensions=(".jl",), - resolve_root=_root_julia, - build_spawn=_spawn_julia, - description="Julia — LanguageServer.jl", - ), - ServerDef( - server_id="clojure-lsp", - extensions=(".clj", ".cljs", ".cljc", ".edn"), - resolve_root=_root_clojure, - build_spawn=_spawn_clojure_lsp, - description="Clojure — clojure-lsp", - ), - ServerDef( - server_id="nixd", - extensions=(".nix",), - resolve_root=_root_nix, - build_spawn=_spawn_nixd, - description="Nix — nixd", - ), - ServerDef( - server_id="zls", - extensions=(".zig", ".zon"), - resolve_root=_root_zig, - build_spawn=_spawn_zls, - description="Zig — zls", - ), - ServerDef( - server_id="gleam", - extensions=(".gleam",), - resolve_root=lambda fp, ws: _root_or_workspace(fp, ws, ["gleam.toml"]), - build_spawn=_spawn_gleam, - description="Gleam — built-in language server", - ), - ServerDef( - server_id="elixir-ls", - extensions=(".ex", ".exs"), - resolve_root=_root_elixir, - build_spawn=_spawn_elixir_ls, - description="Elixir — elixir-ls", - ), - ServerDef( - server_id="prisma", - extensions=(".prisma",), - resolve_root=_root_prisma, - build_spawn=_spawn_prisma, - description="Prisma — built-in language server", - ), - ServerDef( - server_id="kotlin-language-server", - extensions=(".kt", ".kts"), - resolve_root=_root_kotlin, - build_spawn=_spawn_kotlin_ls, - description="Kotlin — kotlin-language-server", - ), - ServerDef( - server_id="jdtls", - extensions=(".java",), - resolve_root=_root_java, - build_spawn=_spawn_jdtls, - description="Java — Eclipse JDT Language Server", - ), - ServerDef( - server_id="powershell", - extensions=(".ps1", ".psm1", ".psd1"), - resolve_root=_root_powershell, - build_spawn=_spawn_powershell_es, - description="PowerShell — PowerShellEditorServices (manual bundle)", - ), + _server("pyright", (".py", ".pyi"), "Python — Microsoft pyright", + markers=["pyproject.toml", "setup.py", "setup.cfg", "requirements.txt", "Pipfile", "pyrightconfig.json"], + build_spawn=_spawn_pyright), + _server("typescript", (".ts", ".tsx", ".js", ".jsx", ".mjs", ".cjs", ".mts", ".cts"), + "JavaScript/TypeScript — typescript-language-server", resolve_root=_root_typescript, + which=("typescript-language-server",), args=("--stdio",), install_pkg="typescript-language-server", seed=True), + _server("vue-language-server", (".vue",), "Vue.js — @vue/language-server", resolve_root=_root_typescript, + args=("--stdio",), install_pkg="@vue/language-server"), + _server("svelte-language-server", (".svelte",), "Svelte — svelte-language-server", resolve_root=_root_typescript, + which=("svelteserver", "svelte-language-server"), args=("--stdio",), install_pkg="svelte-language-server"), + _server("astro-language-server", (".astro",), "Astro — @astrojs/language-server", resolve_root=_root_typescript, + which=("astro-ls", "astro-language-server"), args=("--stdio",), install_pkg="@astrojs/language-server"), + _server("gopls", (".go",), "Go — gopls", markers=["go.work", "go.mod", "go.sum"], install_pkg="gopls"), + _server("rust-analyzer", (".rs",), "Rust — rust-analyzer", markers=["Cargo.toml", "Cargo.lock"], install_pkg="rust-analyzer"), + _server("clangd", (".c", ".cpp", ".cc", ".cxx", ".h", ".hh", ".hpp", ".hxx"), "C/C++ — clangd", + markers=["compile_commands.json", "compile_flags.txt", ".clangd"], + args=("--background-index", "--clang-tidy"), install_pkg="clangd"), + _server("bash-language-server", (".sh", ".bash", ".zsh", ".ksh"), "Bash — bash-language-server", build_spawn=_spawn_bash_ls), + _server("yaml-language-server", (".yaml", ".yml"), "YAML — yaml-language-server", + args=("--stdio",), install_pkg="yaml-language-server"), + _server("lua-language-server", (".lua",), "Lua — lua-language-server", + markers=[".luarc.json", ".luarc.jsonc", ".luacheckrc", ".stylua.toml", "stylua.toml", "selene.toml", "selene.yml"], + install_pkg="lua-language-server"), + _server("intelephense", (".php",), "PHP — intelephense", markers=["composer.json", "composer.lock", ".php-version"], + args=("--stdio",), install_pkg="intelephense", base_init={"telemetry": {"enabled": False}}), + _server("ocaml-lsp", (".ml", ".mli"), "OCaml — ocaml-lsp", markers=["dune-project", "dune-workspace", ".merlin", "opam"], + which=("ocamllsp",)), + _server("dockerfile-ls", (".dockerfile", "Dockerfile"), "Dockerfile — dockerfile-language-server-nodejs", + which=("docker-langserver",), args=("--stdio",), install_pkg="dockerfile-language-server-nodejs"), + # terraform-ls is heavy to auto-install; require the user to provide it. + _server("terraform-ls", (".tf", ".tfvars"), "Terraform — terraform-ls", markers=[".terraform.lock.hcl", "terraform.tfstate"], + args=("serve",), base_init={"experimentalFeatures": {"prefillRequiredFields": True, "validateOnSave": True}}), + _server("dart", (".dart",), "Dart — built-in language server", markers=["pubspec.yaml", "analysis_options.yaml"], + args=("language-server", "--lsp")), + _server("haskell-language-server", (".hs", ".lhs"), "Haskell — haskell-language-server", + markers=["stack.yaml", "cabal.project", "hie.yaml"], + which=("haskell-language-server-wrapper", "haskell-language-server"), args=("--lsp",)), + _server("julia", (".jl",), "Julia — LanguageServer.jl", markers=["Project.toml", "Manifest.toml"], + args=("--startup-file=no", "--history-file=no", "-e", "using LanguageServer; runserver()")), + _server("clojure-lsp", (".clj", ".cljs", ".cljc", ".edn"), "Clojure — clojure-lsp", + markers=["deps.edn", "project.clj", "shadow-cljs.edn", "bb.edn", "build.boot"], args=("listen",)), + _server("nixd", (".nix",), "Nix — nixd", resolve_root=lambda fp, ws: nearest_root(fp, ["flake.nix"]) or ws), + _server("zls", (".zig", ".zon"), "Zig — zls", markers=["build.zig"]), + _server("gleam", (".gleam",), "Gleam — built-in language server", markers=["gleam.toml"], args=("lsp",)), + _server("elixir-ls", (".ex", ".exs"), "Elixir — elixir-ls", markers=["mix.exs", "mix.lock"], + which=("elixir-ls", "language_server.sh")), + _server("prisma", (".prisma",), "Prisma — built-in language server", markers=["schema.prisma", "prisma/schema.prisma"], + args=("language-server",)), + _server("kotlin-language-server", (".kt", ".kts"), "Kotlin — kotlin-language-server", + markers=["settings.gradle", "settings.gradle.kts", "build.gradle", "build.gradle.kts", "pom.xml"]), + # jdtls has a complex install flow; we look for the wrapper script a manual install produces. + _server("jdtls", (".java",), "Java — Eclipse JDT Language Server", + markers=["pom.xml", "build.gradle", "build.gradle.kts", ".project", ".classpath", "settings.gradle"]), + # No universal PowerShell root marker; nearest_root is exact-name only (no globs). + _server("powershell", (".ps1", ".psm1", ".psd1"), "PowerShell — PowerShellEditorServices (manual bundle)", + markers=["PSScriptAnalyzerSettings.psd1"], build_spawn=_spawn_powershell_es), ] @@ -1172,8 +466,7 @@ def find_server_for_file(file_path: str) -> Optional[ServerDef]: def language_id_for(path: str) -> str: """Return the LSP languageId to send in didOpen for ``path``.""" - ext = _file_ext_or_basename(path) - return LANGUAGE_BY_EXT.get(ext, "plaintext") + return LANGUAGE_BY_EXT.get(_file_ext_or_basename(path), "plaintext") __all__ = [ diff --git a/agent/lsp/workspace.py b/agent/lsp/workspace.py index 4f5beacfbb..a33e7908c6 100644 --- a/agent/lsp/workspace.py +++ b/agent/lsp/workspace.py @@ -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() diff --git a/agent/memory_manager.py b/agent/memory_manager.py index 496562c2e8..abb56410a4 100644 --- a/agent/memory_manager.py +++ b/agent/memory_manager.py @@ -1,113 +1,116 @@ """MemoryManager — orchestrates memory providers for the agent. -Single integration point in run_agent.py. Replaces scattered per-backend -code with one manager that delegates to registered providers. - -Only ONE external plugin provider is allowed at a time — attempting to -register a second external provider is rejected with a warning. This -prevents tool schema bloat and conflicting memory backends. +Single integration point (run_agent.py) that fans out to registered providers. +The builtin provider is always allowed; only ONE external plugin provider may be +registered at a time — a second is rejected with a warning to prevent tool +schema bloat and conflicting memory backends. Usage in run_agent.py: self._memory_manager = MemoryManager() - # Only ONE of these: - self._memory_manager.add_provider(plugin_provider) - - # System prompt + self._memory_manager.add_provider(plugin_provider) # at most one external prompt_parts.append(self._memory_manager.build_system_prompt()) - - # Pre-turn - context = self._memory_manager.prefetch_all(user_message) - - # Post-turn - self._memory_manager.sync_all(user_msg, assistant_response) + context = self._memory_manager.prefetch_all(user_message) # pre-turn + self._memory_manager.sync_all(user_msg, assistant_response) # post-turn self._memory_manager.queue_prefetch_all(user_msg) """ from __future__ import annotations +import contextvars +import inspect import json import logging import re -import inspect import threading from concurrent.futures import Future, ThreadPoolExecutor, wait +from functools import partial from typing import Any, Callable, Dict, List, Optional from agent.memory_provider import MemoryProvider, PRE_COMPRESS_CHECKPOINT_API_VERSION from agent.skill_commands import extract_user_instruction_from_skill_message from tools.registry import tool_error +logger = logging.getLogger(__name__) + # Providers that predate the checkpoint-API attribute are implicitly on the # historical best-effort contract (API v1). _LEGACY_PRE_COMPRESS_API_VERSION = 1 +# How long shutdown_all() waits for in-flight background sync/prefetch work to +# drain before abandoning it. Worker threads are daemon, so a wedged provider +# never blocks interpreter exit — it dies with the process past this window. +_SYNC_DRAIN_TIMEOUT_S = 5.0 +_EXTERNAL_PREFETCH_TIMEOUT_S = 8.0 + +_VAR_KEYWORD = inspect.Parameter.VAR_KEYWORD + + +# --------------------------------------------------------------------------- +# Signature introspection (providers are duck-typed; call shapes vary) +# --------------------------------------------------------------------------- + +def _signature_params(fn: Callable[..., Any]): + """Return ``fn``'s parameter mapping, or None when uninspectable (C callables, exotic proxies).""" + try: + return inspect.signature(fn).parameters + except (TypeError, ValueError): + return None + + +def _has_var_kwargs(params) -> bool: + return any(p.kind is _VAR_KEYWORD for p in params.values()) + def _accepts_require_checkpoint(fn: Callable[..., Any]) -> bool: """True if ``fn`` can receive the ``require_checkpoint`` keyword. - Checkpoint (v2) providers written against the original docs example use - the bare ``on_pre_compress(self, messages)`` signature; calling them with - the keyword would raise ``TypeError`` — which, under - ``require_checkpoint=True``, the host would re-raise as a checkpoint - failure even though the provider's durable write succeeded. Inspect the - signature and fall back to the legacy call shape when the keyword (or a - ``**kwargs`` catch-all) is absent. Unreadable signatures (C callables, - exotic proxies) conservatively report False. + Checkpoint (v2) providers written against the original docs example use the + bare ``on_pre_compress(self, messages)`` shape; passing the keyword would + raise TypeError, which under ``require_checkpoint=True`` the host would + re-raise as a checkpoint failure even though the durable write succeeded. + Unreadable signatures conservatively report False. """ - try: - sig = inspect.signature(fn) - except (TypeError, ValueError): + params = _signature_params(fn) + if params is None: return False - for param in sig.parameters.values(): - if param.kind is inspect.Parameter.VAR_KEYWORD: - return True - if ( - param.name == "require_checkpoint" - and param.kind in ( - inspect.Parameter.KEYWORD_ONLY, - inspect.Parameter.POSITIONAL_OR_KEYWORD, - ) - ): - return True - return False + if _has_var_kwargs(params): + return True + param = params.get("require_checkpoint") + return param is not None and param.kind in ( + inspect.Parameter.KEYWORD_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + ) -logger = logging.getLogger(__name__) -# How long shutdown_all() waits for in-flight background sync/prefetch work -# to drain before abandoning it. A wedged provider must never block process -# teardown indefinitely — the worker threads are daemon, so anything still -# running past this window dies with the interpreter. -_SYNC_DRAIN_TIMEOUT_S = 5.0 -_EXTERNAL_PREFETCH_TIMEOUT_S = 8.0 +def _ctx_bound(fn: Callable[[], Any]) -> Callable[[], Any]: + """Bind ``fn`` to the CALLER's contextvars for execution on another thread. + Profile isolation in multi-profile processes (gateway multiplexer, dashboard, + cron) is a ContextVar-scoped HERMES_HOME override; worker threads start with + empty contexts, so an unbound provider resolving config paths or secrets + from a worker would silently land on the default profile. + """ + return partial(contextvars.copy_context().run, fn) + + +# --------------------------------------------------------------------------- +# Tool-schema plumbing +# --------------------------------------------------------------------------- def normalize_tool_schema(schema: Any) -> Optional[Dict[str, Any]]: - """Return a function-tool dict with a resolvable top-level ``name``. + """Return a bare function-tool dict with a resolvable top-level ``name``, else None. - Context engines and memory providers expose tool schemas via - ``get_tool_schemas()``. The expected shape is a bare function schema - (``{"name": ..., "description": ..., "parameters": ...}``) which callers - wrap as ``{"type": "function", "function": schema}``. - - Some providers instead return an entry that is *already* in OpenAI tool - form (``{"type": "function", "function": {"name": ...}}``). Wrapping that - a second time produces ``{"type": "function", "function": {"type": - "function", "function": {...}}}`` whose ``function`` has no top-level - ``name``. Strict providers (e.g. DeepSeek) reject the *entire* request - with ``tools[N].function: missing field name`` (HTTP 400), so one bad - schema disables the whole toolset and breaks every turn (#47707). - - This helper normalizes both shapes to the bare function schema and - returns ``None`` for anything without a resolvable name, so callers can - skip-with-warning rather than appending a nameless tool. + Providers should return ``{"name", "description", "parameters"}`` which callers + wrap as ``{"type": "function", "function": schema}``. Some return the already + wrapped OpenAI form; wrapping that twice yields a ``function`` with no ``name`` + and strict providers (e.g. DeepSeek) reject the ENTIRE request (HTTP 400), + disabling every tool. Both shapes are normalized here so callers can skip + nameless entries with a warning instead of poisoning the request. """ if not isinstance(schema, dict): return None - # Unwrap an already-wrapped OpenAI tool entry. if schema.get("type") == "function" and isinstance(schema.get("function"), dict): schema = schema["function"] - if not isinstance(schema, dict): - return None name = schema.get("name", "") if not name or not isinstance(name, str): return None @@ -123,9 +126,7 @@ def memory_provider_tools_enabled( """Return whether external memory-provider tools should be exposed.""" if disabled_toolsets and "memory" in disabled_toolsets: return False - if memory_tool_present: - return True - if enabled_toolsets is None: + if memory_tool_present or enabled_toolsets is None: return True if not enabled_toolsets: return False @@ -144,19 +145,15 @@ def memory_provider_tools_enabled( def memory_provider_tools_exposed(agent: Any) -> bool: """Whether external memory-provider tools are exposed on ``agent``. - Same gate as ``inject_memory_provider_tools`` so the provider's - ``system_prompt_block()`` and its tool schemas are presented to the - model together — otherwise the system prompt would advertise tools - that don't exist in the tool surface (#81014). + Same gate as ``inject_memory_provider_tools`` so a provider's + ``system_prompt_block()`` and its tool schemas are presented together — + the system prompt must never advertise tools absent from the tool surface. """ tools = getattr(agent, "tools", None) - if isinstance(tools, (list, tuple)): - memory_tool_present = any( - isinstance(tool, dict) and tool.get("function", {}).get("name") == "memory" - for tool in tools - ) - else: - memory_tool_present = False + memory_tool_present = isinstance(tools, (list, tuple)) and any( + isinstance(tool, dict) and tool.get("function", {}).get("name") == "memory" + for tool in tools + ) return memory_provider_tools_enabled( getattr(agent, "enabled_toolsets", None), getattr(agent, "disabled_toolsets", None), @@ -165,7 +162,7 @@ def memory_provider_tools_exposed(agent: Any) -> bool: def inject_memory_provider_tools(agent: Any) -> int: - """Append external memory-provider tool schemas to an agent tool surface.""" + """Append external memory-provider tool schemas to an agent tool surface; return count added.""" memory_manager = getattr(agent, "_memory_manager", None) tools = getattr(agent, "tools", None) if not memory_manager or tools is None: @@ -177,10 +174,8 @@ def inject_memory_provider_tools(agent: Any) -> int: if isinstance(tool, dict) } if not memory_provider_tools_exposed(agent): - # A provider is configured but the memory toolset is gated off - # (platform_toolsets / disabled_toolsets). Say so once — a silent - # return 0 here made #81014 undiagnosable: the provider looked - # "half on" with no clue which config key suppressed its tools. + # Say so once: a silent 0 leaves the provider looking "half on" with no + # clue which config key (platform_toolsets / disabled_toolsets) gated it. _providers = [ p for p in (getattr(memory_manager, "providers", None) or []) if getattr(p, "name", "") != "builtin" @@ -244,56 +239,33 @@ def sanitize_context(text: str) -> str: """Strip fence tags, injected context blocks, and system notes from provider output.""" text = _INTERNAL_CONTEXT_RE.sub('', text) text = _INTERNAL_NOTE_RE.sub('', text) - text = _FENCE_TAG_RE.sub('', text) - return text + return _FENCE_TAG_RE.sub('', text) class StreamingContextScrubber: - """Stateful scrubber for streaming text that may contain split memory-context spans. + """Stateful scrubber for streaming text whose memory-context spans may straddle deltas. - The one-shot ``sanitize_context`` regex cannot survive chunk boundaries: - a ```` opened in one delta and closed in a later delta - leaks its payload to the UI because the non-greedy block regex needs - both tags in one string. This scrubber runs a small state machine - across deltas, holding back partial-tag tails and discarding - everything inside a span (including the system-note line). - - Usage:: - - scrubber = StreamingContextScrubber() - for delta in stream: - visible = scrubber.feed(delta) - if visible: - emit(visible) - trailing = scrubber.flush() # at end of stream - if trailing: - emit(trailing) - - The scrubber is re-entrant per agent instance. Callers building new - top-level responses (new turn) should create a fresh scrubber or call - ``reset()``. + The one-shot ``sanitize_context`` regex needs both tags in one string, so a + span opened in one delta and closed in a later one would leak to the UI. + This state machine holds back partial-tag tails between ``feed()`` calls and + drops everything inside a span (including the system-note line). Create a + fresh scrubber (or ``reset()``) per top-level response; call ``flush()`` at + end of stream. """ _OPEN_TAG = "" _CLOSE_TAG = "" def __init__(self) -> None: + self.reset() + + def reset(self) -> None: self._in_span: bool = False self._buf: str = "" self._at_block_boundary: bool = True - def reset(self) -> None: - self._in_span = False - self._buf = "" - self._at_block_boundary = True - def feed(self, text: str) -> str: - """Return the visible portion of ``text`` after scrubbing. - - Any trailing fragment that could be the start of an open/close tag - is held back in the internal buffer and surfaced on the next - ``feed()`` call or discarded/emitted by ``flush()``. - """ + """Return the visible portion of ``text``; a possible partial tag tail is held for the next call.""" if not text: return "" buf = self._buf + text @@ -304,28 +276,23 @@ class StreamingContextScrubber: if self._in_span: idx = buf.lower().find(self._CLOSE_TAG) if idx == -1: - # Hold back a potential partial close tag; drop the rest + # Hold back a potential partial close tag; drop the rest. held = self._max_partial_suffix(buf, self._CLOSE_TAG) self._buf = buf[-held:] if held else "" return "".join(out) - # Found close — skip span content + tag, continue buf = buf[idx + len(self._CLOSE_TAG):] self._in_span = False else: idx = self._find_boundary_open_tag(buf) if idx == -1: - # No open tag — hold back a potential partial open tag held = ( self._max_pending_open_suffix(buf) or self._max_partial_suffix(buf, self._OPEN_TAG) ) + self._append_visible(out, buf[:-held] if held else buf) if held: - self._append_visible(out, buf[:-held]) self._buf = buf[-held:] - else: - self._append_visible(out, buf) return "".join(out) - # Emit text before the tag, enter span if idx > 0: self._append_visible(out, buf[:idx]) buf = buf[idx + len(self._OPEN_TAG):] @@ -334,12 +301,11 @@ class StreamingContextScrubber: return "".join(out) def flush(self) -> str: - """Emit any held-back buffer at end-of-stream. + """Emit the held-back tail at end-of-stream. - If we're still inside an unterminated span the remaining content is - discarded (safer: leaking partial memory context is worse than a - truncated answer). Otherwise the held-back partial-tag tail is - emitted verbatim (it turned out not to be a real tag). + Inside an unterminated span the remainder is discarded — leaking partial + memory context is worse than a truncated answer. Otherwise the held tail + was not a real tag and is emitted verbatim. """ if self._in_span: self._buf = "" @@ -351,20 +317,16 @@ class StreamingContextScrubber: @staticmethod def _max_partial_suffix(buf: str, tag: str) -> int: - """Return the length of the longest buf-suffix that is a tag-prefix. - - Case-insensitive. Returns 0 if no suffix could start the tag. - """ + """Length of the longest buf-suffix that is a (case-insensitive) prefix of ``tag``, else 0.""" tag_lower = tag.lower() buf_lower = buf.lower() - max_check = min(len(buf_lower), len(tag_lower) - 1) - for i in range(max_check, 0, -1): + for i in range(min(len(buf_lower), len(tag_lower) - 1), 0, -1): if tag_lower.startswith(buf_lower[-i:]): return i return 0 def _find_boundary_open_tag(self, buf: str) -> int: - """Find an opening fence only when it starts a block-like span.""" + """Find an opening fence only when it starts a block-like span (own line, newline after).""" buf_lower = buf.lower() search_start = 0 while True: @@ -376,19 +338,16 @@ class StreamingContextScrubber: search_start = idx + 1 def _max_pending_open_suffix(self, buf: str) -> int: - """Hold a complete boundary tag until the following char confirms it.""" + """Hold a complete boundary tag at the buffer end until the following char confirms it.""" if not buf.lower().endswith(self._OPEN_TAG): return 0 - idx = len(buf) - len(self._OPEN_TAG) - if not self._is_block_boundary(buf, idx): + if not self._is_block_boundary(buf, len(buf) - len(self._OPEN_TAG)): return 0 return len(self._OPEN_TAG) def _has_block_opener_suffix(self, buf: str, idx: int) -> bool: after_idx = idx + len(self._OPEN_TAG) - if after_idx >= len(buf): - return False - return buf[after_idx] in "\r\n" + return after_idx < len(buf) and buf[after_idx] in "\r\n" def _is_block_boundary(self, buf: str, idx: int) -> bool: if idx == 0: @@ -403,9 +362,6 @@ class StreamingContextScrubber: if not text: return out.append(text) - self._update_block_boundary(text) - - def _update_block_boundary(self, text: str) -> None: last_newline = text.rfind("\n") if last_newline != -1: self._at_block_boundary = text[last_newline + 1:].strip() == "" @@ -430,17 +386,22 @@ def build_memory_context_block(raw_context: str) -> str: ) +def _nonblank(text: Any) -> Any: + """Return ``text`` when it has non-whitespace content, else None.""" + return text if text and text.strip() else None + + class MemoryManager: """Orchestrates the built-in provider plus at most one external provider. - The builtin provider is always first. Only one non-builtin (external) - provider is allowed. Failures in one provider never block the other. + The builtin provider is always first. Failures in one provider never block + the other: every fan-out hook logs and swallows per-provider exceptions. """ def __init__(self, *, external_prefetch_timeout: Optional[float] = None) -> None: self._providers: List[MemoryProvider] = [] self._tool_to_provider: Dict[str, MemoryProvider] = {} - self._has_external: bool = False # True once a non-builtin provider is added + self._has_external: bool = False self._external_prefetch_timeout = ( _EXTERNAL_PREFETCH_TIMEOUT_S if external_prefetch_timeout is None @@ -450,15 +411,13 @@ class MemoryManager: raise ValueError("external_prefetch_timeout must be positive") self._external_prefetch_threads: Dict[str, threading.Thread] = {} self._external_prefetch_lock = threading.Lock() - # Background executor for end-of-turn sync/prefetch. Lazily created on - # first use so the common builtin-only path spawns no extra threads. - # A single worker serializes a provider's writes (turn N must land - # before turn N+1) and caps thread growth at one per manager. See - # _submit_background() and the sync_all/queue_prefetch_all rationale. + # Single-worker background executor for end-of-turn sync/prefetch, + # created lazily so the builtin-only path spawns no threads. One worker + # serializes a provider's writes (turn N lands before turn N+1). self._sync_executor: Optional[ThreadPoolExecutor] = None self._sync_executor_lock = threading.Lock() - # Futures are tracked by durability class so shutdown can give writes - # a bounded FIFO drain, then explicitly report anything abandoned. + # Futures tracked by durability class ("write" / "prefetch") so shutdown + # can drain FIFO within a bound, then report exactly what it abandoned. self._background_futures: Dict[Future, str] = {} self._shutting_down = False self._shutdown_drain_state: Dict[str, Any] = { @@ -468,18 +427,38 @@ class MemoryManager: "active_tasks": 0, } + # -- Fan-out helper ------------------------------------------------------ + + def _each_provider( + self, + label: str, + call: Callable[[MemoryProvider], Any], + *, + level: int = logging.DEBUG, + providers: Optional[List[MemoryProvider]] = None, + exc_info: bool = False, + ) -> List[Any]: + """Call ``call(provider)`` for each provider, logging and swallowing failures. + + ``label`` completes the log line ``Memory provider ''