open-swe/agent/middleware/model_fallback.py
seahaven-openswe[bot] e3c7c03c92
feat: port model-fallback resilience from upstream (#1694, #1695) (#161)
* feat: port plan-review & workflow-approval UX (#135)

Port six upstream commits onto dev:

- c03a6be7 (already ported): keep plan guidance high-level
- 546042a4: add workflow approval UI with diff preview, approval URLs,
  web review links, and polling for approval status during active runs
- 216cf181: remove workflow token elevation; approved pushes pass
  through directly without proxy token rewriting
- 3dbc0282: preserve plan redirects after login by accepting relative
  same-origin redirect_to values and rejecting blocked paths
- bb104d93: submit plan comments with cmd+enter
- 90cb6caa: terse Slack replies, shared content via save_plan outside
  plan mode (PLAN_STATUS_SHARED), reject shared-content mutations

Refs: #135

* feat: port durable dispatch hardening and startup latency improvements

Port five upstream PRs onto dev:

- #1621 / #1658: durable dispatch with loopback webhook defense,
  create_durable_run helper, _config_with_prepare_run_id, degradation
  to None for relative/loopback completion webhook URLs
- #1696: run-level completion webhook deduplication (replace
  claim-then-post with post-then-flag per run_id), DeferredErrorModel
  for graph-factory resilience, ToolRetryMiddleware for task subagents,
  TimeoutWrapupMiddleware for all three graphs
- #1697: lazy-load __init__.py for agent.middleware, agent.tools,
  agent.dashboard (PEP 562); defer heavy imports (exa_py in web_search,
  agent.webapp in request_pr_review, deepagents in sandbox.py); add
  ttl_cache.py with stale-while-revalidate for tool loaders

Refs: #137

* fix: restore login page render and clear CI lint/format

The plan-review port removed the authRedirectUrl import from login.tsx
but left its call site, crashing the login page at runtime (blank page,
no 'Sign in to open-swe'). Pass the relative path straight to loginUrl,
matching the plan route and the backend relative-redirect handling.

Also drop an unused os import in the guard test and reformat
workflow_push_guard.py to satisfy ruff.

* feat: port model-fallback resilience from upstream (#1694, #1695)

- Add httpx.TransportError to the transient-exception set so
  incomplete chunked reads on streamed responses trigger a
  fallback instead of cancelling the run (#1694 / c9f6dd86).
- Rewrite fallback to alternate primary/fallback with exponential
  backoff instead of a single failover, so the agent survives
  multi-minute gateway outages spanning both providers (#1695 /
  c9a9a7cd).
- Default backoff schedule (0, 5, 15, 30, 45) reaches past the
  gateway's ~30s recovery window; jittered ±25%.
- On exhaustion, surface a terminal AIMessage explaining the
  outage instead of crashing — progress is checkpointed so the
  user can retrigger to continue.
- Preserve fork conventions: Bedrock ClientError retryability
  check, sync wrap_model_call (using time.sleep instead of
  asyncio.sleep), and existing access-error surfacing for
  Anthropic, OpenAI, and Bedrock (botocore) provider errors.
- Update triage ledger (c9f6dd86, c9a9a7cd → landed) and
  re-render triage.md.

Refs: #139

* fix: align workflow-push-guard tests with dev's transient-elevation impl

The dev merge auto-combined dev's elevation tests with the stale passthrough
tests inherited from the durable-dispatch branch; the passthrough tests
contradict dev's restored _run_with_workflow_token impl. Take dev's test file.

* fix: drop dead ttl_cache module; make fallback backoff jitter two-sided

ttl_cache.py was re-introduced via the dev merge but dev/#160 deliberately
removed it as dead code (no agent importer). Remove it to match dev.
Also make _jittered_delay symmetric (±25%) to match its docstring.

* chore(upstream-sync): triage 4 new upstream commits (#1708-#1713)

Synced ledger to upstream/main (71e3b818). New rows all deferred:
- #1708 add GPT-5.6 OpenAI models (FLAG-HUMAN: fork picker is Bedrock/Fireworks-only)
- #1709 stale admin model defaults after upgrades
- #1710 bump langchain-fireworks 1.4.4
- #1713 align reviewer eval with published findings

---------

Co-authored-by: amoussa1229 <166072409+amoussa1229@users.noreply.github.com>
Co-authored-by: Adam Moussa <adam@seahavenind.com>
2026-07-09 18:28:21 -04:00

284 lines
12 KiB
Python

"""Middleware that retries model calls across a primary and fallback provider.
Wraps the model call. When a model raises a transient provider error (5xx,
429, connection/timeout), the request is retried, alternating between the
primary and the configured fallback model with exponential backoff between
attempts. The fallback is bound to tools by the agent factory on each call,
so swapping ``request.model`` is sufficient.
Why alternate with backoff instead of failing over once: both providers can
be routed through the same LLM Gateway, so a gateway outage takes out the
"cross-provider" fallback too. A single immediate failover cannot ride out
even a short shared outage (the gateway's 502 page literally says "try again
in 30 seconds"), and an unprotected fallback call crashes the whole run.
Alternating with a backoff schedule that reaches past 30s lets a long-running
agent run survive multi-minute provider or gateway blips.
Bidirectional: if the primary is Anthropic the fallback is typically OpenAI,
and vice versa. The middleware itself is provider-agnostic — it inspects the
exception type/status code to decide whether an attempt is retryable.
If every attempt fails, the middleware either raises the last error or (by
default) returns a terminal ``AIMessage`` explaining the outage, so the run
ends with a visible message in Slack/GitHub instead of an abrupt crash. The
turn's progress is checkpointed, so the user can retrigger to continue.
"""
from __future__ import annotations
import asyncio
import logging
import random
import time
from collections.abc import Awaitable, Callable, Sequence
from typing import Any
import anthropic
import httpx
import openai
from botocore.exceptions import ClientError
from langchain.agents.middleware import AgentMiddleware
from langchain.agents.middleware.types import ModelCallResult, ModelRequest, ModelResponse
from langchain_core.language_models import BaseChatModel
from langchain_core.messages import AIMessage
logger = logging.getLogger(__name__)
_RETRYABLE_STATUS_CODES = {408, 409, 425, 429, 500, 502, 503, 504, 529}
_TRANSIENT_EXCEPTIONS: tuple[type[BaseException], ...] = (
anthropic.APIConnectionError,
anthropic.APITimeoutError,
anthropic.RateLimitError,
anthropic.InternalServerError,
openai.APIConnectionError,
openai.APITimeoutError,
openai.RateLimitError,
openai.InternalServerError,
httpx.TransportError,
)
_RETRYABLE_BEDROCK_ERROR_CODES = {
"ThrottlingException",
"ServiceUnavailableException",
"ModelTimeoutException",
"InternalServerException",
}
# Seconds slept before each retry attempt (attempt 0 is the initial call).
# The first failover is immediate: a provider-specific outage should not delay
# the cross-provider retry. Later delays grow past the ~30s the gateway's 502
# page asks for. Each attempt additionally benefits from the SDK's own
# ``max_retries`` backoff, so worst-case wall time before giving up is a few
# minutes — acceptable for a long-running agent, far better than crashing.
DEFAULT_BACKOFF_SCHEDULE: tuple[float, ...] = (0.0, 5.0, 15.0, 30.0, 45.0)
MODEL_OUTAGE_MESSAGE = (
"I wasn't able to reach the language model providers after several retries "
"(both the primary and fallback models returned transient errors, e.g. "
"502/503/overloaded). This is a temporary provider or gateway outage, not a "
"problem with your task. My progress so far has been saved — please retrigger "
"the run in a few minutes to continue."
)
def _should_fallback(exc: BaseException) -> bool:
if isinstance(exc, _TRANSIENT_EXCEPTIONS):
return True
# Catches OverloadedError (529) and other 5xx/429 surfaced as APIStatusError.
if isinstance(exc, (anthropic.APIStatusError, openai.APIStatusError)):
status = getattr(exc, "status_code", None)
if isinstance(status, int) and status in _RETRYABLE_STATUS_CODES:
return True
# Bedrock (Claude) raises botocore ClientError for transient throttling/5xx.
if isinstance(exc, ClientError):
code = exc.response.get("Error", {}).get("Code", "")
if code in _RETRYABLE_BEDROCK_ERROR_CODES:
return True
return False
def _error_body(exc: BaseException) -> dict[str, Any]:
body = getattr(exc, "body", None)
return body if isinstance(body, dict) else {}
def _nested_str(data: dict[str, Any], *keys: str) -> str | None:
current: Any = data
for key in keys:
if not isinstance(current, dict):
return None
current = current.get(key)
return current if isinstance(current, str) and current else None
def _provider_access_error_message(exc: BaseException) -> str | None:
if isinstance(exc, anthropic.BadRequestError):
body = _error_body(exc)
error_code = _nested_str(body, "error", "details", "error_code")
if error_code == "model_not_available":
provider_message = _nested_str(body, "error", "message") or str(exc)
return (
"The selected Anthropic model is not available to this workspace. "
f"Anthropic returned: {provider_message} "
"Choose a different model or update the workspace's Anthropic access and retry."
)
if isinstance(exc, (openai.BadRequestError, openai.NotFoundError)):
body = _error_body(exc)
error_code = _nested_str(body, "error", "code")
if error_code in {"model_not_found", "model_not_available"}:
provider_message = _nested_str(body, "error", "message") or str(exc)
return (
"The selected OpenAI model is not available to this workspace. "
f"OpenAI returned: {provider_message} "
"Choose a different model or update the workspace's OpenAI access and retry."
)
# Bedrock access/lookup failures embed the caller's role ARN and account id in the
# raw botocore message; surface only the error code so identifiers never reach logs
# or the user-facing channel.
if isinstance(exc, ClientError):
code = exc.response.get("Error", {}).get("Code", "")
if code in {"AccessDeniedException", "ResourceNotFoundException"}:
return (
"The selected Bedrock model is not available to this deployment "
f"(Bedrock error: {code}). Verify the model's inference-profile access and "
"IAM permissions, choose a different model, and retry."
)
return None
class ModelFallbackMiddleware(AgentMiddleware):
"""Retry the model call across primary and fallback providers on transient errors.
Args:
fallback_model: Cross-provider model used on odd-numbered attempts.
backoff_schedule: Seconds slept before each retry. ``len(schedule) + 1``
is the total number of attempts. Delays get ±25% jitter.
surface_outage_message: When all attempts fail, return a terminal
``AIMessage`` describing the outage instead of raising, so the run
ends gracefully with a user-visible message rather than a crash.
Set to ``False`` to re-raise the last error (e.g. if platform-level
alerting keys off failed runs).
"""
def __init__(
self,
fallback_model: BaseChatModel,
*,
backoff_schedule: Sequence[float] = DEFAULT_BACKOFF_SCHEDULE,
surface_outage_message: bool = True,
) -> None:
super().__init__()
self._fallback_model = fallback_model
self._backoff_schedule = tuple(backoff_schedule)
self._surface_outage_message = surface_outage_message
def _fallback_name(self) -> str:
return (
getattr(self._fallback_model, "model_name", None)
or getattr(self._fallback_model, "model", None)
or "fallback"
)
def _jittered_delay(self, delay: float) -> float:
return delay + random.uniform(-delay * 0.25, delay * 0.25) if delay > 0 else 0.0
def _log_retry(
self, exc_type: str, use_fallback: bool, attempt: int, total: int, delay: float
) -> None:
logger.warning(
"Model call failed transiently (%s) on %s model "
"(attempt %d/%d); retrying %s model in %.1fs",
exc_type,
"fallback" if use_fallback else "primary",
attempt + 1,
total,
"primary" if use_fallback else f"fallback ({self._fallback_name()})",
delay,
)
def _resolve_request(self, request: ModelRequest, use_fallback: bool) -> ModelRequest:
return request.override(model=self._fallback_model) if use_fallback else request
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelCallResult:
total_attempts = len(self._backoff_schedule) + 1
last_exc: BaseException | None = None
for attempt in range(total_attempts):
use_fallback = attempt % 2 == 1
attempt_request = self._resolve_request(request, use_fallback)
try:
return handler(attempt_request)
except Exception as exc:
access_error_message = _provider_access_error_message(exc)
if access_error_message is not None:
logger.warning("Model access error surfaced to user: %s", type(exc).__name__)
return AIMessage(content=access_error_message)
if not _should_fallback(exc):
raise
last_exc = exc
if attempt + 1 >= total_attempts:
break
delay = self._jittered_delay(self._backoff_schedule[attempt])
self._log_retry(type(exc).__name__, use_fallback, attempt, total_attempts, delay)
if delay > 0:
time.sleep(delay)
assert last_exc is not None
logger.error(
"Model call failed after %d attempts across primary and fallback (%s): %s",
total_attempts,
self._fallback_name(),
last_exc,
)
if self._surface_outage_message:
return AIMessage(content=MODEL_OUTAGE_MESSAGE)
raise last_exc
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> Any:
total_attempts = len(self._backoff_schedule) + 1
last_exc: BaseException | None = None
for attempt in range(total_attempts):
use_fallback = attempt % 2 == 1
attempt_request = self._resolve_request(request, use_fallback)
try:
return await handler(attempt_request)
except Exception as exc:
access_error_message = _provider_access_error_message(exc)
if access_error_message is not None:
logger.warning("Model access error surfaced to user: %s", type(exc).__name__)
return AIMessage(content=access_error_message)
if not _should_fallback(exc):
raise
last_exc = exc
if attempt + 1 >= total_attempts:
break
delay = self._jittered_delay(self._backoff_schedule[attempt])
self._log_retry(type(exc).__name__, use_fallback, attempt, total_attempts, delay)
if delay > 0:
await asyncio.sleep(delay)
assert last_exc is not None
logger.error(
"Model call failed after %d attempts across primary and fallback (%s): %s",
total_attempts,
self._fallback_name(),
last_exc,
)
if self._surface_outage_message:
return AIMessage(content=MODEL_OUTAGE_MESSAGE)
raise last_exc