package com.ankurm.gateway; import java.util.List; import java.util.concurrent.TimeUnit; import io.github.resilience4j.circuitbreaker.CircuitBreaker; import io.github.resilience4j.circuitbreaker.CircuitBreakerRegistry; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.tool.ToolCallingChatOptions; /** * The gateway's {@link ChatModel}: tries each target in order, one circuit breaker per target, * one budget reservation per attempt, and meters every successful call. * *

It sits BELOW the tool-calling advisor on purpose. When a provider fails in the middle of a * tool loop, this class re-sends the conversation (which already holds the tool results) to the * next target. A retry placed above the advisor would start the whole loop again and run the * tools a second time (see {@code ToolLoopPlacementTest}). */ public final class Failover implements ChatModel { private final List targets; private final CircuitBreakerRegistry breakers; private final Budgets budgets; public Failover(List targets, CircuitBreakerRegistry breakers, Budgets budgets) { this.targets = List.copyOf(targets); this.breakers = breakers; this.budgets = budgets; } /** * Must be a {@link ToolCallingChatOptions}: the tool-calling advisor looks at the request's * options and does nothing at all when they are a plain {@code ChatOptions}. */ @Override public ChatOptions getOptions() { return ToolCallingChatOptions.builder().build(); } @Override public ChatResponse call(Prompt prompt) { CallContext ctx = CallContext.CURRENT.isBound() ? CallContext.CURRENT.get() : CallContext.anonymous(); RuntimeException last = null; for (Target t : targets) { CircuitBreaker breaker = breakers.circuitBreaker(t.name()); if (!breaker.tryAcquirePermission()) { ctx.meter().trail(t.name() + ": circuit open, not called"); continue; } long estimate = t.price().micros(estimatePromptTokens(prompt), maxTokens(prompt, t)); Budgets.Reservation reservation; try { reservation = budgets.reserve(ctx.tenant(), estimate); } catch (BudgetExceededException e) { breaker.releasePermission(); throw e; } long started = System.nanoTime(); try { ChatResponse response = t.chat().call(forTarget(prompt, t)); breaker.onSuccess(System.nanoTime() - started, TimeUnit.NANOSECONDS); Usage u = response.getMetadata() == null ? null : response.getMetadata().getUsage(); boolean missing = u == null || u.getTotalTokens() == null || u.getTotalTokens() == 0; long in = missing ? estimatePromptTokens(prompt) : u.getPromptTokens(); long out = missing ? 0 : u.getCompletionTokens(); long actual = t.price().micros(in, out); reservation.settle(actual); ctx.meter().record(t, in, out, actual, missing); ctx.meter().trail(t.name() + ": ok"); return response; } catch (RuntimeException e) { reservation.release(); if (Failures.isTransient(e)) { breaker.onError(System.nanoTime() - started, TimeUnit.NANOSECONDS, e); ctx.meter().trail(t.name() + ": failed with status " + Failures.status(e)); last = e; continue; } breaker.releasePermission(); ctx.meter().trail(t.name() + ": rejected with status " + Failures.status(e) + ", not retried"); throw e; } } throw new AllProvidersUnavailableException(ctx.meter().trail(), last); } /** * The options for ONE target: that provider's own defaults (so the model id is its own, never * another vendor's), overlaid with the portable settings and the tool callbacks the caller set. */ static Prompt forTarget(Prompt prompt, Target t) { ChatOptions in = prompt.getOptions(); ToolCallingChatOptions.Builder b = (ToolCallingChatOptions.Builder) t.chat().getOptions().mutate(); b.model(t.model()); if (in != null) { if (in.getMaxTokens() != null) { b.maxTokens(in.getMaxTokens()); } if (in.getTemperature() != null) { b.temperature(in.getTemperature()); } if (in.getTopP() != null) { b.topP(in.getTopP()); } if (in.getStopSequences() != null) { b.stopSequences(in.getStopSequences()); } if (in instanceof ToolCallingChatOptions tools) { b.toolCallbacks(tools.getToolCallbacks()); b.toolContext(tools.getToolContext()); } } return prompt.mutate().chatOptions(b.build()).build(); } static long estimatePromptTokens(Prompt prompt) { long chars = 0; for (Message m : prompt.getInstructions()) { chars += m.getText() == null ? 0 : m.getText().length(); } return (chars + 3) / 4; } private static long maxTokens(Prompt prompt, Target t) { ChatOptions in = prompt.getOptions(); return in != null && in.getMaxTokens() != null ? in.getMaxTokens() : 1024; } }