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

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