refactor(agent/aux): unify same-provider retry and fallback-candidate prep across sync/async (_prepare_same_provider_retry, _plan_fallback_candidate)

This commit is contained in:
Teknium
2026-09-02 11:31:50 -07:00
parent 127a179ae5
commit dcbf263a09

View File

@@ -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,