refactor(agent): moa_loop guidance-input helper, explicit MoAClient delegation, blank squeeze

This commit is contained in:
Teknium
2026-09-02 19:10:21 -07:00
parent e34096ca19
commit 7c76e5f219
6 changed files with 45 additions and 84 deletions

View File

@@ -49,7 +49,6 @@ def _redact_reference_text(text: Any) -> Any:
if not isinstance(text, str) or not text:
return text
from agent.redact import redact_sensitive_text
text = redact_sensitive_text(text, force=True, code_file=True)
text = _MOA_EMAIL_RE.sub("[redacted email]", text)
return _MOA_PHONE_RE.sub("[redacted phone]", text)
@@ -58,7 +57,6 @@ def _redact_reference_text(text: Any) -> Any:
def _moa_privacy_mode(moa_raw: Any) -> str:
"""Normalized privacy-filter mode from a raw ``moa`` config."""
from hermes_cli.moa_config import coerce_privacy_filter
raw = moa_raw if isinstance(moa_raw, dict) else {}
return coerce_privacy_filter(raw.get("privacy_filter"))
@@ -111,7 +109,6 @@ def _resolve_preset_cached(preset_name: str) -> tuple[dict[str, Any], Any]:
(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
try:
cfg_stamp = get_config_path().stat().st_mtime_ns
except OSError:
@@ -211,7 +208,6 @@ def _slot_reasoning_config(slot: dict[str, Any]) -> dict[str, Any] | None:
effort = slot.get("reasoning_effort")
try:
from hermes_constants import parse_reasoning_effort
return parse_reasoning_effort(effort)
except Exception: # pragma: no cover - bad config must not break MoA
return None
@@ -229,7 +225,6 @@ def _aggregator_reasoning_config(aggregator: dict[str, Any]) -> dict[str, Any] |
try:
from hermes_cli.config import load_config
from hermes_constants import resolve_reasoning_config
return resolve_reasoning_config(load_config() or {}, str(aggregator.get("model") or ""))
except Exception: # pragma: no cover - bad config must not break MoA
return None
@@ -252,7 +247,6 @@ def _slot_runtime(slot: dict[str, Any]) -> dict[str, Any]:
out: dict[str, Any] = {"provider": provider, "model": model}
try:
from hermes_cli.runtime_provider import resolve_runtime_provider
rt = resolve_runtime_provider(requested=provider, target_model=model)
out.update({k: rt[k] for k in ("base_url", "api_key", "api_mode") if rt.get(k)})
request_overrides = rt.get("request_overrides")
@@ -348,7 +342,6 @@ def _price_reference_response(
"""Normalize a reference's usage with the slot's OWN provider/api_mode and price it
at its own rate (hence fan-out cost is summed in dollars). Never raises."""
from agent.usage_pricing import estimate_usage_cost, normalize_usage
usage = CanonicalUsage()
raw_usage = getattr(response, "usage", None)
if raw_usage:
@@ -399,7 +392,6 @@ def _run_reference(
extra_headers = None
# Normalize provider aliases (github, github-copilot, ...) via the canonical table.
from agent.auxiliary_client import _normalize_aux_provider
if _normalize_aux_provider(str(runtime.get("provider") or "")) in ("copilot", "copilot-acp"):
# Copilot gates premium models on request attribution; MoA fan-out serves the
# user's current turn, so mirror the main agent's x-initiator header.
@@ -570,7 +562,6 @@ def _run_references_parallel(
futures: dict[Any, int] = {}
# Propagate the turn's contextvars (approval callbacks, Nous conversation tag).
from tools.thread_context import propagate_context_to_thread
total = len(reference_models)
completed = 0
executor = ThreadPoolExecutor(max_workers=min(_MAX_REFERENCE_WORKERS, total))
@@ -815,6 +806,19 @@ def _slot_labels(slots: list[dict[str, Any]]) -> str:
return ", ".join(_slot_label(slot) for slot in slots)
def _guidance_inputs(
reference_outputs: list[tuple[str, str, Any]], privacy_full: bool, policy: str,
) -> tuple[list[tuple[str, str, Any]], str, bool]:
"""``(advisor outputs for the aggregator, degraded notice, all_failed)``.
'full' privacy mode redacts advisor text reaching the aggregator (applied to a
per-call copy — caches hold raw text). Failed refs are filtered out first.
"""
successful, failed_labels = _split_references(reference_outputs)
agg_refs = _redact_reference_outputs(successful) if privacy_full else successful
return agg_refs, _degraded_notice(failed_labels, policy), bool(reference_outputs) and not successful
def aggregate_moa_context(
*, user_prompt: str, api_messages: list[dict[str, Any]], reference_models: list[dict[str, Any]],
aggregator: dict[str, Any], temperature: float | None = None, aggregator_temperature: float | None = None,
@@ -832,26 +836,19 @@ def aggregate_moa_context(
reference_models, _reference_messages(api_messages), temperature=temperature,
max_tokens=reference_max_tokens, reference_timeout=reference_timeout, agent=agent,
)
successful_outputs, failed_labels = _split_references(reference_outputs)
# 'full' privacy mode also redacts advisor text before it reaches the synthesizer.
privacy_full = False
try:
from hermes_cli.config import load_config as _load_config
if _moa_privacy_mode((_load_config() or {}).get("moa")) == "full":
successful_outputs = _redact_reference_outputs(successful_outputs)
privacy_full = _moa_privacy_mode((_load_config() or {}).get("moa")) == "full"
except Exception: # pragma: no cover - privacy filter must never break a turn
logger.debug("MoA privacy filter check failed", exc_info=True)
degraded = _degraded_notice(failed_labels, degraded_reference_policy)
joined = _join_reference_outputs(successful_outputs, degraded)
agg_refs, degraded, all_failed = _guidance_inputs(reference_outputs, privacy_full, degraded_reference_policy)
joined = _join_reference_outputs(agg_refs, degraded)
# Every reference failed: skip the aggregator (synthesizing over nothing can block
# for the full provider timeout) and return only the sanitized notice.
if reference_outputs and not successful_outputs:
logger.warning(
"MoA: all %d reference(s) failed — skipping aggregator synthesis", len(reference_outputs),
)
if all_failed:
logger.warning("MoA: all %d reference(s) failed — skipping aggregator synthesis", len(reference_outputs))
return (
"[Mixture of Agents context — all reference models failed. "
"Proceeding without aggregated guidance.]\n"
@@ -1059,7 +1056,6 @@ class MoAChatCompletions:
return
try:
from agent.moa_trace import save_moa_turn
agg_slot = pending.get("aggregator_slot") or {}
# Inline capture (non-streaming) beats the caller's streamed text.
agg_output = pending.get("aggregator_output")
@@ -1112,7 +1108,6 @@ class MoAChatCompletions:
"""
try:
from agent.agent_runtime_helpers import plan_cache_sections_for_destination
planning_messages = agg_messages
if guidance:
planning_messages = peel_reference_guidance(agg_messages, str(guidance))
@@ -1298,7 +1293,6 @@ class MoAChatCompletions:
# Derived from the privacy-redacted trace_refs.
try:
from agent.moa_trace import slot_metrics
self._last_reference_metrics = [
slot_metrics(acct, label, output=text) for label, text, acct in trace_refs
]
@@ -1328,20 +1322,15 @@ class MoAChatCompletions:
degraded_reference_policy: str,
) -> 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.
agg_refs = (
_redact_reference_outputs(successful_outputs)
if self._privacy_mode == "full"
else successful_outputs
agg_refs, degraded, all_failed = _guidance_inputs(
reference_outputs, self._privacy_mode == "full", degraded_reference_policy
)
degraded = _degraded_notice(failed_labels, degraded_reference_policy)
header = (
"[Mixture of Agents reference context]\n"
f"Preset: {self.preset_name}\n"
f"Aggregator/acting model: {_slot_label(aggregator)}\n"
)
if reference_outputs and not successful_outputs:
if all_failed:
# Every reference failed: the aggregator acts alone (loud policy → notice).
logger.warning(
"MoA: all %d reference(s) failed — acting aggregator-alone "
@@ -1420,27 +1409,27 @@ 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",
)
"""OpenAI-client-shaped wrapper: ``client.chat.completions`` is a ``MoAChatCompletions``;
the accounting/trace surface below is delegated to that facade."""
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,
self.chat.completions = MoAChatCompletions(preset_name, reference_callback=reference_callback, agent=agent)
def consume_reference_usage(self) -> Any:
return self.chat.completions.consume_reference_usage()
@property
def last_aggregator_slot(self) -> Any:
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:
return self.chat.completions.consume_and_save_trace(
session_id, aggregator_output_fallback=aggregator_output_fallback
)
def __getattr__(self, name: str) -> Any:
if name in MoAClient._DELEGATED:
return getattr(self.chat.completions, name)
raise AttributeError(name)
def last_reference_metrics(self) -> Any:
return self.chat.completions.last_reference_metrics()
# Relay table: event -> (primary kwarg, secondary kwarg or None, {cb kwarg: emit kwarg}).
@@ -1448,11 +1437,7 @@ class MoAClient:
_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"}),
"moa.phase": (
"aggregator",
None,
{"moa_phase": "phase", "moa_refs_done": "refs_done", "moa_refs_total": "refs_total"},
),
"moa.phase": ("aggregator", None, {"moa_phase": "phase", "moa_refs_done": "refs_done", "moa_refs_total": "refs_total"}),
"moa.aggregating": ("aggregator", None, {"moa_ref_count": "ref_count"}),
}
@@ -1482,7 +1467,6 @@ def build_moa_facade(agent, preset_name: Any = None) -> MoAClient:
try:
from hermes_cli.config import load_config
from hermes_cli.moa_config import normalize_moa_config
moa_cfg = normalize_moa_config(load_config().get("moa") or {})
presets = moa_cfg.get("presets") or {}
if resolved_preset not in presets:

View File

@@ -27,7 +27,6 @@ def _traces_enabled_and_dir() -> Optional[Path]:
call (once per cache-MISS turn); ``moa.trace_dir`` overrides the default."""
try:
from hermes_cli.config import load_config
moa_cfg = (load_config() or {}).get("moa") or {}
except Exception: # pragma: no cover - never break a turn over tracing
return None

View File

@@ -112,11 +112,8 @@ class _TrustPolicy:
# The ``allow_*_override`` config keys; each is a same-named ``_TrustPolicy`` field.
_OVERRIDE_FLAGS = (
"allow_provider_override",
"allow_model_override",
"allow_agent_id_override",
"allow_profile_override",
"allow_task_override",
"allow_provider_override", "allow_model_override", "allow_agent_id_override",
"allow_profile_override", "allow_task_override",
)
@@ -217,7 +214,6 @@ def _resolve_task_ownership(plugin_id: str) -> tuple[frozenset, frozenset]:
builtin: set = set()
try:
from hermes_cli.plugins import get_plugin_auxiliary_tasks
owned = {
e.get("key") for e in get_plugin_auxiliary_tasks()
if e.get("plugin") == plugin_id and isinstance(e.get("key"), str) and e.get("key")
@@ -226,7 +222,6 @@ def _resolve_task_ownership(plugin_id: str) -> tuple[frozenset, frozenset]:
pass
try:
from hermes_cli.main import _AUX_TASKS
builtin = {k for k, _name, _desc in _AUX_TASKS}
except Exception: # pragma: no cover — main import failure
pass
@@ -664,13 +659,6 @@ def make_plugin_llm_for_test(
__all__ = [
"PluginLlm",
"PluginLlmTextInput",
"PluginLlmImageInput",
"PluginLlmInput",
"PluginLlmUsage",
"PluginLlmCompleteResult",
"PluginLlmStructuredResult",
"PluginLlmTrustError",
"make_plugin_llm_for_test",
"PluginLlm", "PluginLlmTextInput", "PluginLlmImageInput", "PluginLlmInput", "PluginLlmUsage",
"PluginLlmCompleteResult", "PluginLlmStructuredResult", "PluginLlmTrustError", "make_plugin_llm_for_test",
]

View File

@@ -77,7 +77,6 @@ def _worker(dispatcher: _ConsumerDispatcher) -> None:
def _registered_callbacks(hook_name: str) -> tuple[Callable[..., Any], ...]:
try:
from hermes_cli import plugins
return plugins.iter_hook_callbacks(hook_name)
except Exception:
logger.debug("plugin stream hook callback lookup failed: %s", hook_name, exc_info=True)
@@ -149,7 +148,6 @@ def stream_reasoning_deltas_enabled() -> bool:
"""Return True only when the user opted plugins into reasoning deltas."""
try:
from hermes_cli import config as config_mod
config = config_mod.load_config()
return bool(config_mod.cfg_get(config, "plugins", "stream_reasoning_deltas", default=False))
except Exception:

View File

@@ -19,7 +19,6 @@ from typing import Dict, Tuple
def cache_dir(kind: str) -> Path:
"""Return ``$HERMES_HOME/cache/<kind>/``, creating parents as needed."""
from hermes_constants import get_hermes_home
path = get_hermes_home() / "cache" / kind
path.mkdir(parents=True, exist_ok=True)
return path
@@ -57,7 +56,6 @@ def save_url(
callers can fall back to the bare URL; a partial file is never left behind.
"""
import requests
response = requests.get(url, timeout=timeout, stream=True)
response.raise_for_status()

View File

@@ -20,10 +20,6 @@ from hermes_constants import hermes_home_key
P = TypeVar("P")
def strip_key(name: str) -> str:
return name.strip()
def lower_key(name: str) -> str:
return name.strip().lower()
@@ -31,7 +27,7 @@ def lower_key(name: str) -> str:
class ProviderRegistry(Generic[P]):
"""Global + per-scope provider map with plugin snapshot/restore support.
``normalize`` is ``strip_key`` or ``lower_key`` (case-insensitive registries mirror
``normalize`` is ``str.strip`` 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
@@ -40,7 +36,7 @@ class ProviderRegistry(Generic[P]):
def __init__(
self, *, label: str, provider_cls: type, logger: logging.Logger,
normalize: Callable[[str], str] = strip_key, builtin_names: FrozenSet[str] = frozenset(),
normalize: Callable[[str], str] = str.strip, builtin_names: FrozenSet[str] = frozenset(),
on_builtin_collision: Optional[Callable[[str], None]] = None,
) -> None:
self.label = label
@@ -188,7 +184,6 @@ def configured_provider_name(section: str, logger: logging.Logger) -> Optional[s
configured: Optional[str] = None
try:
from hermes_cli.config import load_config_readonly
cfg = load_config_readonly()
block = cfg.get(section) if isinstance(cfg, dict) else None
raw = block.get("provider") if isinstance(block, dict) else None
@@ -199,7 +194,6 @@ def configured_provider_name(section: str, logger: logging.Logger) -> Optional[s
if configured:
try:
from tools.tool_backend_helpers import NOUS_MANAGED_PROVIDER
if configured.lower() == NOUS_MANAGED_PROVIDER:
configured = "fal"
except Exception: # pragma: no cover — helpers are in-repo