Add llm-gateway module: routing, failover below the tool-calling advisor, per-provider circuit breakers, dollar caps and token limits, tenant-keyed cache; real OpenAI and Anthropic models against a local fake

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
Claude
2026-10-09 10:24:57 +00:00
parent 80cd21f89b
commit cdee85d3f4
51 changed files with 2404 additions and 0 deletions
+1
View File
@@ -21,5 +21,6 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur
| [`multimodal/`](multimodal) | A receipt image through `Media` and `ChatClient.entity(...)` into a Java record, on the real `OpenAiChatModel`, `AnthropicChatModel` and `OllamaChatModel` against a local server that OCRs the image it receives (so accuracy figures describe OCR, not any vision model). The same image on three wire formats, arithmetic validation and a repair retry, accuracy under tilt, shrinking and noise, and an image-token estimate from a documented formula. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Multimodal Spring AI: Extract Structured Data from Images](https://ankurm.com/multimodal-spring-ai-extract-structured-data-from-images-receipts-java-records/) |
| [`text-to-sql/`](text-to-sql) | A question to SQL to rows, safely, on a real PostgreSQL 16: a schema prompt built through the restricted role, a JSqlParser guard (one statement, SELECT only, listed tables and functions), a read-only role with column grants and a statement timeout, a row cap, and evaluation by comparing results. Sixteen queries against four setups (13 harmful: 13 succeed with neither protection, 0 with both). The model is a script, not a language model. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Text-to-SQL with Spring AI, Done Safely](https://ankurm.com/text-to-sql-spring-ai-read-only-roles-query-validation/) |
| [`providers/`](providers) | One ticket-summarising app on the real `OpenAiChatModel`, `AnthropicChatModel` and `GoogleGenAiChatModel` against a local server in all three wire formats: where each puts the system prompt and what it sends by default, portable versus provider-specific options (and the per-call `ChatOptions` that crashes two providers and silently resets the third), prompt caching, a cost table from the vendors' price sheets, and failover with the retry multiplication measured. No vendor API called; token counts and cache hits are simulated from documented rules; latency not measured. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Anthropic Claude vs OpenAI vs Gemini in Spring AI 2.0](https://ankurm.com/spring-ai-claude-vs-openai-vs-gemini-switch-providers-compare-cost/) |
| [`llm-gateway/`](llm-gateway) | A gateway service on Spring AI 2.0: model-hint routing, failover under the tool-calling advisor (a provider failing mid tool loop does not re-run the tool), one circuit breaker per provider, dollar caps that reserve before the call, a token-per-hour limiter, and a tenant-keyed semantic cache. Real OpenAI and Anthropic models against a local fake; no vendor API called. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [The LLM Gateway Pattern for Java Microservices](https://ankurm.com/llm-gateway-pattern-java-microservices/) |
Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/).
+1
View File
@@ -0,0 +1 @@
target/
+66
View File
@@ -0,0 +1,66 @@
# llm-gateway
Companion code for [The LLM Gateway Pattern for Java Microservices](https://ankurm.com/llm-gateway-pattern-java-microservices/) and [LLM Gateway for Java Microservices: Complete Runnable Code](https://ankurm.com/llm-gateway-complete-example/), part of the [Spring AI series](../README.md) on ankurm.com.
A small Spring Boot 4.1 service that sits between your microservices and the LLM providers. A calling service sends one request and a tenant header. The gateway picks a provider chain from a model hint, fails over on transient errors, keeps one circuit breaker per provider, caps each tenant in tokens per hour and in dollars, and runs an allow-listed tool loop.
**No vendor API was called.** The tests use two kinds of stand-in, and each output file says which:
- `Scripted`, a `ChatModel` that plays back a script (answer, tool call, or an HTTP-style failure). Used where the point is the gateway's own logic.
- `FakeVendors`, a local HTTP server that answers in the OpenAI and Anthropic wire formats. The **real** `OpenAiChatModel` and `AnthropicChatModel` talk to it, so the requests it records are the requests those classes would send, and the exceptions are the vendor SDKs' own.
Token counts in the fakes are numbers the test chose, and the semantic-cache test uses a hashed bag-of-words as its embedding model. Nothing here measures a real provider's latency, accuracy or cache behaviour.
## Versions
| Component | Version |
|---|---|
| Spring Boot | 4.1.1 (parent) |
| Spring AI | 2.0.1 (`spring-ai-client-chat`, `spring-ai-openai`, `spring-ai-anthropic`) |
| Resilience4j | 2.4.0 (`resilience4j-circuitbreaker`, used programmatically) |
| Bucket4j | 8.21.0 (`bucket4j_jdk17-core`, in memory) |
| Java | 25 (LTS) |
## Quickstart
```bash
scripts/run-all.sh # runs the suite, regenerates output/
OPENAI_API_KEY=... ANTHROPIC_API_KEY=... mvn spring-boot:run # a live gateway on :8080
curl -X POST localhost:8080/v1/gateway/complete -H 'X-Tenant-Id: acme' -H 'Content-Type: application/json' \
-d '{"featureTag":"support","userMessage":"What is your return policy?","modelHint":"smart","maxTokens":256}'
```
Two consecutive test runs produce byte-identical files. The live run needs real keys and was not done for this article.
## What's here
| File | What it is |
|---|---|
| [`Failover.java`](src/main/java/com/ankurm/gateway/Failover.java) | The gateway's `ChatModel`: ordered targets, a breaker and a budget reservation per attempt |
| [`Gateway.java`](src/main/java/com/ankurm/gateway/Gateway.java) | Rate limit, cache, `ChatClient` with `ToolCallingAdvisor`, metering |
| [`Routes.java`](src/main/java/com/ankurm/gateway/Routes.java) | Hint to provider chain; an unknown hint is an error |
| [`Budgets.java`](src/main/java/com/ankurm/gateway/Budgets.java) | Dollar caps: reserve before the call, settle after |
| [`TokenLimiter.java`](src/main/java/com/ankurm/gateway/TokenLimiter.java) | Tokens per hour per tenant, Bucket4j |
| [`SemanticCache.java`](src/main/java/com/ankurm/gateway/SemanticCache.java) | Embedding cache keyed by tenant and feature |
| [`Failures.java`](src/main/java/com/ankurm/gateway/Failures.java) | HTTP status out of the vendor SDK exceptions |
| [`GatewayConfig.java`](src/main/java/com/ankurm/gateway/GatewayConfig.java), [`application.yml`](src/main/resources/application.yml) | Providers, routes, tenants and breaker settings as data |
## Output files
| File | Written by |
|---|---|
| [`01-routing.txt`](output/01-routing.txt) | `RoutingTest`: hints, one model id per provider, the unknown-hint hazard |
| [`02-failover.txt`](output/02-failover.txt) | `FailoverTest`: which statuses fail over |
| [`03-circuit-breaker.txt`](output/03-circuit-breaker.txt) | `CircuitBreakerTest`: when a breaker opens, isolation, recovery |
| [`04-tool-loop.txt`](output/04-tool-loop.txt) | `ToolLoopPlacementTest`: failover below vs above the tool advisor, usage totals |
| [`05-wire.txt`](output/05-wire.txt) | `WireTest`: the real models failing over mid tool loop, the requests on the wire |
| [`06-real-exceptions.txt`](output/06-real-exceptions.txt) | `WireTest`: what the SDKs throw for each status |
| [`07-budget.txt`](output/07-budget.txt) | `BudgetTest`: caps, estimates, a cap hit mid-loop |
| [`07b-budget-concurrency.txt`](output/07b-budget-concurrency.txt) | `BudgetTest`: 64 simultaneous requests |
| [`08-rate-limit.txt`](output/08-rate-limit.txt) | `RateLimitTest`: pre-consume, settle, refill |
| [`09-cache.txt`](output/09-cache.txt) | `CacheTest`: threshold and tenant key mechanics |
| [`10-http.txt`](output/10-http.txt) | `HttpTest`: the application over HTTP |
## Not covered
Streaming (`ScopedValue` context does not cross Reactor threads), a shared store for breaker or limiter state across replicas (everything here is per instance), real provider behaviour, and anything measured about latency.
+11
View File
@@ -0,0 +1,11 @@
# Routing: hints, per-target options, unknown hints
hint smart -> answered by claude (claude-model)
model id sent to openai: openai-model
model id sent to claude: claude-model
maxTokens sent to both: 256 and 256
hint local -> answered by local; cloud calls made for it: 0
old registry, hint "locla" (typo for local): goes to claude, a cloud provider
this Routes, hint "locla": Unknown model hint 'locla'. Known hints: [local, smart]
+12
View File
@@ -0,0 +1,12 @@
# Failover: which statuses move to the next provider
status second tried? outcome trail
408 true answered [first: failed with status 408, second: ok]
429 true answered [first: failed with status 429, second: ok]
500 true answered [first: failed with status 500, second: ok]
503 true answered [first: failed with status 503, second: ok]
400 false rethrown [first: rejected with status 400, not retried]
401 false rethrown [first: rejected with status 401, not retried]
404 false rethrown [first: rejected with status 404, not retried]
both down: No provider could answer: [a: failed with status 503, b: failed with status 429]
+11
View File
@@ -0,0 +1,11 @@
# Circuit breakers: when they open and how they recover
old yml (window 10, threshold 50%): minimumNumberOfCalls is 100; the first request that skips openai is #11
this module (window 10, minimum 5, threshold 50%): the first request that skips openai is #6
after 8 requests: openai breaker OPEN, claude breaker CLOSED
openai was contacted 5 times of 8; claude answered 8 times
after the wait, one probe goes to openai (it has recovered): breaker CLOSED, openai calls 1
20 requests rejected with 400: breaker CLOSED, calls the breaker recorded: 0
+19
View File
@@ -0,0 +1,19 @@
# Where failover sits relative to the tool loop
A. failover below the advisor (this module)
tool executions: 1
trail: [openai: ok, openai: failed with status 503, claude: ok]
answer: Refunded as refund-A17-1 (from claude)
what claude was sent: UserMessage, AssistantMessage, ToolResponseMessage
B. retry above the advisor (one ChatClient per provider)
tool executions: 2
answer: Refunded as refund-A17-2
C. usage across a two-call tool loop (calls reported 120+15 and 160+12)
usage on the final ChatResponse: prompt=280 completion=27
usage summed per model call: prompt=280 completion=27
D. the same loop with the first call on a cheaper provider (0.75/3.75 then 2.00/10.00)
priced per model call: 587 microdollars
total usage at the last provider's price: 830 microdollars
+8
View File
@@ -0,0 +1,8 @@
# Real OpenAI and Anthropic models: failover in the middle of a tool loop
trail: [openai: ok, openai: failed with status 503, anthropic: ok]
answered by anthropic; tool executions 1; requests: openai 2, anthropic 1
request 1 to OpenAI: model=gpt-6.1-sol, max_tokens=256, tools=[refund_order]
request to Anthropic: model=claude-sonnet-5-5 max_tokens=256
conversation it received: [assistant:tool_use(id=call_1, name=refund_order), user:tool_result(tool_use_id=call_1)]
@@ -0,0 +1,9 @@
# What the vendor SDKs throw, and what the gateway does with it
vendor status exception transient? HTTP requests for one call
openai 400 com.openai.errors.BadRequestException false 1
openai 401 com.openai.errors.UnauthorizedException false 1
openai 429 com.openai.errors.RateLimitException true 1
openai 503 com.openai.errors.InternalServerException true 1
anthropic 400 com.anthropic.errors.BadRequestException false 1
anthropic 529 com.anthropic.errors.InternalServerException true 1
+11
View File
@@ -0,0 +1,11 @@
# Dollar caps: reserve, settle, and the limits of an estimate
cap 5000 microdollars; every call really costs 2000; maxTokens 256 (estimate about 2,570)
call 1: answered, cost 2000, spent so far 2000
call 2: answered, cost 2000, spent so far 4000
call 3: rejected before the provider was called (provider calls so far: 2)
cap 8000; prompt estimated at 100 tokens but the provider counts 3000 (code, other scripts, images do this)
spent after 2 calls: 12200 (cap 8000); the second call was admitted on its estimate
cap 2700, a two-call tool loop: rejected on the SECOND model call; tool executions so far: 1
@@ -0,0 +1,5 @@
# 64 simultaneous requests against a cap that fits 10
cap 10000, each call 1000, 64 threads at once
check then record: admitted 64, spent 64000
reserve then settle: admitted 10
+14
View File
@@ -0,0 +1,14 @@
# Token rate limit: pre-consume, settle, refill
capacity 10000 tokens per hour
take estimate 1500 -> available 8500
settle: prompt 100 + completion 1000 = 1100 -> available 8900
old rule (refund estimate minus PROMPT tokens only): would have refunded 1400 and left 100 charged
take 500, real total 2000 -> available 6900 (the shortfall is charged, not forgiven)
take 6900 -> available 0
take 1000 -> rejected: Token rate limit reached for acme
30 minutes later -> available 5000 (greedy refill: 10000 per hour)
60 more minutes -> available 10000 (capped at capacity)
+10
View File
@@ -0,0 +1,10 @@
# Semantic cache mechanics (stand-in embeddings, threshold 0.92)
EmbeddingModel.embed(String) returns: float[]
probe cosine hit at 0.92?
same words, new order 0.985 true
one word different (A18 for A17) 0.970 true
different question, shares a few words 0.395 false
tenant globex asks the stored question: miss
a cache keyed by feature only, tenant globex asks: 30 days, original packaging.
+10
View File
@@ -0,0 +1,10 @@
# The gateway over HTTP (real beans, real models, fake vendors)
1. normal call -> 200 {"content":"30 days.","provider":"openai","model":"gpt-6.1-sol","promptTokens":47,"completionTokens":52,"costMicros":614,"modelCalls":1,"servedFromCache":false,"trail":["openai: ok"]}
2. openai returns 503 -> 200 {"content":"30 days.","provider":"anthropic","model":"claude-sonnet-5-5","promptTokens":47,"completionTokens":52,"costMicros":614,"modelCalls":1,"servedFromCache":false,"trail":["openai: failed with status 503","anthropic: ok"]}
3. both providers down -> 503 Retry-After=30 No provider could answer: [openai: failed with status 503, anthropic: failed with status 529]
4. unknown hint -> 400 Unknown model hint 'locla'. Known hints: [fast, smart]
5. tool not allow-listed -> 400 Tool 'drop_tables' is not on the gateway allow-list [refund_order]
6. tenant over its cap -> 429 Budget for poor would be exceeded: spent 0 of 1000 microdollars, this call may cost up to 2574
7. body claims tenant acme, header says poor -> 429 (the header wins)
8. no X-Tenant-Id header -> 400
+84
View File
@@ -0,0 +1,84 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-parent</artifactId>
<version>4.1.1</version>
<relativePath/>
</parent>
<groupId>com.ankurm</groupId>
<artifactId>llm-gateway</artifactId>
<version>1.0.0</version>
<name>llm-gateway</name>
<description>An LLM gateway on Spring AI 2.0: routing, failover, circuit breakers, token and dollar caps.</description>
<properties>
<java.version>25</java.version>
<spring-ai.version>2.0.1</spring-ai.version>
<resilience4j.version>2.4.0</resilience4j.version>
<bucket4j.version>8.21.0</bucket4j.version>
</properties>
<dependencyManagement>
<dependencies>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-bom</artifactId>
<version>${spring-ai.version}</version>
<type>pom</type>
<scope>import</scope>
</dependency>
</dependencies>
</dependencyManagement>
<dependencies>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-webmvc</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-client-chat</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-openai</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-anthropic</artifactId>
</dependency>
<dependency>
<groupId>io.github.resilience4j</groupId>
<artifactId>resilience4j-circuitbreaker</artifactId>
<version>${resilience4j.version}</version>
</dependency>
<dependency>
<groupId>com.bucket4j</groupId>
<artifactId>bucket4j_jdk17-core</artifactId>
<version>${bucket4j.version}</version>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
<build>
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-surefire-plugin</artifactId>
<configuration>
<argLine>-Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8</argLine>
</configuration>
</plugin>
</plugins>
</build>
</project>
+8
View File
@@ -0,0 +1,8 @@
#!/usr/bin/env bash
# Regenerates every file under output/ from the test suite. No Docker, no API key, no network
# (Maven needs its usual dependency downloads on the first run). The suite takes under a minute.
set -euo pipefail
cd "$(dirname "$0")/.."
rm -rf target
mvn -q -B test 2>&1 | grep -E "Tests run:|BUILD|FAIL|ERROR" || true
ls output
@@ -0,0 +1,17 @@
package com.ankurm.gateway;
import java.util.List;
public class AllProvidersUnavailableException extends RuntimeException {
private final List<String> trail;
public AllProvidersUnavailableException(List<String> trail, Throwable cause) {
super("No provider could answer: " + trail, cause);
this.trail = List.copyOf(trail);
}
public List<String> trail() {
return trail;
}
}
@@ -0,0 +1,8 @@
package com.ankurm.gateway;
public class BudgetExceededException extends RuntimeException {
public BudgetExceededException(String message) {
super(message);
}
}
@@ -0,0 +1,90 @@
package com.ankurm.gateway;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.function.ToLongFunction;
/**
* Dollar caps per tenant, in microdollars. The point of the class is the order of operations:
* {@link #reserve} takes the money BEFORE the provider is called, and {@link Reservation#settle}
* corrects it afterwards. Checking the balance and recording the spend as two separate steps lets
* concurrent requests all pass the check and all spend (see {@code BudgetTest}).
*/
public final class Budgets {
/** Money held for one call. Settle with the real cost, or release if the call failed. */
public final class Reservation {
private final Account account;
private final long held;
private boolean done;
private Reservation(Account account, long held) {
this.account = account;
this.held = held;
}
public synchronized void settle(long actualMicros) {
if (!done) {
done = true;
account.settle(held, actualMicros);
}
}
public synchronized void release() {
if (!done) {
done = true;
account.settle(held, 0);
}
}
}
private static final class Account {
private long spent;
private long held;
synchronized boolean reserve(long amount, long cap) {
if (spent + held + amount > cap) {
return false;
}
held += amount;
return true;
}
synchronized void settle(long heldAmount, long actual) {
held -= heldAmount;
spent += actual;
}
synchronized long spent() {
return spent;
}
}
private final Map<String, Account> accounts = new ConcurrentHashMap<>();
private final ToLongFunction<String> capMicros;
public Budgets(ToLongFunction<String> capMicros) {
this.capMicros = capMicros;
}
public Reservation reserve(String tenant, long estimateMicros) {
Account a = accounts.computeIfAbsent(tenant, t -> new Account());
long cap = capMicros.applyAsLong(tenant);
if (!a.reserve(estimateMicros, cap)) {
throw new BudgetExceededException("Budget for " + tenant + " would be exceeded: spent "
+ a.spent() + " of " + cap + " microdollars, this call may cost up to " + estimateMicros);
}
return new Reservation(a, estimateMicros);
}
public long spent(String tenant) {
Account a = accounts.get(tenant);
return a == null ? 0 : a.spent();
}
}
@@ -0,0 +1,85 @@
package com.ankurm.gateway;
import java.util.ArrayList;
import java.util.List;
/**
* What the model-level code needs to know about the request it is serving. The tool-calling loop
* runs ABOVE the {@code ChatModel}, so tenant and feature cannot travel as method arguments; they
* travel in a {@link ScopedValue}, which is bound for exactly the duration of one gateway call.
*/
public record CallContext(String tenant, String feature, Meter meter) {
public static final ScopedValue<CallContext> CURRENT = ScopedValue.newInstance();
public static CallContext anonymous() {
return new CallContext("anonymous", "none", new Meter());
}
/** Totals across every model call made for one gateway request, tool rounds and failovers included. */
public static final class Meter {
private final List<String> trail = new ArrayList<>();
private long promptTokens;
private long completionTokens;
private long costMicros;
private int modelCalls;
private String provider = "";
private String model = "";
private boolean usageMissing;
public synchronized void trail(String entry) {
trail.add(entry);
}
public synchronized void record(Target t, long prompt, long completion, long micros, boolean missing) {
promptTokens += prompt;
completionTokens += completion;
costMicros += micros;
modelCalls++;
provider = t.name();
model = t.model();
usageMissing |= missing;
}
public synchronized List<String> 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;
}
}
}
@@ -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.
*
* <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;
}
}
@@ -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;
}
}
@@ -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<List<Target>, Failover> failovers;
private final TokenLimiter limiter;
private final SemanticCache cache;
private final Map<String, ToolCallback[]> toolbox;
public Gateway(Routes routes, java.util.function.Function<List<Target>, Failover> failovers,
TokenLimiter limiter, SemanticCache cache, Map<String, Object> 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<Target> targets = routes.get(req.modelHint());
List<ToolCallback> 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.<String>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<ToolCallback> resolveTools(List<String> names) {
java.util.ArrayList<ToolCallback> 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;
}
}
@@ -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);
}
}
@@ -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<String, Target> 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<String, List<Target>> 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<EmbeddingModel> 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()));
}
}
@@ -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<String> overBudget(BudgetExceededException e) {
return ResponseEntity.status(HttpStatus.TOO_MANY_REQUESTS).body(e.getMessage());
}
@ExceptionHandler(AllProvidersUnavailableException.class)
ResponseEntity<String> down(AllProvidersUnavailableException e) {
return ResponseEntity.status(HttpStatus.SERVICE_UNAVAILABLE).header("Retry-After", "30").body(e.getMessage());
}
@ExceptionHandler({UnknownHintException.class, IllegalArgumentException.class})
ResponseEntity<String> badRequest(RuntimeException e) {
return ResponseEntity.badRequest().body(e.getMessage());
}
}
@@ -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<Provider> providers, Map<String, List<String>> routes,
Map<String, Tenant> 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;
}
}
@@ -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<String> 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 = "";
}
}
}
@@ -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<String> trail) {
}
@@ -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();
}
}
@@ -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();
}
}
@@ -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;
}
}
@@ -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<String, List<Target>> table;
public Routes(Map<String, List<Target>> table) {
this.table = Map.copyOf(table);
}
public List<Target> get(String hint) {
List<Target> t = table.get(hint);
if (t == null) {
throw new UnknownHintException(hint, table.keySet());
}
return t;
}
}
@@ -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<String, List<Entry>> byKey = new ConcurrentHashMap<>();
public SemanticCache(EmbeddingModel embeddings, double threshold) {
this.embeddings = embeddings;
this.threshold = threshold;
}
public Optional<String> 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);
}
}
@@ -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) {
}
@@ -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<String, Bucket> buckets = new ConcurrentHashMap<>();
private final ToLongFunction<String> perHour;
private final TimeMeter clock;
public TokenLimiter(ToLongFunction<String> 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();
}
}
@@ -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<String> known) {
super("Unknown model hint '" + hint + "'. Known hints: " + new TreeSet<>(known));
}
}
@@ -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 }
@@ -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<Void> task) throws Exception {
try (var pool = Executors.newFixedThreadPool(threads)) {
List<Future<Void>> futures = new java.util.ArrayList<>();
for (int i = 0; i < threads; i++) {
futures.add(pool.submit(task));
}
for (Future<Void> f : futures) {
f.get();
}
}
}
}
@@ -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<Embedding> 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<Probe> 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<String, String> 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");
}
}
}
@@ -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);
}
}
}
@@ -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);
}
}
}
@@ -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<String> 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<String> 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<String> 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<String> 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<String> 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<String> 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<String> 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<String> 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<String> none = post(null, body("smart", ""));
t.line("8. no X-Tenant-Id header -> %d", none.statusCode());
assertThat(none.statusCode()).isEqualTo(400);
}
}
}
@@ -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);
}
}
}
@@ -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<String, List<Target>> old = Map.of("smart", List.of(tClaude), "local", List.of(tLocal));
List<Target> 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);
}
}
}
@@ -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);
}
}
}
@@ -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<String> 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<String> 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;
}
}
}
@@ -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<Reply> queue = new ArrayDeque<>();
private final List<String> 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<String> 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("\"", "\\\"") + "\"";
}
}
@@ -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));
}
}
@@ -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<Supplier<ChatResponse>> steps = new ArrayDeque<>();
private final List<Prompt> prompts = new CopyOnWriteArrayList<>();
private final String name;
/** What it does when the script runs out. */
private Supplier<ChatResponse> 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<ChatResponse> s = steps.poll();
return (s != null ? s : otherwise).get();
}
@Override
public ChatOptions getOptions() {
return ToolCallingChatOptions.builder().build();
}
public List<Prompt> prompts() {
return prompts;
}
public int calls() {
return prompts.size();
}
public String modelOfCall(int i) {
return prompts.get(i).getOptions().getModel();
}
}
@@ -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);
}
}