diff --git a/agent/moa_loop.py b/agent/moa_loop.py index c398bfda88..2c8a598df1 100644 --- a/agent/moa_loop.py +++ b/agent/moa_loop.py @@ -63,9 +63,7 @@ def _moa_privacy_mode(moa_raw: Any) -> str: return coerce_privacy_filter(raw.get("privacy_filter")) -def _redact_reference_outputs( - reference_outputs: list[tuple[str, str, Any]], -) -> list[tuple[str, str, Any]]: +def _redact_reference_outputs(reference_outputs: list[tuple[str, str, Any]]) -> list[tuple[str, str, Any]]: """Redact advisor text in reference-output tuples; accounting slot untouched.""" return [(label, _redact_reference_text(text), acct) for label, text, acct in reference_outputs] @@ -99,8 +97,7 @@ def _redact_trace_accounting(acct: Any) -> Any: if not isinstance(acct, _RefAccounting): return acct return replace( - acct, messages=_redact_trace_messages(acct.messages), - output=_redact_reference_text(acct.output), + acct, messages=_redact_trace_messages(acct.messages), output=_redact_reference_text(acct.output), ) @@ -326,8 +323,7 @@ def _maybe_apply_moa_cache_control( # blank_cache_policy_stub is the only sanctioned stub (carries _cache_disabled). should_cache, native_layout = anthropic_prompt_cache_policy( blank_cache_policy_stub(cache_disabled), provider=provider, - base_url=runtime.get("base_url") or "", api_mode=runtime.get("api_mode") or "", - model=model, + base_url=runtime.get("base_url") or "", api_mode=runtime.get("api_mode") or "", model=model, ) if not should_cache: return messages @@ -374,9 +370,8 @@ def _price_reference_response( def _run_reference( slot: dict[str, Any], ref_messages: list[dict[str, Any]], *, temperature: float | None = None, - max_tokens: int | None = None, reference_timeout: float | None = None, - context_length_cache: Any = None, cache_disabled: bool | None = None, - cache_ttl: str | None = None, + max_tokens: int | None = None, reference_timeout: float | None = None, context_length_cache: Any = None, + cache_disabled: bool | None = None, cache_ttl: str | None = None, ) -> tuple[str, str, Any]: """Call one reference model; return ``(label, text, accounting)``. Never raises: a failed reference becomes a labelled ``[failed: …]`` note. Runs in a thread pool.""" @@ -424,9 +419,7 @@ def _run_reference( except Exception as exc: logger.warning("MoA reference model %s failed: %s", label, exc) note = f"[failed: {exc}]" - return label, note, _RefAccounting( - CanonicalUsage(), messages=messages, output=note, **trace_fields, - ) + return label, note, _RefAccounting(CanonicalUsage(), messages=messages, output=note, **trace_fields) # Output headroom reserved in the reference window when reference_max_tokens is unset. @@ -529,8 +522,7 @@ def _placeholder_output(slot: dict[str, Any], note: str) -> tuple[str, str, Any] def _settle_interrupted( - futures: dict[Any, int], results: list, reference_models: list[dict[str, Any]], - late_accounting_sink: Any, + futures: dict[Any, int], results: list, reference_models: list[dict[str, Any]], late_accounting_sink: Any, ) -> None: """Fill every unfinished slot after a user interrupt: cancel never-dispatched futures (nothing billed), keep real output of ones that just finished, and hand @@ -596,10 +588,9 @@ def _run_references_parallel( continue futures[ executor.submit( - propagate_context_to_thread(_run_reference), slot, ref_messages, - temperature=temperature, max_tokens=max_tokens, - reference_timeout=reference_timeout, context_length_cache=ctx_len_cache, - cache_disabled=cache_disabled, cache_ttl=cache_ttl, + propagate_context_to_thread(_run_reference), slot, ref_messages, temperature=temperature, + max_tokens=max_tokens, reference_timeout=reference_timeout, + context_length_cache=ctx_len_cache, cache_disabled=cache_disabled, cache_ttl=cache_ttl, ) ] = idx @@ -826,10 +817,9 @@ def _slot_labels(slots: list[dict[str, Any]]) -> str: 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, reference_max_tokens: int | None = None, - reference_timeout: float | None = None, degraded_reference_policy: str = "loud", - agent: Any = None, + aggregator: dict[str, Any], temperature: float | None = None, aggregator_temperature: float | None = None, + reference_max_tokens: int | None = None, reference_timeout: float | None = None, + degraded_reference_policy: str = "loud", agent: Any = None, ) -> str: """Run configured reference models and synthesize their advice (one-shot /moa). @@ -860,8 +850,7 @@ def aggregate_moa_context( # 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), + "MoA: all %d reference(s) failed — skipping aggregator synthesis", len(reference_outputs), ) return ( "[Mixture of Agents context — all reference models failed. " @@ -1058,9 +1047,7 @@ class MoAChatCompletions: if cost is not None: self._pending_reference_cost = (self._pending_reference_cost or 0) + cost - def consume_and_save_trace( - self, session_id: Any = None, aggregator_output_fallback: Any = None - ) -> None: + def consume_and_save_trace(self, session_id: Any = None, aggregator_output_fallback: Any = None) -> 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 for the @@ -1081,13 +1068,11 @@ class MoAChatCompletions: save_moa_turn( session_id=session_id, preset_name=pending.get("preset", ""), reference_outputs=pending.get("reference_outputs", []), - aggregator_label=pending.get("aggregator_label", ""), - aggregator_model=agg_slot.get("model"), + aggregator_label=pending.get("aggregator_label", ""), aggregator_model=agg_slot.get("model"), aggregator_provider=agg_slot.get("provider"), aggregator_temperature=pending.get("aggregator_temperature"), aggregator_input_messages=pending.get("aggregator_input_messages"), - aggregator_output=agg_output, - aggregator_streamed=bool(pending.get("aggregator_streamed")), + aggregator_output=agg_output, aggregator_streamed=bool(pending.get("aggregator_streamed")), ) except Exception as exc: # pragma: no cover - tracing must never break a turn logger.debug("MoA trace flush failed: %s", exc) @@ -1117,8 +1102,7 @@ class MoAChatCompletions: return {**prepared, "messages": agg_messages} def _plan_aggregator_cache( - self, agg_messages: list[dict[str, Any]], tools: Any, guidance: Any, - agg_runtime: dict[str, Any], + self, agg_messages: list[dict[str, Any]], tools: Any, guidance: Any, agg_runtime: dict[str, Any], ) -> tuple[list[dict[str, Any]], Any]: """Cache-breakpoint the aggregator request for its destination. @@ -1157,9 +1141,7 @@ class MoAChatCompletions: ) return agg_messages, tools - def _call_prepared_aggregator( - self, prepared: dict[str, Any], api_kwargs: dict[str, Any] - ) -> Any: + def _call_prepared_aggregator(self, prepared: dict[str, Any], api_kwargs: dict[str, Any]) -> Any: """Send an already prepared MoA aggregator request exactly once.""" aggregator = prepared["aggregator"] if aggregator.get("provider") == "moa": @@ -1287,8 +1269,7 @@ class MoAChatCompletions: self._emit("moa.progress", refs_done=done, refs_total=total, label=label) reference_outputs = _run_references_parallel( - reference_models, ref_messages, - temperature=_preset_temperature(preset, "reference_temperature"), + reference_models, ref_messages, temperature=_preset_temperature(preset, "reference_temperature"), max_tokens=preset.get("reference_max_tokens"), progress_callback=_progress, reference_timeout=float(raw_reference_timeout) if raw_reference_timeout else None, agent=self._agent, late_accounting_sink=self._record_late_reference_accounting, diff --git a/agent/plugin_llm.py b/agent/plugin_llm.py index d628418c20..0eeab9de89 100644 --- a/agent/plugin_llm.py +++ b/agent/plugin_llm.py @@ -154,9 +154,8 @@ def _resolve_trust_policy(plugin_id: str) -> _TrustPolicy: allowed_models, allow_any_model = _coerce_allowlist(llm_cfg.get("allowed_models")) allowed_providers, allow_any_provider = _coerce_allowlist(llm_cfg.get("allowed_providers")) return _TrustPolicy( - plugin_id=plugin_id, allowed_providers=allowed_providers, - allow_any_provider=allow_any_provider, allowed_models=allowed_models, - allow_any_model=allow_any_model, + plugin_id=plugin_id, allowed_providers=allowed_providers, allow_any_provider=allow_any_provider, + allowed_models=allowed_models, allow_any_model=allow_any_model, **{name: bool(llm_cfg.get(name, False)) for name in _OVERRIDE_FLAGS}, ) @@ -188,8 +187,7 @@ def _gate_ref_override(policy: _TrustPolicy, kind: str, requested: str) -> str: # Overrides gated by a bare trust flag (no allowlist): ``kind`` -> denial wording. _FLAG_ONLY_OVERRIDES = { - "agent_id": "run completions against a non-default agent id", - "profile": "override the auth profile", + "agent_id": "run completions against a non-default agent id", "profile": "override the auth profile", } @@ -235,9 +233,7 @@ def _resolve_task_ownership(plugin_id: str) -> tuple[frozenset, frozenset]: return frozenset(owned), frozenset(builtin) -def _check_task( - policy: _TrustPolicy, *, plugin_id: str, requested_task: Optional[str], -) -> Optional[str]: +def _check_task(policy: _TrustPolicy, *, plugin_id: str, requested_task: Optional[str]) -> Optional[str]: """Validate a plugin's requested auxiliary ``task`` key. unset / ``""`` / ``"auto"`` → ``None`` (main-model path); a key the plugin @@ -382,9 +378,7 @@ def _parse_structured_text( except ImportError: logger.debug("jsonschema unavailable; skipping schema validation") except jsonschema.ValidationError as exc: # type: ignore[attr-defined] - raise ValueError( - f"Plugin LLM structured output did not match schema: {exc.message}" - ) from exc + raise ValueError(f"Plugin LLM structured output did not match schema: {exc.message}") from exc return parsed, "json" @@ -512,10 +506,9 @@ class PluginLlm: # -- public API ----------------------------------------------------------- def complete( - self, messages: List[Dict[str, Any]], *, provider: Optional[str] = None, - model: Optional[str] = None, temperature: Optional[float] = None, - max_tokens: Optional[int] = None, timeout: Optional[float] = None, - agent_id: Optional[str] = None, profile: Optional[str] = None, + self, messages: List[Dict[str, Any]], *, provider: Optional[str] = None, model: Optional[str] = None, + temperature: Optional[float] = None, max_tokens: Optional[int] = None, + timeout: Optional[float] = None, agent_id: Optional[str] = None, profile: Optional[str] = None, purpose: Optional[str] = None, task: Optional[str] = None, ) -> PluginLlmCompleteResult: """Run a host-owned chat completion against the user's active model. @@ -527,12 +520,10 @@ class PluginLlm: return self._finish("complete", agent, kw, self._invoke_sync(kw), purpose) def complete_structured( - self, *, instructions: str, input: Sequence[PluginLlmInput], - json_schema: Optional[Any] = None, json_mode: bool = False, - schema_name: Optional[str] = None, system_prompt: Optional[str] = None, - provider: Optional[str] = None, model: Optional[str] = None, - temperature: Optional[float] = None, max_tokens: Optional[int] = None, - timeout: Optional[float] = None, agent_id: Optional[str] = None, + self, *, instructions: str, input: Sequence[PluginLlmInput], json_schema: Optional[Any] = None, + json_mode: bool = False, schema_name: Optional[str] = None, system_prompt: Optional[str] = None, + provider: Optional[str] = None, model: Optional[str] = None, temperature: Optional[float] = None, + max_tokens: Optional[int] = None, timeout: Optional[float] = None, agent_id: Optional[str] = None, profile: Optional[str] = None, purpose: Optional[str] = None, task: Optional[str] = None, ) -> PluginLlmStructuredResult: """Run a bounded host-owned structured completion. @@ -545,10 +536,9 @@ class PluginLlm: return self._finish("complete_structured", agent, kw, self._invoke_sync(kw), purpose, spec) async def acomplete( - self, messages: List[Dict[str, Any]], *, provider: Optional[str] = None, - model: Optional[str] = None, temperature: Optional[float] = None, - max_tokens: Optional[int] = None, timeout: Optional[float] = None, - agent_id: Optional[str] = None, profile: Optional[str] = None, + self, messages: List[Dict[str, Any]], *, provider: Optional[str] = None, model: Optional[str] = None, + temperature: Optional[float] = None, max_tokens: Optional[int] = None, + timeout: Optional[float] = None, agent_id: Optional[str] = None, profile: Optional[str] = None, purpose: Optional[str] = None, task: Optional[str] = None, ) -> PluginLlmCompleteResult: """Async sibling of :meth:`complete`.""" @@ -556,12 +546,10 @@ class PluginLlm: return self._finish("acomplete", agent, kw, await self._invoke_async(kw), purpose) async def acomplete_structured( - self, *, instructions: str, input: Sequence[PluginLlmInput], - json_schema: Optional[Any] = None, json_mode: bool = False, - schema_name: Optional[str] = None, system_prompt: Optional[str] = None, - provider: Optional[str] = None, model: Optional[str] = None, - temperature: Optional[float] = None, max_tokens: Optional[int] = None, - timeout: Optional[float] = None, agent_id: Optional[str] = None, + self, *, instructions: str, input: Sequence[PluginLlmInput], json_schema: Optional[Any] = None, + json_mode: bool = False, schema_name: Optional[str] = None, system_prompt: Optional[str] = None, + provider: Optional[str] = None, model: Optional[str] = None, temperature: Optional[float] = None, + max_tokens: Optional[int] = None, timeout: Optional[float] = None, agent_id: Optional[str] = None, profile: Optional[str] = None, purpose: Optional[str] = None, task: Optional[str] = None, ) -> PluginLlmStructuredResult: """Async sibling of :meth:`complete_structured`.""" @@ -572,10 +560,9 @@ class PluginLlm: # -- shared core ---------------------------------------------------------- def _gate( - self, provider: Optional[str], model: Optional[str], agent_id: Optional[str], - profile: Optional[str], task: Optional[str], messages: Optional[List[Dict[str, Any]]], - temperature: Optional[float], max_tokens: Optional[int], timeout: Optional[float], - spec: Optional[Dict[str, Any]] = None, + self, provider: Optional[str], model: Optional[str], agent_id: Optional[str], profile: Optional[str], + task: Optional[str], messages: Optional[List[Dict[str, Any]]], temperature: Optional[float], + max_tokens: Optional[int], timeout: Optional[float], spec: Optional[Dict[str, Any]] = None, ) -> tuple[Optional[str], Dict[str, Any]]: """Trust gate (task first, then overrides), then — for a structured ``spec`` — build messages/response_format (input-shape errors surface only after trust diff --git a/agent/provider_registry.py b/agent/provider_registry.py index 549670ce95..21d4b8e483 100644 --- a/agent/provider_registry.py +++ b/agent/provider_registry.py @@ -172,8 +172,7 @@ class ProviderRegistry(Generic[P]): def is_available_safe( - provider: Any, logger: logging.Logger, fmt: str, *, level: int = logging.DEBUG, - exc_info: bool = False, + provider: Any, logger: logging.Logger, fmt: str, *, level: int = logging.DEBUG, exc_info: bool = False, ) -> bool: """``bool(provider.is_available())`` that treats a raising provider as unavailable.""" try: