Packages

ExClaw — OpenClaw rebuilt on ADK Elixir

Current section

Files

Jump to
ex_claw lib ex_claw llm cascade.ex
Raw

lib/ex_claw/llm/cascade.ex

defmodule ExClaw.LLM.Cascade do
@moduledoc """
LLM provider cascade with exponential backoff.
Global tracking of provider health, per-agent cascade configs.
When a provider fails, automatically tries the next in the cascade
with sticky backoff (don't hammer a broken provider).
## Per-Agent Config
Each agent can define its own cascade strategy:
%{
providers: ["anthropic/claude-sonnet-4", "google/gemini-2.0-flash", "openai/gpt-4o"],
strategy: :ordered, # :ordered | :least_latency | :cheapest
max_retries: 2
}
## Global Health Tracking
Provider health is tracked globally — if Claude is down, ALL agents
see it and skip to the next provider until backoff expires.
"""
use GenServer
require Logger
@type provider_state :: %{
name: String.t(),
healthy: boolean(),
consecutive_failures: non_neg_integer(),
last_failure: DateTime.t() | nil,
last_success: DateTime.t() | nil,
backoff_until: DateTime.t() | nil,
avg_latency_ms: float(),
success_count: non_neg_integer(),
failure_count: non_neg_integer()
}
@type cascade_config :: %{
providers: [String.t()],
strategy: :ordered | :least_latency | :cheapest,
max_retries: non_neg_integer()
}
# 30 minutes
@max_backoff_ms 30 * 60 * 1000
# 5 seconds
@base_backoff_ms 5_000
# --- Public API ---
def start_link(opts \\ []) do
GenServer.start_link(__MODULE__, opts, name: __MODULE__)
end
@doc "Get the next provider to try for a given agent cascade config."
@spec next_provider(cascade_config()) :: {:ok, String.t()} | {:error, :all_providers_down}
def next_provider(config) do
GenServer.call(__MODULE__, {:next_provider, config})
end
@doc "Report a successful call to a provider."
@spec report_success(String.t(), non_neg_integer()) :: :ok
def report_success(provider, latency_ms) do
GenServer.call(__MODULE__, {:success, provider, latency_ms})
end
@doc "Report a failed call to a provider."
@spec report_failure(String.t(), term()) :: :ok
def report_failure(provider, reason) do
GenServer.call(__MODULE__, {:failure, provider, reason})
end
@doc "Get health status of all known providers."
@spec provider_health() :: %{String.t() => provider_state()}
def provider_health do
GenServer.call(__MODULE__, :health)
end
@doc "Manually reset a provider's backoff (e.g. after fixing an API key)."
@spec reset_provider(String.t()) :: :ok
def reset_provider(provider) do
GenServer.call(__MODULE__, {:reset, provider})
end
# --- GenServer ---
@impl true
def init(_opts) do
{:ok, %{providers: %{}}}
end
@impl true
def handle_call({:next_provider, config}, _from, state) do
now = DateTime.utc_now()
available =
config.providers
|> Enum.filter(fn name ->
case Map.get(state.providers, name) do
# unknown = assume healthy
nil -> true
%{backoff_until: nil} -> true
%{backoff_until: until} -> DateTime.compare(now, until) != :lt
end
end)
# Apply strategy
chosen =
case {config.strategy, available} do
{_, []} ->
nil
{:ordered, providers} ->
List.first(providers)
{:least_latency, providers} ->
Enum.min_by(providers, fn name ->
case Map.get(state.providers, name) do
%{avg_latency_ms: lat} when lat > 0 -> lat
_ -> 999_999
end
end)
{:cheapest, providers} ->
# For now, prefer order (cost estimation TBD)
List.first(providers)
end
if chosen do
{:reply, {:ok, chosen}, state}
else
{:reply, {:error, :all_providers_down}, state}
end
end
def handle_call(:health, _from, state) do
{:reply, state.providers, state}
end
def handle_call({:reset, provider}, _from, state) do
updated =
Map.update(state.providers, provider, new_provider(provider), fn ps ->
%{ps | healthy: true, consecutive_failures: 0, backoff_until: nil}
end)
{:reply, :ok, %{state | providers: updated}}
end
def handle_call({:success, provider, latency_ms}, _from, state) do
ps = Map.get(state.providers, provider, new_provider(provider))
new_avg =
if ps.success_count > 0 do
(ps.avg_latency_ms * ps.success_count + latency_ms) / (ps.success_count + 1)
else
latency_ms * 1.0
end
updated_ps = %{
ps
| healthy: true,
consecutive_failures: 0,
last_success: DateTime.utc_now(),
backoff_until: nil,
avg_latency_ms: new_avg,
success_count: ps.success_count + 1
}
updated = Map.put(state.providers, provider, updated_ps)
{:reply, :ok, %{state | providers: updated}}
end
def handle_call({:failure, provider, _reason}, _from, state) do
ps = Map.get(state.providers, provider, new_provider(provider))
failures = ps.consecutive_failures + 1
backoff_ms = min((@base_backoff_ms * :math.pow(2, failures - 1)) |> trunc(), @max_backoff_ms)
backoff_until = DateTime.add(DateTime.utc_now(), backoff_ms, :millisecond)
Logger.warning(
"[Cascade] #{provider} failed (#{failures} consecutive). " <>
"Backing off until #{DateTime.to_iso8601(backoff_until)}"
)
updated_ps = %{
ps
| healthy: false,
consecutive_failures: failures,
last_failure: DateTime.utc_now(),
backoff_until: backoff_until,
failure_count: ps.failure_count + 1
}
updated = Map.put(state.providers, provider, updated_ps)
{:reply, :ok, %{state | providers: updated}}
end
defp new_provider(name) do
%{
name: name,
healthy: true,
consecutive_failures: 0,
last_failure: nil,
last_success: nil,
backoff_until: nil,
avg_latency_ms: 0.0,
success_count: 0,
failure_count: 0
}
end
end