refactor(agent): pack multi-line calls/signatures in r3-06-C files (AST-identical)

This commit is contained in:
Teknium
2026-09-02 18:44:36 -07:00
parent 23ae37b4bc
commit 28a974aef5
7 changed files with 124 additions and 336 deletions
+58 -149
View File
@@ -98,8 +98,7 @@ def _redact_trace_accounting(acct: Any) -> Any:
if not isinstance(acct, _RefAccounting):
return acct
return replace(
acct,
messages=_redact_trace_messages(acct.messages),
acct, messages=_redact_trace_messages(acct.messages),
output=_redact_reference_text(acct.output),
)
@@ -297,10 +296,7 @@ def _with_cache_disabled(runtime: dict[str, Any], cache_disabled: Any) -> dict[s
def _maybe_apply_moa_cache_control(
messages: list[dict[str, Any]],
runtime: dict[str, Any],
*,
cache_disabled: bool | None = None,
messages: list[dict[str, Any]], runtime: dict[str, Any], *, cache_disabled: bool | None = None,
cache_ttl: str | None = None,
) -> list[dict[str, Any]]:
"""Apply cache_control to an advisor/aggregator request when its route honors it.
@@ -328,10 +324,8 @@ def _maybe_apply_moa_cache_control(
model = runtime.get("model") or ""
# 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 "",
blank_cache_policy_stub(cache_disabled), provider=provider,
base_url=runtime.get("base_url") or "", api_mode=runtime.get("api_mode") or "",
model=model,
)
if not should_cache:
@@ -369,11 +363,8 @@ def _price_reference_response(
usage = CanonicalUsage()
try:
cost = estimate_usage_cost(
slot.get("model") or "",
usage,
provider=runtime.get("provider"),
base_url=runtime.get("base_url"),
api_key=runtime.get("api_key"),
slot.get("model") or "", usage, provider=runtime.get("provider"),
base_url=runtime.get("base_url"), api_key=runtime.get("api_key"),
)
return usage, cost.amount_usd, cost.status, cost.source
except Exception: # pragma: no cover - defensive
@@ -381,14 +372,9 @@ 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,
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,
) -> tuple[str, str, Any]:
"""Call one reference model; return ``(label, text, accounting)``. Never raises:
@@ -396,8 +382,7 @@ def _run_reference(
label = _slot_label(slot)
runtime = _slot_runtime(slot)
trace_fields = {
"model": slot.get("model"),
"provider": runtime.get("provider") or slot.get("provider"),
"model": slot.get("model"), "provider": runtime.get("provider") or slot.get("provider"),
"temperature": temperature,
}
# The advisory view already stripped the agent's system prompt; this is the only one.
@@ -405,10 +390,7 @@ def _run_reference(
try:
# Trim to THIS model's window (advisors may be smaller than the aggregator).
trimmed = _trim_messages_for_reference(
messages,
slot,
runtime,
reserve_output_tokens=max_tokens,
messages, slot, runtime, reserve_output_tokens=max_tokens,
context_length_cache=context_length_cache,
)
# The advisory view is append-only across iterations, so cache_control lets
@@ -427,21 +409,15 @@ def _run_reference(
# user's current turn, so mirror the main agent's x-initiator header.
extra_headers = {"x-initiator": "user"}
response = call_llm(
task="moa_reference",
messages=trimmed,
temperature=temperature,
task="moa_reference", messages=trimmed, temperature=temperature,
max_tokens=slot_max_tokens if slot_max_tokens is not None else max_tokens,
timeout=reference_timeout,
reasoning_config=_slot_reasoning_config(slot),
extra_headers=extra_headers,
**runtime,
timeout=reference_timeout, reasoning_config=_slot_reasoning_config(slot),
extra_headers=extra_headers, **runtime,
)
output_text = _extract_text(response) or "(empty response)"
acct = _RefAccounting(
*_price_reference_response(response, slot, runtime),
messages=trimmed,
output=output_text,
**trace_fields,
*_price_reference_response(response, slot, runtime), messages=trimmed,
output=output_text, **trace_fields,
)
return label, output_text, acct
except Exception as exc:
@@ -460,12 +436,8 @@ _REFERENCE_TRIM_SAFETY_FRACTION = 0.10
def _trim_messages_for_reference(
messages: list[dict[str, Any]],
slot: dict[str, str],
runtime: dict[str, Any],
*,
reserve_output_tokens: int | None = None,
context_length_cache: Any = None,
messages: list[dict[str, Any]], slot: dict[str, str], runtime: dict[str, Any], *,
reserve_output_tokens: int | None = None, context_length_cache: Any = None,
) -> list[dict[str, Any]]:
"""Trim an advisory request to fit a reference model's context window.
@@ -495,10 +467,8 @@ def _trim_messages_for_reference(
else:
try:
context_length = get_model_context_length(
model=model,
base_url=str(runtime.get("base_url") or ""),
api_key=str(runtime.get("api_key") or ""),
provider=provider,
model=model, base_url=str(runtime.get("base_url") or ""),
api_key=str(runtime.get("api_key") or ""), provider=provider,
)
except Exception:
logger.debug("MoA reference context-length resolution failed for %s", _slot_label(slot))
@@ -560,9 +530,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]],
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
@@ -592,15 +560,9 @@ def _settle_interrupted(
def _run_references_parallel(
reference_models: list[dict[str, Any]],
ref_messages: list[dict[str, Any]],
*,
temperature: float | None = None,
max_tokens: int | None = None,
progress_callback: Any = None,
reference_timeout: float | None = None,
agent: Any = None,
late_accounting_sink: Any = None,
reference_models: list[dict[str, Any]], ref_messages: list[dict[str, Any]], *,
temperature: float | None = None, max_tokens: int | None = None, progress_callback: Any = None,
reference_timeout: float | None = None, agent: Any = None, late_accounting_sink: Any = None,
) -> list[tuple[str, str, Any]]:
"""Fan out all reference models in parallel; ``(label, text, _RefAccounting)`` per
slot in ``reference_models`` order.
@@ -635,15 +597,10 @@ 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
@@ -871,16 +828,10 @@ 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",
*, 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,
) -> str:
"""Run configured reference models and synthesize their advice (one-shot /moa).
@@ -891,12 +842,8 @@ def aggregate_moa_context(
"""
reference_models = [slot for slot in reference_models if slot.get("enabled", True)]
reference_outputs = _run_references_parallel(
reference_models,
_reference_messages(api_messages),
temperature=temperature,
max_tokens=reference_max_tokens,
reference_timeout=reference_timeout,
agent=agent,
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)
@@ -944,15 +891,11 @@ def aggregate_moa_context(
# a third independent MoA call path that otherwise re-bills its full input.
agg_messages = _maybe_apply_moa_cache_control(
[{"role": "user", "content": synth_prompt}],
_with_cache_disabled(agg_runtime, cache_disabled),
cache_ttl=cache_ttl,
_with_cache_disabled(agg_runtime, cache_disabled), cache_ttl=cache_ttl,
)
response = call_llm(
task="moa_aggregator",
messages=agg_messages,
temperature=aggregator_temperature,
reasoning_config=_aggregator_reasoning_config(aggregator),
**agg_runtime,
task="moa_aggregator", messages=agg_messages, temperature=aggregator_temperature,
reasoning_config=_aggregator_reasoning_config(aggregator), **agg_runtime,
)
synthesis = _extract_text(response)
except Exception as exc:
@@ -990,21 +933,17 @@ def _completed_response_as_stream_chunk(response: Any) -> Any:
for index, tc in enumerate(raw_tool_calls)
]
delta = SimpleNamespace(
content=getattr(message, "content", None),
tool_calls=tool_call_deltas,
content=getattr(message, "content", None), tool_calls=tool_call_deltas,
reasoning_content=getattr(message, "reasoning_content", None),
reasoning=getattr(message, "reasoning", None),
reasoning_details=getattr(message, "reasoning_details", None),
)
choice = SimpleNamespace(
index=getattr(first_choice, "index", 0),
delta=delta,
index=getattr(first_choice, "index", 0), delta=delta,
finish_reason=getattr(first_choice, "finish_reason", None) or "stop",
)
return SimpleNamespace(
id=getattr(response, "id", None),
model=getattr(response, "model", None),
choices=[choice],
id=getattr(response, "id", None), model=getattr(response, "model", None), choices=[choice],
usage=getattr(response, "usage", None),
)
@@ -1029,10 +968,7 @@ def _attach_reference_guidance(agg_messages: list[dict[str, Any]], guidance: str
agg_messages.append({"role": "user", "content": guidance})
def peel_reference_guidance(
messages: list[dict[str, Any]],
guidance: Any,
) -> list[dict[str, Any]]:
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), so a
cache breakpoint never lands on the turn-varying guidance. Inputs are not mutated."""
if not guidance or not messages:
@@ -1150,8 +1086,7 @@ class MoAChatCompletions:
if agg_output is None and aggregator_output_fallback:
agg_output = aggregator_output_fallback
save_moa_turn(
session_id=session_id,
preset_name=pending.get("preset", ""),
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"),
@@ -1189,10 +1124,7 @@ class MoAChatCompletions:
return {**prepared, "messages": agg_messages}
def _plan_aggregator_cache(
self,
agg_messages: list[dict[str, Any]],
tools: Any,
guidance: 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.
@@ -1294,9 +1226,7 @@ class MoAChatCompletions:
return agg_response
def _fanout_cache_key(
self,
preset: dict[str, Any],
ref_messages: list[dict[str, Any]],
self, preset: dict[str, Any], ref_messages: list[dict[str, Any]],
reference_models: list[dict[str, Any]],
) -> tuple:
"""Turn-scoped reference cache key per the preset's fan-out cadence.
@@ -1353,13 +1283,9 @@ class MoAChatCompletions:
)
def _run_fanout(
self,
preset: dict[str, Any],
ref_messages: list[dict[str, Any]],
reference_models: list[dict[str, Any]],
aggregator: dict[str, Any],
aggregator_temperature: Any,
cache_key: tuple,
self, preset: dict[str, Any], ref_messages: list[dict[str, Any]],
reference_models: list[dict[str, Any]], aggregator: dict[str, Any],
aggregator_temperature: Any, cache_key: tuple,
) -> list[tuple[str, str, Any]]:
"""Cache-MISS path of ``create``: run the advisors, account, trace and emit.
@@ -1372,14 +1298,11 @@ class MoAChatCompletions:
self._emit("moa.progress", refs_done=done, refs_total=total, label=label)
reference_outputs = _run_references_parallel(
reference_models,
ref_messages,
reference_models, ref_messages,
temperature=_preset_temperature(preset, "reference_temperature"),
max_tokens=preset.get("reference_max_tokens"),
progress_callback=_progress,
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,
agent=self._agent, late_accounting_sink=self._record_late_reference_accounting,
)
if any(text == _INTERRUPTED_REFERENCE_NOTE for _lbl, text, _acct in reference_outputs):
# An interrupted fan-out is a partial snapshot: never cache it (a HIT
@@ -1402,10 +1325,8 @@ class MoAChatCompletions:
else:
trace_refs = list(reference_outputs)
self._pending_trace = {
"preset": self.preset_name,
"reference_outputs": trace_refs,
"aggregator_slot": aggregator,
"aggregator_temperature": aggregator_temperature,
"preset": self.preset_name, "reference_outputs": trace_refs,
"aggregator_slot": aggregator, "aggregator_temperature": aggregator_temperature,
}
# Derived from the privacy-redacted trace_refs.
try:
@@ -1422,29 +1343,21 @@ class MoAChatCompletions:
ref_count = len(reference_outputs)
for idx, (label, text, _accounting) in enumerate(reference_outputs, start=1):
self._emit(
"moa.reference",
index=idx,
count=ref_count,
label=label,
"moa.reference", index=idx, count=ref_count, label=label,
text=_redact_reference_text(text) if privacy_mode else text,
)
if ref_count:
# Phase transition: fan-out complete, aggregator about to act.
agg_label = _slot_label(aggregator)
self._emit(
"moa.phase",
phase="aggregator",
refs_done=ref_count,
refs_total=ref_count,
"moa.phase", phase="aggregator", refs_done=ref_count, refs_total=ref_count,
aggregator=agg_label,
)
self._emit("moa.aggregating", aggregator=agg_label, ref_count=ref_count)
return reference_outputs
def _build_guidance(
self,
reference_outputs: list[tuple[str, str, Any]],
aggregator: dict[str, Any],
self, reference_outputs: list[tuple[str, str, Any]], aggregator: dict[str, Any],
degraded_reference_policy: str,
) -> str | None:
"""Render the reference block attached to the aggregator prompt (None = nothing)."""
@@ -1531,9 +1444,7 @@ class MoAChatCompletions:
_attach_reference_guidance(agg_messages, guidance)
prepared_request = {
"messages": agg_messages,
"guidance": guidance,
"aggregator": aggregator,
"messages": agg_messages, "guidance": guidance, "aggregator": aggregator,
"aggregator_temperature": aggregator_temperature,
}
if api_kwargs.pop("_moa_prepare_only", False):
@@ -1591,10 +1502,8 @@ def build_moa_facade(agent, preset_name: Any = None) -> MoAClient:
primary, secondary, extra_map = spec
try:
cb(
event,
str(kwargs.get(primary) or ""),
str(kwargs.get(secondary) or "") if secondary else None,
None,
event, str(kwargs.get(primary) or ""),
str(kwargs.get(secondary) or "") if secondary else None, None,
**{out: kwargs.get(src) for out, src in extra_map.items()},
)
except Exception:
+9 -22
View File
@@ -55,10 +55,8 @@ def _slot_trace(acct: Any, label: str) -> dict[str, Any]:
input messages and output (not the truncated display preview)."""
usage = getattr(acct, "usage", None)
return {
"label": label,
**{f: getattr(acct, f, None) for f in _ACCT_FIELDS},
"input_messages": getattr(acct, "messages", None),
"output": getattr(acct, "output", None),
"label": label, **{f: getattr(acct, f, None) for f in _ACCT_FIELDS},
"input_messages": getattr(acct, "messages", None), "output": getattr(acct, "output", None),
"usage": {f: getattr(usage, f, 0) for f in _USAGE_FIELDS} if usage is not None else {},
**{f: getattr(acct, f, None) for f in _COST_FIELDS},
}
@@ -76,16 +74,9 @@ def slot_metrics(acct: Any, label: str, output: Any = None) -> dict[str, Any]:
def save_moa_turn(
*,
session_id: Optional[str],
preset_name: str,
reference_outputs: list[tuple[str, str, Any]],
aggregator_label: str,
aggregator_model: Optional[str],
aggregator_provider: Optional[str],
aggregator_temperature: Any,
aggregator_input_messages: Any,
aggregator_output: Optional[str],
*, session_id: Optional[str], preset_name: str, reference_outputs: list[tuple[str, str, Any]],
aggregator_label: str, aggregator_model: Optional[str], aggregator_provider: Optional[str],
aggregator_temperature: Any, aggregator_input_messages: Any, aggregator_output: Optional[str],
aggregator_streamed: bool,
) -> None:
"""Append one full MoA turn record to the session's trace JSONL, if enabled.
@@ -113,14 +104,10 @@ def save_moa_turn(
"preset": preset_name,
"references": [_slot_trace(acct, label) for label, _text, acct in reference_outputs],
"aggregator": {
"label": aggregator_label,
"model": aggregator_model,
"provider": aggregator_provider,
"temperature": aggregator_temperature,
"input_messages": aggregator_input_messages,
"output": aggregator_output,
"streamed": aggregator_streamed,
"output_location": output_location,
"label": aggregator_label, "model": aggregator_model,
"provider": aggregator_provider, "temperature": aggregator_temperature,
"input_messages": aggregator_input_messages, "output": aggregator_output,
"streamed": aggregator_streamed, "output_location": output_location,
},
}
with path.open("a", encoding="utf-8") as f:
+3 -9
View File
@@ -54,9 +54,7 @@ def _parse_reset_seconds(headers: Optional[Mapping[str, str]]) -> Optional[float
def record_nous_rate_limit(
*,
headers: Optional[Mapping[str, str]] = None,
error_context: Optional[dict[str, Any]] = None,
*, headers: Optional[Mapping[str, str]] = None, error_context: Optional[dict[str, Any]] = None,
default_cooldown: float = 300.0,
) -> None:
"""Record that Nous Portal is rate-limited in the shared state file.
@@ -121,9 +119,7 @@ def _is_exhausted(remaining: Optional[int], reset: Optional[float]) -> bool:
def is_genuine_nous_rate_limit(
*,
headers: Optional[Mapping[str, str]] = None,
last_known_state: Optional[Any] = None,
*, headers: Optional[Mapping[str, str]] = None, last_known_state: Optional[Any] = None,
) -> bool:
"""Decide whether a 429 from Nous Portal is a real account rate limit.
@@ -154,9 +150,7 @@ def _parse_buckets_from_headers(
return result
def _has_exhausted_bucket(
buckets: Mapping[str, tuple[Optional[int], Optional[float]]],
) -> bool:
def _has_exhausted_bucket(buckets: Mapping[str, tuple[Optional[int], Optional[float]]]) -> bool:
return any(_is_exhausted(remaining, reset) for remaining, reset in buckets.values())
+43 -120
View File
@@ -154,10 +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,
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},
)
@@ -238,10 +236,7 @@ def _resolve_task_ownership(plugin_id: str) -> tuple[frozenset, frozenset]:
def _check_task(
policy: _TrustPolicy,
*,
plugin_id: str,
requested_task: Optional[str],
policy: _TrustPolicy, *, plugin_id: str, requested_task: Optional[str],
) -> Optional[str]:
"""Validate a plugin's requested auxiliary ``task`` key.
@@ -321,13 +316,8 @@ def _image_part(norm: Dict[str, Any]) -> Dict[str, Any]:
def _build_structured_messages(
*,
instructions: str,
inputs: Sequence[PluginLlmInput],
json_mode: bool,
json_schema: Optional[Any],
schema_name: Optional[str],
system_prompt: Optional[str],
*, instructions: str, inputs: Sequence[PluginLlmInput], json_mode: bool,
json_schema: Optional[Any], schema_name: Optional[str], system_prompt: Optional[str],
) -> List[Dict[str, Any]]:
"""OpenAI-style messages for a structured call: optional system message (prompt +
JSON-only directive), then a user message whose first text part is the
@@ -453,10 +443,7 @@ def _main_config_value(reader: str, default: str) -> str:
def _resolve_attribution(
*,
provider_override: Optional[str],
model_override: Optional[str],
response: Any,
*, provider_override: Optional[str], model_override: Optional[str], response: Any,
route_info: Optional[Dict[str, str]] = None,
) -> tuple[str, str]:
"""``(provider, model)`` to record on the result.
@@ -511,10 +498,7 @@ class PluginLlm:
caller) → ``_finish`` (result + audit log)."""
def __init__(
self,
*,
plugin_id: str,
policy_loader: Optional[Callable[[str], _TrustPolicy]] = None,
self, *, plugin_id: str, policy_loader: Optional[Callable[[str], _TrustPolicy]] = None,
sync_caller: Optional[Callable[..., Any]] = None,
async_caller: Optional[Callable[..., Awaitable[Any]]] = None,
) -> None:
@@ -526,18 +510,11 @@ 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,
purpose: Optional[str] = None,
task: 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.
@@ -548,23 +525,13 @@ 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,
profile: Optional[str] = None,
purpose: Optional[str] = None,
task: 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.
@@ -576,41 +543,24 @@ 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,
purpose: Optional[str] = None,
task: 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`."""
agent, kw = self._gate(provider, model, agent_id, profile, task, messages, temperature, max_tokens, timeout)
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,
profile: Optional[str] = None,
purpose: Optional[str] = None,
task: 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`."""
spec = _structured_spec("acomplete_structured", instructions, input, system_prompt, json_mode, json_schema, schema_name)
@@ -620,16 +570,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],
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`` —
@@ -648,25 +591,14 @@ class PluginLlm:
messages = _build_structured_messages(**spec)
extra_body = _json_response_format(json_mode=spec["json_mode"], json_schema=spec["json_schema"])
return eff_agent, dict(
messages=messages,
provider_override=eff_provider,
model_override=eff_model,
profile_override=eff_profile,
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
extra_body=extra_body,
task=eff_task,
messages=messages, provider_override=eff_provider, model_override=eff_model,
profile_override=eff_profile, temperature=temperature, max_tokens=max_tokens,
timeout=timeout, extra_body=extra_body, task=eff_task,
)
def _finish(
self,
name: str,
agent_id: Optional[str],
kw: Dict[str, Any],
invoked: tuple[str, str, Any],
purpose: Optional[str],
spec: Optional[Dict[str, Any]] = None,
self, name: str, agent_id: Optional[str], kw: Dict[str, Any], invoked: tuple[str, str, Any],
purpose: Optional[str], spec: Optional[Dict[str, Any]] = None,
) -> Any:
"""Build the result object + audit dict and emit the INFO audit line."""
real_provider, real_model, response = invoked
@@ -701,15 +633,9 @@ class PluginLlm:
merged_extra.setdefault("metadata", {})["auth_profile"] = kw["profile_override"]
route_info: Optional[Dict[str, str]] = {} if kw["task"] else None
return dict(
task=kw["task"],
provider=kw["provider_override"],
model=kw["model_override"],
messages=kw["messages"],
temperature=kw["temperature"],
max_tokens=kw["max_tokens"],
timeout=kw["timeout"],
extra_body=merged_extra or None,
route_info=route_info,
task=kw["task"], provider=kw["provider_override"], model=kw["model_override"],
messages=kw["messages"], temperature=kw["temperature"], max_tokens=kw["max_tokens"],
timeout=kw["timeout"], extra_body=merged_extra or None, route_info=route_info,
), route_info
@staticmethod
@@ -740,10 +666,7 @@ class PluginLlm:
def make_plugin_llm_for_test(
*,
plugin_id: str,
policy: _TrustPolicy,
sync_caller: Optional[Callable[..., Any]] = None,
*, plugin_id: str, policy: _TrustPolicy, 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
+3 -12
View File
@@ -44,18 +44,9 @@ def save_b64(kind: str, b64_data: str, *, prefix: str, extension: str) -> Path:
def save_url(
kind: str,
url: str,
*,
prefix: str,
timeout: float,
max_bytes: int,
chunk_size: int,
content_types: Dict[str, str],
url_extensions: Tuple[str, ...],
default_extension: str,
label: str,
empty_error: str,
kind: str, url: str, *, prefix: str, timeout: float, max_bytes: int, chunk_size: int,
content_types: Dict[str, str], url_extensions: Tuple[str, ...], default_extension: str,
label: str, empty_error: str,
) -> Path:
"""Stream-download *url* into the cache with a size cap.
+1 -3
View File
@@ -22,9 +22,7 @@ logger = logging.getLogger(__name__)
__all__ = ["splice_provider_projection"]
def splice_provider_projection(
agent: Any, response: Any, messages: list[dict[str, Any]]
) -> int:
def splice_provider_projection(agent: Any, response: Any, messages: list[dict[str, Any]]) -> int:
"""Append the provider's projected history rows and tick the nudge counter.
Returns the number of rows spliced. Tolerates absent/garbage attributes so a
+7 -21
View File
@@ -39,13 +39,8 @@ 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(),
self, *, label: str, provider_cls: type, logger: logging.Logger,
normalize: Callable[[str], str] = strip_key, builtin_names: FrozenSet[str] = frozenset(),
on_builtin_collision: Optional[Callable[[str], None]] = None,
) -> None:
self.label = label
@@ -168,25 +163,16 @@ class ProviderRegistry(Generic[P]):
"""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,
_lock=self._lock,
register_provider=self.register,
list_providers=self.list_providers,
get_provider=self.get_provider,
snapshot_registration=self.snapshot_registration,
_providers=self._providers, _scoped_providers=self._scoped_providers, _lock=self._lock,
register_provider=self.register, list_providers=self.list_providers,
get_provider=self.get_provider, snapshot_registration=self.snapshot_registration,
restore_registration=self.restore_registration,
registry_generation=self.registry_generation,
_reset_for_tests=self.reset_for_tests,
registry_generation=self.registry_generation, _reset_for_tests=self.reset_for_tests,
)
def is_available_safe(
provider: Any,
logger: logging.Logger,
fmt: str,
*,
level: int = logging.DEBUG,
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."""