From dcbf263a09259b120a76372cfc4e27e988c8f508 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:31:50 -0700 Subject: [PATCH] refactor(agent/aux): unify same-provider retry and fallback-candidate prep across sync/async (_prepare_same_provider_retry, _plan_fallback_candidate) --- agent/auxiliary_client.py | 408 +++++++++++++++----------------------- 1 file changed, 164 insertions(+), 244 deletions(-) diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index 6bd0419c71..6e2e73b270 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -4405,7 +4405,7 @@ def _recover_provider_pool(provider: str, exc: Exception, *, failed_api_key: str return False -def _retry_same_provider_sync( +def _prepare_same_provider_retry( *, task: Optional[str], resolved_provider: str, @@ -4422,20 +4422,23 @@ def _retry_same_provider_sync( effective_timeout: float, effective_extra_body: dict, reasoning_config: Optional[dict], + async_mode: bool, extra_headers: Optional[Dict[str, str]] = None, -) -> Any: +) -> Tuple[Any, Dict[str, Any]]: + """Rebuild (client, request kwargs) for a same-provider retry after credential recovery.""" if task == "vision": effective_provider, retry_client, retry_model = resolve_vision_provider_client( provider=resolved_provider, model=final_model, base_url=resolved_base_url, api_key=resolved_api_key, - async_mode=False, + async_mode=async_mode, ) else: retry_client, retry_model = _get_cached_client( resolved_provider, resolved_model, + async_mode=async_mode, base_url=resolved_base_url, api_key=resolved_api_key, api_mode=resolved_api_mode, @@ -4469,90 +4472,37 @@ def _retry_same_provider_sync( retry_kwargs["extra_headers"] = dict(extra_headers) if _is_anthropic_compat_endpoint(resolved_provider, retry_base): retry_kwargs["messages"] = _convert_openai_images_to_anthropic(retry_kwargs["messages"]) + return retry_client, retry_kwargs + + +def _retry_same_provider_sync(*, resolved_provider: str, resolved_api_mode: Optional[str], task: Optional[str], **prep) -> Any: + retry_client, retry_kwargs = _prepare_same_provider_retry( + task=task, resolved_provider=resolved_provider, resolved_api_mode=resolved_api_mode, + async_mode=False, **prep, + ) return _validate_llm_response( _relay_sync_completion( - retry_client, - retry_kwargs, - provider=resolved_provider, - api_mode=resolved_api_mode, + retry_client, retry_kwargs, provider=resolved_provider, api_mode=resolved_api_mode, ), task, ) -async def _retry_same_provider_async( - *, - task: Optional[str], - resolved_provider: str, - resolved_model: Optional[str], - resolved_base_url: Optional[str], - resolved_api_key: Optional[str], - resolved_api_mode: Optional[str], - final_model: Optional[str], - messages: list, - temperature: Optional[float], - max_tokens: Optional[int], - tools: Optional[list], - effective_timeout: float, - effective_extra_body: dict, - reasoning_config: Optional[dict], - extra_headers: Optional[Dict[str, str]] = None, -) -> Any: - if task == "vision": - effective_provider, retry_client, retry_model = resolve_vision_provider_client( - provider=resolved_provider, - model=final_model, - base_url=resolved_base_url, - api_key=resolved_api_key, - async_mode=True, - ) - else: - retry_client, retry_model = _get_cached_client( - resolved_provider, - resolved_model, - async_mode=True, - base_url=resolved_base_url, - api_key=resolved_api_key, - api_mode=resolved_api_mode, - ) - effective_provider = _effective_provider_for_client( - retry_client, resolved_provider, - ) - if retry_client is None: - raise RuntimeError( - f"Auxiliary {task or 'call'}: provider {resolved_provider} could not be rebuilt after recovery" - ) - - retry_base = str(getattr(retry_client, "base_url", "") or "") - retry_kwargs = _build_call_kwargs( - effective_provider or resolved_provider, - retry_model or final_model, - messages, - temperature=temperature, - max_tokens=max_tokens, - tools=tools, - timeout=effective_timeout, - extra_body=effective_extra_body, - reasoning_config=reasoning_config, - base_url=retry_base or resolved_base_url, - task=task, +async def _retry_same_provider_async(*, resolved_provider: str, resolved_api_mode: Optional[str], task: Optional[str], **prep) -> Any: + retry_client, retry_kwargs = _prepare_same_provider_retry( + task=task, resolved_provider=resolved_provider, resolved_api_mode=resolved_api_mode, + async_mode=True, **prep, ) - # Preserve attribution headers across the retry — see the sync variant. - if extra_headers: - retry_kwargs["extra_headers"] = dict(extra_headers) - if _is_anthropic_compat_endpoint(resolved_provider, retry_base): - retry_kwargs["messages"] = _convert_openai_images_to_anthropic(retry_kwargs["messages"]) return _validate_llm_response( await _relay_async_completion( - retry_client, - retry_kwargs, - provider=resolved_provider, - api_mode=resolved_api_mode, + retry_client, retry_kwargs, provider=resolved_provider, api_mode=resolved_api_mode, ), task, ) + + def _refresh_provider_credentials(provider: str) -> bool: """Refresh short-lived credentials for OAuth-backed auxiliary providers.""" normalized = _normalize_aux_provider(provider) @@ -4801,6 +4751,108 @@ def _replan_synchronous_cache_sections( ) +def _fallback_request_kwargs( + destination: _FallbackDestination, + *, + task: Optional[str], + messages: list, + tools: Optional[list], + temperature: Optional[float], + max_tokens: Optional[int], + effective_timeout: float, + effective_extra_body: dict, + reasoning_config: Optional[dict], + fallback_entry: dict, + task_config: dict, + apply_fast_lane: bool, +) -> Dict[str, Any]: + """Build request kwargs for one fallback destination (cache-section replan + fast-lane cap).""" + fallback_max_tokens, fallback_extra_body = max_tokens, effective_extra_body + if apply_fast_lane: + fallback_max_tokens, fallback_extra_body = _compression_fast_lane_controls( + task, + actual_provider=destination.provider, + actual_model=destination.model, + requested_provider=fallback_entry.get("provider"), + requested_model=fallback_entry.get("model"), + route_config=fallback_entry, + leak_guard_config=task_config, + max_tokens=max_tokens, + extra_body=effective_extra_body, + ) + fallback_messages, fallback_tools = _replan_synchronous_cache_sections( + messages, + tools, + destination=destination, + ) + fb_kwargs = _build_call_kwargs( + destination.provider, destination.model, fallback_messages, + temperature=temperature, max_tokens=fallback_max_tokens, + tools=fallback_tools, timeout=effective_timeout, + extra_body=fallback_extra_body, reasoning_config=reasoning_config, + base_url=destination.base_url, task=task) + if apply_fast_lane and fallback_max_tokens is not None and max_tokens is None: + fb_kwargs.update( + auxiliary_max_tokens_param(fallback_max_tokens, model=destination.model) + ) + return fb_kwargs + + +def _plan_fallback_candidate( + fb_client: Any, + fb_model: Optional[str], + fb_label: str, + *, + task: Optional[str], + effective_timeout: float, + apply_fast_lane: bool, + **request, +) -> Tuple[_FallbackDestination, Dict[str, Any], Callable[[str, Any, Optional[str]], Dict[str, Any]]]: + """Resolve the destination + first-attempt kwargs for a fallback candidate. + + Returns ``(destination, kwargs, rebuild)`` where ``rebuild(provider, client, model)`` + produces kwargs for the credential-refreshed retry destination. A configured-chain + entry's own ``timeout`` overrides ``effective_timeout``. + """ + fb_timeout = _fallback_entry_timeout(task, fb_label) + if fb_timeout is not None and fb_timeout != effective_timeout: + logger.info( + "Auxiliary %s: %s using its configured timeout %.0fs " + "(task-level was %.0fs)", + task or "call", fb_label, fb_timeout, effective_timeout, + ) + effective_timeout = fb_timeout + destination = _fallback_destination(task, fb_client, fb_model, fb_label) + task_config = _get_auxiliary_task_config(task) if task == "compression" else {} + fallback_entry = _fallback_chain_entry(task, fb_label) or {} + common = dict( + task=task, effective_timeout=effective_timeout, + fallback_entry=fallback_entry, task_config=task_config, + apply_fast_lane=apply_fast_lane, **request, + ) + + def _rebuild(provider: str, client: Any, model: Optional[str]) -> Tuple[_FallbackDestination, Dict[str, Any]]: + retry_destination = _FallbackDestination( + provider, + destination.base_url or str(getattr(client, "base_url", "") or ""), + destination.api_mode, + model or destination.model, + ) + return retry_destination, _fallback_request_kwargs(retry_destination, **common) + + return destination, _fallback_request_kwargs(destination, **common), _rebuild + + +def _quarantine_fallback_candidate(task: Optional[str], fb_label: str, fb_provider: str, fb_err: Exception, *, tag: str = "") -> None: + """Refresh unavailable or still 401s: token is dead. Quarantine the candidate so the caller moves on.""" + _mark_provider_unhealthy(fb_provider or fb_label) + logger.warning( + "Auxiliary %s%s: fallback candidate %s has a stale/unrefreshable " + "credential (%s) — skipping to next fallback", + task or "call", tag, fb_label, fb_err, + ) + + def _call_fallback_candidate_sync( fb_client: Any, fb_model: Optional[str], @@ -4820,64 +4872,34 @@ def _call_fallback_candidate_sync( On an auth error: refresh the candidate's credentials and retry once with a rebuilt client; if that also auth-fails, mark the provider unhealthy and return ``None`` so the caller moves to the next layer instead of aborting - the task. Non-auth errors raise. A configured-chain entry's own ``timeout`` - overrides ``effective_timeout``. + the task. Non-auth errors raise. """ - fb_timeout = _fallback_entry_timeout(task, fb_label) - if fb_timeout is not None and fb_timeout != effective_timeout: - logger.info( - "Auxiliary %s: %s using its configured timeout %.0fs " - "(task-level was %.0fs)", - task or "call", fb_label, fb_timeout, effective_timeout, - ) - effective_timeout = fb_timeout - destination = _fallback_destination(task, fb_client, fb_model, fb_label) - task_config = _get_auxiliary_task_config(task) if task == "compression" else {} - fallback_entry = _fallback_chain_entry(task, fb_label) or {} - fallback_max_tokens, fallback_extra_body = _compression_fast_lane_controls( - task, - actual_provider=destination.provider, - actual_model=destination.model, - requested_provider=fallback_entry.get("provider"), - requested_model=fallback_entry.get("model"), - route_config=fallback_entry, - leak_guard_config=task_config, - max_tokens=max_tokens, - extra_body=effective_extra_body, + destination, fb_kwargs, rebuild = _plan_fallback_candidate( + fb_client, fb_model, fb_label, task=task, effective_timeout=effective_timeout, + apply_fast_lane=True, messages=messages, tools=tools, temperature=temperature, + max_tokens=max_tokens, effective_extra_body=effective_extra_body, + reasoning_config=reasoning_config, ) - fallback_messages, fallback_tools = _replan_synchronous_cache_sections( - messages, - tools, - destination=destination, - ) - fb_kwargs = _build_call_kwargs( - destination.provider, destination.model, fallback_messages, - temperature=temperature, max_tokens=fallback_max_tokens, - tools=fallback_tools, timeout=effective_timeout, - extra_body=fallback_extra_body, reasoning_config=reasoning_config, - base_url=destination.base_url, task=task) - if fallback_max_tokens is not None and max_tokens is None: - fb_kwargs.update( - auxiliary_max_tokens_param(fallback_max_tokens, model=destination.model) - ) - try: + + def _send(client: Any, request_kwargs: Dict[str, Any], dest: _FallbackDestination) -> Any: return _validate_llm_response( _relay_sync_completion( - fb_client, - fb_kwargs, - provider=destination.provider, - api_mode=destination.api_mode, + client, + request_kwargs, + provider=dest.provider, + api_mode=dest.api_mode, create=lambda request: _create_with_progress( - fb_client, + client, request, task, - force_stream=_provider_requires_stream( - destination.provider, destination.base_url - ), + force_stream=_provider_requires_stream(dest.provider, dest.base_url), ), ), task, ) + + try: + return _send(fb_client, fb_kwargs, destination) except Exception as fb_err: if not _is_auth_error(fb_err): raise @@ -4892,74 +4914,13 @@ def _call_fallback_candidate_sync( api_mode=destination.api_mode, ) if retry_client is not None: - retry_destination = _FallbackDestination( - fb_provider, - destination.base_url - or str(getattr(retry_client, "base_url", "") or ""), - destination.api_mode, - retry_model or destination.model, - ) - retry_messages, retry_tools = _replan_synchronous_cache_sections( - messages, - tools, - destination=retry_destination, - ) - retry_max_tokens, retry_extra_body = _compression_fast_lane_controls( - task, - actual_provider=retry_destination.provider, - actual_model=retry_destination.model, - requested_provider=fallback_entry.get("provider"), - requested_model=fallback_entry.get("model"), - route_config=fallback_entry, - leak_guard_config=task_config, - max_tokens=max_tokens, - extra_body=effective_extra_body, - ) - retry_kwargs = _build_call_kwargs( - retry_destination.provider, - retry_destination.model, - retry_messages, - temperature=temperature, max_tokens=retry_max_tokens, - tools=retry_tools, timeout=effective_timeout, - extra_body=retry_extra_body, - reasoning_config=reasoning_config, - base_url=retry_destination.base_url, task=task) - if retry_max_tokens is not None and max_tokens is None: - retry_kwargs.update( - auxiliary_max_tokens_param( - retry_max_tokens, model=retry_destination.model - ) - ) + retry_destination, retry_kwargs = rebuild(fb_provider, retry_client, retry_model) try: - return _validate_llm_response( - _relay_sync_completion( - retry_client, - retry_kwargs, - provider=retry_destination.provider, - api_mode=retry_destination.api_mode, - create=lambda request: _create_with_progress( - retry_client, - request, - task, - force_stream=_provider_requires_stream( - retry_destination.provider, - retry_destination.base_url, - ), - ), - ), - task, - ) + return _send(retry_client, retry_kwargs, retry_destination) except Exception as retry_err: if not _is_auth_error(retry_err): raise - # Refresh unavailable or still 401s: token is dead. Quarantine the - # candidate and let the caller move on. - _mark_provider_unhealthy(fb_provider or fb_label) - logger.warning( - "Auxiliary %s: fallback candidate %s has a stale/unrefreshable " - "credential (%s) — skipping to next fallback", - task or "call", fb_label, fb_err, - ) + _quarantine_fallback_candidate(task, fb_label, fb_provider, fb_err) return None @@ -4977,37 +4938,27 @@ async def _call_fallback_candidate_async( effective_extra_body: dict, reasoning_config: Optional[dict], ) -> Optional[Any]: - """Async mirror of :func:`_call_fallback_candidate_sync`.""" - fb_timeout = _fallback_entry_timeout(task, fb_label) - if fb_timeout is not None and fb_timeout != effective_timeout: - logger.info( - "Auxiliary %s: %s using its configured timeout %.0fs " - "(task-level was %.0fs)", - task or "call", fb_label, fb_timeout, effective_timeout, - ) - effective_timeout = fb_timeout - destination = _fallback_destination(task, fb_client, fb_model, fb_label) - fallback_messages, fallback_tools = _replan_synchronous_cache_sections( - messages, - tools, - destination=destination, + """Async mirror of :func:`_call_fallback_candidate_sync` (no fast-lane cap on this wire).""" + destination, fb_kwargs, rebuild = _plan_fallback_candidate( + fb_client, fb_model, fb_label, task=task, effective_timeout=effective_timeout, + apply_fast_lane=False, messages=messages, tools=tools, temperature=temperature, + max_tokens=max_tokens, effective_extra_body=effective_extra_body, + reasoning_config=reasoning_config, ) - fb_kwargs = _build_call_kwargs( - destination.provider, destination.model, fallback_messages, - temperature=temperature, max_tokens=max_tokens, - tools=fallback_tools, timeout=effective_timeout, - extra_body=effective_extra_body, reasoning_config=reasoning_config, - base_url=destination.base_url, task=task) - try: + + async def _send(client: Any, request_kwargs: Dict[str, Any], dest: _FallbackDestination) -> Any: return _validate_llm_response( await _relay_async_completion( - fb_client, - fb_kwargs, - provider=destination.provider, - api_mode=destination.api_mode, + client, + request_kwargs, + provider=dest.provider, + api_mode=dest.api_mode, ), task, ) + + try: + return await _send(fb_client, fb_kwargs, destination) except Exception as fb_err: if not _is_auth_error(fb_err): raise @@ -5023,49 +4974,18 @@ async def _call_fallback_candidate_async( api_mode=destination.api_mode, ) if retry_client is not None: - retry_destination = _FallbackDestination( - fb_provider, - destination.base_url - or str(getattr(retry_client, "base_url", "") or ""), - destination.api_mode, - retry_model or destination.model, - ) - retry_messages, retry_tools = _replan_synchronous_cache_sections( - messages, - tools, - destination=retry_destination, - ) - retry_kwargs = _build_call_kwargs( - retry_destination.provider, - retry_destination.model, - retry_messages, - temperature=temperature, max_tokens=max_tokens, - tools=retry_tools, timeout=effective_timeout, - extra_body=effective_extra_body, - reasoning_config=reasoning_config, - base_url=retry_destination.base_url, task=task) + retry_destination, retry_kwargs = rebuild(fb_provider, retry_client, retry_model) try: - return _validate_llm_response( - await _relay_async_completion( - retry_client, - retry_kwargs, - provider=retry_destination.provider, - api_mode=retry_destination.api_mode, - ), - task, - ) + return await _send(retry_client, retry_kwargs, retry_destination) except Exception as retry_err: if not _is_auth_error(retry_err): raise - _mark_provider_unhealthy(fb_provider or fb_label) - logger.warning( - "Auxiliary %s (async): fallback candidate %s has a stale/unrefreshable " - "credential (%s) — skipping to next fallback", - task or "call", fb_label, fb_err, - ) + _quarantine_fallback_candidate(task, fb_label, fb_provider, fb_err, tag=" (async)") return None + + def _try_payment_fallback( failed_provider: str, task: str = None,