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:
@@ -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