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

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
Claude
2026-10-09 10:24:57 +00:00
parent 80cd21f89b
commit cdee85d3f4
51 changed files with 2404 additions and 0 deletions
@@ -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);
}
}