diff --git a/agent/moa_loop.py b/agent/moa_loop.py index 631e726cf7..4a980fb604 100644 --- a/agent/moa_loop.py +++ b/agent/moa_loop.py @@ -25,10 +25,9 @@ from agent.usage_pricing import CanonicalUsage logger = logging.getLogger(__name__) -# Privacy filter (moa.privacy_filter: '' | display | full). Secret shapes are handled -# by agent.redact; these add the PII classes it leaves alone (emails, formatted NA -# phones). The phone pattern requires explicit delimiters so line numbers, dates, -# times, SHAs, IPs and versions never match. +# Privacy filter (moa.privacy_filter: '' | display | full): PII classes agent.redact +# leaves alone. The phone pattern requires explicit delimiters so line numbers, +# dates, times, SHAs, IPs and versions never match. _MOA_EMAIL_RE = re.compile(r"\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}\b") _MOA_PHONE_RE = re.compile( r"(? Any: """Redact secrets (central redactor) then MoA PII patterns. force=True: the privacy filter is its own opt-in, independent of the global - log-redaction toggle. code_file=True: advisory text is prose/code, so the - ENV/JSON assignment heuristics that mangle source snippets stay off. + log-redaction toggle. code_file=True: keeps the ENV/JSON assignment heuristics + (which mangle source snippets) off advisory prose/code. """ if not isinstance(text, str) or not text: return text @@ -105,18 +104,14 @@ def _redact_trace_accounting(acct: Any) -> Any: ) -# Cold-start caches: the preset and each (provider, model) runtime are immutable -# for a turn, so avoid re-resolving them on every create() call. +# Cold-start caches: preset and per-(provider, model) runtime are immutable for a turn. _preset_cache_lock = threading.Lock() _preset_cache: dict[tuple, Any] = {} def _resolve_preset_cached(preset_name: str) -> tuple[dict[str, Any], Any]: - """Return ``(preset, raw moa config)``, caching the resolved preset per config mtime. - - load_config() is (mtime_ns, size)-cached upstream; the saving here is skipping - resolve_moa_preset's full validation of the moa block on every create(). - """ + """``(preset, raw moa config)``; the resolved preset is cached per config mtime + (skips resolve_moa_preset's full validation of the moa block on every create()).""" from hermes_cli.config import get_config_path, load_config from hermes_cli.moa_config import resolve_moa_preset @@ -153,9 +148,8 @@ _MAX_REFERENCE_WORKERS = 8 class _RefAccounting: """Per-reference usage, cost and full trace (third slot of a reference-output tuple). - Advisors may run on a different model than the aggregator, so cost is priced at - the advisor's OWN rate and summed in dollars. The trace fields (``messages``, - ``output``, ``model``, ``provider``, ``temperature``) are only populated when + Cost is priced at the advisor's OWN rate and summed in dollars (advisors may run + on a different model than the aggregator). Trace fields are only populated when tracing is on. """ @@ -227,11 +221,10 @@ def _slot_reasoning_config(slot: dict[str, Any]) -> dict[str, Any] | None: def _aggregator_reasoning_config(aggregator: dict[str, Any]) -> dict[str, Any] | None: - """Resolve the aggregator's reasoning config: slot > per-model > global. + """Aggregator reasoning config: slot > per-model > global (shared chokepoint). - The aggregator is the ACTING model, so it falls back through the shared - ``resolve_reasoning_config`` chokepoint. References deliberately do not: - inheriting a global ``xhigh`` into every advisor would multiply cost. + References deliberately do NOT fall back: inheriting a global ``xhigh`` into + every advisor would multiply cost. """ cfg = _slot_reasoning_config(aggregator) if cfg is not None: @@ -246,12 +239,10 @@ def _aggregator_reasoning_config(aggregator: dict[str, Any]) -> dict[str, Any] | def _slot_runtime(slot: dict[str, Any]) -> dict[str, Any]: - """Resolve a slot to ``call_llm`` kwargs via ``resolve_runtime_provider``. + """Slot → ``call_llm`` kwargs with the provider's real api_mode/base_url/api_key. - Gives the slot its provider's real api_mode/base_url/api_key instead of letting - the auxiliary auto-detector guess. Falls back to bare provider/model on error - (never cached: a transient error would pin bare kwargs for a TTL). Cached per - (provider, model) with a short TTL. + Cached per (provider, model) with a short TTL. Falls back to bare provider/model + on error — never cached, or a transient error would pin bare kwargs for a TTL. """ provider = str(slot.get("provider") or "").strip() model = str(slot.get("model") or "").strip() @@ -314,11 +305,10 @@ def _maybe_apply_moa_cache_control( ) -> list[dict[str, Any]]: """Apply cache_control to an advisor/aggregator request when its route honors it. - Same policy function and marker helper as the main loop; MoA has no static - prefix so the legacy system-and-3 fallback is used. ``cache_disabled`` is - stamped onto the stub so ``cache_ttl: off`` is honored; ``cache_ttl`` threads - the agent's tier, clamped per destination. Returns the messages unchanged on - any error. + Same policy/marker helpers as the main loop; MoA has no static prefix so the + legacy system-and-3 fallback is used. ``cache_disabled`` is stamped onto the + stub so ``cache_ttl: off`` is honored; ``cache_ttl`` is clamped per destination. + Returns the messages unchanged on any error. """ try: from agent.agent_runtime_helpers import ( @@ -401,12 +391,8 @@ def _run_reference( cache_disabled: bool | None = None, cache_ttl: str | None = None, ) -> tuple[str, str, Any]: - """Call one reference model; return ``(label, text, accounting)``. Never raises. - - The slot is resolved to its provider's real runtime and called through the - same ``call_llm`` path any model uses. A failed reference becomes a labelled - ``[failed: …]`` note. Runs inside a thread pool (call_llm is blocking). - """ + """Call one reference model; return ``(label, text, accounting)``. Never raises: + a failed reference becomes a labelled ``[failed: …]`` note. Runs in a thread pool.""" label = _slot_label(slot) runtime = _slot_runtime(slot) trace_fields = { @@ -414,11 +400,10 @@ def _run_reference( "provider": runtime.get("provider") or slot.get("provider"), "temperature": temperature, } - # The trimmed view already stripped the agent's system prompt; this is the only one. + # The advisory view already stripped the agent's system prompt; this is the only one. messages = [{"role": "system", "content": _REFERENCE_SYSTEM_PROMPT}, *ref_messages] try: - # Trim to THIS model's window (advisors may be smaller than the aggregator); - # estimated after the system prompt is prepended so it counts too. + # Trim to THIS model's window (advisors may be smaller than the aggregator). trimmed = _trim_messages_for_reference( messages, slot, @@ -426,9 +411,8 @@ def _run_reference( reserve_output_tokens=max_tokens, context_length_cache=context_length_cache, ) - # Anthropic-style caching is opt-in per request; the advisory view is append-only - # across iterations, so decorating lets iteration N+1 replay N's cached prefix. - # The live agent disable is pinned onto the runtime (not a fresh config read). + # The advisory view is append-only across iterations, so cache_control lets + # iteration N+1 replay N's cached prefix. trimmed = _maybe_apply_moa_cache_control( trimmed, _with_cache_disabled(runtime, cache_disabled), cache_ttl=cache_ttl ) @@ -485,12 +469,11 @@ def _trim_messages_for_reference( ) -> list[dict[str, Any]]: """Trim an advisory request to fit a reference model's context window. - ``messages`` is the full request (advisory system prompt included). Budget = - window minus ``reserve_output_tokens`` (or a default) minus a safety fraction. - Drops the OLDEST frames after the system prompt while keeping: the system - prompt, a user-first body, and the trailing user turn plus one preceding turn - (even if still over budget). ``context_length_cache`` memoizes the window per - (provider, model) for the fan-out; unresolvable windows leave messages unchanged. + Budget = window − ``reserve_output_tokens`` (or a default) − a safety fraction. + Drops the OLDEST frames after the system prompt, always keeping a user-first + body and the trailing user turn plus one preceding turn (even if still over + budget). ``context_length_cache`` memoizes the window per (provider, model); + unresolvable windows leave messages unchanged. """ if not messages: return messages @@ -582,13 +565,9 @@ def _settle_interrupted( reference_models: list[dict[str, Any]], late_accounting_sink: Any, ) -> None: - """Fill every unfinished slot after a user interrupt. - - Never-dispatched futures are cancelled (nothing billed). Ones that finished - between the interrupt check and now keep their real output. Running ones cannot - be killed and WILL be billed on completion, so their eventual accounting is - handed to ``late_accounting_sink``. - """ + """Fill every unfinished slot after a user interrupt: cancel never-dispatched + futures (nothing billed), keep real output of ones that just finished, and hand + running ones (cannot be killed, WILL bill) to ``late_accounting_sink``.""" for future, idx in futures.items(): if results[idx] is not None: continue @@ -623,23 +602,20 @@ def _run_references_parallel( agent: Any = None, late_accounting_sink: Any = None, ) -> list[tuple[str, str, Any]]: - """Fan out all reference models in parallel; outputs are in ``reference_models`` order. + """Fan out all reference models in parallel; ``(label, text, _RefAccounting)`` per + slot in ``reference_models`` order. - Slots with ``provider == "moa"`` are skipped with a labelled note (recursion - guard). ``progress_callback(refs_done, refs_total, label)`` fires per completion - (best-effort). Each element is ``(label, text, _RefAccounting)``. - - With *agent*, the wait polls every ``_REFERENCE_POLL_INTERVAL_S`` so a user - interrupt can abort it (like tool_executor's batch). In-flight calls cannot be - killed; their eventual accounting goes to ``late_accounting_sink``. + ``provider == "moa"`` slots are skipped with a note (recursion guard). + ``progress_callback(refs_done, refs_total, label)`` fires per completion. With + *agent*, the wait polls every ``_REFERENCE_POLL_INTERVAL_S`` so a user interrupt + can abort it; in-flight calls cannot be killed and bill via ``late_accounting_sink``. """ if not reference_models: return [] results: list[tuple[str, str, Any] | None] = [None] * len(reference_models) futures: dict[Any, int] = {} - # Executor threads start with an empty contextvars.Context; propagate the turn's - # (approval callbacks + Nous Portal conversation tag) into each worker. + # Propagate the turn's contextvars (approval callbacks, Nous conversation tag). from tools.thread_context import propagate_context_to_thread total = len(reference_models) @@ -748,32 +724,28 @@ _ADVISORY_INSTRUCTION = ( def _reference_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: """Build the advisory (reference-model) view of the conversation. - Flattens the transcript to plain user/assistant TEXT turns: the system prompt - is dropped, tool_calls are rendered inline, and tool results are folded into - the preceding assistant turn as ``[tool result: ...]`` previews. Zero tool-role - messages / tool_calls arrays are emitted, so strict providers do not 400. - The view always ends on a ``user`` turn (Anthropic treats a trailing assistant - turn as prefill): a synthetic advisory request is APPENDED rather than - deleting context. The aggregator always receives the full transcript. + Plain user/assistant TEXT turns only: system prompt dropped, tool_calls rendered + inline, tool results folded into the preceding assistant turn as previews (no + tool-role messages / tool_calls arrays, so strict providers do not 400). Always + ends on a ``user`` turn (Anthropic treats a trailing assistant turn as prefill) + by APPENDING a synthetic request. The aggregator always gets the full transcript. """ rendered: list[dict[str, Any]] = [] last_user_content: str | None = None for msg in messages: role = msg.get("role") content = msg.get("content") - # Content may be a list (cache_control-decorated text parts, multimodal turns); - # flatten_message_text extracts text so decorated and undecorated transcripts - # yield a byte-identical view (stable advisory prefix for advisor caching). + # Decorated (cache_control parts) and undecorated transcripts must yield a + # byte-identical view so the advisory prefix stays cache-stable. text = flatten_message_text(content) if role == "user": if not text.strip() and isinstance(content, list) and content: - # Structured content with no text (e.g. image-only): an empty user - # message is rejected by strict providers and skipping breaks alternation. + # Image-only turn: empty user messages are rejected by strict providers + # and skipping would break alternation. text = "[user sent non-text content (e.g. an image attachment)]" if not text.strip(): - # Genuinely empty user turn: strict providers (Kimi, ZAI) 400 on it; dropping - # is safe because the advisory view is not strictly alternating anyway. + # Genuinely empty user turn: strict providers 400 on it; safe to drop. continue last_user_content = text rendered.append({"role": "user", "content": text}) @@ -795,8 +767,7 @@ def _reference_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: rendered.append({"role": "assistant", "content": block}) # system and any other role are ignored. - # End on a user turn by appending a synthetic advisory request (Anthropic - # rejects trailing assistant prefill); an existing trailing user turn is left as is. + # Anthropic rejects trailing assistant prefill: end on a synthetic user request. if rendered and rendered[-1].get("role") == "assistant": rendered.append({"role": "user", "content": _ADVISORY_INSTRUCTION}) @@ -915,9 +886,8 @@ def aggregate_moa_context( """Run configured reference models and synthesize their advice (one-shot /moa). Failures become model-specific notes instead of aborting the loop. - ``reference_max_tokens`` caps ONLY the fan-out — capping the aggregator - truncated long syntheses. ``temperature`` / ``aggregator_temperature`` - default to None (provider default). ``agent`` makes the fan-out interruptible. + ``reference_max_tokens`` caps ONLY the fan-out (capping the aggregator truncated + long syntheses). ``agent`` makes the fan-out interruptible. """ reference_models = [slot for slot in reference_models if slot.get("enabled", True)] reference_outputs = _run_references_parallel( @@ -930,8 +900,7 @@ def aggregate_moa_context( ) successful_outputs, failed_labels = _split_references(reference_outputs) - # 'full' privacy mode also redacts advisor text before it reaches this synthesizer - # ('display' has no surface here). Failed refs are already filtered out. + # 'full' privacy mode also redacts advisor text before it reaches the synthesizer. try: from hermes_cli.config import load_config as _load_config @@ -1001,10 +970,7 @@ def aggregate_moa_context( def _completed_response_as_stream_chunk(response: Any) -> Any: - """Adapt a completed response (``choices[0].message``) into one delta stream chunk. - - Done at the MoA facade boundary so transports stay untouched. - """ + """Adapt a completed response into one delta stream chunk (facade boundary only).""" choices = getattr(response, "choices", None) first_choice = choices[0] if isinstance(choices, (list, tuple)) and choices else None message = getattr(first_choice, "message", None) @@ -1046,12 +1012,10 @@ def _completed_response_as_stream_chunk(response: Any) -> Any: def _attach_reference_guidance(agg_messages: list[dict[str, Any]], guidance: str) -> None: """Attach the per-turn reference block at the END of the aggregator prompt. - The block varies per iteration; merging it into the (early) original user - message would diverge the prompt prefix and re-prefill the whole conversation - each step. Appending keeps ``[system][task][tool-history]`` cache-stable. A - trailing user turn is merged in place (string or content-part list — a new text - part rides AFTER the cache_control-marked part); otherwise a user message is - appended (two consecutive user turns would be rejected by strict providers). + The block varies per iteration; appending keeps ``[system][task][tool-history]`` + cache-stable. A trailing user turn is merged in place (string, or a new text part + AFTER the cache_control-marked part); otherwise a user message is appended (two + consecutive user turns would be rejected by strict providers). """ last = agg_messages[-1] if agg_messages else None if last is not None and last.get("role") == "user": @@ -1069,11 +1033,8 @@ def peel_reference_guidance( messages: list[dict[str, Any]], guidance: Any, ) -> list[dict[str, Any]]: - """Exact inverse of ``_attach_reference_guidance`` (the three attach shapes). - - Used by the failover redecoration chokepoint so a cache breakpoint never lands - on the turn-varying guidance. Returns a new list; inputs are not mutated. - """ + """Exact inverse of ``_attach_reference_guidance`` (the three attach shapes), so a + cache breakpoint never lands on the turn-varying guidance. Inputs are not mutated.""" if not guidance or not messages: return messages guidance_text = str(guidance) @@ -1107,39 +1068,33 @@ def peel_reference_guidance( class MoAChatCompletions: """OpenAI-chat-compatible facade where the aggregator is the acting model. - ``reference_callback(event, **kwargs)`` is an optional best-effort display hook: - "moa.reference" index, count, label, text - "moa.progress" refs_done, refs_total, label (per reference completion) - "moa.phase" phase, refs_done, refs_total, aggregator - "moa.aggregating" aggregator (label), ref_count - ``agent`` is the owning AIAgent; it lets the fan-out check ``_interrupt_requested``. + ``reference_callback(event, **kwargs)`` is an optional best-effort display hook + (events: ``moa.reference``, ``moa.progress``, ``moa.phase``, ``moa.aggregating``; + kwargs per ``_RELAY_EVENTS``). ``agent`` is the owning AIAgent; it lets the + fan-out check ``_interrupt_requested``. """ def __init__(self, preset_name: str, reference_callback: Any = None, agent: Any = None): self.preset_name = preset_name or "default" self.reference_callback = reference_callback self._agent = agent - # State-scoped reference cache keyed on the advisory-view signature: a new - # user/tool message is a MISS (references re-run), a redundant create() with - # identical state is a HIT (no re-run, no re-emit). + # Reference cache keyed on the advisory-view signature: new state = MISS + # (references re-run), identical state = HIT (no re-run, no re-emit). self._ref_cache_key: tuple | None = None self._ref_cache_outputs: list[tuple[str, str, Any]] = [] - # Fan-out usage/cost from the latest cache-MISS create(), awaiting - # consume_reference_usage (zero deposited on a HIT so spend counts once). - # The lock guards them against late-accounting callbacks on worker threads. + # Fan-out spend awaiting consume_reference_usage (nothing deposited on a HIT so + # spend counts once); the lock guards late-accounting callbacks on worker threads. self._pending_reference_usage: Any = CanonicalUsage() self._pending_reference_cost: Any = None self._accounting_lock = threading.Lock() - # Resolved aggregator slot from the latest create(); cost accounting prices the - # acting turn at its real model instead of the virtual preset name. + # Real aggregator slot so cost accounting prices the acting turn at its model. self.last_aggregator_slot: Any = None # Full-turn trace parts from a cache-MISS create(), flushed by consume_and_save_trace. self._pending_trace: Any = None # Per-advisor metrics for observability hooks; NOT consumed (post_api_request # fires on a different branch than consume_and_save_trace). self._last_reference_metrics: Any = None - # every_n cadence state, scoped to a single USER TURN (resets on a new user - # message) so iteration 1 of every turn is on-cadence. + # every_n cadence state, scoped to one USER TURN so iteration 1 is on-cadence. self._fanout_iteration_count = 0 self._fanout_turn_sig: str | None = None self._fanout_last_state_sig: str | None = None @@ -1147,10 +1102,8 @@ class MoAChatCompletions: self._privacy_mode: str = "" def consume_reference_usage(self) -> tuple[Any, Any]: - """Pop pending fan-out ``(CanonicalUsage, cost_usd_or_None)`` and reset both. - - Clearing prevents a streaming retry re-entering accounting from double-counting. - """ + """Pop pending fan-out ``(CanonicalUsage, cost_usd_or_None)`` and reset both + (so a streaming retry re-entering accounting cannot double-count).""" with self._accounting_lock: usage = self._pending_reference_usage or CanonicalUsage() cost = self._pending_reference_cost @@ -1163,10 +1116,7 @@ class MoAChatCompletions: return self._last_reference_metrics def _record_late_reference_accounting(self, label: str, accounting: Any) -> None: - """Fold a late-completing interrupted reference's real spend into pending totals. - - Registered as a done-callback on abandoned futures (they still bill). Thread-safe. - """ + """Done-callback for abandoned (still billing) futures: fold their real spend in.""" if not isinstance(accounting, _RefAccounting): return self._fold_pending_accounting(*_sum_reference_accounting([(label, "", accounting)])) @@ -1184,9 +1134,8 @@ class MoAChatCompletions: ) -> None: """Flush the pending full-turn trace to disk (no-op when nothing is pending). - ``aggregator_output_fallback`` is the caller's resolved acting text: on the - streaming path the output could not be captured at ``create()`` time, so it - is folded in here. Clears the pending trace; never raises. + ``aggregator_output_fallback`` is the caller's resolved acting text for the + streaming path (not capturable at ``create()`` time). Never raises. """ pending = self._pending_trace self._pending_trace = None @@ -1225,11 +1174,8 @@ class MoAChatCompletions: logger.debug("MoA reference_callback failed for %s: %s", event, exc) def prepare(self, messages: list[dict[str, Any]]) -> dict[str, Any]: - """Run the advisor fan-out and return the exact aggregator request. - - The loop measures this augmented prompt before its compression gate, then - hands the object back to ``create()`` so the fan-out is not repeated. - """ + """Run the advisor fan-out and return the exact aggregator request, which the + loop measures before its compression gate and hands back to ``create()``.""" return self.create(messages=messages, _moa_prepare_only=True) def rebase_prepared_request( @@ -1249,11 +1195,11 @@ class MoAChatCompletions: guidance: Any, agg_runtime: dict[str, Any], ) -> tuple[list[dict[str, Any]], Any]: - """Decorate the aggregator request with cache breakpoints for its destination. + """Cache-breakpoint the aggregator request for its destination. - The guidance is peeled before planning and re-attached after so a breakpoint - never lands on the turn-varying block. On any error the undecorated request - is returned (warning, not debug: this is the aggregator's ONLY decoration path). + Guidance is peeled before planning and re-attached after so a breakpoint never + lands on the turn-varying block. Any error → undecorated request (warning, not + debug: this is the aggregator's ONLY decoration path). """ try: from agent.agent_runtime_helpers import plan_cache_sections_for_destination @@ -1261,7 +1207,6 @@ class MoAChatCompletions: planning_messages = agg_messages if guidance: planning_messages = peel_reference_guidance(agg_messages, str(guidance)) - # plan_cache_sections_for_destination returns request-local copies. # Tri-state cache_disabled: facades built via __new__ have no _agent; forcing # False would suppress the planner's config fallback. _agent = getattr(self, "_agent", None) @@ -1298,8 +1243,7 @@ class MoAChatCompletions: agg_messages, tools = self._plan_aggregator_cache( prepared["messages"], api_kwargs.get("tools"), prepared.get("guidance"), agg_runtime ) - # Record the exact aggregator INPUT into the pending trace; the persisted COPY - # is redacted under any privacy mode while the live input stays raw. + # Trace the exact aggregator INPUT (persisted copy redacted; live input raw). if self._pending_trace is not None: self._pending_trace["aggregator_input_messages"] = ( _redact_trace_messages([dict(m) for m in agg_messages]) @@ -1307,8 +1251,6 @@ class MoAChatCompletions: else agg_messages ) self._pending_trace["aggregator_label"] = _slot_label(aggregator) - # The aggregator is the acting model: call it through the same request path - # any model uses, with max_tokens passed through (None → model maximum). # stream=True returns the RAW token stream (consumer reassembles + retries); # the non-streaming path forwards no stream/stream_options/timeout. stream = bool(api_kwargs.get("stream")) @@ -1319,8 +1261,7 @@ class MoAChatCompletions: # The consumer's stream-read timeout must govern the aggregator stream. if api_kwargs.get("timeout") is not None: stream_kwargs["timeout"] = api_kwargs["timeout"] - # Pop the runtime's extra_body and merge with the caller's (caller wins) so the - # explicit kwarg never collides with **agg_runtime. + # Pop the runtime's extra_body so the explicit kwarg never collides with **agg_runtime. agg_extra_body = _merge_slot_extra_body( agg_runtime.pop("extra_body", None), api_kwargs.get("extra_body"), ) @@ -1336,8 +1277,7 @@ class MoAChatCompletions: **stream_kwargs, **agg_runtime, ) - # Non-streaming: capture the aggregator output inline. Streaming: the output - # lands as the turn's assistant message; the trace marks it streamed. + # Streaming output lands as the turn's assistant message; the trace marks it. if self._pending_trace is not None: self._pending_trace["aggregator_streamed"] = stream output = None @@ -1359,14 +1299,12 @@ class MoAChatCompletions: ref_messages: list[dict[str, Any]], reference_models: list[dict[str, Any]], ) -> tuple: - """Compute the turn-scoped reference cache key per the preset's fan-out cadence. + """Turn-scoped reference cache key per the preset's fan-out cadence. - "user_turn" (default): advisors run once per user turn — the signature hashes - only the prefix up to the LAST USER message, so later tool iterations are - cache HITs. "per_iteration": re-run whenever the advisory view changes. - "every_n:": iteration 1 of a turn, then every Nth; in-between iterations - reuse the last on-cadence guidance (key pinned to that run so the lookup is a - HIT: no advisor calls, no double accounting, no re-emit). + "user_turn" (default) hashes only the prefix up to the LAST USER message, so + later tool iterations are HITs. "per_iteration" re-runs whenever the advisory + view changes. "every_n:": iteration 1 of a turn, then every Nth; in-between + iterations return the pinned last on-cadence key (HIT: no calls, no re-emit). """ fanout_mode = str(preset.get("fanout") or "user_turn").strip().lower() every_n = 0 @@ -1423,10 +1361,11 @@ class MoAChatCompletions: aggregator_temperature: Any, cache_key: tuple, ) -> list[tuple[str, str, Any]]: - """Cache-MISS path of ``create``: run the advisors, account, trace and emit.""" - # No output caps by default (None → call_llm omits max_tokens). A preset MAY cap - # ADVISOR output (dominant MoA latency); the acting aggregator is never capped. - # None reference_timeout = inherit auxiliary.moa_reference.timeout via call_llm. + """Cache-MISS path of ``create``: run the advisors, account, trace and emit. + + A preset MAY cap ADVISOR output (dominant MoA latency); the acting aggregator + is never capped. None timeout = inherit auxiliary.moa_reference.timeout. + """ raw_reference_timeout = preset.get("reference_timeout") def _progress(done: int, total: int, label: str) -> None: @@ -1452,9 +1391,8 @@ class MoAChatCompletions: self._ref_cache_outputs = list(reference_outputs) # Fold advisor spend into accounting exactly once per turn. self._fold_pending_accounting(*_sum_reference_accounting(reference_outputs)) - # Stash the fan-out for trace persistence (aggregator input/label filled in - # later; output stitched in by consume_and_save_trace). Traces are persisted, - # so ANY active privacy mode redacts advisor text and per-advisor input/output. + # Stash the fan-out for trace persistence (aggregator parts filled in later). + # Traces are persisted, so ANY active privacy mode redacts them. privacy_mode = self._privacy_mode if privacy_mode: trace_refs = [ @@ -1480,8 +1418,7 @@ class MoAChatCompletions: logger.debug("MoA reference metrics render failed: %s", exc) self._last_reference_metrics = None - # Surface each reference's answer BEFORE the aggregator acts (once per turn). - # The cache keeps RAW text; redaction happens at each consuming surface. + # Surface each answer BEFORE the aggregator acts; the cache keeps RAW text. ref_count = len(reference_outputs) for idx, (label, text, _accounting) in enumerate(reference_outputs, start=1): self._emit( @@ -1512,8 +1449,7 @@ class MoAChatCompletions: ) -> str | None: """Render the reference block attached to the aggregator prompt (None = nothing).""" successful_outputs, failed_labels = _split_references(reference_outputs) - # 'full' privacy mode redacts advisor text reaching the AGGREGATOR too; 'display' - # leaves it raw. Applied to a per-call copy — the cache holds raw text. + # 'full' privacy mode redacts advisor text reaching the AGGREGATOR too. agg_refs = ( _redact_reference_outputs(successful_outputs) if self._privacy_mode == "full" @@ -1526,8 +1462,7 @@ class MoAChatCompletions: f"Aggregator/acting model: {_slot_label(aggregator)}\n" ) if reference_outputs and not successful_outputs: - # Every reference failed: the aggregator acts alone. Under the loud policy it - # still gets the sanitized unavailability notice; under silent, nothing. + # Every reference failed: the aggregator acts alone (loud policy → notice). logger.warning( "MoA: all %d reference(s) failed — acting aggregator-alone " "without reference guidance", @@ -1579,9 +1514,8 @@ class MoAChatCompletions: ref_messages = _reference_messages(messages) cache_key = self._fanout_cache_key(preset, ref_messages, reference_models) if cache_key == self._ref_cache_key and self._ref_cache_outputs: - # Cache HIT: references already ran and were accounted this turn. Deposit - # nothing, but do NOT zero pending totals (a late interrupted reference may - # have deposited real spend). No trace either — a repeat iteration is not a turn. + # HIT: already ran and accounted. Do NOT zero pending totals (a late + # interrupted reference may have deposited) and no trace (not a new turn). reference_outputs = list(self._ref_cache_outputs) self._pending_trace = None else: @@ -1608,38 +1542,31 @@ class MoAChatCompletions: class MoAClient: + """OpenAI-client-shaped wrapper: ``client.chat.completions`` is a ``MoAChatCompletions``. + + The accounting/trace surface (``consume_reference_usage``, ``last_aggregator_slot``, + ``consume_and_save_trace``, ``last_reference_metrics``) is delegated to the facade. + """ + + _DELEGATED = ( + "consume_reference_usage", "last_aggregator_slot", + "consume_and_save_trace", "last_reference_metrics", + ) + def __init__(self, preset_name: str, reference_callback: Any = None, agent: Any = None): self.chat = type("_MoAChat", (), {})() self.chat.completions = MoAChatCompletions( preset_name, reference_callback=reference_callback, agent=agent, ) - def consume_reference_usage(self) -> Any: - """Pop pending reference-fan-out usage + cost from the completions facade.""" - return self.chat.completions.consume_reference_usage() - - @property - def last_aggregator_slot(self) -> Any: - """Resolved aggregator slot from the most recent create(), or None.""" - return getattr(self.chat.completions, "last_aggregator_slot", None) - - def consume_and_save_trace( - self, session_id: Any = None, aggregator_output_fallback: Any = None - ) -> None: - """Flush the pending full-turn MoA trace via the completions facade.""" - return self.chat.completions.consume_and_save_trace( - session_id, aggregator_output_fallback=aggregator_output_fallback - ) - - def last_reference_metrics(self) -> Any: - """Per-advisor metrics from the most recent fan-out, or None (read-only).""" - return self.chat.completions.last_reference_metrics() + def __getattr__(self, name: str) -> Any: + if name in MoAClient._DELEGATED: + return getattr(self.chat.completions, name) + raise AttributeError(name) -# Display-event relay table for build_moa_facade: event -> (primary kwarg, secondary -# kwarg or None, {tool_progress_callback kwarg: emit kwarg}). The callback signature is -# ``cb(event, label, text, None, **moa_*)``; "moa.progress" is rendered by frontends as -# a status-bar ``MOA: N/M refs done``. +# Relay table: event -> (primary kwarg, secondary kwarg or None, {cb kwarg: emit kwarg}). +# The callback signature is ``cb(event, label, text, None, **moa_*)``. _RELAY_EVENTS: dict[str, tuple[str, str | None, dict[str, str]]] = { "moa.reference": ("label", "text", {"moa_index": "index", "moa_count": "count"}), "moa.progress": ("label", None, {"moa_refs_done": "refs_done", "moa_refs_total": "refs_total"}), @@ -1653,13 +1580,9 @@ _RELAY_EVENTS: dict[str, tuple[str, str | None, dict[str, str]]] = { def build_moa_facade(agent, preset_name: Any = None) -> MoAClient: - """Build the MoA facade client for ``agent``, wiring the reference relay. - - Single construction point for ``MoAClient`` (agent_init, fallback restore, - transport recovery, switch_model): a bare ``MoAClient(preset)`` would drop the - ``reference_callback`` relay and silence display events for the session. - The relay reads ``agent.tool_progress_callback`` at emit time. - """ + """Single construction point for ``MoAClient``: a bare ``MoAClient(preset)`` would + drop the ``reference_callback`` relay and silence display events for the session. + The relay reads ``agent.tool_progress_callback`` at emit time.""" def _moa_reference_relay(event: str, **kwargs: Any) -> None: cb = getattr(agent, "tool_progress_callback", None) spec = _RELAY_EVENTS.get(event) diff --git a/agent/moa_trace.py b/agent/moa_trace.py index 964d72de41..b4ccade104 100644 --- a/agent/moa_trace.py +++ b/agent/moa_trace.py @@ -1,12 +1,9 @@ """Full MoA turn trace persistence (opt-in via config ``moa.save_traces``). -When enabled, every Mixture-of-Agents turn that runs the reference fan-out (a -cache MISS in ``MoAChatCompletions.create``) appends one JSON line to -``/moa-traces/.jsonl``: the exact messages each -reference received, each reference's full output, and the exact aggregator input -plus its output when available — what every model saw, said, and cost. - -Side-channel only: never enters the ``messages`` table, history or replay +Every MoA turn that runs the reference fan-out (a cache MISS in +``MoAChatCompletions.create``) appends one JSON line to +``/moa-traces/.jsonl``: what every model saw, said, and +cost. Side-channel only: never enters the ``messages`` table, history or replay (references are advisory side-calls whose rows would corrupt role alternation). Off by default; when off the only overhead is the config read. """ @@ -32,7 +29,7 @@ def _traces_enabled_and_dir() -> Optional[Path]: from hermes_cli.config import load_config moa_cfg = (load_config() or {}).get("moa") or {} - except Exception: # pragma: no cover - defensive: never break a turn over tracing + except Exception: # pragma: no cover - never break a turn over tracing return None if not moa_cfg.get("save_traces"): return None @@ -54,8 +51,8 @@ _COST_FIELDS = ("cost_usd", "cost_status", "cost_source") def _slot_trace(acct: Any, label: str) -> dict[str, Any]: - """Render one reference's _RefAccounting into a full trace dict, including - the FULL input messages and output (not the truncated display preview).""" + """One reference's _RefAccounting as a full trace dict, including the FULL + input messages and output (not the truncated display preview).""" usage = getattr(acct, "usage", None) return { "label": label, diff --git a/agent/moonshot_schema.py b/agent/moonshot_schema.py index ab7319f2cc..079280cfbb 100644 --- a/agent/moonshot_schema.py +++ b/agent/moonshot_schema.py @@ -1,16 +1,11 @@ -"""Helpers for translating OpenAI-style tool schemas to Moonshot's schema subset. +"""Translate OpenAI-style tool schemas to Moonshot's (Kimi) stricter JSON Schema subset. -Moonshot (Kimi) accepts a stricter subset of JSON Schema than OpenAI tool -calling; violations fail with HTTP 400 "tools.function.parameters is not a -valid moonshot flavored json schema". Rules applied here: - -1. Every property schema must carry a ``type`` (JSON Schema allows omitting it). -2. With ``anyOf``, ``type`` belongs on the children, never the parent. -3. Enum arrays under scalar types may not contain null / empty string. -4. Every object schema must carry a ``required`` array, even an empty one. - -The ``#/definitions/...`` → ``#/$defs/...`` rewrite for draft-07 refs lives in -``tools/mcp_tool._normalize_mcp_input_schema`` so it applies to all providers. +Violations fail with HTTP 400 "tools.function.parameters is not a valid moonshot +flavored json schema". Rules: (1) every property schema carries a ``type``; +(2) with ``anyOf``, ``type`` belongs on the children, never the parent; (3) enum +arrays under scalar types may not contain null / empty string; (4) every object +schema carries a ``required`` array, even an empty one. The ``#/definitions/`` → +``#/$defs/`` rewrite lives in ``tools/mcp_tool`` so it applies to all providers. """ from __future__ import annotations @@ -26,6 +21,10 @@ _SCHEMA_LIST_KEYS = frozenset({"anyOf", "oneOf", "allOf", "prefixItems"}) # Values are a single nested schema (additionalProperties may also be a bool). _SCHEMA_NODE_KEYS = frozenset({"items", "contains", "not", "additionalProperties", "propertyNames"}) +_SCALAR_TYPES = frozenset({"string", "integer", "number", "boolean"}) +# bool before int: bool is an int subclass. +_ENUM_SAMPLE_TYPES = ((bool, "boolean"), (int, "integer"), (float, "number")) + def _empty_object_schema() -> Dict[str, Any]: return {"type": "object", "properties": {}, "required": []} @@ -42,9 +41,9 @@ def _repair_schema(node: Any) -> Any: for key, value in node.items(): if key in _SCHEMA_MAP_KEYS and isinstance(value, dict): repaired[key] = {sub_key: _repair_schema(sub_val) for sub_key, sub_val in value.items()} - elif key in _SCHEMA_LIST_KEYS and isinstance(value, list): - repaired[key] = [_repair_schema(v) for v in value] - elif key in _SCHEMA_NODE_KEYS and isinstance(value, dict): + elif (key in _SCHEMA_LIST_KEYS and isinstance(value, list)) or ( + key in _SCHEMA_NODE_KEYS and isinstance(value, dict) + ): repaired[key] = _repair_schema(value) else: repaired[key] = value @@ -60,9 +59,7 @@ def _repair_schema(node: Any) -> Any: if len(non_null) > 1: repaired["anyOf"] = non_null return repaired - merge = {k: v for k, v in repaired.items() if k != "anyOf"} - merge.update(non_null[0]) - repaired = merge + repaired = {**{k: v for k, v in repaired.items() if k != "anyOf"}, **non_null[0]} # Moonshot also rejects the non-standard ``nullable`` keyword. repaired.pop("nullable", None) @@ -73,7 +70,7 @@ def _repair_schema(node: Any) -> Any: repaired = _fill_missing_type(repaired) # Rule 3: drop null/"" enum values under scalar types; drop an emptied enum. - if isinstance(repaired.get("enum"), list) and repaired.get("type") in {"string", "integer", "number", "boolean"}: + if isinstance(repaired.get("enum"), list) and repaired.get("type") in _SCALAR_TYPES: cleaned = [v for v in repaired["enum"] if v is not None and v != ""] if cleaned: repaired["enum"] = cleaned @@ -88,9 +85,8 @@ def _repair_schema(node: Any) -> Any: def _ensure_required_array(node: Dict[str, Any]) -> Dict[str, Any]: - """Guarantee an object schema carries a ``required`` list (Moonshot 400s - otherwise), pruning names that don't exist in ``properties`` — Moonshot - also rejects dangling names. Mutates and returns ``node``.""" + """Guarantee an object schema carries a ``required`` list, pruning names not in + ``properties`` (Moonshot also rejects dangling names). Mutates and returns ``node``.""" props = node.get("properties") req = node.get("required") if isinstance(req, list): @@ -102,7 +98,7 @@ def _ensure_required_array(node: Dict[str, Any]) -> Dict[str, Any]: def _fill_missing_type(node: Dict[str, Any]) -> Dict[str, Any]: - """Infer a reasonable ``type`` if this schema node has none. + """Infer a ``type`` if this schema node has none. A type list collapses to its first concrete member; otherwise ``properties``/``required``/``additionalProperties`` → object, @@ -111,10 +107,7 @@ def _fill_missing_type(node: Dict[str, Any]) -> Dict[str, Any]: """ node_type = node.get("type") if isinstance(node_type, list): - concrete = next( - (t for t in node_type if isinstance(t, str) and t not in {"", "null"}), - "string", - ) + concrete = next((t for t in node_type if isinstance(t, str) and t not in {"", "null"}), "string") return {**node, "type": concrete} if "type" in node and node_type not in {None, ""}: return node @@ -124,24 +117,20 @@ def _fill_missing_type(node: Dict[str, Any]) -> Dict[str, Any]: elif "items" in node or "prefixItems" in node: inferred = "array" elif isinstance(node.get("enum"), list) and node["enum"]: - sample = node["enum"][0] # bool before int: bool is an int subclass - scalar_types = ((bool, "boolean"), (int, "integer"), (float, "number")) - inferred = next((t for cls, t in scalar_types if isinstance(sample, cls)), "string") + sample = node["enum"][0] + inferred = next((t for cls, t in _ENUM_SAMPLE_TYPES if isinstance(sample, cls)), "string") else: inferred = "string" - return {**node, "type": inferred} def sanitize_moonshot_tool_parameters(parameters: Any) -> Dict[str, Any]: - """Return a deep-copied, Moonshot-compatible object schema; input is not mutated.""" + """Deep-copied, Moonshot-compatible object schema; input is not mutated.""" if not isinstance(parameters, dict): return _empty_object_schema() - repaired = _repair_schema(copy.deepcopy(parameters)) if not isinstance(repaired, dict): return _empty_object_schema() - # Top-level must be an object schema. repaired["type"] = "object" repaired.setdefault("properties", {}) @@ -155,22 +144,17 @@ def sanitize_moonshot_tools(tools: List[Dict[str, Any]]) -> List[Dict[str, Any]] """ if not tools: return tools - sanitized: List[Dict[str, Any]] = [] any_change = False for tool in tools: fn = tool.get("function") if isinstance(tool, dict) else None - if not isinstance(fn, dict): - sanitized.append(tool) - continue - params = fn.get("parameters") - repaired = sanitize_moonshot_tool_parameters(params) - if repaired is not params: - any_change = True - sanitized.append({**tool, "function": {**fn, "parameters": repaired}}) - else: - sanitized.append(tool) - + if isinstance(fn, dict): + params = fn.get("parameters") + repaired = sanitize_moonshot_tool_parameters(params) + if repaired is not params: + any_change = True + tool = {**tool, "function": {**fn, "parameters": repaired}} + sanitized.append(tool) return sanitized if any_change else tools diff --git a/agent/nous_rate_guard.py b/agent/nous_rate_guard.py index b5acb048b3..2e3e95f214 100644 --- a/agent/nous_rate_guard.py +++ b/agent/nous_rate_guard.py @@ -1,9 +1,9 @@ """Cross-session rate limit guard for Nous Portal. -Writes rate limit state to a shared file so all sessions (CLI, gateway, -cron, auxiliary) can check whether Nous Portal is currently rate-limited -before making requests. Without it each 429 fans out into up to 9 calls per -turn (3 SDK retries x 3 Hermes retries), all counted against RPH. +Writes rate limit state to a shared file so all sessions (CLI, gateway, cron, +auxiliary) can check whether Nous Portal is currently rate-limited before making +requests. Without it each 429 fans out into up to 9 calls per turn (3 SDK +retries x 3 Hermes retries), all counted against RPH. """ from __future__ import annotations @@ -12,10 +12,9 @@ import contextlib import json import logging import os -import tempfile import time from typing import Any, Mapping, Optional -from utils import atomic_replace +from utils import atomic_write_text from agent.rate_limit_tracker import ( _BUCKET_TAGS, _fmt_seconds, @@ -35,7 +34,7 @@ format_remaining = _fmt_seconds def _state_path() -> str: - """Return the path to the Nous rate limit state file.""" + """Path to the Nous rate limit state file.""" try: from hermes_constants import get_hermes_home base = get_hermes_home() @@ -47,11 +46,7 @@ def _state_path() -> str: def _parse_reset_seconds(headers: Optional[Mapping[str, str]]) -> Optional[float]: """Best reset estimate (seconds from now) from hourly, per-minute, then retry-after headers.""" lowered = lower_headers(headers) - for key in ( - "x-ratelimit-reset-requests-1h", - "x-ratelimit-reset-requests", - "retry-after", - ): + for key in ("x-ratelimit-reset-requests-1h", "x-ratelimit-reset-requests", "retry-after"): val = _safe_float(lowered.get(key), 0.0) if val > 0: return val @@ -71,44 +66,20 @@ def record_nous_rate_limit( """ now = time.time() reset_at = None - header_seconds = _parse_reset_seconds(headers) if header_seconds is not None: reset_at = now + header_seconds - if reset_at is None and isinstance(error_context, dict): ctx_reset = error_context.get("reset_at") if isinstance(ctx_reset, (int, float)) and ctx_reset > now: reset_at = float(ctx_reset) - if reset_at is None: reset_at = now + default_cooldown - path = _state_path() + state = {"reset_at": reset_at, "recorded_at": now, "reset_seconds": reset_at - now} try: - state_dir = os.path.dirname(path) - os.makedirs(state_dir, exist_ok=True) - - state = { - "reset_at": reset_at, - "recorded_at": now, - "reset_seconds": reset_at - now, - } - - fd, tmp_path = tempfile.mkstemp(dir=state_dir, suffix=".tmp") - try: - with os.fdopen(fd, "w", encoding="utf-8") as f: - json.dump(state, f) - atomic_replace(tmp_path, path) - except Exception: - with contextlib.suppress(OSError): - os.unlink(tmp_path) - raise - - logger.info( - "Nous rate limit recorded: resets in %.0fs (at %.0f)", - reset_at - now, reset_at, - ) + atomic_write_text(_state_path(), json.dumps(state)) + logger.info("Nous rate limit recorded: resets in %.0fs (at %.0f)", reset_at - now, reset_at) except Exception as exc: logger.debug("Failed to write Nous rate limit state: %s", exc) @@ -170,11 +141,10 @@ def is_genuine_nous_rate_limit( def _parse_buckets_from_headers( headers: Optional[Mapping[str, str]], ) -> dict[str, tuple[Optional[int], Optional[float]]]: - """Extract (remaining, reset_seconds) per bucket from x-ratelimit-* headers ({} if none).""" + """(remaining, reset_seconds) per bucket from x-ratelimit-* headers ({} if none).""" lowered = lower_headers(headers) if not has_rate_limit_headers(lowered): return {} - result: dict[str, tuple[Optional[int], Optional[float]]] = {} for _attr, tag in _BUCKET_TAGS: remaining = _safe_int(lowered.get(f"x-ratelimit-remaining-{tag}"), None) diff --git a/agent/plugin_llm.py b/agent/plugin_llm.py index c0ea2d8528..14ab45b15b 100644 --- a/agent/plugin_llm.py +++ b/agent/plugin_llm.py @@ -1,39 +1,12 @@ -""" -Plugin LLM facade — host-owned LLM access for trusted plugins. -============================================================== +"""Plugin LLM facade — host-owned LLM access for trusted plugins (``ctx.llm``). -Plugins that need their own out-of-band model call (rewrite a tool error, -translate inbound text, summarise a paste, score a scheduled job) get -``ctx.llm`` on :class:`~hermes_cli.plugins.PluginContext`: ``complete`` / -``complete_structured`` (text + image inputs, JSON schema validation) and their -async siblings ``acomplete`` / ``acomplete_structured``. - -Provider/model/agent_id/profile are explicit keyword arguments mirroring the -host config shape (``model.provider`` + ``model.model``) — no embedded slugs. -The host owns routing, auth, timeouts, and fallback; the plugin never sees raw -tokens or keys. Every override knob is gated by per-plugin trust flags:: - - plugins: - entries: - my-plugin: - llm: - allow_provider_override: true - allow_model_override: true - allowed_providers: [openrouter, anthropic] # optional - allowed_models: [openai/gpt-4o-mini] # optional - allow_agent_id_override: false - allow_profile_override: false - allow_task_override: false # borrow the host's built-in aux tasks - -The gate is fail-closed: a missing config block means "no overrides". - -``task=`` routes a call through a plugin-registered auxiliary model slot -(``ctx.register_auxiliary_task``). A plugin may always name a slot it -registered itself; ``allow_task_override`` additionally lets it use the host's -*built-in* auxiliary tasks. A foreign or unknown key is rejected loudly -(error + logged warning), never silently downgraded to the main model. - -Backed by :func:`agent.auxiliary_client.call_llm`. +``complete`` / ``complete_structured`` (text + image inputs, JSON schema validation) +and their async siblings. Provider/model/agent_id/profile are explicit keyword +arguments mirroring the host config shape; the host owns routing, auth, timeouts +and fallback, so the plugin never sees raw tokens or keys. Every override knob is +gated by the per-plugin ``plugins.entries..llm.allow_*_override`` trust flags +(fail-closed: a missing block means "no overrides"). Backed by +:func:`agent.auxiliary_client.call_llm`. """ from __future__ import annotations @@ -48,9 +21,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional, Sequence, Uni logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Public dataclasses -# --------------------------------------------------------------------------- +# -- public dataclasses ------------------------------------------------------- @dataclass @@ -63,7 +34,7 @@ class PluginLlmTextInput: @dataclass class PluginLlmImageInput: - """Image block. Provide ``data`` (raw bytes) or ``url`` (http(s)/data: URL). + """Image block: ``data`` (raw bytes) or ``url`` (http(s)/data: URL). ``mime_type`` is required for non-PNG bytes to render across providers.""" data: Optional[bytes] = None @@ -74,16 +45,12 @@ class PluginLlmImageInput: PluginLlmInput = Union[PluginLlmTextInput, PluginLlmImageInput, Dict[str, Any]] -"""A single structured input block: one of the dataclasses above or a plain dict -of the same shape (``{"type": "text", "text": ...}`` / -``{"type": "image", "data": , "mime_type": ..., "file_name": ...}`` / -``{"type": "image", "url": ...}``).""" +"""One structured input block: a dataclass above or a plain dict of the same shape.""" @dataclass class PluginLlmUsage: - """Token + cost usage. All fields optional — providers differ on what they - return. ``cost_usd`` is the host's best estimate.""" + """Token + cost usage; every field optional. ``cost_usd`` is the host's best estimate.""" input_tokens: int = 0 output_tokens: int = 0 @@ -109,9 +76,8 @@ class PluginLlmCompleteResult: class PluginLlmStructuredResult: """Result of :meth:`PluginLlm.complete_structured`. - ``parsed`` is set only when JSON output was requested (``json_mode`` or - ``json_schema``) AND the response was valid JSON; ``content_type`` is then - ``"json"``, otherwise ``"text"``.""" + ``parsed`` is set only when JSON output was requested AND the response was + valid JSON; ``content_type`` is then ``"json"``, otherwise ``"text"``.""" text: str provider: str @@ -123,9 +89,7 @@ class PluginLlmStructuredResult: audit: Dict[str, Any] = field(default_factory=dict) -# --------------------------------------------------------------------------- -# Trust gate -# --------------------------------------------------------------------------- +# -- trust gate --------------------------------------------------------------- @dataclass(frozen=True) @@ -157,7 +121,6 @@ _OVERRIDE_FLAGS = ( def _normalize_ref(raw: str) -> str: - """Lower-case + strip whitespace. Used for allowlist matching.""" return (raw or "").strip().lower() @@ -167,16 +130,12 @@ def _coerce_allowlist(raw: Any) -> tuple[Optional[frozenset], bool]: if not isinstance(raw, list): return None, False normalized = [_normalize_ref(item) for item in raw if isinstance(item, str)] - allow_any = "*" in normalized - cleaned = {item for item in normalized if item and item != "*"} - return frozenset(cleaned), allow_any + return frozenset(item for item in normalized if item and item != "*"), "*" in normalized def _resolve_trust_policy(plugin_id: str) -> _TrustPolicy: - """Read ``plugins.entries..llm`` from config.yaml. - - Missing config → fully restrictive policy. Resolved per call (not cached) - so config edits take effect without restarting the agent.""" + """Read ``plugins.entries..llm`` from config.yaml (missing → fully + restrictive). Resolved per call so config edits apply without a restart.""" if not plugin_id: return _TrustPolicy(plugin_id="") @@ -209,7 +168,6 @@ class PluginLlmTrustError(PermissionError): def _denied(plugin_id: str, what: str, flag: str) -> PluginLlmTrustError: - """Uniform "flag not set" trust error.""" return PluginLlmTrustError( f"Plugin {plugin_id!r} cannot {what} " f"(set plugins.entries.{plugin_id}.llm.{flag} to true to allow)." @@ -217,8 +175,7 @@ def _denied(plugin_id: str, what: str, flag: str) -> PluginLlmTrustError: def _gate_ref_override(policy: _TrustPolicy, kind: str, requested: str) -> str: - """Gate a ``provider`` / ``model`` override: trust flag, then optional - allowlist. Returns the stripped value or raises.""" + """Gate a ``provider`` / ``model`` override: trust flag, then optional allowlist.""" if not getattr(policy, f"allow_{kind}_override"): raise _denied(policy.plugin_id, f"override the {kind}", f"allow_{kind}_override") allowed = getattr(policy, f"allowed_{kind}s") @@ -242,9 +199,8 @@ def _check_overrides( policy: _TrustPolicy, *, requested_provider: Optional[str], requested_model: Optional[str], requested_agent_id: Optional[str], requested_profile: Optional[str], ) -> tuple[Optional[str], Optional[str], Optional[str], Optional[str]]: - """Apply the trust gate; each override is gated independently, in the order - provider, model, agent_id, profile. Returns ``(provider, model, agent_id, - profile)`` (agent_id unstripped) or raises :class:`PluginLlmTrustError`.""" + """Gate each override independently, in the order provider, model, agent_id, + profile. Returns ``(provider, model, agent_id, profile)`` (agent_id unstripped).""" final_provider = _gate_ref_override(policy, "provider", requested_provider) if requested_provider else None final_model = _gate_ref_override(policy, "model", requested_model) if requested_model else None for kind, requested in (("agent_id", requested_agent_id), ("profile", requested_profile)): @@ -255,12 +211,12 @@ def _check_overrides( def _resolve_task_ownership(plugin_id: str) -> tuple[frozenset, frozenset]: - """Return ``(owned_keys, builtin_keys)`` for the task trust gate. + """``(owned_keys, builtin_keys)`` for the task trust gate. - Imports are lazy (circular import at plugin discovery). An unreadable - registry yields empty sets, failing the gate closed. Ownership matches on - the canonical id ``ctx.llm`` is bound to (``manifest.key or manifest.name``), - which is what ``register_auxiliary_task`` stores as the entry's ``plugin``.""" + Imports are lazy (circular import at plugin discovery); an unreadable registry + yields empty sets, failing the gate closed. Ownership matches on the canonical id + ``ctx.llm`` is bound to (``manifest.key or manifest.name``), which is what + ``register_auxiliary_task`` stores as the entry's ``plugin``.""" owned: set = set() builtin: set = set() try: @@ -289,15 +245,12 @@ def _check_task( ) -> Optional[str]: """Validate a plugin's requested auxiliary ``task`` key. - * unset / ``""`` / ``"auto"`` → ``None`` (main-model path). - * a key the plugin registered itself → allowed. - * a built-in key → allowed only with ``allow_task_override``. - * anything else → raises + logs a warning. Never silently downgraded to - ``auto``: that would mask the misconfiguration and could route to a main - model the user steered elsewhere on purpose.""" - if not requested_task: - return None - task = requested_task.strip() + unset / ``""`` / ``"auto"`` → ``None`` (main-model path); a key the plugin + registered itself → allowed; a built-in key → only with ``allow_task_override``; + anything else raises + logs. Never silently downgraded to ``auto``: that would + mask the misconfiguration and could route to a main model the user steered + elsewhere on purpose.""" + task = (requested_task or "").strip() if not task or task.lower() == "auto": return None @@ -323,9 +276,7 @@ def _check_task( ) -# --------------------------------------------------------------------------- -# Input normalization -# --------------------------------------------------------------------------- +# -- input normalization ------------------------------------------------------ def _normalize_input_block(block: PluginLlmInput) -> Dict[str, Any]: @@ -378,10 +329,9 @@ def _build_structured_messages( schema_name: Optional[str], system_prompt: Optional[str], ) -> List[Dict[str, Any]]: - """Build OpenAI-style messages for a structured call: optional system - message (prompt + JSON-only directive), then a user message whose first - text part is the instructions (+ schema name / JSON schema) followed by - the input blocks.""" + """OpenAI-style messages for a structured call: optional system message (prompt + + JSON-only directive), then a user message whose first text part is the + instructions (+ schema name / JSON schema) followed by the input blocks.""" messages: List[Dict[str, Any]] = [] sys_parts: List[str] = [system_prompt.strip()] if system_prompt else [] if json_mode or json_schema is not None: @@ -402,7 +352,6 @@ def _build_structured_messages( schema_text = str(json_schema) header = f"{header}\n\nJSON schema:\n{schema_text}" user_parts: List[Dict[str, Any]] = [{"type": "text", "text": header}] - for block in inputs: norm = _normalize_input_block(block) # always "text" or "image" user_parts.append({"type": "text", "text": norm["text"]} if norm["type"] == "text" else _image_part(norm)) @@ -410,16 +359,14 @@ def _build_structured_messages( return messages -# --------------------------------------------------------------------------- -# JSON parsing -# --------------------------------------------------------------------------- +# -- JSON parsing / response extraction -------------------------------------- _FENCE_RE = re.compile(r"```(?:json)?\s*(.+?)```", re.DOTALL | re.IGNORECASE) def _strip_code_fences(text: str) -> str: - """Return the first fenced code block's body, or the stripped text when unfenced.""" + """The first fenced code block's body, or the stripped text when unfenced.""" match = _FENCE_RE.search(text) return match.group(1).strip() if match else text.strip() @@ -427,10 +374,9 @@ def _strip_code_fences(text: str) -> str: def _parse_structured_text( *, text: str, json_mode: bool, json_schema: Optional[Any] ) -> tuple[Optional[Any], str]: - """Return ``(parsed, content_type)``: ``"json"`` when parsing (and schema - validation, if a schema was given) succeeded, ``"text"`` otherwise. - Schema violations raise ``ValueError``; a missing ``jsonschema`` package - skips validation with a debug log.""" + """``(parsed, content_type)``: ``"json"`` when parsing (and schema validation, if + given) succeeded, ``"text"`` otherwise. Schema violations raise ``ValueError``; + a missing ``jsonschema`` package skips validation with a debug log.""" if not (json_mode or json_schema is not None) or not text: return None, "text" @@ -453,15 +399,10 @@ def _parse_structured_text( return parsed, "json" -# --------------------------------------------------------------------------- -# Response extraction -# --------------------------------------------------------------------------- - - def _extract_usage(response: Any) -> PluginLlmUsage: - """Pull token usage out of an OpenAI-shaped response, tolerating provider - naming differences (Anthropic via the aux adapter: ``prompt_tokens`` / - ``completion_tokens``; direct OpenAI adds ``cache_read_input_tokens``).""" + """Token usage from an OpenAI-shaped response, tolerating provider naming + (``prompt_tokens``/``completion_tokens`` vs ``input_tokens``/``output_tokens``, + ``cache_read_input_tokens`` vs ``cache_read_tokens``).""" usage = PluginLlmUsage() raw = getattr(response, "usage", None) if raw is None: @@ -485,7 +426,7 @@ def _extract_usage(response: Any) -> PluginLlmUsage: def _extract_text(response: Any) -> str: - """Pull the assistant text out of an OpenAI-shaped response object.""" + """Assistant text of an OpenAI-shaped response (string or text-part list content).""" try: content = getattr(response.choices[0].message, "content", None) if isinstance(content, str): @@ -518,12 +459,11 @@ def _resolve_attribution( response: Any, route_info: Optional[Dict[str, str]] = None, ) -> tuple[str, str]: - """Decide what to record as ``result.provider`` / ``result.model``. + """``(provider, model)`` to record on the result. - Provider: route selected by ``auxiliary_client`` > explicit override > - current main provider > ``"auto"``. Model: ``response.model`` (providers - return the canonical id that actually ran, e.g. ``gpt-4o-2024-08-06``) > - route > override > current main model > ``"default"``.""" + Provider: route selected by ``auxiliary_client`` > explicit override > current + main provider > ``"auto"``. Model: ``response.model`` (the canonical id that + actually ran) > route > override > current main model > ``"default"``.""" route_info = route_info or {} provider = route_info.get("provider") or provider_override or _main_config_value("_read_main_provider", "auto") response_model = getattr(response, "model", None) @@ -532,14 +472,12 @@ def _resolve_attribution( return provider, route_info.get("model") or model_override or _main_config_value("_read_main_model", "default") -# --------------------------------------------------------------------------- -# PluginLlm facade -# --------------------------------------------------------------------------- +# -- PluginLlm facade --------------------------------------------------------- def _json_response_format(*, json_mode: bool, json_schema: Optional[Any]) -> Optional[Dict[str, Any]]: - """``extra_body.response_format`` for the request; falls back to - ``json_object`` without a schema so schema-blind providers still get a hint.""" + """``extra_body.response_format``; falls back to ``json_object`` without a + schema so schema-blind providers still get a hint.""" if json_schema is not None: schema = {"name": "plugin_structured_output", "schema": json_schema, "strict": False} return {"response_format": {"type": "json_schema", "json_schema": schema}} @@ -552,8 +490,7 @@ def _structured_spec( name: str, instructions: str, input: Sequence[PluginLlmInput], system_prompt: Optional[str], json_mode: bool, json_schema: Optional[Any], schema_name: Optional[str], ) -> Dict[str, Any]: - """Argument check for the structured methods (runs before the trust gate); - returns the spec ``_gate`` / ``_finish`` consume.""" + """Argument check for the structured methods (runs before the trust gate).""" if not instructions or not instructions.strip(): raise ValueError(f"{name} requires non-empty instructions") if not input: @@ -568,13 +505,10 @@ class PluginLlm: """Host-owned LLM access for one trusted plugin. Constructed by :class:`hermes_cli.plugins.PluginContext` and exposed as - ``ctx.llm``; the constructor binds plugin identity for trust enforcement, - so plugins should not instantiate it directly. - - Every public method is ``_gate`` (trust checks → call kwargs) → - ``_invoke_*`` (host ``call_llm`` or injected caller) → ``_finish`` - (result + audit log); the sync/async and plain/structured variants differ - only in which pieces they pass through.""" + ``ctx.llm``; the constructor binds plugin identity for trust enforcement, so + plugins should not instantiate it directly. Every public method is ``_gate`` + (trust checks → call kwargs) → ``_invoke_*`` (host ``call_llm`` or injected + caller) → ``_finish`` (result + audit log).""" def __init__( self, @@ -589,7 +523,7 @@ class PluginLlm: self._sync_caller = sync_caller self._async_caller = async_caller - # -- public sync API ---------------------------------------------------- + # -- public API ----------------------------------------------------------- def complete( self, @@ -607,10 +541,9 @@ class PluginLlm: ) -> PluginLlmCompleteResult: """Run a host-owned chat completion against the user's active model. - ``messages`` is the standard OpenAI shape. ``provider``/``model``/ - ``agent_id``/``profile`` are each gated by - ``plugins.entries..llm.allow_*_override``. ``task`` routes through - a plugin-registered auxiliary slot (see :func:`_check_task`).""" + ``provider``/``model``/``agent_id``/``profile`` are each gated by + ``plugins.entries..llm.allow_*_override``. ``task`` routes through a + plugin-registered auxiliary slot (see :func:`_check_task`).""" agent, kw = self._gate(provider, model, agent_id, profile, task, messages, temperature, max_tokens, timeout) return self._finish("complete", agent, kw, self._invoke_sync(kw), purpose) @@ -637,14 +570,11 @@ class PluginLlm: ``input`` accepts text and image blocks. With ``json_mode=True`` or a ``json_schema`` the response is parsed (and validated when the optional - ``jsonschema`` package is installed) into ``result.parsed``. - ``task`` routes as in :meth:`complete`.""" + ``jsonschema`` package is installed) into ``result.parsed``.""" spec = _structured_spec("complete_structured", instructions, input, system_prompt, json_mode, json_schema, schema_name) agent, kw = self._gate(provider, model, agent_id, profile, task, None, temperature, max_tokens, timeout, spec) return self._finish("complete_structured", agent, kw, self._invoke_sync(kw), purpose, spec) - # -- public async API --------------------------------------------------- - async def acomplete( self, messages: List[Dict[str, Any]], @@ -687,7 +617,7 @@ class PluginLlm: agent, kw = self._gate(provider, model, agent_id, profile, task, None, temperature, max_tokens, timeout, spec) return self._finish("acomplete_structured", agent, kw, await self._invoke_async(kw), purpose, spec) - # -- shared core -------------------------------------------------------- + # -- shared core ---------------------------------------------------------- def _gate( self, @@ -702,13 +632,11 @@ class PluginLlm: timeout: Optional[float], spec: Optional[Dict[str, Any]] = None, ) -> tuple[Optional[str], Dict[str, Any]]: - """Run the trust gate (task first, then overrides), then — for a - structured ``spec`` — build messages/response_format (input-shape errors - surface only after trust passes). Returns the effective agent id - (result-only) and the call kwargs handed to ``_invoke_*`` / an injected - caller, in the documented order: messages, provider_override, - model_override, profile_override, temperature, max_tokens, timeout, - extra_body, task.""" + """Trust gate (task first, then overrides), then — for a structured ``spec`` — + build messages/response_format (input-shape errors surface only after trust + passes). Returns the effective agent id and the call kwargs, in the documented + order: messages, provider_override, model_override, profile_override, + temperature, max_tokens, timeout, extra_body, task.""" policy = self._policy_loader(self._plugin_id) eff_task = _check_task(policy, plugin_id=self._plugin_id, requested_task=task) eff_provider, eff_model, eff_agent, eff_profile = _check_overrides( @@ -761,13 +689,13 @@ class PluginLlm: logger.info(fmt + "tokens=%d", *log_args, usage.total_tokens) return cls(**fields, audit=audit) - # -- host invocation --------------------------------------------------- + # -- host invocation ------------------------------------------------------ @staticmethod def _host_kwargs(kw: Dict[str, Any]) -> tuple[Dict[str, Any], Optional[Dict[str, str]]]: - """Translate call kwargs into ``call_llm`` kwargs. The auth profile - rides in ``extra_body.metadata.auth_profile``; ``route_info`` is only - requested when routing through a task slot.""" + """Call kwargs → ``call_llm`` kwargs. The auth profile rides in + ``extra_body.metadata.auth_profile``; ``route_info`` is only requested when + routing through a task slot.""" merged_extra = dict(kw["extra_body"] or {}) if kw["profile_override"]: merged_extra.setdefault("metadata", {})["auth_profile"] = kw["profile_override"] @@ -793,10 +721,9 @@ class PluginLlm: return provider, model, response def _invoke_sync(self, kw: Dict[str, Any]) -> tuple[str, str, Any]: - """Invoke the host's ``call_llm`` (lazy import: circular deps at plugin - discovery) and return ``(provider, model, response)``. ``task`` is - already trust-checked; ``None`` keeps the main model. An injected - ``sync_caller`` replaces the whole path and receives the call kwargs.""" + """Host ``call_llm`` (lazy import: circular deps at plugin discovery) → + ``(provider, model, response)``. An injected ``sync_caller`` replaces the + whole path and receives the call kwargs.""" if self._sync_caller is not None: return self._sync_caller(**kw) from agent.auxiliary_client import call_llm @@ -812,11 +739,6 @@ class PluginLlm: return self._attributed(kw, await async_call_llm(**call_kw), route_info) -# --------------------------------------------------------------------------- -# Test helpers -# --------------------------------------------------------------------------- - - def make_plugin_llm_for_test( *, plugin_id: str, @@ -824,8 +746,8 @@ def make_plugin_llm_for_test( sync_caller: Optional[Callable[..., Any]] = None, async_caller: Optional[Callable[..., Awaitable[Any]]] = None, ) -> PluginLlm: - """:class:`PluginLlm` with an injected policy and caller (no config.yaml, - no provider). Not part of the public plugin API.""" + """:class:`PluginLlm` with an injected policy and caller (no config.yaml, no + provider). Not part of the public plugin API.""" return PluginLlm(plugin_id=plugin_id, policy_loader=lambda _pid: policy, sync_caller=sync_caller, async_caller=async_caller) diff --git a/agent/provider_base.py b/agent/provider_base.py index 470b1c0a6c..68ab819e28 100644 --- a/agent/provider_base.py +++ b/agent/provider_base.py @@ -1,12 +1,10 @@ """Shared base classes for the pluggable-backend provider ABCs. -Every tool-provider ABC (browser, TTS, image/video gen, transcription, web -search, terminal env) shares the same identity + ``hermes tools`` picker -surface; the previous per-ABC copies of these defaults were byte-identical. -Concrete ABCs subclass :class:`ProviderBase` (or :class:`CatalogProviderBase` -when the backend also exposes a model catalog and is available by default) and -add only their domain methods. Plugins keep subclassing the concrete ABC, so -``isinstance`` checks and abstract-method sets are unchanged. +Every tool-provider ABC shares the same identity + ``hermes tools`` picker surface. +Concrete ABCs subclass :class:`ProviderBase` (or :class:`CatalogProviderBase` when +the backend also exposes a model catalog and is available by default) and add only +their domain methods. Plugins keep subclassing the concrete ABC, so ``isinstance`` +checks and abstract-method sets are unchanged. """ from __future__ import annotations @@ -36,16 +34,10 @@ class ProviderBase(abc.ABC): """Provider row for the ``hermes tools`` picker. Shape: ``{"name", "badge", "tag", "env_vars": [{"key", "prompt", "url"}, ...]}`` - (browser providers may add ``"post_setup"``). Default: a minimal entry - derived from ``display_name`` with no env vars — override to expose API - key prompts and badges. + (browser providers may add ``"post_setup"``). Override to expose API key + prompts and badges. """ - return { - "name": self.display_name, - "badge": "", - "tag": "", - "env_vars": [], - } + return {"name": self.display_name, "badge": "", "tag": "", "env_vars": []} class CatalogProviderBase(ProviderBase): @@ -59,19 +51,16 @@ class CatalogProviderBase(ProviderBase): def is_available(self) -> bool: """True when this provider can service calls (API key present, SDK importable). - Default True. Must NOT raise and must NOT make network calls — the picker - and ``hermes setup`` call it on every paint. + Must NOT raise and must NOT make network calls — the picker and + ``hermes setup`` call it on every paint. """ return True def list_models(self) -> List[Dict[str, Any]]: - """Model catalog entries (``{"id": ..., "display": ...}`` plus optional - provider-specific keys). Default: empty (no user-selectable models).""" + """Model catalog entries (``{"id": ..., "display": ...}`` + provider-specific keys).""" return [] def default_model(self) -> Optional[str]: """Id of the first catalog entry, or None when the catalog is empty.""" models = self.list_models() - if models: - return models[0].get("id") - return None + return models[0].get("id") if models else None diff --git a/agent/provider_media.py b/agent/provider_media.py index ba69a003fd..d3dfffe553 100644 --- a/agent/provider_media.py +++ b/agent/provider_media.py @@ -1,10 +1,10 @@ -"""Shared ``$HERMES_HOME/cache//`` materialisation helpers for the -image/video generation provider ABCs. +"""``$HERMES_HOME/cache//`` materialisation helpers for the image/video +generation provider ABCs. -Several backends (xAI, OpenAI, DeepInfra, FAL) return *ephemeral* delivery URLs -that expire before a downstream consumer (Telegram ``send_photo``, browser -fetch) can resolve them, so providers materialise the bytes locally at -tool-completion time. Filenames are ``__.``. +Several backends return *ephemeral* delivery URLs that expire before a downstream +consumer (Telegram ``send_photo``, browser fetch) can resolve them, so providers +materialise the bytes locally at tool-completion time. Filenames are +``__.``. """ from __future__ import annotations @@ -59,12 +59,11 @@ def save_url( ) -> Path: """Stream-download *url* into the cache with a size cap. - The extension comes from the response ``Content-Type`` (a small explicit - table — never inherit a type that points at HTML/JSON from a degenerate - response), then the URL suffix (some CDNs return - ``application/octet-stream``), then *default_extension*. Raises on any - network / HTTP / oversize / empty error so callers can fall back to the bare - URL; a partial file is never left behind. + The extension comes from the response ``Content-Type`` (an explicit table — + never inherit a type pointing at HTML/JSON from a degenerate response), then + the URL suffix (some CDNs return ``application/octet-stream``), then + *default_extension*. Raises on any network / HTTP / oversize / empty error so + callers can fall back to the bare URL; a partial file is never left behind. """ import requests @@ -75,15 +74,11 @@ def save_url( extension = content_types.get(content_type) if extension is None: url_path = url.split("?", 1)[0].lower() - for ext in url_extensions: - if url_path.endswith(f".{ext}"): - extension = "jpg" if ext == "jpeg" else ext - break - if extension is None: - extension = default_extension - + extension = next( + ("jpg" if ext == "jpeg" else ext for ext in url_extensions if url_path.endswith(f".{ext}")), + default_extension, + ) path = cache_path(kind, prefix, extension) - bytes_written = 0 with path.open("wb") as fh: for chunk in response.iter_content(chunk_size=chunk_size): diff --git a/agent/provider_projection.py b/agent/provider_projection.py index b3788941e5..dc56b98530 100644 --- a/agent/provider_projection.py +++ b/agent/provider_projection.py @@ -1,23 +1,13 @@ """Fold an agent-as-provider's own activity back into Hermes' turn state. -Some providers are *agents* (an ACP CLI behind a client shim; the codex -app-server takes an analogous path in ``agent/codex_runtime.py``): they run -their own tools inside their own session, so by the time Hermes sees the -response that work is done. Those calls must never come back as pending -``tool_calls`` (Hermes would re-run finished work), but two subsystems go blind -if they are merely summarised into ``reasoning``: - -* the **self-improvement loop**, which replays ``messages`` to distil memories - and skills; -* the **skill-review nudge**, whose ``_iters_since_skill`` counter only moves on - Hermes tool iterations. - -So the client hands both back on the completion object — -``hermes_projected_messages`` (completed ``assistant(tool_calls=[…])`` + -``tool(result)`` rows) and ``hermes_provider_tool_iterations`` — and this helper -applies them. Ordinary OpenAI-compatible clients set neither and are unaffected. -The splice is append-only through ``append_message`` so rows carry a timestamp -and persist like any other live-transcript append. +Agent providers (ACP CLI shims, the codex app-server) run their own tools, so that +work must never come back as pending ``tool_calls`` (Hermes would re-run it) — but +the self-improvement loop (replays ``messages``) and the skill-review nudge +(``_iters_since_skill`` counter) go blind if it is merely summarised into +``reasoning``. The client hands back ``hermes_projected_messages`` (completed +assistant/tool rows) and ``hermes_provider_tool_iterations`` on the completion +object; this helper applies them append-only via ``append_message`` (timestamped, +persisted). Ordinary OpenAI-compatible clients set neither and are unaffected. """ from __future__ import annotations @@ -46,9 +36,7 @@ def splice_provider_projection( append_message(messages, row) if rows: logger.debug( - "spliced %d provider-projected transcript row(s) from %s", - len(rows), - getattr(agent, "provider", "?"), + "spliced %d provider-projected transcript row(s) from %s", len(rows), getattr(agent, "provider", "?"), ) try: diff --git a/agent/provider_registry.py b/agent/provider_registry.py index f8c00ce0a4..27b859d820 100644 --- a/agent/provider_registry.py +++ b/agent/provider_registry.py @@ -1,14 +1,12 @@ """Shared engine behind the ``agent.*_registry`` provider registries. -Every pluggable-backend registry (browser, TTS, image/video gen, transcription, -web search, terminal env) has the same shape: a global name->provider map plus -per-profile *scoped* maps (multiplexed gateways), a lock, registration with -re-registration logging, and the snapshot/restore pair that -:mod:`hermes_cli.plugins` uses to unwind a plugin's registrations. Each -``*_registry`` module instantiates one :class:`ProviderRegistry` and re-exports -its bound methods under the historical module-level names via -:meth:`ProviderRegistry.export`, so call sites, ``patch("agent.x_registry.get_provider")`` -targets, and the ``_providers`` / ``_scoped_providers`` / ``_lock`` test hooks are unchanged. +Every pluggable-backend registry has the same shape: a global name->provider map +plus per-profile *scoped* maps (multiplexed gateways), a lock, registration with +re-registration logging, and the snapshot/restore pair :mod:`hermes_cli.plugins` +uses to unwind a plugin. Each ``*_registry`` module instantiates one +:class:`ProviderRegistry` and re-exports its bound methods under the historical +module-level names via :meth:`ProviderRegistry.export`, so ``patch("agent.x_registry.get_provider")`` +targets and the ``_providers`` / ``_scoped_providers`` / ``_lock`` test hooks are unchanged. """ from __future__ import annotations @@ -33,14 +31,11 @@ def lower_key(name: str) -> str: class ProviderRegistry(Generic[P]): """Global + per-scope provider map with plugin snapshot/restore support. - Args: - label: Human label used in log/error strings (``"Browser"``, ``"TTS"``). - provider_cls: ABC every registered instance must satisfy (TypeError otherwise). - logger: The owning module's logger, so record names stay per-registry. - normalize: Key normalizer — ``strip_key`` or ``lower_key`` (case-insensitive - registries mirror how their dispatcher normalizes the configured name). - builtin_names: Reserved names owned by in-tree implementations; a collision - calls ``on_builtin_collision(key)`` and, if that returns, skips registration. + ``normalize`` is ``strip_key`` or ``lower_key`` (case-insensitive registries mirror + how their dispatcher normalizes the configured name). ``builtin_names`` are reserved + for in-tree implementations; a collision calls ``on_builtin_collision(key)`` and, if + that returns, skips registration. ``logger`` is the owning module's so record names + stay per-registry. """ def __init__( @@ -67,8 +62,6 @@ class ProviderRegistry(Generic[P]): # "TTS provider" but "Registered browser provider": acronyms keep their case. self._log_label = label if label.isupper() else label[0].lower() + label[1:] - # -- internal helpers (caller holds the lock) --------------------------- - def _target(self, scope: Optional[str], *, create: bool) -> Dict[str, P]: if scope is None: return self._providers @@ -82,8 +75,6 @@ class ProviderRegistry(Generic[P]): else: self._scoped_generations[scope] = self._scoped_generations.get(scope, 0) + 1 - # -- registration ------------------------------------------------------- - def register(self, provider: P, *, scope: Optional[str] = None) -> None: """Register a provider; same-name re-registration overwrites (hot reload).""" if not isinstance(provider, self.provider_cls): @@ -107,17 +98,13 @@ class ProviderRegistry(Generic[P]): self._bump(scope) if existing is not None: self.logger.debug( - f"{self.label} provider '%s' re-registered (was %r)", - key, type(existing).__name__, + f"{self.label} provider '%s' re-registered (was %r)", key, type(existing).__name__, ) else: self.logger.debug( - f"Registered {self._log_label} provider '%s' (%s)", - key, type(provider).__name__, + f"Registered {self._log_label} provider '%s' (%s)", key, type(provider).__name__, ) - # -- lookup --------------------------------------------------------------- - def merged(self, scope: Optional[str] = None) -> Dict[str, P]: """Global map overlaid with the active profile's scoped map (a copy).""" with self._lock: @@ -146,8 +133,6 @@ class ProviderRegistry(Generic[P]): with self._lock: return self._generation, self._scoped_generations.get(active_scope, 0) - # -- plugin unload support (hermes_cli.plugins) ----------------------------- - def snapshot_registration(self, name: str, *, scope: Optional[str] = None) -> Optional[P]: """Exact-slot lookup (no global fallback) used to detect plugin ownership.""" with self._lock: @@ -180,14 +165,8 @@ class ProviderRegistry(Generic[P]): self._generation += 1 def export(self, namespace: Dict[str, Any]) -> None: - """Bind the historical module-level API into a ``*_registry`` module. - - Installs ``register_provider``/``list_providers``/``get_provider``/ - ``snapshot_registration``/``restore_registration``/``registry_generation``/ - ``_reset_for_tests`` plus the ``_providers``/``_scoped_providers``/``_lock`` - test hooks, so ``patch("agent.x_registry.get_provider")`` and direct - ``_providers`` manipulation in tests keep working unchanged. - """ + """Bind the historical module-level API (+ ``_providers``/``_scoped_providers``/ + ``_lock`` test hooks) into a ``*_registry`` module namespace.""" namespace.update( _providers=self._providers, _scoped_providers=self._scoped_providers, @@ -227,10 +206,9 @@ def configured_provider_name(section: str, logger: logging.Logger) -> Optional[s cfg = load_config_readonly() block = cfg.get(section) if isinstance(cfg, dict) else None - if isinstance(block, dict): - raw = block.get("provider") - if isinstance(raw, str) and raw.strip(): - configured = raw.strip() + raw = block.get("provider") if isinstance(block, dict) else None + if isinstance(raw, str) and raw.strip(): + configured = raw.strip() except Exception as exc: logger.debug("Could not read %s.provider from config: %s", section, exc) if configured: