refactor(agent): pack multi-line calls/signatures in r3-06-C files (AST-identical)
This commit is contained in:
+58
-149
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user