refactor(agent): pack remaining multi-line signatures (AST-identical)
This commit is contained in:
+20
-39
@@ -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,
|
||||
|
||||
+22
-35
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user