Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
139 lines
5.7 KiB
Java
139 lines
5.7 KiB
Java
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.
|
|
*
|
|
* <p>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<Target> targets;
|
|
|
|
private final CircuitBreakerRegistry breakers;
|
|
|
|
private final Budgets budgets;
|
|
|
|
public Failover(List<Target> 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;
|
|
}
|
|
}
|