trail() {
+ return List.copyOf(trail);
+ }
+
+ public synchronized long promptTokens() {
+ return promptTokens;
+ }
+
+ public synchronized long completionTokens() {
+ return completionTokens;
+ }
+
+ public synchronized long costMicros() {
+ return costMicros;
+ }
+
+ public synchronized int modelCalls() {
+ return modelCalls;
+ }
+
+ /** The provider that answered last. */
+ public synchronized String provider() {
+ return provider;
+ }
+
+ public synchronized String model() {
+ return model;
+ }
+
+ public synchronized boolean usageMissing() {
+ return usageMissing;
+ }
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Failover.java b/llm-gateway/src/main/java/com/ankurm/gateway/Failover.java
new file mode 100644
index 0000000..2fda99a
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Failover.java
@@ -0,0 +1,138 @@
+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;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Failures.java b/llm-gateway/src/main/java/com/ankurm/gateway/Failures.java
new file mode 100644
index 0000000..f17e64f
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Failures.java
@@ -0,0 +1,48 @@
+package com.ankurm.gateway;
+
+import java.io.IOException;
+
+import com.anthropic.errors.AnthropicIoException;
+import com.anthropic.errors.AnthropicServiceException;
+import com.openai.errors.OpenAIIoException;
+import com.openai.errors.OpenAIServiceException;
+
+/**
+ * Decides whether a failed model call is worth sending to another provider. Same idea as the
+ * {@code providers} module: the vendor SDKs throw unrelated exception families, so walk the cause
+ * chain for an HTTP status.
+ */
+public final class Failures {
+
+ private Failures() {
+ }
+
+ public static int status(Throwable t) {
+ for (Throwable c = t; c != null; c = c.getCause()) {
+ if (c instanceof ProviderFailure e) {
+ return e.status();
+ }
+ if (c instanceof OpenAIServiceException e) {
+ return e.statusCode();
+ }
+ if (c instanceof AnthropicServiceException e) {
+ return e.statusCode();
+ }
+ }
+ return -1;
+ }
+
+ /** True for overload, rate limit, 5xx and network failures. A 400 or 401 is the caller's problem. */
+ public static boolean isTransient(Throwable t) {
+ int status = status(t);
+ if (status != -1) {
+ return status == 408 || status == 429 || status >= 500;
+ }
+ for (Throwable c = t; c != null; c = c.getCause()) {
+ if (c instanceof OpenAIIoException || c instanceof AnthropicIoException || c instanceof IOException) {
+ return true;
+ }
+ }
+ return false;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Gateway.java b/llm-gateway/src/main/java/com/ankurm/gateway/Gateway.java
new file mode 100644
index 0000000..3332f91
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Gateway.java
@@ -0,0 +1,98 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+import java.util.Map;
+
+import org.springframework.ai.chat.client.ChatClient;
+import org.springframework.ai.chat.client.advisor.ToolCallingAdvisor;
+import org.springframework.ai.chat.model.ChatResponse;
+import org.springframework.ai.support.ToolCallbacks;
+import org.springframework.ai.tool.ToolCallback;
+
+/**
+ * The one method every calling service uses: rate limit, cache, route, record. Each request gets
+ * its own {@link Failover} over the targets for its hint, wrapped in a {@code ChatClient} whose
+ * tool loop is the standard {@link ToolCallingAdvisor}.
+ */
+public final class Gateway {
+
+ private final Routes routes;
+
+ private final java.util.function.Function, Failover> failovers;
+
+ private final TokenLimiter limiter;
+
+ private final SemanticCache cache;
+
+ private final Map toolbox;
+
+ public Gateway(Routes routes, java.util.function.Function, Failover> failovers,
+ TokenLimiter limiter, SemanticCache cache, Map toolObjects) {
+ this.routes = routes;
+ this.failovers = failovers;
+ this.limiter = limiter;
+ this.cache = cache;
+ this.toolbox = new java.util.LinkedHashMap<>();
+ toolObjects.forEach((name, obj) -> toolbox.put(name, ToolCallbacks.from(obj)));
+ }
+
+ public GatewayResponse complete(String tenant, GatewayRequest req) {
+ List targets = routes.get(req.modelHint());
+ List callbacks = resolveTools(req.tools());
+
+ long estimate = (req.systemPrompt().length() + req.userMessage().length() + 3) / 4 + req.maxTokens();
+ limiter.take(tenant, estimate);
+ CallContext.Meter meter = new CallContext.Meter();
+ try {
+ var cached = cache == null ? java.util.Optional.empty()
+ : cache.lookup(tenant, req.featureTag(), req.userMessage());
+ if (cached.isPresent()) {
+ limiter.settle(tenant, estimate, 0);
+ return new GatewayResponse(cached.get(), "cache", "cache", 0, 0, 0, 0, true, List.of("cache: hit"));
+ }
+ ChatClient client = ChatClient.builder(failovers.apply(targets))
+ .defaultAdvisors(ToolCallingAdvisor.builder().build())
+ .build();
+ CallContext ctx = new CallContext(tenant, req.featureTag(), meter);
+ ChatResponse response = ScopedValue.where(CallContext.CURRENT, ctx).call(() -> {
+ var spec = client.prompt().user(req.userMessage())
+ .options(org.springframework.ai.model.tool.ToolCallingChatOptions.builder()
+ .maxTokens(req.maxTokens()));
+ if (!req.systemPrompt().isBlank()) {
+ spec = spec.system(req.systemPrompt());
+ }
+ if (!callbacks.isEmpty()) {
+ spec = spec.toolCallbacks(callbacks);
+ }
+ return spec.call().chatResponse();
+ });
+ String content = response.getResult().getOutput().getText();
+ limiter.settle(tenant, estimate, meter.promptTokens() + meter.completionTokens());
+ if (cache != null) {
+ cache.put(tenant, req.featureTag(), req.userMessage(), content);
+ }
+ return new GatewayResponse(content, meter.provider(), meter.model(), meter.promptTokens(),
+ meter.completionTokens(), meter.costMicros(), meter.modelCalls(), false, meter.trail());
+ }
+ catch (RuntimeException e) {
+ limiter.settle(tenant, estimate, meter.promptTokens() + meter.completionTokens());
+ throw e;
+ }
+ catch (Exception e) {
+ throw new IllegalStateException(e);
+ }
+ }
+
+ private List resolveTools(List names) {
+ java.util.ArrayList out = new java.util.ArrayList<>();
+ for (String n : names) {
+ ToolCallback[] cb = toolbox.get(n);
+ if (cb == null) {
+ throw new IllegalArgumentException("Tool '" + n + "' is not on the gateway allow-list "
+ + toolbox.keySet());
+ }
+ out.addAll(List.of(cb));
+ }
+ return out;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayApplication.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayApplication.java
new file mode 100644
index 0000000..d767459
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayApplication.java
@@ -0,0 +1,12 @@
+package com.ankurm.gateway;
+
+import org.springframework.boot.SpringApplication;
+import org.springframework.boot.autoconfigure.SpringBootApplication;
+
+@SpringBootApplication
+public class GatewayApplication {
+
+ public static void main(String[] args) {
+ SpringApplication.run(GatewayApplication.class, args);
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayConfig.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayConfig.java
new file mode 100644
index 0000000..054560c
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayConfig.java
@@ -0,0 +1,82 @@
+package com.ankurm.gateway;
+
+import java.time.Duration;
+import java.util.LinkedHashMap;
+import java.util.List;
+import java.util.Map;
+
+import io.github.bucket4j.TimeMeter;
+import io.github.resilience4j.circuitbreaker.CircuitBreakerConfig;
+import io.github.resilience4j.circuitbreaker.CircuitBreakerRegistry;
+import org.springframework.ai.anthropic.AnthropicChatModel;
+import org.springframework.ai.anthropic.AnthropicChatOptions;
+import org.springframework.ai.chat.model.ChatModel;
+import org.springframework.ai.embedding.EmbeddingModel;
+import org.springframework.ai.openai.OpenAiChatModel;
+import org.springframework.ai.openai.OpenAiChatOptions;
+import org.springframework.beans.factory.ObjectProvider;
+import org.springframework.boot.context.properties.EnableConfigurationProperties;
+import org.springframework.context.annotation.Bean;
+import org.springframework.context.annotation.Configuration;
+
+@Configuration
+@EnableConfigurationProperties(GatewayProperties.class)
+public class GatewayConfig {
+
+ @Bean
+ Routes routes(GatewayProperties p) {
+ Map byName = new LinkedHashMap<>();
+ for (GatewayProperties.Provider c : p.providers()) {
+ byName.put(c.name(), new Target(c.name(), c.model(), chatModel(c), Price.of(c.inputPrice(), c.outputPrice())));
+ }
+ Map> table = new LinkedHashMap<>();
+ p.routes().forEach((hint, names) -> table.put(hint, names.stream().map(byName::get).toList()));
+ return new Routes(table);
+ }
+
+ /** No retries inside the SDK: the gateway's own failover is the retry policy. */
+ public static ChatModel chatModel(GatewayProperties.Provider c) {
+ return switch (c.type()) {
+ case "openai" -> OpenAiChatModel.builder().options(OpenAiChatOptions.builder()
+ .baseUrl(c.baseUrl()).apiKey(c.apiKey()).model(c.model()).maxRetries(0).build()).build();
+ case "anthropic" -> AnthropicChatModel.builder().options(AnthropicChatOptions.builder()
+ .baseUrl(c.baseUrl()).apiKey(c.apiKey()).model(c.model()).maxTokens(1024).maxRetries(0).build())
+ .build();
+ default -> throw new IllegalArgumentException("Unknown provider type " + c.type());
+ };
+ }
+
+ @Bean
+ CircuitBreakerRegistry breakers(GatewayProperties p) {
+ return registry(p.breaker());
+ }
+
+ /** One breaker per provider name, all built from the same settings. */
+ public static CircuitBreakerRegistry registry(GatewayProperties.Breaker b) {
+ return CircuitBreakerRegistry.of(CircuitBreakerConfig.custom()
+ .slidingWindowSize(b.slidingWindowSize())
+ .minimumNumberOfCalls(b.minimumCalls())
+ .failureRateThreshold(b.failureRateThreshold())
+ .waitDurationInOpenState(Duration.ofMillis(b.waitOpenMillis()))
+ .permittedNumberOfCallsInHalfOpenState(1)
+ .build());
+ }
+
+ @Bean
+ Budgets budgets(GatewayProperties p) {
+ return new Budgets(t -> p.tenant(t).budgetMicros());
+ }
+
+ @Bean
+ TokenLimiter tokenLimiter(GatewayProperties p) {
+ return new TokenLimiter(t -> p.tenant(t).tokensPerHour(), TimeMeter.SYSTEM_MILLISECONDS);
+ }
+
+ @Bean
+ Gateway gateway(Routes routes, CircuitBreakerRegistry breakers, Budgets budgets, TokenLimiter limiter,
+ ObjectProvider embeddings) {
+ EmbeddingModel em = embeddings.getIfAvailable();
+ return new Gateway(routes, targets -> new Failover(targets, breakers, budgets), limiter,
+ em == null ? null : new SemanticCache(em, 0.92), Map.of("refund_order", new OrderTools()));
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayController.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayController.java
new file mode 100644
index 0000000..0cef3e2
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayController.java
@@ -0,0 +1,45 @@
+package com.ankurm.gateway;
+
+import org.springframework.http.HttpStatus;
+import org.springframework.http.ResponseEntity;
+import org.springframework.web.bind.annotation.ExceptionHandler;
+import org.springframework.web.bind.annotation.PostMapping;
+import org.springframework.web.bind.annotation.RequestBody;
+import org.springframework.web.bind.annotation.RequestHeader;
+import org.springframework.web.bind.annotation.RequestMapping;
+import org.springframework.web.bind.annotation.RestController;
+
+/**
+ * Deploy inside the private network only. The tenant comes from a header set by the mesh or the
+ * API gateway upstream; the request body has no tenant field to spoof.
+ */
+@RestController
+@RequestMapping("/v1/gateway")
+public class GatewayController {
+
+ private final Gateway gateway;
+
+ GatewayController(Gateway gateway) {
+ this.gateway = gateway;
+ }
+
+ @PostMapping("/complete")
+ GatewayResponse complete(@RequestHeader("X-Tenant-Id") String tenant, @RequestBody GatewayRequest request) {
+ return gateway.complete(tenant, request);
+ }
+
+ @ExceptionHandler(BudgetExceededException.class)
+ ResponseEntity overBudget(BudgetExceededException e) {
+ return ResponseEntity.status(HttpStatus.TOO_MANY_REQUESTS).body(e.getMessage());
+ }
+
+ @ExceptionHandler(AllProvidersUnavailableException.class)
+ ResponseEntity down(AllProvidersUnavailableException e) {
+ return ResponseEntity.status(HttpStatus.SERVICE_UNAVAILABLE).header("Retry-After", "30").body(e.getMessage());
+ }
+
+ @ExceptionHandler({UnknownHintException.class, IllegalArgumentException.class})
+ ResponseEntity badRequest(RuntimeException e) {
+ return ResponseEntity.badRequest().body(e.getMessage());
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayProperties.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayProperties.java
new file mode 100644
index 0000000..200db67
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayProperties.java
@@ -0,0 +1,26 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+import java.util.Map;
+
+import org.springframework.boot.context.properties.ConfigurationProperties;
+
+@ConfigurationProperties("gateway")
+public record GatewayProperties(List providers, Map> routes,
+ Map tenants, Tenant defaults, Breaker breaker) {
+
+ /** {@code type} is {@code openai} or {@code anthropic}; prices are dollars per million tokens. */
+ public record Provider(String name, String type, String baseUrl, String apiKey, String model,
+ String inputPrice, String outputPrice) {
+ }
+
+ public record Tenant(long budgetMicros, long tokensPerHour) {
+ }
+
+ public record Breaker(int slidingWindowSize, int minimumCalls, int failureRateThreshold, long waitOpenMillis) {
+ }
+
+ public Tenant tenant(String id) {
+ return tenants != null && tenants.containsKey(id) ? tenants.get(id) : defaults;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayRequest.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayRequest.java
new file mode 100644
index 0000000..c5c4474
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayRequest.java
@@ -0,0 +1,26 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+
+/**
+ * What a calling service sends. The tenant is NOT here: it comes from a header set upstream.
+ * {@code tools} are names from the gateway's own allow-list, never code from the caller.
+ */
+public record GatewayRequest(String featureTag, String systemPrompt, String userMessage, String modelHint,
+ Integer maxTokens, List tools) {
+
+ public GatewayRequest {
+ if (modelHint == null) {
+ modelHint = "smart";
+ }
+ if (maxTokens == null || maxTokens <= 0) {
+ maxTokens = 1024;
+ }
+ if (tools == null) {
+ tools = List.of();
+ }
+ if (systemPrompt == null) {
+ systemPrompt = "";
+ }
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/GatewayResponse.java b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayResponse.java
new file mode 100644
index 0000000..1442314
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/GatewayResponse.java
@@ -0,0 +1,7 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+
+public record GatewayResponse(String content, String provider, String model, long promptTokens,
+ long completionTokens, long costMicros, int modelCalls, boolean servedFromCache, List trail) {
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/OrderTools.java b/llm-gateway/src/main/java/com/ankurm/gateway/OrderTools.java
new file mode 100644
index 0000000..a450e66
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/OrderTools.java
@@ -0,0 +1,17 @@
+package com.ankurm.gateway;
+
+import java.util.concurrent.atomic.AtomicInteger;
+
+import org.springframework.ai.tool.annotation.Tool;
+import org.springframework.ai.tool.annotation.ToolParam;
+
+/** A tool with a side effect that can be counted: every execution is one more refund issued. */
+public class OrderTools {
+
+ public final AtomicInteger executions = new AtomicInteger();
+
+ @Tool(name = "refund_order", description = "Refunds an order and returns the refund id")
+ public String refund(@ToolParam(description = "The order id") String orderId) {
+ return "refund-" + orderId + "-" + executions.incrementAndGet();
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Price.java b/llm-gateway/src/main/java/com/ankurm/gateway/Price.java
new file mode 100644
index 0000000..d2a9b03
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Price.java
@@ -0,0 +1,22 @@
+package com.ankurm.gateway;
+
+import java.math.BigDecimal;
+import java.math.RoundingMode;
+
+/**
+ * Dollars per million tokens. One dollar per million tokens is exactly one microdollar per token,
+ * so the arithmetic stays in whole microdollars and never touches a double.
+ */
+public record Price(BigDecimal inputPerMillion, BigDecimal outputPerMillion) {
+
+ public static Price of(String input, String output) {
+ return new Price(new BigDecimal(input), new BigDecimal(output));
+ }
+
+ /** Cost in microdollars, rounded up so a cap is never undercounted. */
+ public long micros(long inputTokens, long outputTokens) {
+ return inputPerMillion.multiply(BigDecimal.valueOf(inputTokens))
+ .add(outputPerMillion.multiply(BigDecimal.valueOf(outputTokens)))
+ .setScale(0, RoundingMode.CEILING).longValueExact();
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/ProviderFailure.java b/llm-gateway/src/main/java/com/ankurm/gateway/ProviderFailure.java
new file mode 100644
index 0000000..e13dad9
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/ProviderFailure.java
@@ -0,0 +1,16 @@
+package com.ankurm.gateway;
+
+/** A provider call that failed with an HTTP status. Stand-in for the vendor SDK exceptions in tests. */
+public class ProviderFailure extends RuntimeException {
+
+ private final int status;
+
+ public ProviderFailure(int status, String message) {
+ super(message);
+ this.status = status;
+ }
+
+ public int status() {
+ return status;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Routes.java b/llm-gateway/src/main/java/com/ankurm/gateway/Routes.java
new file mode 100644
index 0000000..3bf932e
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Routes.java
@@ -0,0 +1,22 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+import java.util.Map;
+
+/** Maps a model hint to an ordered list of targets. An unknown hint is an error, not a default. */
+public final class Routes {
+
+ private final Map> table;
+
+ public Routes(Map> table) {
+ this.table = Map.copyOf(table);
+ }
+
+ public List get(String hint) {
+ List t = table.get(hint);
+ if (t == null) {
+ throw new UnknownHintException(hint, table.keySet());
+ }
+ return t;
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/SemanticCache.java b/llm-gateway/src/main/java/com/ankurm/gateway/SemanticCache.java
new file mode 100644
index 0000000..61744cc
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/SemanticCache.java
@@ -0,0 +1,59 @@
+package com.ankurm.gateway;
+
+import java.util.List;
+import java.util.Map;
+import java.util.Optional;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.CopyOnWriteArrayList;
+
+import org.springframework.ai.embedding.EmbeddingModel;
+
+/**
+ * Serves a stored answer when a new question embeds close enough to an old one. The key includes
+ * the tenant: a cache shared across tenants answers tenant B with tenant A's data.
+ */
+public final class SemanticCache {
+
+ private record Entry(float[] vector, String answer) {
+ }
+
+ private final EmbeddingModel embeddings;
+
+ private final double threshold;
+
+ private final Map> byKey = new ConcurrentHashMap<>();
+
+ public SemanticCache(EmbeddingModel embeddings, double threshold) {
+ this.embeddings = embeddings;
+ this.threshold = threshold;
+ }
+
+ public Optional lookup(String tenant, String feature, String question) {
+ float[] q = embeddings.embed(question);
+ Entry best = null;
+ double bestScore = -1;
+ for (Entry e : byKey.getOrDefault(tenant + "|" + feature, List.of())) {
+ double s = cosine(q, e.vector());
+ if (s >= threshold && s > bestScore) {
+ best = e;
+ bestScore = s;
+ }
+ }
+ return Optional.ofNullable(best).map(Entry::answer);
+ }
+
+ public void put(String tenant, String feature, String question, String answer) {
+ byKey.computeIfAbsent(tenant + "|" + feature, k -> new CopyOnWriteArrayList<>())
+ .add(new Entry(embeddings.embed(question), answer));
+ }
+
+ static double cosine(float[] a, float[] b) {
+ double dot = 0, na = 0, nb = 0;
+ for (int i = 0; i < a.length; i++) {
+ dot += a[i] * b[i];
+ na += a[i] * a[i];
+ nb += b[i] * b[i];
+ }
+ return dot / (Math.sqrt(na) * Math.sqrt(nb) + 1e-10);
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/Target.java b/llm-gateway/src/main/java/com/ankurm/gateway/Target.java
new file mode 100644
index 0000000..ed6bcb8
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/Target.java
@@ -0,0 +1,7 @@
+package com.ankurm.gateway;
+
+import org.springframework.ai.chat.model.ChatModel;
+
+/** One provider the gateway can send a call to: a name for logs and breakers, a model, a price. */
+public record Target(String name, String model, ChatModel chat, Price price) {
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/TokenLimiter.java b/llm-gateway/src/main/java/com/ankurm/gateway/TokenLimiter.java
new file mode 100644
index 0000000..d3f9edd
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/TokenLimiter.java
@@ -0,0 +1,60 @@
+package com.ankurm.gateway;
+
+import java.time.Duration;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.function.ToLongFunction;
+
+import io.github.bucket4j.Bandwidth;
+import io.github.bucket4j.Bucket;
+import io.github.bucket4j.TimeMeter;
+
+/**
+ * Tokens-per-hour per tenant (Bucket4j, in memory). The unit is LLM tokens, not requests, because
+ * one summarisation can cost fifty lookups. Pre-consume an estimate, then settle with the real
+ * total: refund the surplus, or charge the shortfall. The old post refunded against prompt tokens
+ * only and never charged the completion, so every answer was free of the limit.
+ */
+public final class TokenLimiter {
+
+ private final Map buckets = new ConcurrentHashMap<>();
+
+ private final ToLongFunction perHour;
+
+ private final TimeMeter clock;
+
+ public TokenLimiter(ToLongFunction perHour, TimeMeter clock) {
+ this.perHour = perHour;
+ this.clock = clock;
+ }
+
+ private Bucket bucket(String tenant) {
+ return buckets.computeIfAbsent(tenant, t -> {
+ long cap = perHour.applyAsLong(t);
+ return Bucket.builder().withCustomTimePrecision(clock)
+ .addLimit(Bandwidth.builder().capacity(cap).refillGreedy(cap, Duration.ofHours(1)).build())
+ .build();
+ });
+ }
+
+ public void take(String tenant, long estimate) {
+ if (!bucket(tenant).tryConsume(estimate)) {
+ throw new BudgetExceededException("Token rate limit reached for " + tenant);
+ }
+ }
+
+ /** Correct the estimate once the real token count is known. */
+ public void settle(String tenant, long estimate, long actual) {
+ Bucket b = bucket(tenant);
+ if (actual < estimate) {
+ b.addTokens(estimate - actual);
+ }
+ else if (actual > estimate) {
+ b.consumeIgnoringRateLimits(actual - estimate);
+ }
+ }
+
+ public long available(String tenant) {
+ return bucket(tenant).getAvailableTokens();
+ }
+}
diff --git a/llm-gateway/src/main/java/com/ankurm/gateway/UnknownHintException.java b/llm-gateway/src/main/java/com/ankurm/gateway/UnknownHintException.java
new file mode 100644
index 0000000..4e548be
--- /dev/null
+++ b/llm-gateway/src/main/java/com/ankurm/gateway/UnknownHintException.java
@@ -0,0 +1,11 @@
+package com.ankurm.gateway;
+
+import java.util.Set;
+import java.util.TreeSet;
+
+public class UnknownHintException extends RuntimeException {
+
+ public UnknownHintException(String hint, Set known) {
+ super("Unknown model hint '" + hint + "'. Known hints: " + new TreeSet<>(known));
+ }
+}
diff --git a/llm-gateway/src/main/resources/application.yml b/llm-gateway/src/main/resources/application.yml
new file mode 100644
index 0000000..d393df8
--- /dev/null
+++ b/llm-gateway/src/main/resources/application.yml
@@ -0,0 +1,12 @@
+# Providers are plain data. Prices are dollars per million tokens (read 2026-10-09, see CostTest in the providers module).
+gateway:
+ providers:
+ - { name: openai, type: openai, base-url: "https://api.openai.com/v1", api-key: "${OPENAI_API_KEY:unset}", model: gpt-6.1-sol, input-price: "2.00", output-price: "10.00" }
+ - { name: anthropic, type: anthropic, base-url: "https://api.anthropic.com", api-key: "${ANTHROPIC_API_KEY:unset}", model: claude-sonnet-5-5, input-price: "2.00", output-price: "10.00" }
+ routes:
+ smart: [openai, anthropic]
+ fast: [anthropic, openai]
+ defaults: { budget-micros: 5000000, tokens-per-hour: 100000 }
+ tenants:
+ free-tier-tenant: { budget-micros: 50000, tokens-per-hour: 10000 }
+ breaker: { sliding-window-size: 10, minimum-calls: 5, failure-rate-threshold: 50, wait-open-millis: 30000 }
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/BudgetTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/BudgetTest.java
new file mode 100644
index 0000000..71d90fb
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/BudgetTest.java
@@ -0,0 +1,130 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+import static org.assertj.core.api.Assertions.assertThatThrownBy;
+
+import java.util.List;
+import java.util.concurrent.CyclicBarrier;
+import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicLong;
+
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Scripted;
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.Test;
+
+/** Dollar caps: reserve before the call, settle after, and where the number can still be wrong. */
+class BudgetTest {
+
+ private static GatewayRequest ask(String text, int maxTokens) {
+ return new GatewayRequest("support", "", text, "smart", maxTokens, null);
+ }
+
+ @Test
+ void capsHoldBeforeTheCallAndAreOnlyAsGoodAsTheEstimate() {
+ try (Transcript t = new Transcript("07-budget.txt", "Dollar caps: reserve, settle, and the limits of an estimate")) {
+ // A. sequential calls against a 5,000 microdollar cap; the model reports 500 in / 100 out = 2,000 each
+ Scripted model = new Scripted("openai").answerAlways("ok", 500, 100);
+ Budgets budgets = new Budgets(tenant -> 5_000);
+ Gateway g = Fixtures.gateway(new OrderTools(), budgets, Fixtures.target("openai", model));
+ t.line("cap 5000 microdollars; every call really costs 2000; maxTokens 256 (estimate about 2,570)");
+ for (int i = 1; i <= 3; i++) {
+ try {
+ GatewayResponse r = g.complete("acme", ask("hello", 256));
+ t.line(" call %d: answered, cost %d, spent so far %d", i, r.costMicros(), budgets.spent("acme"));
+ }
+ catch (BudgetExceededException e) {
+ t.line(" call %d: rejected before the provider was called (provider calls so far: %d)", i, model.calls());
+ }
+ }
+ assertThat(model.calls()).isEqualTo(2);
+ assertThat(budgets.spent("acme")).isEqualTo(4_000);
+
+ // B. the estimate is chars/4. A request whose real prompt is bigger than that overshoots the cap.
+ Scripted heavy = new Scripted("openai").answerAlways("ok", 3_000, 10);
+ Budgets b2 = new Budgets(tenant -> 8_000);
+ Gateway g2 = Fixtures.gateway(new OrderTools(), b2, Fixtures.target("openai", heavy));
+ String text = "x".repeat(400); // estimated at 100 tokens; the provider reports 3,000
+ g2.complete("acme", ask(text, 32));
+ g2.complete("acme", ask(text, 32));
+ t.blank();
+ t.line("cap 8000; prompt estimated at 100 tokens but the provider counts 3000 (code, other scripts, images do this)");
+ t.line(" spent after 2 calls: %d (cap 8000); the second call was admitted on its estimate", b2.spent("acme"));
+ assertThat(b2.spent("acme")).isGreaterThan(8_000);
+
+ // C. a cap hit in the middle of a tool loop: the tool has already run
+ OrderTools tools = new OrderTools();
+ Scripted looping = new Scripted("openai").callTool("c1", "refund_order", "{\"orderId\":\"A17\"}", 120, 15)
+ .answer("done", 160, 12);
+ Budgets b3 = new Budgets(tenant -> 2_700);
+ Gateway g3 = Fixtures.gateway(tools, b3, Fixtures.target("openai", looping));
+ t.blank();
+ try {
+ g3.complete("acme", new GatewayRequest("support", "", "Refund order A17", "smart", 256, List.of("refund_order")));
+ }
+ catch (BudgetExceededException e) {
+ t.line("cap 2700, a two-call tool loop: rejected on the SECOND model call; tool executions so far: %d",
+ tools.executions.get());
+ }
+ assertThat(tools.executions.get()).isEqualTo(1);
+ assertThat(looping.calls()).isEqualTo(1);
+ }
+ }
+
+ @Test
+ void reserveFirstAdmitsExactlyTheCapAndCheckThenRecordDoesNot() throws Exception {
+ try (Transcript t = new Transcript("07b-budget-concurrency.txt", "64 simultaneous requests against a cap that fits 10")) {
+ int threads = 64;
+ long each = 1_000;
+ long cap = 10 * each;
+
+ // check-then-record: every thread passes the check before any records its spend
+ AtomicLong spent = new AtomicLong();
+ AtomicInteger admittedNaive = new AtomicInteger();
+ CyclicBarrier gap = new CyclicBarrier(threads);
+ run(threads, () -> {
+ if (spent.get() + each <= cap) {
+ gap.await(); // in real life this gap is the duration of the provider call
+ spent.addAndGet(each);
+ admittedNaive.incrementAndGet();
+ }
+ return null;
+ });
+
+ // reserve-then-settle
+ Budgets budgets = new Budgets(tenant -> cap);
+ AtomicInteger admittedReserved = new AtomicInteger();
+ run(threads, () -> {
+ try {
+ budgets.reserve("acme", each);
+ admittedReserved.incrementAndGet(); // never settled, so the hold stays: the worst case
+ }
+ catch (BudgetExceededException rejected) {
+ // expected for 54 of them
+ }
+ return null;
+ });
+
+ t.line("cap %d, each call %d, %d threads at once", cap, each, threads);
+ t.line("check then record: admitted %d, spent %d", admittedNaive.get(), spent.get());
+ t.line("reserve then settle: admitted %d", admittedReserved.get());
+ assertThat(admittedNaive.get()).isEqualTo(64);
+ assertThat(admittedReserved.get()).isEqualTo(10);
+ assertThatThrownBy(() -> budgets.reserve("acme", each)).isInstanceOf(BudgetExceededException.class);
+ }
+ }
+
+ private static void run(int threads, java.util.concurrent.Callable task) throws Exception {
+ try (var pool = Executors.newFixedThreadPool(threads)) {
+ List> futures = new java.util.ArrayList<>();
+ for (int i = 0; i < threads; i++) {
+ futures.add(pool.submit(task));
+ }
+ for (Future f : futures) {
+ f.get();
+ }
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/CacheTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/CacheTest.java
new file mode 100644
index 0000000..6139ab1
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/CacheTest.java
@@ -0,0 +1,103 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.util.ArrayList;
+import java.util.HashMap;
+import java.util.List;
+import java.util.Locale;
+import java.util.Map;
+import java.util.Set;
+
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.document.Document;
+import org.springframework.ai.embedding.Embedding;
+import org.springframework.ai.embedding.EmbeddingModel;
+import org.springframework.ai.embedding.EmbeddingRequest;
+import org.springframework.ai.embedding.EmbeddingResponse;
+
+/**
+ * The cache mechanics, with a STAND-IN embedding model: hashed bag of words, no meaning at all. It
+ * proves how the threshold, the tenant key and the API behave. It says nothing about how any real
+ * embedding model scores paraphrases, and the article says so.
+ */
+class CacheTest {
+
+ /** Each word hashes to one of 256 slots; the vector is the normalised word count. */
+ static final class BagOfWords implements EmbeddingModel {
+
+ @Override
+ public EmbeddingResponse call(EmbeddingRequest request) {
+ List out = new ArrayList<>();
+ int i = 0;
+ for (String text : request.getInstructions()) {
+ out.add(new Embedding(vector(text), i++));
+ }
+ return new EmbeddingResponse(out);
+ }
+
+ @Override
+ public float[] embed(Document document) {
+ return vector(document.getText());
+ }
+
+ static float[] vector(String text) {
+ float[] v = new float[256];
+ for (String w : text.toLowerCase(Locale.ROOT).split("[^a-z0-9]+")) {
+ if (!w.isEmpty()) {
+ v[Math.floorMod(w.hashCode(), 256)] += 1;
+ }
+ }
+ double norm = 0;
+ for (float x : v) {
+ norm += x * x;
+ }
+ norm = Math.sqrt(norm);
+ for (int i = 0; i < v.length; i++) {
+ v[i] /= (float) norm;
+ }
+ return v;
+ }
+ }
+
+ @Test
+ void thresholdTenantKeyAndTheEmbedApi() throws Exception {
+ BagOfWords model = new BagOfWords();
+ SemanticCache cache = new SemanticCache(model, 0.92);
+ try (Transcript t = new Transcript("09-cache.txt", "Semantic cache mechanics (stand-in embeddings, threshold 0.92)")) {
+ t.line("EmbeddingModel.embed(String) returns: %s",
+ EmbeddingModel.class.getMethod("embed", String.class).getReturnType().getSimpleName());
+ assertThat(EmbeddingModel.class.getMethod("embed", String.class).getReturnType()).isEqualTo(float[].class);
+
+ String stored = "Please tell me the return policy for electronics bought in our online store during the "
+ + "holiday season for order A17 including refunds exchanges and store credit options";
+ cache.put("acme", "support", stored, "30 days, original packaging.");
+
+ record Probe(String label, String question) {
+ }
+ List probes = List.of(
+ new Probe("same words, new order", "For order A17 please tell me the electronics return policy bought in our online store during the "
+ + "holiday season including refunds exchanges and store credit options"),
+ new Probe("one word different (A18 for A17)", stored.replace("A17", "A18")),
+ new Probe("different question, shares a few words", "What is the return policy for software"));
+ t.line("%-38s %-8s %s", "probe", "cosine", "hit at 0.92?");
+ for (Probe p : probes) {
+ double score = SemanticCache.cosine(BagOfWords.vector(stored), BagOfWords.vector(p.question()));
+ boolean hit = cache.lookup("acme", "support", p.question()).isPresent();
+ t.line("%-38s %-8.3f %s", p.label(), score, hit);
+ }
+ assertThat(cache.lookup("acme", "support", stored.replace("A17", "A18"))).isPresent();
+ assertThat(cache.lookup("acme", "support", "What is the return policy for software")).isEmpty();
+
+ t.blank();
+ t.line("tenant globex asks the stored question: %s", cache.lookup("globex", "support", stored).isPresent() ? "hit" : "miss");
+ assertThat(cache.lookup("globex", "support", stored)).isEmpty();
+
+ Map keyedByFeatureOnly = new HashMap<>();
+ keyedByFeatureOnly.put("support", "30 days, original packaging.");
+ t.line("a cache keyed by feature only, tenant globex asks: %s", keyedByFeatureOnly.get("support"));
+ assertThat(Set.copyOf(keyedByFeatureOnly.keySet())).containsExactly("support");
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/CircuitBreakerTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/CircuitBreakerTest.java
new file mode 100644
index 0000000..0a90519
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/CircuitBreakerTest.java
@@ -0,0 +1,96 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.time.Duration;
+import java.util.List;
+
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Scripted;
+import com.ankurm.gateway.support.Transcript;
+import io.github.resilience4j.circuitbreaker.CircuitBreaker;
+import io.github.resilience4j.circuitbreaker.CircuitBreakerConfig;
+import io.github.resilience4j.circuitbreaker.CircuitBreakerRegistry;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.prompt.Prompt;
+
+/** One breaker per provider: when it opens, who is isolated, and how it recovers. */
+class CircuitBreakerTest {
+
+ private static int requestsUntilSkipped(CircuitBreakerRegistry registry) {
+ Scripted openai = new Scripted("openai").failAlways(503);
+ Scripted claude = new Scripted("claude").answerAlways("ok", 10, 5);
+ Failover f = new Failover(List.of(Fixtures.target("openai", openai), Fixtures.target("claude", claude)),
+ registry, Fixtures.unlimited());
+ for (int i = 1; i <= 120; i++) {
+ int before = openai.calls();
+ f.call(new Prompt("hi"));
+ if (openai.calls() == before) {
+ return i;
+ }
+ }
+ return -1;
+ }
+
+ @Test
+ void whenItOpensWhoIsIsolatedAndHowItRecovers() throws Exception {
+ try (Transcript t = new Transcript("03-circuit-breaker.txt", "Circuit breakers: when they open and how they recover")) {
+ // A. the old post's yml: slidingWindowSize 10, failureRateThreshold 50, nothing about minimumNumberOfCalls
+ CircuitBreakerConfig oldYml = CircuitBreakerConfig.custom().slidingWindowSize(10).failureRateThreshold(50)
+ .waitDurationInOpenState(Duration.ofSeconds(30)).build();
+ int oldAt = requestsUntilSkipped(CircuitBreakerRegistry.of(oldYml));
+ t.line("old yml (window 10, threshold 50%%): minimumNumberOfCalls is %d; the first request that skips openai is #%d",
+ oldYml.getMinimumNumberOfCalls(), oldAt);
+
+ // B. this module's yml: minimum-calls 5
+ int newAt = requestsUntilSkipped(Fixtures.breakers());
+ t.line("this module (window 10, minimum 5, threshold 50%%): the first request that skips openai is #%d", newAt);
+ assertThat(newAt).isEqualTo(6);
+ assertThat(oldAt).isEqualTo(11);
+
+ // C. isolation and recovery
+ CircuitBreakerRegistry registry = Fixtures.breakers();
+ Scripted openai = new Scripted("openai").failAlways(503);
+ Scripted claude = new Scripted("claude").answerAlways("ok", 10, 5);
+ Failover f = new Failover(List.of(Fixtures.target("openai", openai), Fixtures.target("claude", claude)),
+ registry, Fixtures.unlimited());
+ for (int i = 0; i < 8; i++) {
+ f.call(new Prompt("hi"));
+ }
+ t.blank();
+ t.line("after 8 requests: openai breaker %s, claude breaker %s",
+ registry.circuitBreaker("openai").getState(), registry.circuitBreaker("claude").getState());
+ t.line("openai was contacted %d times of 8; claude answered %d times", openai.calls(), claude.calls());
+ assertThat(registry.circuitBreaker("openai").getState()).isEqualTo(CircuitBreaker.State.OPEN);
+ assertThat(registry.circuitBreaker("claude").getState()).isEqualTo(CircuitBreaker.State.CLOSED);
+
+ Thread.sleep(250); // waitDurationInOpenState is 200 ms in the test registry
+ Scripted healed = new Scripted("openai").answerAlways("back", 10, 5);
+ Failover g = new Failover(List.of(Fixtures.target("openai", healed), Fixtures.target("claude", claude)),
+ registry, Fixtures.unlimited());
+ g.call(new Prompt("hi"));
+ t.blank();
+ t.line("after the wait, one probe goes to openai (it has recovered): breaker %s, openai calls %d",
+ registry.circuitBreaker("openai").getState(), healed.calls());
+ assertThat(registry.circuitBreaker("openai").getState()).isEqualTo(CircuitBreaker.State.CLOSED);
+
+ // D. 4xx are the caller's fault and must not count against the provider
+ CircuitBreakerRegistry r2 = Fixtures.breakers();
+ Scripted picky = new Scripted("openai").failAlways(400);
+ Failover h = new Failover(List.of(Fixtures.target("openai", picky)), r2, Fixtures.unlimited());
+ for (int i = 0; i < 20; i++) {
+ try {
+ h.call(new Prompt("hi"));
+ }
+ catch (ProviderFailure expected) {
+ // rethrown, not counted
+ }
+ }
+ t.blank();
+ t.line("20 requests rejected with 400: breaker %s, calls the breaker recorded: %d",
+ r2.circuitBreaker("openai").getState(), r2.circuitBreaker("openai").getMetrics().getNumberOfBufferedCalls());
+ assertThat(r2.circuitBreaker("openai").getMetrics().getNumberOfBufferedCalls()).isZero();
+ assertThat(r2.circuitBreaker("openai").getState()).isEqualTo(CircuitBreaker.State.CLOSED);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/FailoverTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/FailoverTest.java
new file mode 100644
index 0000000..dfe69f5
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/FailoverTest.java
@@ -0,0 +1,52 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+import static org.assertj.core.api.Assertions.assertThatThrownBy;
+
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Scripted;
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.prompt.Prompt;
+
+/** Which failures move on to the next provider, and which stop the call. */
+class FailoverTest {
+
+ @Test
+ void onlyTransientFailuresFailOver() {
+ try (Transcript t = new Transcript("02-failover.txt", "Failover: which statuses move to the next provider")) {
+ t.line("%-8s %-14s %-12s %s", "status", "second tried?", "outcome", "trail");
+ for (int status : new int[] {408, 429, 500, 503, 400, 401, 404}) {
+ Scripted first = new Scripted("first").fail(status);
+ Scripted second = new Scripted("second").answerAlways("ok", 10, 5);
+ Failover f = Fixtures.failover(Fixtures.target("first", first), Fixtures.target("second", second));
+ CallContext.Meter meter = new CallContext.Meter();
+ String outcome;
+ try {
+ ScopedValue.where(CallContext.CURRENT, new CallContext("acme", "x", meter))
+ .run(() -> f.call(new Prompt("hi")));
+ outcome = "answered";
+ }
+ catch (ProviderFailure e) {
+ outcome = "rethrown";
+ }
+ t.line("%-8d %-14s %-12s %s", status, second.calls() > 0, outcome, meter.trail());
+ boolean transientFailure = status == 408 || status == 429 || status >= 500;
+ assertThat(second.calls() > 0).isEqualTo(transientFailure);
+ }
+
+ Scripted a = new Scripted("a").failAlways(503);
+ Scripted b = new Scripted("b").failAlways(429);
+ Failover f = Fixtures.failover(Fixtures.target("a", a), Fixtures.target("b", b));
+ t.blank();
+ try {
+ f.call(new Prompt("hi"));
+ }
+ catch (AllProvidersUnavailableException e) {
+ t.line("both down: %s", e.getMessage());
+ assertThat(e.trail()).containsExactly("a: failed with status 503", "b: failed with status 429");
+ }
+ assertThatThrownBy(() -> f.call(new Prompt("hi"))).isInstanceOf(AllProvidersUnavailableException.class);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/HttpTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/HttpTest.java
new file mode 100644
index 0000000..f07f711
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/HttpTest.java
@@ -0,0 +1,133 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.net.URI;
+import java.net.http.HttpClient;
+import java.net.http.HttpRequest;
+import java.net.http.HttpResponse;
+
+import com.ankurm.gateway.support.FakeVendors;
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.AfterAll;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+import org.springframework.boot.test.context.SpringBootTest;
+import org.springframework.boot.test.web.server.LocalServerPort;
+import org.springframework.test.context.DynamicPropertyRegistry;
+import org.springframework.test.context.DynamicPropertySource;
+
+/** The whole application over HTTP: real beans, real models, a fake at the far end. */
+@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT)
+class HttpTest {
+
+ static final FakeVendors VENDORS;
+
+ static {
+ try {
+ VENDORS = new FakeVendors();
+ }
+ catch (java.io.IOException e) {
+ throw new ExceptionInInitializerError(e);
+ }
+ }
+
+ @DynamicPropertySource
+ static void providers(DynamicPropertyRegistry r) {
+ r.add("gateway.providers[0].name", () -> "openai");
+ r.add("gateway.providers[0].type", () -> "openai");
+ r.add("gateway.providers[0].base-url", () -> VENDORS.url() + "/v1");
+ r.add("gateway.providers[0].api-key", () -> "test-key");
+ r.add("gateway.providers[0].model", () -> "gpt-6.1-sol");
+ r.add("gateway.providers[0].input-price", () -> "2.00");
+ r.add("gateway.providers[0].output-price", () -> "10.00");
+ r.add("gateway.providers[1].name", () -> "anthropic");
+ r.add("gateway.providers[1].type", () -> "anthropic");
+ r.add("gateway.providers[1].base-url", () -> VENDORS.url());
+ r.add("gateway.providers[1].api-key", () -> "test-key");
+ r.add("gateway.providers[1].model", () -> "claude-sonnet-5-5");
+ r.add("gateway.providers[1].input-price", () -> "2.00");
+ r.add("gateway.providers[1].output-price", () -> "10.00");
+ r.add("gateway.tenants.poor.budget-micros", () -> "1000");
+ r.add("gateway.tenants.poor.tokens-per-hour", () -> "100000");
+ }
+
+ @LocalServerPort
+ int port;
+
+ private final HttpClient http = HttpClient.newHttpClient();
+
+ @BeforeEach
+ void reset() {
+ VENDORS.openai.reset();
+ VENDORS.anthropic.reset();
+ }
+
+ @AfterAll
+ static void stop() {
+ VENDORS.close();
+ }
+
+ private HttpResponse post(String tenant, String json) throws Exception {
+ HttpRequest.Builder b = HttpRequest.newBuilder(URI.create("http://127.0.0.1:" + port + "/v1/gateway/complete"))
+ .header("Content-Type", "application/json").POST(HttpRequest.BodyPublishers.ofString(json));
+ if (tenant != null) {
+ b.header("X-Tenant-Id", tenant);
+ }
+ return http.send(b.build(), HttpResponse.BodyHandlers.ofString());
+ }
+
+ private static String body(String hint, String extra) {
+ return "{\"featureTag\":\"support\",\"userMessage\":\"What is your return policy?\",\"modelHint\":\"" + hint
+ + "\",\"maxTokens\":256" + extra + "}";
+ }
+
+ @Test
+ void statusCodesAndBodiesOverHttp() throws Exception {
+ try (Transcript t = new Transcript("10-http.txt", "The gateway over HTTP (real beans, real models, fake vendors)")) {
+ VENDORS.openai.then(200, FakeVendors.openaiText("30 days.", 47, 52));
+ HttpResponse ok = post("acme", body("smart", ""));
+ t.line("1. normal call -> %d %s", ok.statusCode(), ok.body());
+ assertThat(ok.statusCode()).isEqualTo(200);
+ assertThat(ok.body()).contains("\"provider\":\"openai\"").contains("\"costMicros\":614");
+
+ reset();
+ VENDORS.openai.then(503, FakeVendors.error("server_error", "overloaded"));
+ VENDORS.anthropic.then(200, FakeVendors.anthropicText("30 days.", 47, 52));
+ HttpResponse failover = post("acme", body("smart", ""));
+ t.line("2. openai returns 503 -> %d %s", failover.statusCode(), failover.body());
+ assertThat(failover.body()).contains("\"provider\":\"anthropic\"");
+
+ reset();
+ VENDORS.openai.always(503, FakeVendors.error("server_error", "overloaded"));
+ VENDORS.anthropic.always(529, FakeVendors.error("overloaded_error", "overloaded"));
+ HttpResponse down = post("acme", body("smart", ""));
+ t.line("3. both providers down -> %d Retry-After=%s %s", down.statusCode(),
+ down.headers().firstValue("Retry-After").orElse("-"), down.body());
+ assertThat(down.statusCode()).isEqualTo(503);
+
+ reset();
+ HttpResponse hint = post("acme", body("locla", ""));
+ t.line("4. unknown hint -> %d %s", hint.statusCode(), hint.body());
+ assertThat(hint.statusCode()).isEqualTo(400);
+ assertThat(VENDORS.openai.requests() + VENDORS.anthropic.requests()).isZero();
+
+ HttpResponse tool = post("acme", body("smart", ",\"tools\":[\"drop_tables\"]"));
+ t.line("5. tool not allow-listed -> %d %s", tool.statusCode(), tool.body());
+ assertThat(tool.statusCode()).isEqualTo(400);
+
+ HttpResponse poor = post("poor", body("smart", ""));
+ t.line("6. tenant over its cap -> %d %s", poor.statusCode(), poor.body());
+ assertThat(poor.statusCode()).isEqualTo(429);
+ assertThat(VENDORS.openai.requests() + VENDORS.anthropic.requests()).isZero();
+
+ HttpResponse spoof = post("poor", body("smart", ",\"tenantId\":\"acme\""));
+ t.line("7. body claims tenant acme, header says poor -> %d (the header wins)", spoof.statusCode());
+ assertThat(spoof.statusCode()).isEqualTo(429);
+
+ HttpResponse none = post(null, body("smart", ""));
+ t.line("8. no X-Tenant-Id header -> %d", none.statusCode());
+ assertThat(none.statusCode()).isEqualTo(400);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/RateLimitTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/RateLimitTest.java
new file mode 100644
index 0000000..0142170
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/RateLimitTest.java
@@ -0,0 +1,67 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+import static org.assertj.core.api.Assertions.assertThatThrownBy;
+
+import java.time.Duration;
+import java.util.concurrent.atomic.AtomicLong;
+
+import com.ankurm.gateway.support.Transcript;
+import io.github.bucket4j.TimeMeter;
+import org.junit.jupiter.api.Test;
+
+/** Tokens per hour per tenant: estimate, settle against the real total, refill with a fake clock. */
+class RateLimitTest {
+
+ @Test
+ void settleChargesTheRealTotalIncludingTheCompletion() {
+ AtomicLong now = new AtomicLong();
+ TimeMeter clock = new TimeMeter() {
+ @Override
+ public long currentTimeNanos() {
+ return now.get();
+ }
+
+ @Override
+ public boolean isWallClockBased() {
+ return false;
+ }
+ };
+ TokenLimiter limiter = new TokenLimiter(t -> 10_000, clock);
+ try (Transcript t = new Transcript("08-rate-limit.txt", "Token rate limit: pre-consume, settle, refill")) {
+ t.line("capacity 10000 tokens per hour");
+ limiter.take("acme", 1_500);
+ t.line("take estimate 1500 -> available %d", limiter.available("acme"));
+ limiter.settle("acme", 1_500, 1_100); // prompt 100 + completion 1000
+ t.line("settle: prompt 100 + completion 1000 = 1100 -> available %d", limiter.available("acme"));
+ assertThat(limiter.available("acme")).isEqualTo(8_900);
+
+ // the old post refunded estimate - promptTokens, so the completion was never charged
+ long oldRefund = 1_500 - 100;
+ t.line("old rule (refund estimate minus PROMPT tokens only): would have refunded %d and left %d charged",
+ oldRefund, 1_500 - oldRefund);
+
+ limiter.take("acme", 500);
+ limiter.settle("acme", 500, 2_000); // the answer was longer than estimated
+ t.blank();
+ t.line("take 500, real total 2000 -> available %d (the shortfall is charged, not forgiven)",
+ limiter.available("acme"));
+ assertThat(limiter.available("acme")).isEqualTo(6_900);
+
+ limiter.take("acme", 6_900);
+ t.blank();
+ t.line("take 6900 -> available %d", limiter.available("acme"));
+ assertThatThrownBy(() -> limiter.take("acme", 1_000)).isInstanceOf(BudgetExceededException.class)
+ .hasMessageContaining("Token rate limit reached");
+ t.line("take 1000 -> rejected: Token rate limit reached for acme");
+
+ now.addAndGet(Duration.ofMinutes(30).toNanos());
+ t.blank();
+ t.line("30 minutes later -> available %d (greedy refill: 10000 per hour)", limiter.available("acme"));
+ assertThat(limiter.available("acme")).isEqualTo(5_000);
+ now.addAndGet(Duration.ofMinutes(60).toNanos());
+ t.line("60 more minutes -> available %d (capped at capacity)", limiter.available("acme"));
+ assertThat(limiter.available("acme")).isEqualTo(10_000);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/RoutingTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/RoutingTest.java
new file mode 100644
index 0000000..c7c4cdf
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/RoutingTest.java
@@ -0,0 +1,71 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+import static org.assertj.core.api.Assertions.assertThatThrownBy;
+
+import java.util.List;
+import java.util.Map;
+
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Scripted;
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.Test;
+
+/** Routing table: hints, one model id per provider, and what an unknown hint does. */
+class RoutingTest {
+
+ @Test
+ void hintsOptionsAndUnknownHints() {
+ Scripted openai = new Scripted("openai").failAlways(503);
+ Scripted claude = new Scripted("claude").answerAlways("from claude", 100, 20);
+ Scripted local = new Scripted("local").answerAlways("from local", 100, 20);
+ Target tOpenai = Fixtures.target("openai", openai);
+ Target tClaude = Fixtures.target("claude", claude);
+ Target tLocal = Fixtures.target("local", local);
+ Routes routes = new Routes(Map.of("smart", List.of(tOpenai, tClaude), "local", List.of(tLocal)));
+
+ OrderTools tools = new OrderTools();
+ var budgets = Fixtures.unlimited();
+ var breakers = Fixtures.breakers();
+ Gateway gateway = new Gateway(routes, t -> new Failover(t, breakers, budgets),
+ new TokenLimiter(t -> 1_000_000, io.github.bucket4j.TimeMeter.SYSTEM_MILLISECONDS), null,
+ Map.of("refund_order", tools));
+
+ try (Transcript t = new Transcript("01-routing.txt", "Routing: hints, per-target options, unknown hints")) {
+ // 1. the same request, hint "smart": openai is down, claude answers, each got its OWN model id
+ GatewayResponse r = gateway.complete("acme", new GatewayRequest("support", "Be brief.", "Hi", "smart", 256, null));
+ t.line("hint smart -> answered by %s (%s)", r.provider(), r.model());
+ t.line(" model id sent to openai: %s", openai.modelOfCall(0));
+ t.line(" model id sent to claude: %s", claude.modelOfCall(0));
+ t.line(" maxTokens sent to both: %s and %s", openai.prompts().get(0).getOptions().getMaxTokens(),
+ claude.prompts().get(0).getOptions().getMaxTokens());
+ assertThat(openai.modelOfCall(0)).isEqualTo("openai-model");
+ assertThat(claude.modelOfCall(0)).isEqualTo("claude-model");
+
+ // 2. hint "local" never touches the cloud targets
+ int cloudBefore = openai.calls() + claude.calls();
+ GatewayResponse l = gateway.complete("acme", new GatewayRequest("support", "", "Hi", "local", 256, null));
+ t.blank();
+ t.line("hint local -> answered by %s; cloud calls made for it: %d", l.provider(),
+ openai.calls() + claude.calls() - cloudBefore);
+ assertThat(openai.calls() + claude.calls()).isEqualTo(cloudBefore);
+
+ // 3. a typo: the old registry did getOrDefault(hint, smart)
+ t.blank();
+ Map> old = Map.of("smart", List.of(tClaude), "local", List.of(tLocal));
+ List oldChoice = old.getOrDefault("locla", old.get("smart"));
+ t.line("old registry, hint \"locla\" (typo for local): goes to %s, a cloud provider", oldChoice.get(0).name());
+ assertThat(oldChoice.get(0).name()).isEqualTo("claude");
+
+ assertThatThrownBy(() -> gateway.complete("acme", new GatewayRequest("support", "", "Hi", "locla", 256, null)))
+ .isInstanceOf(UnknownHintException.class);
+ try {
+ gateway.complete("acme", new GatewayRequest("support", "", "Hi", "locla", 256, null));
+ }
+ catch (UnknownHintException e) {
+ t.line("this Routes, hint \"locla\": %s", e.getMessage());
+ }
+ assertThat(claude.calls()).isEqualTo(1);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/ToolLoopPlacementTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/ToolLoopPlacementTest.java
new file mode 100644
index 0000000..0a81f57
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/ToolLoopPlacementTest.java
@@ -0,0 +1,107 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.util.List;
+import java.util.stream.Collectors;
+
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Scripted;
+import com.ankurm.gateway.support.Transcript;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.client.ChatClient;
+import org.springframework.ai.chat.client.advisor.ToolCallingAdvisor;
+import org.springframework.ai.chat.model.ChatResponse;
+import org.springframework.ai.support.ToolCallbacks;
+
+/**
+ * In Spring AI 2.0 the tool loop is an advisor ABOVE the chat model. So where failover sits decides
+ * whether a provider failing mid-loop re-runs the tools.
+ */
+class ToolLoopPlacementTest {
+
+ private static final String ARGS = "{\"orderId\":\"A17\"}";
+
+ @Test
+ void failoverBelowTheAdvisorDoesNotReRunTools() {
+ try (Transcript t = new Transcript("04-tool-loop.txt", "Where failover sits relative to the tool loop")) {
+ // A. the gateway as built: Failover is the ChatModel, ToolCallingAdvisor is above it.
+ OrderTools toolsA = new OrderTools();
+ Scripted openaiA = new Scripted("openai").callTool("call-1", "refund_order", ARGS, 120, 15).fail(503);
+ Scripted claudeA = new Scripted("claude").answer("Refunded as refund-A17-1", 160, 12);
+ Gateway gateway = Fixtures.gateway(toolsA, Fixtures.unlimited(), Fixtures.target("openai", openaiA),
+ Fixtures.target("claude", claudeA));
+ GatewayResponse r = gateway.complete("acme", new GatewayRequest("support", "", "Refund order A17", "smart",
+ 256, List.of("refund_order")));
+
+ t.line("A. failover below the advisor (this module)");
+ t.line(" tool executions: %d", toolsA.executions.get());
+ t.line(" trail: %s", r.trail());
+ t.line(" answer: %s (from %s)", r.content(), r.provider());
+ t.line(" what claude was sent: %s", claudeA.prompts().get(0).getInstructions().stream()
+ .map(m -> m.getClass().getSimpleName()).collect(Collectors.joining(", ")));
+ assertThat(toolsA.executions.get()).isEqualTo(1);
+ assertThat(claudeA.calls()).isEqualTo(1);
+
+ // B. the intuitive alternative: one ChatClient per provider, retry the whole call on failure.
+ OrderTools toolsB = new OrderTools();
+ Scripted openaiB = new Scripted("openai").callTool("call-1", "refund_order", ARGS, 120, 15).fail(503);
+ Scripted claudeB = new Scripted("claude").callTool("call-9", "refund_order", ARGS, 120, 15)
+ .answer("Refunded as refund-A17-2", 160, 12);
+ String answer = null;
+ for (Scripted model : List.of(openaiB, claudeB)) {
+ ChatClient client = ChatClient.builder(model).defaultAdvisors(ToolCallingAdvisor.builder().build())
+ .build();
+ try {
+ answer = client.prompt().user("Refund order A17").toolCallbacks(ToolCallbacks.from(toolsB)).call()
+ .content();
+ break;
+ }
+ catch (ProviderFailure e) {
+ // retry the whole call on the next provider
+ }
+ }
+ t.blank();
+ t.line("B. retry above the advisor (one ChatClient per provider)");
+ t.line(" tool executions: %d", toolsB.executions.get());
+ t.line(" answer: %s", answer);
+ assertThat(toolsB.executions.get()).isEqualTo(2);
+
+ // C. what the response says about usage, and what the meter says
+ Scripted usageModel = new Scripted("openai").callTool("call-1", "refund_order", ARGS, 120, 15)
+ .answer("done", 160, 12);
+ ChatClient client = ChatClient.builder(Fixtures.failover(Fixtures.target("openai", usageModel)))
+ .defaultAdvisors(ToolCallingAdvisor.builder().build()).build();
+ CallContext.Meter meter = new CallContext.Meter();
+ ChatResponse response = ScopedValue.where(CallContext.CURRENT, new CallContext("acme", "x", meter))
+ .call(() -> client.prompt().user("Refund order A17")
+ .toolCallbacks(ToolCallbacks.from(new OrderTools())).call().chatResponse());
+ var usage = response.getMetadata().getUsage();
+ t.blank();
+ t.line("C. usage across a two-call tool loop (calls reported 120+15 and 160+12)");
+ t.line(" usage on the final ChatResponse: prompt=%d completion=%d", usage.getPromptTokens(),
+ usage.getCompletionTokens());
+ t.line(" usage summed per model call: prompt=%d completion=%d", meter.promptTokens(),
+ meter.completionTokens());
+ assertThat(meter.promptTokens()).isEqualTo(280);
+ assertThat(meter.completionTokens()).isEqualTo(27);
+
+ // D. a loop that changes provider halfway: one total, two prices
+ Target cheap = new Target("openai", "openai-model",
+ new Scripted("openai").callTool("call-1", "refund_order", ARGS, 120, 15).fail(503),
+ Price.of("0.75", "3.75"));
+ Target dear = new Target("claude", "claude-model", new Scripted("claude").answer("done", 160, 12),
+ Price.of("2.00", "10.00"));
+ GatewayResponse mixed = Fixtures.gateway(new OrderTools(), Fixtures.unlimited(), cheap, dear)
+ .complete("acme", new GatewayRequest("support", "", "Refund order A17", "smart", 256,
+ List.of("refund_order")));
+ long oneRate = dear.price().micros(280, 27);
+ t.blank();
+ t.line("D. the same loop with the first call on a cheaper provider (0.75/3.75 then 2.00/10.00)");
+ t.line(" priced per model call: %d microdollars", mixed.costMicros());
+ t.line(" total usage at the last provider's price: %d microdollars", oneRate);
+ assertThat(mixed.costMicros()).isEqualTo(587);
+ assertThat(oneRate).isEqualTo(830);
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/WireTest.java b/llm-gateway/src/test/java/com/ankurm/gateway/WireTest.java
new file mode 100644
index 0000000..20efd6a
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/WireTest.java
@@ -0,0 +1,127 @@
+package com.ankurm.gateway;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.Map;
+
+import com.ankurm.gateway.support.FakeVendors;
+import com.ankurm.gateway.support.Fixtures;
+import com.ankurm.gateway.support.Transcript;
+import io.github.bucket4j.TimeMeter;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.prompt.Prompt;
+import tools.jackson.databind.JsonNode;
+import tools.jackson.databind.json.JsonMapper;
+
+/** The REAL OpenAI and Anthropic Spring AI models behind the gateway, against a local fake of both wire formats. */
+class WireTest {
+
+ private static final JsonMapper JSON = JsonMapper.builder().build();
+
+ private static Target target(String name, String type, String base, String model) {
+ GatewayProperties.Provider p = new GatewayProperties.Provider(name, type, base, "test-key", model, "2.00", "10.00");
+ return new Target(name, model, GatewayConfig.chatModel(p), Price.of("2.00", "10.00"));
+ }
+
+ private static Gateway gateway(FakeVendors v, OrderTools tools) {
+ var breakers = Fixtures.breakers();
+ var budgets = Fixtures.unlimited();
+ Target openai = target("openai", "openai", v.url() + "/v1", "gpt-6.1-sol");
+ Target claude = target("anthropic", "anthropic", v.url(), "claude-sonnet-5-5");
+ return new Gateway(new Routes(Map.of("smart", List.of(openai, claude))), t -> new Failover(t, breakers, budgets),
+ new TokenLimiter(x -> 1_000_000, TimeMeter.SYSTEM_MILLISECONDS), null, Map.of("refund_order", tools));
+ }
+
+ @Test
+ void realModelsFailOverInTheMiddleOfAToolLoop() throws Exception {
+ try (Transcript t = new Transcript("05-wire.txt", "Real OpenAI and Anthropic models: failover in the middle of a tool loop");
+ FakeVendors v = new FakeVendors()) {
+ v.openai.then(200, FakeVendors.openaiToolCall("call_1", "refund_order", "{\"orderId\":\"A17\"}", 120, 15))
+ .then(503, FakeVendors.error("server_error", "overloaded"));
+ v.anthropic.then(200, FakeVendors.anthropicText("Refunded as refund-A17-1", 160, 12));
+ OrderTools tools = new OrderTools();
+
+ GatewayResponse r = gateway(v, tools).complete("acme",
+ new GatewayRequest("support", "Be brief.", "Refund order A17", "smart", 256, List.of("refund_order")));
+
+ t.line("trail: %s", r.trail());
+ t.line("answered by %s; tool executions %d; requests: openai %d, anthropic %d", r.provider(),
+ tools.executions.get(), v.openai.requests(), v.anthropic.requests());
+ assertThat(r.provider()).isEqualTo("anthropic");
+ assertThat(tools.executions.get()).isEqualTo(1);
+ assertThat(v.openai.requests()).isEqualTo(2);
+ assertThat(v.anthropic.requests()).isEqualTo(1);
+
+ JsonNode first = JSON.readTree(v.openai.bodies().get(0));
+ List toolNames = new ArrayList<>();
+ for (JsonNode n : first.path("tools")) {
+ toolNames.add(n.path("function").path("name").asString());
+ }
+ t.blank();
+ t.line("request 1 to OpenAI: model=%s, %s, tools=%s", first.path("model").asString(),
+ first.has("max_completion_tokens") ? "max_completion_tokens=" + first.path("max_completion_tokens").asInt()
+ : "max_tokens=" + first.path("max_tokens").asInt(),
+ toolNames);
+
+ JsonNode body = JSON.readTree(v.anthropic.bodies().get(0));
+ List blocks = new ArrayList<>();
+ for (JsonNode m : body.path("messages")) {
+ for (JsonNode b : m.path("content")) {
+ String type = b.path("type").asString();
+ String detail = switch (type) {
+ case "tool_use" -> "tool_use(id=" + b.path("id").asString() + ", name=" + b.path("name").asString() + ")";
+ case "tool_result" -> "tool_result(tool_use_id=" + b.path("tool_use_id").asString() + ")";
+ default -> type;
+ };
+ blocks.add(m.path("role").asString() + ":" + detail);
+ }
+ }
+ t.line("request to Anthropic: model=%s max_tokens=%d", body.path("model").asString(), body.path("max_tokens").asInt());
+ t.line(" conversation it received: %s", blocks);
+ assertThat(body.path("model").asString()).isEqualTo("claude-sonnet-5-5");
+ assertThat(blocks).anyMatch(b -> b.startsWith("assistant:tool_use(id=call_1"));
+ assertThat(blocks).anyMatch(b -> b.equals("user:tool_result(tool_use_id=call_1)"));
+ }
+ }
+
+ @Test
+ void realExceptionsAreClassifiedAndSdkRetriesAreOff() throws Exception {
+ try (Transcript t = new Transcript("06-real-exceptions.txt", "What the vendor SDKs throw, and what the gateway does with it");
+ FakeVendors v = new FakeVendors()) {
+ Target openai = target("openai", "openai", v.url() + "/v1", "gpt-6.1-sol");
+ Target claude = target("anthropic", "anthropic", v.url(), "claude-sonnet-5-5");
+ t.line("%-10s %-8s %-52s %-11s %s", "vendor", "status", "exception", "transient?", "HTTP requests for one call");
+ for (int status : new int[] {400, 401, 429, 503}) {
+ v.openai.bodies().clear();
+ v.openai.always(status, FakeVendors.error("x", "x"));
+ Throwable e = thrown(openai);
+ t.line("%-10s %-8d %-52s %-11s %d", "openai", status, e.getClass().getName(), Failures.isTransient(e),
+ v.openai.requests());
+ assertThat(Failures.status(e)).isEqualTo(status);
+ assertThat(v.openai.requests()).isEqualTo(1);
+ }
+ for (int status : new int[] {400, 529}) {
+ v.anthropic.bodies().clear();
+ v.anthropic.always(status, FakeVendors.error("x", "x"));
+ Throwable e = thrown(claude);
+ t.line("%-10s %-8d %-52s %-11s %d", "anthropic", status, e.getClass().getName(), Failures.isTransient(e),
+ v.anthropic.requests());
+ assertThat(Failures.status(e)).isEqualTo(status);
+ assertThat(Failures.isTransient(e)).isEqualTo(status == 529);
+ assertThat(v.anthropic.requests()).isEqualTo(1);
+ }
+ }
+ }
+
+ private static Throwable thrown(Target target) {
+ try {
+ target.chat().call(new Prompt("hi"));
+ throw new AssertionError("expected a failure");
+ }
+ catch (RuntimeException e) {
+ return e;
+ }
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/support/FakeVendors.java b/llm-gateway/src/test/java/com/ankurm/gateway/support/FakeVendors.java
new file mode 100644
index 0000000..16e43b0
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/support/FakeVendors.java
@@ -0,0 +1,127 @@
+package com.ankurm.gateway.support;
+
+import java.io.IOException;
+import java.io.OutputStream;
+import java.net.InetSocketAddress;
+import java.nio.charset.StandardCharsets;
+import java.util.ArrayDeque;
+import java.util.Deque;
+import java.util.List;
+import java.util.concurrent.CopyOnWriteArrayList;
+
+import com.sun.net.httpserver.HttpExchange;
+import com.sun.net.httpserver.HttpServer;
+
+/**
+ * One local HTTP server that speaks two vendor wire formats, so the REAL Spring AI OpenAI and
+ * Anthropic models can talk to it. It is not a language model: each vendor has a queue of canned
+ * replies (status and body), and every request body is recorded.
+ */
+public final class FakeVendors implements AutoCloseable {
+
+ public record Reply(int status, String body) {
+ }
+
+ public static final class Vendor {
+
+ private final Deque queue = new ArrayDeque<>();
+
+ private final List bodies = new CopyOnWriteArrayList<>();
+
+ private volatile Reply otherwise;
+
+ public synchronized Vendor then(int status, String body) {
+ queue.add(new Reply(status, body));
+ return this;
+ }
+
+ public Vendor always(int status, String body) {
+ otherwise = new Reply(status, body);
+ return this;
+ }
+
+ public synchronized Vendor reset() {
+ queue.clear();
+ bodies.clear();
+ otherwise = null;
+ return this;
+ }
+
+ synchronized Reply next() {
+ Reply r = queue.poll();
+ return r != null ? r : otherwise != null ? otherwise : new Reply(500, "{\"error\":\"script exhausted\"}");
+ }
+
+ public List bodies() {
+ return bodies;
+ }
+
+ public int requests() {
+ return bodies.size();
+ }
+ }
+
+ public final Vendor openai = new Vendor();
+
+ public final Vendor anthropic = new Vendor();
+
+ private final HttpServer server;
+
+ public FakeVendors() throws IOException {
+ server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
+ server.createContext("/v1/chat/completions", ex -> serve(ex, openai));
+ server.createContext("/v1/messages", ex -> serve(ex, anthropic));
+ server.start();
+ }
+
+ public String url() {
+ return "http://127.0.0.1:" + server.getAddress().getPort();
+ }
+
+ private static void serve(HttpExchange ex, Vendor v) throws IOException {
+ v.bodies.add(new String(ex.getRequestBody().readAllBytes(), StandardCharsets.UTF_8));
+ Reply r = v.next();
+ byte[] out = r.body().getBytes(StandardCharsets.UTF_8);
+ ex.getResponseHeaders().add("Content-Type", "application/json");
+ ex.sendResponseHeaders(r.status(), out.length);
+ try (OutputStream os = ex.getResponseBody()) {
+ os.write(out);
+ }
+ }
+
+ @Override
+ public void close() {
+ server.stop(0);
+ }
+
+ // ---- canned bodies ----
+
+ public static String openaiText(String text, int in, int out) {
+ return "{\"id\":\"chatcmpl-1\",\"object\":\"chat.completion\",\"created\":1700000000,\"model\":\"gpt-6.1-sol\","
+ + "\"choices\":[{\"index\":0,\"message\":{\"role\":\"assistant\",\"content\":\"" + text
+ + "\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":" + in + ",\"completion_tokens\":"
+ + out + ",\"total_tokens\":" + (in + out) + "}}";
+ }
+
+ public static String openaiToolCall(String id, String tool, String argsJson, int in, int out) {
+ return "{\"id\":\"chatcmpl-2\",\"object\":\"chat.completion\",\"created\":1700000000,\"model\":\"gpt-6.1-sol\","
+ + "\"choices\":[{\"index\":0,\"message\":{\"role\":\"assistant\",\"content\":null,\"tool_calls\":[{\"id\":\""
+ + id + "\",\"type\":\"function\",\"function\":{\"name\":\"" + tool + "\",\"arguments\":"
+ + quote(argsJson) + "}}]},\"finish_reason\":\"tool_calls\"}],\"usage\":{\"prompt_tokens\":" + in
+ + ",\"completion_tokens\":" + out + ",\"total_tokens\":" + (in + out) + "}}";
+ }
+
+ public static String anthropicText(String text, int in, int out) {
+ return "{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"model\":\"claude-sonnet-5-5\","
+ + "\"content\":[{\"type\":\"text\",\"text\":\"" + text + "\"}],\"stop_reason\":\"end_turn\","
+ + "\"stop_sequence\":null,\"usage\":{\"input_tokens\":" + in + ",\"output_tokens\":" + out + "}}";
+ }
+
+ public static String error(String type, String message) {
+ return "{\"error\":{\"type\":\"" + type + "\",\"message\":\"" + message + "\"}}";
+ }
+
+ private static String quote(String s) {
+ return "\"" + s.replace("\\", "\\\\").replace("\"", "\\\"") + "\"";
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/support/Fixtures.java b/llm-gateway/src/test/java/com/ankurm/gateway/support/Fixtures.java
new file mode 100644
index 0000000..97847c5
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/support/Fixtures.java
@@ -0,0 +1,49 @@
+package com.ankurm.gateway.support;
+
+import java.util.List;
+import java.util.Map;
+
+import com.ankurm.gateway.Budgets;
+import com.ankurm.gateway.Failover;
+import com.ankurm.gateway.Gateway;
+import com.ankurm.gateway.GatewayConfig;
+import com.ankurm.gateway.GatewayProperties;
+import com.ankurm.gateway.OrderTools;
+import com.ankurm.gateway.Price;
+import com.ankurm.gateway.Routes;
+import com.ankurm.gateway.Target;
+import com.ankurm.gateway.TokenLimiter;
+import io.github.bucket4j.TimeMeter;
+import io.github.resilience4j.circuitbreaker.CircuitBreakerRegistry;
+
+public final class Fixtures {
+
+ /** Dollars per million tokens, the "smart tier" sheet from the providers module (read 2026-10-09). */
+ public static final Price PRICE = Price.of("2.00", "10.00");
+
+ private Fixtures() {
+ }
+
+ public static Target target(String name, Scripted model) {
+ return new Target(name, name + "-model", model, PRICE);
+ }
+
+ public static Budgets unlimited() {
+ return new Budgets(t -> Long.MAX_VALUE / 4);
+ }
+
+ /** The same settings application.yml uses, with a short open-state wait so tests can see recovery. */
+ public static CircuitBreakerRegistry breakers() {
+ return GatewayConfig.registry(new GatewayProperties.Breaker(10, 5, 50, 200));
+ }
+
+ public static Failover failover(Target... targets) {
+ return new Failover(List.of(targets), breakers(), unlimited());
+ }
+
+ public static Gateway gateway(OrderTools tools, Budgets budgets, Target... targets) {
+ CircuitBreakerRegistry r = breakers();
+ return new Gateway(new Routes(Map.of("smart", List.of(targets))), t -> new Failover(t, r, budgets),
+ new TokenLimiter(t -> 1_000_000, TimeMeter.SYSTEM_MILLISECONDS), null, Map.of("refund_order", tools));
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/support/Scripted.java b/llm-gateway/src/test/java/com/ankurm/gateway/support/Scripted.java
new file mode 100644
index 0000000..9923e1e
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/support/Scripted.java
@@ -0,0 +1,107 @@
+package com.ankurm.gateway.support;
+
+import java.util.ArrayDeque;
+import java.util.Deque;
+import java.util.List;
+import java.util.concurrent.CopyOnWriteArrayList;
+import java.util.function.Supplier;
+
+import com.ankurm.gateway.ProviderFailure;
+import org.springframework.ai.chat.messages.AssistantMessage;
+import org.springframework.ai.chat.metadata.ChatResponseMetadata;
+import org.springframework.ai.chat.metadata.DefaultUsage;
+import org.springframework.ai.chat.model.ChatModel;
+import org.springframework.ai.chat.model.ChatResponse;
+import org.springframework.ai.chat.model.Generation;
+import org.springframework.ai.chat.prompt.ChatOptions;
+import org.springframework.ai.chat.prompt.Prompt;
+import org.springframework.ai.model.tool.ToolCallingChatOptions;
+
+/**
+ * A {@link ChatModel} that plays back a script: each call takes the next step, which either
+ * answers (with usage) or throws an HTTP-style failure. It is a stand-in for a provider, so it
+ * records every prompt it receives, including the options the gateway built for it.
+ */
+public final class Scripted implements ChatModel {
+
+ private final Deque> steps = new ArrayDeque<>();
+
+ private final List prompts = new CopyOnWriteArrayList<>();
+
+ private final String name;
+
+ /** What it does when the script runs out. */
+ private Supplier otherwise;
+
+ public Scripted(String name) {
+ this.name = name;
+ this.otherwise = () -> {
+ throw new IllegalStateException(name + " ran out of script");
+ };
+ }
+
+ public static ChatResponse text(String text, int in, int out) {
+ return new ChatResponse(List.of(new Generation(new AssistantMessage(text))),
+ ChatResponseMetadata.builder().usage(new DefaultUsage(in, out)).build());
+ }
+
+ public static ChatResponse toolCall(String id, String tool, String json, int in, int out) {
+ AssistantMessage m = AssistantMessage.builder().content("")
+ .toolCalls(List.of(new AssistantMessage.ToolCall(id, "function", tool, json))).build();
+ return new ChatResponse(List.of(new Generation(m)),
+ ChatResponseMetadata.builder().usage(new DefaultUsage(in, out)).build());
+ }
+
+ public Scripted answer(String text, int in, int out) {
+ steps.add(() -> text(text, in, out));
+ return this;
+ }
+
+ public Scripted callTool(String id, String tool, String json, int in, int out) {
+ steps.add(() -> toolCall(id, tool, json, in, out));
+ return this;
+ }
+
+ public Scripted fail(int status) {
+ steps.add(() -> {
+ throw new ProviderFailure(status, name + " HTTP " + status);
+ });
+ return this;
+ }
+
+ public Scripted failAlways(int status) {
+ otherwise = () -> {
+ throw new ProviderFailure(status, name + " HTTP " + status);
+ };
+ return this;
+ }
+
+ public Scripted answerAlways(String text, int in, int out) {
+ otherwise = () -> text(text, in, out);
+ return this;
+ }
+
+ @Override
+ public ChatResponse call(Prompt prompt) {
+ prompts.add(prompt);
+ Supplier s = steps.poll();
+ return (s != null ? s : otherwise).get();
+ }
+
+ @Override
+ public ChatOptions getOptions() {
+ return ToolCallingChatOptions.builder().build();
+ }
+
+ public List prompts() {
+ return prompts;
+ }
+
+ public int calls() {
+ return prompts.size();
+ }
+
+ public String modelOfCall(int i) {
+ return prompts.get(i).getOptions().getModel();
+ }
+}
diff --git a/llm-gateway/src/test/java/com/ankurm/gateway/support/Transcript.java b/llm-gateway/src/test/java/com/ankurm/gateway/support/Transcript.java
new file mode 100644
index 0000000..4287522
--- /dev/null
+++ b/llm-gateway/src/test/java/com/ankurm/gateway/support/Transcript.java
@@ -0,0 +1,47 @@
+package com.ankurm.gateway.support;
+
+import java.io.IOException;
+import java.io.PrintWriter;
+import java.io.StringWriter;
+import java.nio.file.Files;
+import java.nio.file.Path;
+
+/**
+ * Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
+ * echoes it to the console. Every console block quoted in the article comes out of one of these
+ * files verbatim.
+ */
+public final class Transcript implements AutoCloseable {
+
+ private final Path path;
+ private final StringWriter buffer = new StringWriter();
+ private final PrintWriter out = new PrintWriter(buffer);
+
+ public Transcript(String fileName, String title) {
+ this.path = Path.of("output", fileName);
+ out.println("# " + title);
+ out.println();
+ }
+
+ public Transcript line(String format, Object... args) {
+ out.println(args.length == 0 ? format : String.format(format, args));
+ return this;
+ }
+
+ public Transcript blank() {
+ out.println();
+ return this;
+ }
+
+ @Override
+ public void close() {
+ out.flush();
+ try {
+ Files.createDirectories(path.getParent());
+ Files.writeString(path, buffer.toString());
+ } catch (IOException e) {
+ throw new IllegalStateException("could not write " + path, e);
+ }
+ System.out.print(buffer);
+ }
+}