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:
@@ -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/).
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
target/
|
||||
@@ -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.
|
||||
@@ -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]
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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>
|
||||
Executable
+8
@@ -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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user