refactor(agent): plugin_llm/nous_rate_guard/moonshot_schema/provider_* — unify atomic write, compact docstrings, collapse defensive layers
This commit is contained in:
+123
-200
@@ -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"(?<![\w.+-])" # no leading word char / dot / + / - (kills IPs, IDs, versions)
|
||||
@@ -43,8 +42,8 @@ def _redact_reference_text(text: Any) -> 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:<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:<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)
|
||||
|
||||
+7
-10
@@ -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
|
||||
``<hermes_home>/moa-traces/<session_id>.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
|
||||
``<hermes_home>/moa-traces/<session_id>.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,
|
||||
|
||||
+30
-46
@@ -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
|
||||
|
||||
|
||||
|
||||
+11
-41
@@ -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)
|
||||
|
||||
+77
-155
@@ -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.<id>.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": <bytes>, "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.<plugin_id>.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.<plugin_id>.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.<id>.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.<id>.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)
|
||||
|
||||
|
||||
|
||||
+12
-23
@@ -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
|
||||
|
||||
+15
-20
@@ -1,10 +1,10 @@
|
||||
"""Shared ``$HERMES_HOME/cache/<kind>/`` materialisation helpers for the
|
||||
image/video generation provider ABCs.
|
||||
"""``$HERMES_HOME/cache/<kind>/`` materialisation helpers for the image/video
|
||||
generation provider ABCs.
|
||||
|
||||
Several backends (xAI, OpenAI, DeepInfra, FAL) return *ephemeral* delivery URLs
|
||||
that expire before a downstream consumer (Telegram ``send_photo``, browser
|
||||
fetch) can resolve them, so providers materialise the bytes locally at
|
||||
tool-completion time. Filenames are ``<prefix>_<YYYYMMDD_HHMMSS>_<uuid8>.<ext>``.
|
||||
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
|
||||
``<prefix>_<YYYYMMDD_HHMMSS>_<uuid8>.<ext>``.
|
||||
"""
|
||||
|
||||
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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
+19
-41
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user