diff --git a/agent/moa_loop.py b/agent/moa_loop.py index 4a980fb604..0f02b39f40 100644 --- a/agent/moa_loop.py +++ b/agent/moa_loop.py @@ -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: diff --git a/agent/moa_trace.py b/agent/moa_trace.py index b4ccade104..329b28ab7a 100644 --- a/agent/moa_trace.py +++ b/agent/moa_trace.py @@ -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: diff --git a/agent/nous_rate_guard.py b/agent/nous_rate_guard.py index 2e3e95f214..15654598a4 100644 --- a/agent/nous_rate_guard.py +++ b/agent/nous_rate_guard.py @@ -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()) diff --git a/agent/plugin_llm.py b/agent/plugin_llm.py index 14ab45b15b..5acc3ae3d4 100644 --- a/agent/plugin_llm.py +++ b/agent/plugin_llm.py @@ -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 diff --git a/agent/provider_media.py b/agent/provider_media.py index d3dfffe553..eed99fc5d2 100644 --- a/agent/provider_media.py +++ b/agent/provider_media.py @@ -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. diff --git a/agent/provider_projection.py b/agent/provider_projection.py index dc56b98530..36c3765266 100644 --- a/agent/provider_projection.py +++ b/agent/provider_projection.py @@ -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 diff --git a/agent/provider_registry.py b/agent/provider_registry.py index 27b859d820..549670ce95 100644 --- a/agent/provider_registry.py +++ b/agent/provider_registry.py @@ -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."""