Add providers module: one app on the real OpenAI, Anthropic and Gemini Spring AI models against a local server in three wire formats; options, prompt caching, cost from price sheets, failover with retry layers measured
Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
@@ -18,7 +18,8 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur
|
||||
| [`evaluation/`](evaluation) | Testing an LLM app: `RelevancyEvaluator` and `FactCheckingEvaluator` (exactly what they send and which judge replies they accept), a 12-case golden dataset with a pass-rate gate, a deterministic judge for CI, simulated judge noise, a 1-5 graded evaluator and a composite. Stub models only; the one live-judge test is skipped without a key. Spring Boot 4.1.1, Spring AI 2.0.1, JUnit 6, Java 25. | [Testing LLM Apps in Java: Spring AI Evaluators and LLM-as-Judge in JUnit 6](https://ankurm.com/spring-ai-2-0-testing-llm-apps-evaluators-llm-as-judge-junit-6/) |
|
||||
| [`observability/`](observability) | What Spring AI 2.0.1 records on its own (model, chat client, advisor and tool meters, spans for a tool call), token usage turned into cost per endpoint, a Grafana dashboard checked against live Prometheus and Grafana, and the traps: histograms are opt-in, a response with no usage looks like a free call, prompt text is logged only if switched on. The real `OpenAiChatModel` against a local fake server, so counts are approximate and prices illustrative. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Observability for Spring AI: Tokens, Latency and Cost with Micrometer and OpenTelemetry](https://ankurm.com/spring-ai-2-0-observability-micrometer-opentelemetry-tokens-cost/) |
|
||||
| [`guardrails/`](guardrails) | Prompt injection against a Spring AI assistant with tools: a poisoned document, a poisoned tool result, a markdown-image leak and a system prompt leak, run against a document filter, a tool allow-list with argument policies, and output validation, alone and together (6 of 6 attacks succeed with no defence, 0 of 6 with all three). A deliberately gullible stub model, so it measures what each defence stops when the model *is* fooled, not how often a real model is. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Prompt Injection Defense in Spring AI](https://ankurm.com/prompt-injection-defense-spring-ai-guardrails-tool-allow-lists-output-validation/) |
|
||||
|
||||
Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/).
|
||||
| [`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/) |
|
||||
|
||||
Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/).
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
target/
|
||||
@@ -0,0 +1,44 @@
|
||||
# providers
|
||||
|
||||
Companion code for [Anthropic Claude vs OpenAI vs Gemini in Spring AI 2.0: Switching Providers and Comparing Cost](https://ankurm.com/spring-ai-claude-vs-openai-vs-gemini-switch-providers-compare-cost/), part of the [Spring AI series](../README.md) on ankurm.com.
|
||||
|
||||
One small application (`TicketService`, which summarises a support ticket) runs on the real `OpenAiChatModel`, `AnthropicChatModel` and `GoogleGenAiChatModel`, and the tests look at what each one does differently.
|
||||
|
||||
**No vendor API was called, and no claim is made about any vendor's behaviour.** `FakeProviderServer` is a local HTTP server that answers in the three wire formats. It records every request the real Spring AI classes send, which is the evidence for the "what is different on the wire" results. Three things in it are simulation: token counts are characters / 4, the answer is a fixed JSON string, and cache hits follow the rules the vendors document (thresholds and prefix matching, read on 2026-10-09) applied to the request that really arrived. So the cost table is a worked example of the price sheets, not a bill, and **latency is not measured at all**.
|
||||
|
||||
## Versions
|
||||
|
||||
| Component | Version |
|
||||
|---|---|
|
||||
| Spring Boot | 4.1.1 (parent) |
|
||||
| Spring AI | 2.0.1 (`spring-ai-openai`, `spring-ai-anthropic`, `spring-ai-google-genai`) |
|
||||
| Models named in the examples | `gpt-6.1-sol`, `claude-sonnet-5-5`, `gemini-3.8-flash` |
|
||||
| Java | 25 (LTS) |
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
scripts/run-all.sh # runs the suite, regenerates output/01 .. 05
|
||||
```
|
||||
|
||||
Two consecutive runs produce byte-identical files.
|
||||
|
||||
## What's here
|
||||
|
||||
| File | What it is |
|
||||
|---|---|
|
||||
| [`ProviderModels.java`](src/main/java/com/ankurm/providers/ProviderModels.java) | The one class that knows which vendor is behind a `ChatModel` |
|
||||
| [`TicketService.java`](src/main/java/com/ankurm/providers/TicketService.java) | The application; sees only `ChatModel` |
|
||||
| [`FallbackChatModel.java`](src/main/java/com/ankurm/providers/FallbackChatModel.java) | Tries providers in order on transient failures |
|
||||
| [`Failures.java`](src/main/java/com/ankurm/providers/Failures.java) | Reads an HTTP status out of three unrelated exception families |
|
||||
| [`Cost.java`](src/main/java/com/ankurm/providers/Cost.java) | Usage to dollars, with the per-provider meaning of "prompt tokens" |
|
||||
|
||||
## Output files
|
||||
|
||||
| File | Written by |
|
||||
|---|---|
|
||||
| [`01-wire.txt`](output/01-wire.txt) | `WireTest`: the same call as each provider's HTTP request |
|
||||
| [`02-options.txt`](output/02-options.txt) | `OptionsTest`: portable and provider-specific options, and the per-call trap |
|
||||
| [`03-caching.txt`](output/03-caching.txt) | `CachingTest`: breakpoints sent, cache tokens reported, text order |
|
||||
| [`04-cost.txt`](output/04-cost.txt) | `CostTest`: 100 requests priced on three price sheets |
|
||||
| [`05-fallback.txt`](output/05-fallback.txt) | `FallbackTest`: which failures fail over, and requests per plan |
|
||||
@@ -0,0 +1,20 @@
|
||||
# One TicketService.raw() call, three providers, as the HTTP server saw it
|
||||
|
||||
OPENAI POST /v1/chat/completions
|
||||
top-level keys : messages, model
|
||||
model sent : gpt-6.1-sol
|
||||
system prompt : messages[0], role "system"
|
||||
sampling sent : nothing (the vendor's own defaults apply)
|
||||
|
||||
ANTHROPIC POST /v1/messages
|
||||
top-level keys : max_tokens, messages, model, system
|
||||
model sent : claude-sonnet-5-5
|
||||
system prompt : top-level "system" (a string)
|
||||
sampling sent : max_tokens=1024
|
||||
|
||||
GEMINI POST /v1beta/models/gemini-3.8-flash:generateContent
|
||||
top-level keys : contents, systemInstruction, generationConfig
|
||||
model sent : gemini-3.8-flash
|
||||
system prompt : top-level "systemInstruction".parts[0]
|
||||
sampling sent : temperature=0.7, topP=1.0
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# Options: portable, provider-specific, and the per-call trap
|
||||
|
||||
A. ChatClient.options(ChatOptions.builder().temperature(0.2).maxTokens(200)) on each provider
|
||||
OPENAI model=gpt-6.1-sol, max_tokens=200, temperature=0.2
|
||||
ANTHROPIC max_tokens=200, model=claude-sonnet-5-5, temperature=0.2
|
||||
GEMINI temperature=0.2, topP=1.0, maxOutputTokens=200
|
||||
|
||||
B. Provider-specific options through the same ChatClient call
|
||||
OPENAI reasoningEffort + promptCacheKey -> prompt_cache_key=tickets-v1, reasoning_effort=low
|
||||
ANTHROPIC topK(5) -> top_k=5
|
||||
GEMINI thinkingBudget(0) -> generationConfig {"temperature":0.7,"topP":1.0,"thinkingConfig":{"thinkingBudget":0}}
|
||||
|
||||
C. The trap: one built ChatOptions object handed to Prompt, then model.call(prompt)
|
||||
OPENAI ClassCastException (DefaultChatOptions cannot be cast to OpenAiChatOptions)
|
||||
ANTHROPIC no error; model sent "claude-haiku-4-5", max_tokens 4096, temperature absent
|
||||
GEMINI ClassCastException (DefaultChatOptions cannot be cast to GoogleGenAiChatOptions)
|
||||
|
||||
D. The other trap: OpenAI-specific options sent to the other two models
|
||||
ANTHROPIC no error; model sent "claude-haiku-4-5"
|
||||
GEMINI ClassCastException (OpenAiChatOptions cannot be cast to GoogleGenAiChatOptions)
|
||||
@@ -0,0 +1,18 @@
|
||||
# Prompt caching: what each client sends and what Usage reports
|
||||
|
||||
A. Anthropic: the strategy decides whether the system prompt carries a cache breakpoint
|
||||
strategy NONE system is a plain string
|
||||
call 1: promptTokens=6014 cacheRead=0 cacheWrite=0
|
||||
call 2: promptTokens=6014 cacheRead=0 cacheWrite=0
|
||||
strategy SYSTEM_ONLY system is a list of 1 block(s), cache_control on block 0: true
|
||||
call 1: promptTokens=3 cacheRead=0 cacheWrite=6011
|
||||
call 2: promptTokens=3 cacheRead=6011 cacheWrite=0
|
||||
|
||||
B. OpenAI and Gemini: nothing to switch on, but the order of the text decides the hit
|
||||
OPENAI timestamp AFTER the policy: call 1 cacheRead=0 call 2 cacheRead=6016 of 6019 prompt tokens
|
||||
OPENAI timestamp BEFORE the policy: call 1 cacheRead=0 call 2 cacheRead=0 of 6019 prompt tokens
|
||||
GEMINI timestamp AFTER the policy: call 1 cacheRead=0 call 2 cacheRead=6016 of 6019 prompt tokens
|
||||
GEMINI timestamp BEFORE the policy: call 1 cacheRead=0 call 2 cacheRead=0 of 6019 prompt tokens
|
||||
|
||||
C. A prompt below the vendor's minimum: the client still asks, the vendor does not cache
|
||||
Anthropic SYSTEM_ONLY, 847-character policy: cache_control sent=true, cacheRead=0 cacheWrite=0
|
||||
@@ -0,0 +1,15 @@
|
||||
# Cost of 100 ticket summaries with a ~6,000-token policy
|
||||
|
||||
Workload per request: system policy 24044 characters, ticket about 12 characters, answer 62 characters.
|
||||
Token counts are the fake server's estimate (characters / 4). Prices are per million tokens as read on 2026-10-09.
|
||||
|
||||
provider, model, price sheet cache miss cache hit saving
|
||||
OpenAI gpt-6.1-sol $1.22 $0.09 92.7%
|
||||
Anthropic claude-sonnet-5-5 $1.22 $0.09 92.5%
|
||||
Anthropic, clock inside the block $1.22 $1.52 -24.7%
|
||||
Gemini gemini-3.8-flash (to 2026) $0.46 $0.06 87.9%
|
||||
Gemini gemini-3.8-flash (2027) $0.92 $0.11 87.9%
|
||||
|
||||
"cache miss": OpenAI and Gemini with the changing clock text at the START of the system prompt; Anthropic with caching off.
|
||||
"cache hit": the clock text moved into the user message, so the system prompt never changes; Anthropic with SYSTEM_ONLY.
|
||||
The Anthropic "clock inside the block" row keeps the clock at the END of the one cached system block: it changes every request, so every request pays the cache write.
|
||||
@@ -0,0 +1,22 @@
|
||||
# Fallback: OpenAI first, then Anthropic, then Gemini
|
||||
|
||||
A. OpenAI fails with each status (no client retries). Does the call move on?
|
||||
OpenAI 400 -> thrown requests: openai=1 anthropic=0 gemini=0 trail: [openai: failed with status 400]
|
||||
OpenAI 401 -> thrown requests: openai=1 anthropic=0 gemini=0 trail: [openai: failed with status 401]
|
||||
OpenAI 429 -> answered requests: openai=1 anthropic=1 gemini=0 trail: [openai: failed with status 429, anthropic: ok]
|
||||
OpenAI 500 -> answered requests: openai=1 anthropic=1 gemini=0 trail: [openai: failed with status 500, anthropic: ok]
|
||||
OpenAI 503 -> answered requests: openai=1 anthropic=1 gemini=0 trail: [openai: failed with status 503, anthropic: ok]
|
||||
|
||||
B. The same 503 with the client's own retries switched on (maxRetries 2)
|
||||
OpenAI 503 -> answered requests: openai=3 anthropic=1 gemini=0
|
||||
|
||||
C. Gemini retries in two layers: the Google SDK (HttpRetryOptions.attempts) and Spring AI's RetryTemplate
|
||||
SDK attempts=1, RetryTemplate retries=0 -> 1 request(s) for one call
|
||||
SDK attempts=1, RetryTemplate retries=2 -> 3 request(s) for one call
|
||||
SDK attempts=3, RetryTemplate retries=0 -> 3 request(s) for one call
|
||||
SDK attempts=3, RetryTemplate retries=2 -> 9 request(s) for one call
|
||||
nothing configured at all -> 5 requests for one call, and it took more than 5 seconds: true
|
||||
|
||||
D. Everything down: the last provider's failure is the one you see
|
||||
all three 503 -> thrown requests: openai=1 anthropic=1 gemini=1
|
||||
trail: [openai: failed with status 503, anthropic: failed with status 503, gemini: failed with status 503]
|
||||
@@ -0,0 +1,76 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project xmlns="http://maven.apache.org/POM/4.0.0"
|
||||
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
|
||||
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
|
||||
<modelVersion>4.0.0</modelVersion>
|
||||
|
||||
<parent>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-parent</artifactId>
|
||||
<version>4.1.1</version>
|
||||
<relativePath/>
|
||||
</parent>
|
||||
|
||||
<groupId>com.ankurm</groupId>
|
||||
<artifactId>providers</artifactId>
|
||||
<version>1.0.0</version>
|
||||
<name>providers</name>
|
||||
<description>One Spring AI app on OpenAI, Anthropic and Gemini: switching providers, provider options, prompt caching, cost and fallback.</description>
|
||||
|
||||
<properties>
|
||||
<java.version>25</java.version>
|
||||
<spring-ai.version>2.0.1</spring-ai.version>
|
||||
</properties>
|
||||
|
||||
<dependencyManagement>
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-bom</artifactId>
|
||||
<version>${spring-ai.version}</version>
|
||||
<type>pom</type>
|
||||
<scope>import</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
</dependencyManagement>
|
||||
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-client-chat</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-openai</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-anthropic</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-google-genai</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>tools.jackson.core</groupId>
|
||||
<artifactId>jackson-databind</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-test</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
|
||||
<build>
|
||||
<plugins>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-surefire-plugin</artifactId>
|
||||
<configuration>
|
||||
<argLine>-Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8</argLine>
|
||||
</configuration>
|
||||
</plugin>
|
||||
</plugins>
|
||||
</build>
|
||||
</project>
|
||||
Executable
+9
@@ -0,0 +1,9 @@
|
||||
#!/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 about a minute and a half
|
||||
# (one test waits out Gemini's default retry backoff on purpose).
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.."
|
||||
rm -rf target
|
||||
mvn -q -B test 2>&1 | grep -E "Tests run:|BUILD|FAIL|ERROR" || true
|
||||
ls output
|
||||
@@ -0,0 +1,44 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.RoundingMode;
|
||||
|
||||
import org.springframework.ai.chat.metadata.Usage;
|
||||
|
||||
/**
|
||||
* Turns one response's {@link Usage} into dollars. The reason this is a class and not a one-liner:
|
||||
* Spring AI's {@code getPromptTokens()} means different things per provider. OpenAI and Gemini
|
||||
* count cached tokens INSIDE it; Anthropic counts only the tokens after the cache breakpoint.
|
||||
*/
|
||||
public final class Cost {
|
||||
|
||||
/** Token counts split by how they are billed, and the price. */
|
||||
public record Billed(long uncachedInput, long cacheRead, long cacheWrite, long output, BigDecimal usd) {
|
||||
}
|
||||
|
||||
private static final BigDecimal MILLION = BigDecimal.valueOf(1_000_000);
|
||||
|
||||
private Cost() {
|
||||
}
|
||||
|
||||
public static Billed of(Provider provider, Usage u, Prices p) {
|
||||
long prompt = u.getPromptTokens();
|
||||
long read = orZero(u.getCacheReadInputTokens());
|
||||
long write = orZero(u.getCacheWriteInputTokens());
|
||||
long uncached = switch (provider) {
|
||||
case ANTHROPIC -> prompt;
|
||||
case OPENAI, GEMINI -> prompt - read;
|
||||
};
|
||||
long out = u.getCompletionTokens();
|
||||
BigDecimal usd = p.input().multiply(BigDecimal.valueOf(uncached))
|
||||
.add(p.cacheRead().multiply(BigDecimal.valueOf(read)))
|
||||
.add(p.cacheWrite().multiply(BigDecimal.valueOf(write)))
|
||||
.add(p.output().multiply(BigDecimal.valueOf(out)))
|
||||
.divide(MILLION, 8, RoundingMode.HALF_UP);
|
||||
return new Billed(uncached, read, write, out, usd);
|
||||
}
|
||||
|
||||
private static long orZero(Long v) {
|
||||
return v == null ? 0 : v;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import java.io.IOException;
|
||||
|
||||
import com.anthropic.errors.AnthropicIoException;
|
||||
import com.anthropic.errors.AnthropicServiceException;
|
||||
import com.google.genai.errors.ApiException;
|
||||
import com.openai.errors.OpenAIIoException;
|
||||
import com.openai.errors.OpenAIServiceException;
|
||||
|
||||
/**
|
||||
* Decides whether a failed model call is worth sending to another provider. The three vendors'
|
||||
* SDKs throw three unrelated exception families, and Gemini's arrives wrapped in a plain
|
||||
* RuntimeException, so the cause chain is walked.
|
||||
*/
|
||||
public final class Failures {
|
||||
|
||||
private Failures() {
|
||||
}
|
||||
|
||||
/** The HTTP status found anywhere in the cause chain, or -1. */
|
||||
public static int status(Throwable t) {
|
||||
for (Throwable c = t; c != null; c = c.getCause()) {
|
||||
if (c instanceof OpenAIServiceException e) {
|
||||
return e.statusCode();
|
||||
}
|
||||
if (c instanceof AnthropicServiceException e) {
|
||||
return e.statusCode();
|
||||
}
|
||||
if (c instanceof ApiException e) {
|
||||
return e.code();
|
||||
}
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
/** True for overload, rate limit, server errors and network failures; false for 4xx the caller caused. */
|
||||
public static boolean isTransient(Throwable t) {
|
||||
int status = status(t);
|
||||
if (status != -1) {
|
||||
return status == 408 || status == 429 || status >= 500;
|
||||
}
|
||||
for (Throwable c = t; c != null; c = c.getCause()) {
|
||||
if (c instanceof OpenAIIoException || c instanceof AnthropicIoException || c instanceof IOException) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.function.Predicate;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
|
||||
/**
|
||||
* Tries each model in order and moves on when the failure is one {@code shouldFallBack} accepts.
|
||||
* A failure it does not accept (a bad request, a wrong key) is rethrown at once: another provider
|
||||
* would only hide it.
|
||||
*/
|
||||
public class FallbackChatModel implements ChatModel {
|
||||
|
||||
public record Target(String name, ChatModel model) {
|
||||
}
|
||||
|
||||
private final List<Target> targets;
|
||||
|
||||
private final Predicate<Throwable> shouldFallBack;
|
||||
|
||||
private final List<String> trail = new ArrayList<>();
|
||||
|
||||
public FallbackChatModel(List<Target> targets, Predicate<Throwable> shouldFallBack) {
|
||||
this.targets = List.copyOf(targets);
|
||||
this.shouldFallBack = shouldFallBack;
|
||||
}
|
||||
|
||||
/** Which targets were tried by the most recent call, in order, with how each ended. */
|
||||
public synchronized List<String> trail() {
|
||||
return List.copyOf(trail);
|
||||
}
|
||||
|
||||
@Override
|
||||
public synchronized ChatResponse call(Prompt prompt) {
|
||||
trail.clear();
|
||||
RuntimeException last = null;
|
||||
for (Target t : targets) {
|
||||
try {
|
||||
ChatResponse r = t.model().call(prompt);
|
||||
trail.add(t.name() + ": ok");
|
||||
return r;
|
||||
}
|
||||
catch (RuntimeException e) {
|
||||
trail.add(t.name() + ": failed with status " + Failures.status(e));
|
||||
last = e;
|
||||
if (!shouldFallBack.test(e)) {
|
||||
throw e;
|
||||
}
|
||||
}
|
||||
}
|
||||
throw last;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
|
||||
/** Dollars per million tokens. The numbers live in the test that uses them, with the date they were read. */
|
||||
public record Prices(BigDecimal input, BigDecimal cacheRead, BigDecimal cacheWrite, BigDecimal output) {
|
||||
|
||||
public static Prices of(String input, String cacheRead, String cacheWrite, String output) {
|
||||
return new Prices(new BigDecimal(input), new BigDecimal(cacheRead), new BigDecimal(cacheWrite),
|
||||
new BigDecimal(output));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
/** The three vendors this module talks to, with the model id used for each in the examples. */
|
||||
public enum Provider {
|
||||
|
||||
OPENAI("gpt-6.1-sol"),
|
||||
ANTHROPIC("claude-sonnet-5-5"),
|
||||
GEMINI("gemini-3.8-flash");
|
||||
|
||||
private final String defaultModel;
|
||||
|
||||
Provider(String defaultModel) {
|
||||
this.defaultModel = defaultModel;
|
||||
}
|
||||
|
||||
public String defaultModel() {
|
||||
return defaultModel;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import com.google.genai.Client;
|
||||
import java.time.Duration;
|
||||
|
||||
import com.google.genai.types.HttpOptions;
|
||||
import com.google.genai.types.HttpRetryOptions;
|
||||
import org.springframework.ai.anthropic.AnthropicCacheOptions;
|
||||
import org.springframework.ai.anthropic.AnthropicCacheStrategy;
|
||||
import org.springframework.ai.anthropic.AnthropicChatModel;
|
||||
import org.springframework.ai.anthropic.AnthropicChatOptions;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.google.genai.GoogleGenAiChatModel;
|
||||
import org.springframework.ai.google.genai.GoogleGenAiChatOptions;
|
||||
import org.springframework.ai.openai.OpenAiChatModel;
|
||||
import org.springframework.ai.openai.OpenAiChatOptions;
|
||||
import org.springframework.core.retry.RetryPolicy;
|
||||
import org.springframework.core.retry.RetryTemplate;
|
||||
|
||||
/**
|
||||
* Builds the REAL Spring AI chat model for each provider. This is the only class that knows which
|
||||
* vendor is behind a {@link ChatModel}; everything else in the application sees the interface.
|
||||
*/
|
||||
public final class ProviderModels {
|
||||
|
||||
private ProviderModels() {
|
||||
}
|
||||
|
||||
/** Settings that differ per deployment, not per call. */
|
||||
public record Settings(String baseUrl, String apiKey, int maxRetries, AnthropicCacheStrategy anthropicCache) {
|
||||
|
||||
public static Settings of(String baseUrl) {
|
||||
return new Settings(baseUrl, "test-key", 0, AnthropicCacheStrategy.NONE);
|
||||
}
|
||||
|
||||
public Settings withRetries(int n) {
|
||||
return new Settings(baseUrl, apiKey, n, anthropicCache);
|
||||
}
|
||||
|
||||
public Settings withAnthropicCache(AnthropicCacheStrategy s) {
|
||||
return new Settings(baseUrl, apiKey, maxRetries, s);
|
||||
}
|
||||
}
|
||||
|
||||
public static ChatModel create(Provider p, Settings s) {
|
||||
return create(p, s, p.defaultModel());
|
||||
}
|
||||
|
||||
public static ChatModel create(Provider p, Settings s, String model) {
|
||||
return switch (p) {
|
||||
case OPENAI -> OpenAiChatModel.builder()
|
||||
.options(OpenAiChatOptions.builder().baseUrl(s.baseUrl() + "/v1").apiKey(s.apiKey())
|
||||
.model(model).maxRetries(s.maxRetries()).build())
|
||||
.build();
|
||||
case ANTHROPIC -> AnthropicChatModel.builder()
|
||||
.options(AnthropicChatOptions.builder().baseUrl(s.baseUrl()).apiKey(s.apiKey())
|
||||
.model(model).maxTokens(1024).maxRetries(s.maxRetries())
|
||||
.cacheOptions(AnthropicCacheOptions.builder().strategy(s.anthropicCache()).build())
|
||||
.build())
|
||||
.build();
|
||||
case GEMINI -> GoogleGenAiChatModel.builder()
|
||||
.genAiClient(Client.builder().apiKey(s.apiKey())
|
||||
.httpOptions(HttpOptions.builder().baseUrl(s.baseUrl())
|
||||
.retryOptions(HttpRetryOptions.builder().attempts(s.maxRetries() + 1).initialDelay(0.001).build()).build()).build())
|
||||
.options(GoogleGenAiChatOptions.builder().model(model).build())
|
||||
.retryTemplate(new RetryTemplate(RetryPolicy.builder().maxRetries(0)
|
||||
.delay(Duration.ofMillis(1)).build()))
|
||||
.build();
|
||||
};
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
/** What the application wants back, whichever vendor answers. */
|
||||
public record Summary(String title, String priority) {
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
|
||||
/** The application: summarise a support ticket. It depends on {@link ChatModel} only. */
|
||||
public class TicketService {
|
||||
|
||||
private final ChatClient client;
|
||||
|
||||
private final String policy;
|
||||
|
||||
public TicketService(ChatModel model, String policy) {
|
||||
this.client = ChatClient.create(model);
|
||||
this.policy = policy;
|
||||
}
|
||||
|
||||
public Summary summarise(String ticket) {
|
||||
return client.prompt().system(policy).user(ticket).call().entity(Summary.class);
|
||||
}
|
||||
|
||||
public String raw(String ticket) {
|
||||
return client.prompt().system(policy).user(ticket).call().content();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.providers.support.FakeProviderServer;
|
||||
import com.ankurm.providers.support.Fixtures;
|
||||
import com.ankurm.providers.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.anthropic.AnthropicCacheStrategy;
|
||||
import org.springframework.ai.chat.messages.SystemMessage;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.metadata.Usage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
|
||||
/**
|
||||
* Prompt caching. What is REAL here: the request each Spring AI client builds (does it mark a cache
|
||||
* breakpoint? where does the changing text sit?) and which Usage getters report cache tokens.
|
||||
* What is SIMULATED: whether the "vendor" hits its cache. The fake server applies each vendor's
|
||||
* documented rule (see FakeProviderServer) to the request it received.
|
||||
*/
|
||||
class CachingTest {
|
||||
|
||||
static final String POLICY = Fixtures.policy(24000);
|
||||
|
||||
@Test
|
||||
void caching() throws Exception {
|
||||
try (FakeProviderServer fake = new FakeProviderServer();
|
||||
Transcript t = new Transcript("03-caching.txt", "Prompt caching: what each client sends and what Usage reports")) {
|
||||
|
||||
t.line("A. Anthropic: the strategy decides whether the system prompt carries a cache breakpoint");
|
||||
for (AnthropicCacheStrategy s : new AnthropicCacheStrategy[] {AnthropicCacheStrategy.NONE, AnthropicCacheStrategy.SYSTEM_ONLY}) {
|
||||
fake.reset();
|
||||
ChatModel m = ProviderModels.create(Provider.ANTHROPIC, ProviderModels.Settings.of(fake.url()).withAnthropicCache(s));
|
||||
Usage u1 = call(m, POLICY, "ticket 1");
|
||||
Usage u2 = call(m, POLICY, "ticket 2");
|
||||
JsonNode sys = fake.seen().get(0).body().path("system");
|
||||
t.line(" strategy %-12s system is %s", s, sys.isString() ? "a plain string"
|
||||
: "a list of " + sys.size() + " block(s), cache_control on block 0: " + sys.get(0).has("cache_control"));
|
||||
t.line(" call 1: promptTokens=%d cacheRead=%d cacheWrite=%d", u1.getPromptTokens(), u1.getCacheReadInputTokens(), u1.getCacheWriteInputTokens());
|
||||
t.line(" call 2: promptTokens=%d cacheRead=%d cacheWrite=%d", u2.getPromptTokens(), u2.getCacheReadInputTokens(), u2.getCacheWriteInputTokens());
|
||||
if (s == AnthropicCacheStrategy.SYSTEM_ONLY) {
|
||||
assertThat(sys.get(0).has("cache_control")).isTrue();
|
||||
assertThat(u1.getCacheWriteInputTokens()).isGreaterThan(5000);
|
||||
assertThat(u2.getCacheReadInputTokens()).isEqualTo(u1.getCacheWriteInputTokens());
|
||||
assertThat(u2.getPromptTokens()).isLessThan(10);
|
||||
}
|
||||
else {
|
||||
assertThat(sys.isString()).isTrue();
|
||||
assertThat(u2.getCacheReadInputTokens()).isZero();
|
||||
}
|
||||
}
|
||||
|
||||
t.blank();
|
||||
t.line("B. OpenAI and Gemini: nothing to switch on, but the order of the text decides the hit");
|
||||
for (Provider p : new Provider[] {Provider.OPENAI, Provider.GEMINI}) {
|
||||
for (boolean volatileFirst : new boolean[] {false, true}) {
|
||||
fake.reset();
|
||||
ChatModel m = ProviderModels.create(p, ProviderModels.Settings.of(fake.url()));
|
||||
Usage u1 = call(m, system(volatileFirst, "10:00:01"), "ticket 1");
|
||||
Usage u2 = call(m, system(volatileFirst, "10:00:02"), "ticket 2");
|
||||
t.line(" %-7s %-34s call 1 cacheRead=%-5d call 2 cacheRead=%d of %d prompt tokens", p,
|
||||
volatileFirst ? "timestamp BEFORE the policy:" : "timestamp AFTER the policy:",
|
||||
u1.getCacheReadInputTokens(), u2.getCacheReadInputTokens(), u2.getPromptTokens());
|
||||
if (volatileFirst) {
|
||||
assertThat(u2.getCacheReadInputTokens()).isZero();
|
||||
}
|
||||
else {
|
||||
assertThat(u2.getCacheReadInputTokens()).isGreaterThan(5000);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
t.blank();
|
||||
t.line("C. A prompt below the vendor's minimum: the client still asks, the vendor does not cache");
|
||||
String shortPolicy = Fixtures.policy(800);
|
||||
fake.reset();
|
||||
ChatModel am = ProviderModels.create(Provider.ANTHROPIC,
|
||||
ProviderModels.Settings.of(fake.url()).withAnthropicCache(AnthropicCacheStrategy.SYSTEM_ONLY));
|
||||
call(am, shortPolicy, "ticket 1");
|
||||
Usage s2 = call(am, shortPolicy, "ticket 2");
|
||||
JsonNode sys = fake.seen().get(0).body().path("system");
|
||||
t.line(" Anthropic SYSTEM_ONLY, %d-character policy: cache_control sent=%s, cacheRead=%d cacheWrite=%d",
|
||||
shortPolicy.length(), sys.isArray() && sys.get(0).has("cache_control"), s2.getCacheReadInputTokens(), s2.getCacheWriteInputTokens());
|
||||
assertThat(sys.isArray() && sys.get(0).has("cache_control")).isTrue();
|
||||
assertThat(s2.getCacheReadInputTokens()).isZero();
|
||||
assertThat(s2.getCacheWriteInputTokens()).isZero();
|
||||
}
|
||||
}
|
||||
|
||||
static String system(boolean volatileFirst, String clock) {
|
||||
return volatileFirst ? "Current time: " + clock + "\n" + POLICY : POLICY + "Current time: " + clock + "\n";
|
||||
}
|
||||
|
||||
static Usage call(ChatModel m, String system, String user) {
|
||||
return m.call(new Prompt(List.of(new SystemMessage(system), new UserMessage(user)))).getMetadata().getUsage();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.math.RoundingMode;
|
||||
|
||||
import com.ankurm.providers.support.FakeProviderServer;
|
||||
import com.ankurm.providers.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.anthropic.AnthropicCacheStrategy;
|
||||
import org.springframework.ai.chat.metadata.Usage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
|
||||
/**
|
||||
* What 100 requests cost. The arithmetic and the Usage plumbing are real; the token counts are
|
||||
* the fake server's estimate (characters / 4) and the cache hits follow documented rules, so the
|
||||
* dollar figures are a worked example of the price sheets, not a measured bill.
|
||||
*/
|
||||
class CostTest {
|
||||
|
||||
// Prices in USD per million tokens, read from the vendors' pricing pages on 2026-10-09 (standard tier, short context).
|
||||
// OpenAI gpt-6.1-sol https://developers.openai.com/api/docs/pricing input 2.00, cached input 0.10, output 10.00
|
||||
// Anthropic claude-sonnet-5-5 https://platform.claude.com/docs/en/about-claude/pricing input 2, 5-minute cache write 2.50, read 0.10, output 10
|
||||
// Gemini gemini-3.8-flash https://ai.google.dev/gemini-api/docs/pricing input 0.75, cached 0.075, output 3.75 (through 2026-12-31)
|
||||
// from 2027-01-01: input 1.50, cached 0.15, output 7.50
|
||||
static final Prices OPENAI = Prices.of("2.00", "0.10", "0", "10.00");
|
||||
|
||||
static final Prices ANTHROPIC = Prices.of("2.00", "0.10", "2.50", "10.00");
|
||||
|
||||
static final Prices GEMINI_NOW = Prices.of("0.75", "0.075", "0", "3.75");
|
||||
|
||||
static final Prices GEMINI_2027 = Prices.of("1.50", "0.15", "0", "7.50");
|
||||
|
||||
static final int REQUESTS = 100;
|
||||
|
||||
@Test
|
||||
void hundredRequests() throws Exception {
|
||||
try (FakeProviderServer fake = new FakeProviderServer();
|
||||
Transcript t = new Transcript("04-cost.txt", "Cost of 100 ticket summaries with a ~6,000-token policy")) {
|
||||
t.line("Workload per request: system policy %d characters, ticket about 12 characters, answer %d characters.",
|
||||
CachingTest.POLICY.length(), "{\"title\":\"Login fails after password reset\",\"priority\":\"high\"}".length());
|
||||
t.line("Token counts are the fake server's estimate (characters / 4). Prices are per million tokens as read on 2026-10-09.");
|
||||
t.blank();
|
||||
t.line("%-34s %12s %12s %9s", "provider, model, price sheet", "cache miss", "cache hit", "saving");
|
||||
|
||||
BigDecimal oMiss = run(fake, Provider.OPENAI, Where.FIRST, AnthropicCacheStrategy.NONE, OPENAI);
|
||||
BigDecimal oHit = run(fake, Provider.OPENAI, Where.USER, AnthropicCacheStrategy.NONE, OPENAI);
|
||||
row(t, "OpenAI gpt-6.1-sol", oMiss, oHit);
|
||||
BigDecimal aMiss = run(fake, Provider.ANTHROPIC, Where.USER, AnthropicCacheStrategy.NONE, ANTHROPIC);
|
||||
BigDecimal aHit = run(fake, Provider.ANTHROPIC, Where.USER, AnthropicCacheStrategy.SYSTEM_ONLY, ANTHROPIC);
|
||||
row(t, "Anthropic claude-sonnet-5-5", aMiss, aHit);
|
||||
BigDecimal aTrap = run(fake, Provider.ANTHROPIC, Where.LAST, AnthropicCacheStrategy.SYSTEM_ONLY, ANTHROPIC);
|
||||
row(t, "Anthropic, clock inside the block", aMiss, aTrap);
|
||||
BigDecimal gMiss = run(fake, Provider.GEMINI, Where.FIRST, AnthropicCacheStrategy.NONE, GEMINI_NOW);
|
||||
BigDecimal gHit = run(fake, Provider.GEMINI, Where.USER, AnthropicCacheStrategy.NONE, GEMINI_NOW);
|
||||
row(t, "Gemini gemini-3.8-flash (to 2026)", gMiss, gHit);
|
||||
BigDecimal g27Miss = run(fake, Provider.GEMINI, Where.FIRST, AnthropicCacheStrategy.NONE, GEMINI_2027);
|
||||
BigDecimal g27Hit = run(fake, Provider.GEMINI, Where.USER, AnthropicCacheStrategy.NONE, GEMINI_2027);
|
||||
row(t, "Gemini gemini-3.8-flash (2027)", g27Miss, g27Hit);
|
||||
|
||||
t.blank();
|
||||
t.line("\"cache miss\": OpenAI and Gemini with the changing clock text at the START of the system prompt; Anthropic with caching off.");
|
||||
t.line("\"cache hit\": the clock text moved into the user message, so the system prompt never changes; Anthropic with SYSTEM_ONLY.");
|
||||
t.line("The Anthropic \"clock inside the block\" row keeps the clock at the END of the one cached system block: it changes every request, so every request pays the cache write.");
|
||||
|
||||
assertThat(oHit).isLessThan(oMiss);
|
||||
assertThat(aHit).isLessThan(aMiss);
|
||||
assertThat(gHit).isLessThan(gMiss);
|
||||
assertThat(aTrap).isGreaterThan(aMiss);
|
||||
// the doubling of the Gemini price sheet doubles the bill
|
||||
assertThat(g27Miss).isEqualByComparingTo(gMiss.multiply(BigDecimal.valueOf(2)));
|
||||
}
|
||||
}
|
||||
|
||||
private static void row(Transcript t, String name, BigDecimal miss, BigDecimal hit) {
|
||||
BigDecimal saving = BigDecimal.ONE.subtract(hit.divide(miss, 6, RoundingMode.HALF_UP)).multiply(BigDecimal.valueOf(100));
|
||||
t.line("%-34s %12s %12s %8s%%", name, "$" + miss.setScale(2, RoundingMode.HALF_UP), "$" + hit.setScale(2, RoundingMode.HALF_UP),
|
||||
saving.setScale(1, RoundingMode.HALF_UP));
|
||||
}
|
||||
|
||||
/** Where the text that changes on every request goes. */
|
||||
enum Where { FIRST, LAST, USER }
|
||||
|
||||
private static BigDecimal run(FakeProviderServer fake, Provider p, Where where, AnthropicCacheStrategy cache, Prices prices) {
|
||||
fake.reset();
|
||||
ChatModel m = ProviderModels.create(p, ProviderModels.Settings.of(fake.url()).withAnthropicCache(cache));
|
||||
BigDecimal total = BigDecimal.ZERO;
|
||||
for (int i = 0; i < REQUESTS; i++) {
|
||||
String clock = String.format("10:%02d:%02d.%d", i / 3600, i / 60 % 60, i);
|
||||
Usage u = switch (where) {
|
||||
case FIRST -> CachingTest.call(m, CachingTest.system(true, clock), "ticket " + i);
|
||||
case LAST -> CachingTest.call(m, CachingTest.system(false, clock), "ticket " + i);
|
||||
case USER -> CachingTest.call(m, CachingTest.POLICY, "ticket " + i + " at " + clock);
|
||||
};
|
||||
total = total.add(Cost.of(p, u, prices).usd());
|
||||
}
|
||||
return total;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.providers.support.FakeProviderServer;
|
||||
import com.ankurm.providers.support.Transcript;
|
||||
import com.google.genai.Client;
|
||||
import com.google.genai.types.HttpOptions;
|
||||
import com.google.genai.types.HttpRetryOptions;
|
||||
import java.time.Duration;
|
||||
import org.springframework.core.retry.RetryPolicy;
|
||||
import org.springframework.core.retry.RetryTemplate;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.google.genai.GoogleGenAiChatModel;
|
||||
import org.springframework.ai.google.genai.GoogleGenAiChatOptions;
|
||||
|
||||
/** Failover: which failures move to the next provider, and how many HTTP requests each plan really costs. */
|
||||
class FallbackTest {
|
||||
|
||||
@Test
|
||||
void fallback() throws Exception {
|
||||
try (FakeProviderServer fake = new FakeProviderServer();
|
||||
Transcript t = new Transcript("05-fallback.txt", "Fallback: OpenAI first, then Anthropic, then Gemini")) {
|
||||
|
||||
t.line("A. OpenAI fails with each status (no client retries). Does the call move on?");
|
||||
for (int status : new int[] {400, 401, 429, 500, 503}) {
|
||||
fake.reset();
|
||||
fake.failNext(Provider.OPENAI, status, 100);
|
||||
FallbackChatModel f = chain(fake, 0);
|
||||
String result = attempt(f);
|
||||
t.line(" OpenAI %d -> %-8s requests: %s trail: %s", status, result, counts(fake), f.trail());
|
||||
boolean transientStatus = status == 429 || status >= 500;
|
||||
assertThat(result).isEqualTo(transientStatus ? "answered" : "thrown");
|
||||
}
|
||||
|
||||
t.blank();
|
||||
t.line("B. The same 503 with the client's own retries switched on (maxRetries 2)");
|
||||
fake.reset();
|
||||
fake.failNext(Provider.OPENAI, 503, 100);
|
||||
FallbackChatModel f2 = chain(fake, 2);
|
||||
t.line(" OpenAI 503 -> %-8s requests: %s", attempt(f2), counts(fake));
|
||||
assertThat(fake.seen(Provider.OPENAI)).hasSize(3);
|
||||
|
||||
t.blank();
|
||||
t.line("C. Gemini retries in two layers: the Google SDK (HttpRetryOptions.attempts) and Spring AI's RetryTemplate");
|
||||
for (int attempts : new int[] {1, 3}) {
|
||||
for (int templateRetries : new int[] {0, 2}) {
|
||||
fake.reset();
|
||||
fake.failNext(Provider.GEMINI, 503, 100);
|
||||
attempt(gemini(fake, attempts, templateRetries));
|
||||
t.line(" SDK attempts=%d, RetryTemplate retries=%d -> %d request(s) for one call", attempts, templateRetries,
|
||||
fake.seen(Provider.GEMINI).size());
|
||||
assertThat(fake.seen(Provider.GEMINI)).hasSize(attempts * (1 + templateRetries));
|
||||
}
|
||||
}
|
||||
fake.reset();
|
||||
fake.failNext(Provider.GEMINI, 503, 100);
|
||||
long t0 = System.nanoTime();
|
||||
attempt(GoogleGenAiChatModel.builder()
|
||||
.genAiClient(Client.builder().apiKey("k").httpOptions(HttpOptions.builder().baseUrl(fake.url()).build()).build())
|
||||
.options(GoogleGenAiChatOptions.builder().model("gemini-3.8-flash").build()).build());
|
||||
long seconds = (System.nanoTime() - t0) / 1_000_000_000L;
|
||||
t.line(" nothing configured at all -> %d requests for one call, and it took more than 5 seconds: %s",
|
||||
fake.seen(Provider.GEMINI).size(), seconds > 5);
|
||||
assertThat(seconds).isGreaterThan(5);
|
||||
|
||||
t.blank();
|
||||
t.line("D. Everything down: the last provider's failure is the one you see");
|
||||
fake.reset();
|
||||
for (Provider p : Provider.values()) {
|
||||
fake.failNext(p, 503, 100);
|
||||
}
|
||||
FallbackChatModel f3 = chain(fake, 0);
|
||||
t.line(" all three 503 -> %s requests: %s", attempt(f3), counts(fake));
|
||||
t.line(" trail: %s", f3.trail());
|
||||
assertThat(f3.trail()).hasSize(3);
|
||||
}
|
||||
}
|
||||
|
||||
private static ChatModel gemini(FakeProviderServer fake, int attempts, int templateRetries) {
|
||||
return GoogleGenAiChatModel.builder()
|
||||
.genAiClient(Client.builder().apiKey("k").httpOptions(HttpOptions.builder().baseUrl(fake.url())
|
||||
.retryOptions(HttpRetryOptions.builder().attempts(attempts).initialDelay(0.001).build()).build()).build())
|
||||
.options(GoogleGenAiChatOptions.builder().model("gemini-3.8-flash").build())
|
||||
.retryTemplate(new RetryTemplate(RetryPolicy.builder().maxRetries(templateRetries).delay(Duration.ofMillis(1)).build()))
|
||||
.build();
|
||||
}
|
||||
|
||||
private static String attempt(ChatModel m) {
|
||||
try {
|
||||
m.call(new Prompt("x"));
|
||||
return "answered";
|
||||
}
|
||||
catch (RuntimeException e) {
|
||||
return "thrown";
|
||||
}
|
||||
}
|
||||
|
||||
private static FallbackChatModel chain(FakeProviderServer fake, int retries) {
|
||||
var s = ProviderModels.Settings.of(fake.url()).withRetries(retries);
|
||||
return new FallbackChatModel(List.of(
|
||||
new FallbackChatModel.Target("openai", ProviderModels.create(Provider.OPENAI, s)),
|
||||
new FallbackChatModel.Target("anthropic", ProviderModels.create(Provider.ANTHROPIC, s)),
|
||||
new FallbackChatModel.Target("gemini", ProviderModels.create(Provider.GEMINI, s))), Failures::isTransient);
|
||||
}
|
||||
|
||||
private static String counts(FakeProviderServer fake) {
|
||||
return "openai=" + fake.seen(Provider.OPENAI).size() + " anthropic=" + fake.seen(Provider.ANTHROPIC).size()
|
||||
+ " gemini=" + fake.seen(Provider.GEMINI).size();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.providers.support.FakeProviderServer;
|
||||
import com.ankurm.providers.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.google.genai.GoogleGenAiChatOptions;
|
||||
import org.springframework.ai.openai.OpenAiChatOptions;
|
||||
import org.springframework.ai.anthropic.AnthropicChatOptions;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
|
||||
/** Portable options, provider-specific options, and the two ways to get them wrong. */
|
||||
class OptionsTest {
|
||||
|
||||
@Test
|
||||
void options() throws Exception {
|
||||
try (FakeProviderServer fake = new FakeProviderServer();
|
||||
Transcript t = new Transcript("02-options.txt", "Options: portable, provider-specific, and the per-call trap")) {
|
||||
t.line("A. ChatClient.options(ChatOptions.builder().temperature(0.2).maxTokens(200)) on each provider");
|
||||
for (Provider p : Provider.values()) {
|
||||
fake.reset();
|
||||
ChatClient.create(model(p, fake)).prompt().user("x")
|
||||
.options(ChatOptions.builder().temperature(0.2).maxTokens(200)).call().content();
|
||||
JsonNode b = fake.seen().get(0).body();
|
||||
JsonNode src = p == Provider.GEMINI ? b.path("generationConfig") : b;
|
||||
List<String> fields = new ArrayList<>();
|
||||
for (String f : src.propertyNames()) {
|
||||
if (!f.equals("messages") && !f.equals("contents")) {
|
||||
fields.add(f + "=" + src.path(f).asString());
|
||||
}
|
||||
}
|
||||
t.line(" %-9s %s", p, String.join(", ", fields));
|
||||
assertThat(src.path("temperature").asDouble()).isEqualTo(0.2);
|
||||
}
|
||||
assertThat(fake.seen().get(0).model()).isNotEmpty();
|
||||
|
||||
t.blank();
|
||||
t.line("B. Provider-specific options through the same ChatClient call");
|
||||
fake.reset();
|
||||
ChatClient.create(model(Provider.OPENAI, fake)).prompt().user("x")
|
||||
.options(OpenAiChatOptions.builder().reasoningEffort("low").promptCacheKey("tickets-v1")).call().content();
|
||||
t.line(" %-9s reasoningEffort + promptCacheKey -> %s", Provider.OPENAI, keysExcept(fake.seen().get(0).body(), "messages", "model"));
|
||||
fake.reset();
|
||||
ChatClient.create(model(Provider.ANTHROPIC, fake)).prompt().user("x")
|
||||
.options(AnthropicChatOptions.builder().topK(5)).call().content();
|
||||
t.line(" %-9s topK(5) -> %s", Provider.ANTHROPIC, keysExcept(fake.seen().get(0).body(), "messages", "model", "max_tokens"));
|
||||
fake.reset();
|
||||
ChatClient.create(model(Provider.GEMINI, fake)).prompt().user("x")
|
||||
.options(GoogleGenAiChatOptions.builder().thinkingBudget(0)).call().content();
|
||||
t.line(" %-9s thinkingBudget(0) -> generationConfig %s", Provider.GEMINI,
|
||||
fake.seen().get(0).body().path("generationConfig").toString());
|
||||
|
||||
t.blank();
|
||||
t.line("C. The trap: one built ChatOptions object handed to Prompt, then model.call(prompt)");
|
||||
for (Provider p : Provider.values()) {
|
||||
fake.reset();
|
||||
String outcome;
|
||||
try {
|
||||
model(p, fake).call(new Prompt("x", ChatOptions.builder().temperature(0.2).maxTokens(200).build()));
|
||||
JsonNode b = fake.seen().get(0).body();
|
||||
outcome = "no error; model sent \"" + b.path("model").asString() + "\", max_tokens "
|
||||
+ b.path("max_tokens").asString() + ", temperature "
|
||||
+ (b.has("temperature") ? b.path("temperature").asString() : "absent");
|
||||
assertThat(p).isEqualTo(Provider.ANTHROPIC);
|
||||
assertThat(b.path("model").asString()).isNotEqualTo(p.defaultModel());
|
||||
}
|
||||
catch (ClassCastException e) {
|
||||
outcome = "ClassCastException (" + simple(e.getMessage()) + ")";
|
||||
}
|
||||
t.line(" %-9s %s", p, outcome);
|
||||
}
|
||||
|
||||
t.blank();
|
||||
t.line("D. The other trap: OpenAI-specific options sent to the other two models");
|
||||
for (Provider p : new Provider[] {Provider.ANTHROPIC, Provider.GEMINI}) {
|
||||
fake.reset();
|
||||
String outcome;
|
||||
try {
|
||||
model(p, fake).call(new Prompt("x", OpenAiChatOptions.builder().reasoningEffort("low").build()));
|
||||
JsonNode b = fake.seen().get(0).body();
|
||||
outcome = "no error; model sent \"" + b.path("model").asString() + "\"";
|
||||
}
|
||||
catch (ClassCastException e) {
|
||||
outcome = "ClassCastException (" + simple(e.getMessage()) + ")";
|
||||
}
|
||||
t.line(" %-9s %s", p, outcome);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private static ChatModel model(Provider p, FakeProviderServer fake) {
|
||||
return ProviderModels.create(p, ProviderModels.Settings.of(fake.url()));
|
||||
}
|
||||
|
||||
private static String keysExcept(JsonNode b, String... skip) {
|
||||
List<String> out = new ArrayList<>();
|
||||
outer:
|
||||
for (String k : b.propertyNames()) {
|
||||
for (String s : skip) {
|
||||
if (k.equals(s)) {
|
||||
continue outer;
|
||||
}
|
||||
}
|
||||
out.add(k + "=" + b.path(k).asString());
|
||||
}
|
||||
return String.join(", ", out);
|
||||
}
|
||||
|
||||
private static String simple(String message) {
|
||||
// "class a.b.C cannot be cast to class x.y.D (...)" -> "C cannot be cast to D"
|
||||
String[] w = message.split(" ");
|
||||
return w[1].substring(w[1].lastIndexOf('.') + 1) + " cannot be cast to " + w[7].substring(w[7].lastIndexOf('.') + 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package com.ankurm.providers;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.providers.support.FakeProviderServer;
|
||||
import com.ankurm.providers.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
|
||||
/** Same application call, three vendors: what does each real Spring AI model put on the wire? */
|
||||
class WireTest {
|
||||
|
||||
@Test
|
||||
void sameCallThreeWireFormats() throws Exception {
|
||||
try (FakeProviderServer fake = new FakeProviderServer();
|
||||
Transcript t = new Transcript("01-wire.txt", "One TicketService.raw() call, three providers, as the HTTP server saw it")) {
|
||||
for (Provider p : Provider.values()) {
|
||||
new TicketService(ProviderModels.create(p, ProviderModels.Settings.of(fake.url())),
|
||||
"You triage tickets.").raw("Cannot log in after reset");
|
||||
}
|
||||
for (Provider p : Provider.values()) {
|
||||
FakeProviderServer.Seen s = fake.seen(p).get(0);
|
||||
JsonNode b = s.body();
|
||||
List<String> keys = new ArrayList<>(b.propertyNames());
|
||||
t.line("%-9s POST %s", p, s.path());
|
||||
t.line(" top-level keys : %s", String.join(", ", keys));
|
||||
t.line(" model sent : %s", s.model());
|
||||
t.line(" system prompt : %s", where(p, b));
|
||||
t.line(" sampling sent : %s", sampling(p, b));
|
||||
t.blank();
|
||||
}
|
||||
assertThat(fake.seen(Provider.OPENAI).get(0).body().path("messages").get(0).path("role").asString())
|
||||
.isEqualTo("system");
|
||||
assertThat(fake.seen(Provider.ANTHROPIC).get(0).body().path("system").asString()).isEqualTo("You triage tickets.");
|
||||
assertThat(fake.seen(Provider.GEMINI).get(0).body().path("systemInstruction").path("parts").get(0)
|
||||
.path("text").asString()).isEqualTo("You triage tickets.");
|
||||
assertThat(fake.seen(Provider.GEMINI).get(0).body().path("generationConfig").path("temperature").asDouble())
|
||||
.isEqualTo(0.7);
|
||||
assertThat(fake.seen(Provider.OPENAI).get(0).body().has("temperature")).isFalse();
|
||||
assertThat(fake.seen(Provider.ANTHROPIC).get(0).body().has("temperature")).isFalse();
|
||||
}
|
||||
}
|
||||
|
||||
private static String where(Provider p, JsonNode b) {
|
||||
return switch (p) {
|
||||
case OPENAI -> "messages[0], role \"" + b.path("messages").get(0).path("role").asString() + "\"";
|
||||
case ANTHROPIC -> "top-level \"system\" (" + (b.path("system").isString() ? "a string" : "a list of blocks") + ")";
|
||||
case GEMINI -> "top-level \"systemInstruction\".parts[0]";
|
||||
};
|
||||
}
|
||||
|
||||
private static String sampling(Provider p, JsonNode b) {
|
||||
List<String> parts = new ArrayList<>();
|
||||
JsonNode src = p == Provider.GEMINI ? b.path("generationConfig") : b;
|
||||
for (String f : new String[] {"temperature", "top_p", "topP", "max_tokens", "max_completion_tokens", "maxOutputTokens"}) {
|
||||
if (src.has(f)) {
|
||||
parts.add(f + "=" + src.path(f).asString());
|
||||
}
|
||||
}
|
||||
return parts.isEmpty() ? "nothing (the vendor's own defaults apply)" : String.join(", ", parts);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,287 @@
|
||||
package com.ankurm.providers.support;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.OutputStream;
|
||||
import java.net.InetSocketAddress;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import com.ankurm.providers.Provider;
|
||||
import com.sun.net.httpserver.HttpExchange;
|
||||
import com.sun.net.httpserver.HttpServer;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
import tools.jackson.databind.json.JsonMapper;
|
||||
import tools.jackson.databind.node.ObjectNode;
|
||||
|
||||
/**
|
||||
* One local HTTP server that answers like OpenAI chat completions, Anthropic messages and Gemini
|
||||
* generateContent. The REAL Spring AI model classes talk to it, so the requests it records are the
|
||||
* requests those classes would send to the vendors.
|
||||
*
|
||||
* <p>It is not a language model and not a vendor. Three things it does are SIMULATION, not
|
||||
* measurement, and every number derived from them says so in the article:
|
||||
* <ul>
|
||||
* <li>token counts are {@code ceil(characters / 4)};</li>
|
||||
* <li>the answer is a fixed JSON string;</li>
|
||||
* <li>cache hits follow the rules the vendors DOCUMENT (see {@link #cacheRules}), applied to the
|
||||
* request the client really sent. Whether a real vendor hits its cache is not tested here.</li>
|
||||
* </ul>
|
||||
*/
|
||||
public class FakeProviderServer implements AutoCloseable {
|
||||
|
||||
/** What one request carried, reduced to the things the article compares. */
|
||||
public record Seen(Provider provider, String path, JsonNode body, String system, String user, String model) {
|
||||
}
|
||||
|
||||
/** Vendor cache thresholds, in tokens, as documented on 2026-10-09 for the models named in {@link Provider}. */
|
||||
public record CacheRules(int openAiMinTokens, int anthropicMinTokens, int geminiMinTokens) {
|
||||
public static CacheRules documented() {
|
||||
return new CacheRules(1024, 512, 4096);
|
||||
}
|
||||
}
|
||||
|
||||
private static final JsonMapper JSON = JsonMapper.builder().build();
|
||||
|
||||
private final HttpServer server;
|
||||
|
||||
private final List<Seen> seen = new CopyOnWriteArrayList<>();
|
||||
|
||||
private final Map<Provider, AtomicInteger> failuresLeft = new ConcurrentHashMap<>();
|
||||
|
||||
private final Map<Provider, Integer> failureStatus = new ConcurrentHashMap<>();
|
||||
|
||||
private final Map<Provider, List<String>> previousPrompts = new ConcurrentHashMap<>();
|
||||
|
||||
private final Map<Provider, Set<String>> anthropicWritten = new ConcurrentHashMap<>();
|
||||
|
||||
private volatile CacheRules cacheRules = CacheRules.documented();
|
||||
|
||||
private volatile String answer = "{\"title\":\"Login fails after password reset\",\"priority\":\"high\"}";
|
||||
|
||||
public FakeProviderServer() throws IOException {
|
||||
server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
|
||||
server.createContext("/", this::handle);
|
||||
server.start();
|
||||
}
|
||||
|
||||
public String url() {
|
||||
return "http://127.0.0.1:" + server.getAddress().getPort();
|
||||
}
|
||||
|
||||
public List<Seen> seen() {
|
||||
return seen;
|
||||
}
|
||||
|
||||
public List<Seen> seen(Provider p) {
|
||||
return seen.stream().filter(s -> s.provider() == p).toList();
|
||||
}
|
||||
|
||||
public FakeProviderServer answer(String json) {
|
||||
this.answer = json;
|
||||
return this;
|
||||
}
|
||||
|
||||
public FakeProviderServer cacheRules(CacheRules r) {
|
||||
this.cacheRules = r;
|
||||
return this;
|
||||
}
|
||||
|
||||
/** The next {@code count} requests to this provider fail with {@code status}. */
|
||||
public FakeProviderServer failNext(Provider p, int status, int count) {
|
||||
failureStatus.put(p, status);
|
||||
failuresLeft.put(p, new AtomicInteger(count));
|
||||
return this;
|
||||
}
|
||||
|
||||
public void reset() {
|
||||
seen.clear();
|
||||
failuresLeft.clear();
|
||||
previousPrompts.clear();
|
||||
anthropicWritten.clear();
|
||||
}
|
||||
|
||||
/** Estimated tokens: one per four characters, rounded up. Documented as a simulation. */
|
||||
public static int tokens(String s) {
|
||||
return (s.length() + 3) / 4;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
server.stop(0);
|
||||
}
|
||||
|
||||
private void handle(HttpExchange ex) throws IOException {
|
||||
String path = ex.getRequestURI().getPath();
|
||||
Provider p = path.contains("/chat/completions") ? Provider.OPENAI
|
||||
: path.contains("/messages") ? Provider.ANTHROPIC
|
||||
: path.contains(":generateContent") ? Provider.GEMINI : null;
|
||||
String raw = new String(ex.getRequestBody().readAllBytes(), StandardCharsets.UTF_8);
|
||||
if (p == null) {
|
||||
reply(ex, 404, "{\"error\":\"unknown path " + path + "\"}");
|
||||
return;
|
||||
}
|
||||
JsonNode body = JSON.readTree(raw);
|
||||
Seen s = extract(p, path, body);
|
||||
seen.add(s);
|
||||
AtomicInteger left = failuresLeft.get(p);
|
||||
if (left != null && left.getAndDecrement() > 0) {
|
||||
reply(ex, failureStatus.get(p), errorBody(p, failureStatus.get(p)));
|
||||
return;
|
||||
}
|
||||
reply(ex, 200, success(p, s));
|
||||
}
|
||||
|
||||
private static Seen extract(Provider p, String path, JsonNode b) {
|
||||
StringBuilder system = new StringBuilder();
|
||||
StringBuilder user = new StringBuilder();
|
||||
String model = b.path("model").asString("");
|
||||
switch (p) {
|
||||
case OPENAI -> {
|
||||
for (JsonNode m : b.path("messages")) {
|
||||
StringBuilder into = m.path("role").asString().equals("user") ? user : system;
|
||||
text(m.path("content"), into);
|
||||
}
|
||||
}
|
||||
case ANTHROPIC -> {
|
||||
text(b.path("system"), system);
|
||||
for (JsonNode m : b.path("messages")) {
|
||||
text(m.path("content"), user);
|
||||
}
|
||||
}
|
||||
case GEMINI -> {
|
||||
for (JsonNode part : b.path("systemInstruction").path("parts")) {
|
||||
system.append(part.path("text").asString(""));
|
||||
}
|
||||
for (JsonNode c : b.path("contents")) {
|
||||
for (JsonNode part : c.path("parts")) {
|
||||
user.append(part.path("text").asString(""));
|
||||
}
|
||||
}
|
||||
String[] seg = path.split("/models/");
|
||||
model = seg.length > 1 ? seg[1].replace(":generateContent", "") : "";
|
||||
}
|
||||
}
|
||||
return new Seen(p, path, b, system.toString(), user.toString(), model);
|
||||
}
|
||||
|
||||
private static void text(JsonNode content, StringBuilder into) {
|
||||
if (content.isString()) {
|
||||
into.append(content.asString());
|
||||
} else {
|
||||
for (JsonNode part : content) {
|
||||
into.append(part.path("text").asString(""));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private String errorBody(Provider p, int status) {
|
||||
return switch (p) {
|
||||
case OPENAI -> "{\"error\":{\"message\":\"simulated " + status + "\",\"type\":\"server_error\",\"code\":null}}";
|
||||
case ANTHROPIC -> "{\"type\":\"error\",\"error\":{\"type\":\"overloaded_error\",\"message\":\"simulated " + status + "\"}}";
|
||||
case GEMINI -> "{\"error\":{\"code\":" + status + ",\"message\":\"simulated " + status + "\",\"status\":\"UNAVAILABLE\"}}";
|
||||
};
|
||||
}
|
||||
|
||||
private String success(Provider p, Seen s) {
|
||||
String prompt = s.system() + "\n" + s.user();
|
||||
int in = tokens(prompt);
|
||||
int out = tokens(answer);
|
||||
ObjectNode r = JSON.createObjectNode();
|
||||
switch (p) {
|
||||
case OPENAI -> {
|
||||
int cached = prefixCached(p, prompt, cacheRules.openAiMinTokens());
|
||||
r.put("id", "chatcmpl-fake").put("object", "chat.completion").put("created", 1700000000)
|
||||
.put("model", s.model());
|
||||
ObjectNode msg = r.putArray("choices").addObject().put("index", 0).put("finish_reason", "stop")
|
||||
.putObject("message").put("role", "assistant").put("content", answer);
|
||||
msg.putNull("refusal");
|
||||
ObjectNode u = r.putObject("usage").put("prompt_tokens", in).put("completion_tokens", out)
|
||||
.put("total_tokens", in + out);
|
||||
u.putObject("prompt_tokens_details").put("cached_tokens", cached);
|
||||
}
|
||||
case ANTHROPIC -> {
|
||||
int created = 0;
|
||||
int read = 0;
|
||||
String prefix = anthropicCachedPrefix(s.body());
|
||||
if (prefix != null && tokens(prefix) >= cacheRules.anthropicMinTokens()) {
|
||||
Set<String> written = anthropicWritten.computeIfAbsent(p, k -> new HashSet<>());
|
||||
if (written.contains(prefix)) {
|
||||
read = tokens(prefix);
|
||||
} else {
|
||||
written.add(prefix);
|
||||
created = tokens(prefix);
|
||||
}
|
||||
}
|
||||
r.put("id", "msg_fake").put("type", "message").put("role", "assistant").put("model", s.model())
|
||||
.put("stop_reason", "end_turn").putNull("stop_sequence");
|
||||
r.putArray("content").addObject().put("type", "text").put("text", answer);
|
||||
r.putObject("usage").put("input_tokens", in - created - read).put("output_tokens", out)
|
||||
.put("cache_creation_input_tokens", created).put("cache_read_input_tokens", read);
|
||||
}
|
||||
case GEMINI -> {
|
||||
int cached = prefixCached(p, prompt, cacheRules.geminiMinTokens());
|
||||
ObjectNode cand = r.putArray("candidates").addObject().put("finishReason", "STOP").put("index", 0);
|
||||
cand.putObject("content").put("role", "model").putArray("parts").addObject().put("text", answer);
|
||||
r.putObject("usageMetadata").put("promptTokenCount", in).put("candidatesTokenCount", out)
|
||||
.put("totalTokenCount", in + out).put("cachedContentTokenCount", cached);
|
||||
r.put("modelVersion", s.model());
|
||||
}
|
||||
}
|
||||
return JSON.writeValueAsString(r);
|
||||
}
|
||||
|
||||
/** OpenAI and Gemini: documented automatic prefix caching. Cached = longest shared prefix with one of the last four prompts. */
|
||||
private int prefixCached(Provider p, String prompt, int minTokens) {
|
||||
List<String> prev = previousPrompts.computeIfAbsent(p, k -> new ArrayList<>());
|
||||
int best = 0;
|
||||
synchronized (prev) {
|
||||
for (String old : prev) {
|
||||
int n = 0;
|
||||
int max = Math.min(old.length(), prompt.length());
|
||||
while (n < max && old.charAt(n) == prompt.charAt(n)) {
|
||||
n++;
|
||||
}
|
||||
best = Math.max(best, n);
|
||||
}
|
||||
prev.add(prompt);
|
||||
if (prev.size() > 4) {
|
||||
prev.remove(0);
|
||||
}
|
||||
}
|
||||
int t = best / 4;
|
||||
return t >= minTokens ? t : 0;
|
||||
}
|
||||
|
||||
/** Anthropic: text of tools-then-system blocks up to and including the last block carrying cache_control, else null. */
|
||||
private static String anthropicCachedPrefix(JsonNode body) {
|
||||
JsonNode sys = body.path("system");
|
||||
if (!sys.isArray()) {
|
||||
return null;
|
||||
}
|
||||
StringBuilder prefix = new StringBuilder();
|
||||
String marked = null;
|
||||
for (JsonNode block : sys) {
|
||||
prefix.append(block.path("text").asString(""));
|
||||
if (block.has("cache_control")) {
|
||||
marked = prefix.toString();
|
||||
}
|
||||
}
|
||||
return marked;
|
||||
}
|
||||
|
||||
private void reply(HttpExchange ex, int status, String json) throws IOException {
|
||||
byte[] out = json.getBytes(StandardCharsets.UTF_8);
|
||||
ex.getResponseHeaders().add("Content-Type", "application/json");
|
||||
ex.sendResponseHeaders(status, out.length);
|
||||
try (OutputStream os = ex.getResponseBody()) {
|
||||
os.write(out);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package com.ankurm.providers.support;
|
||||
|
||||
/** Deterministic text of a chosen size, so the prompt length (and so the cost) is known exactly. */
|
||||
public final class Fixtures {
|
||||
|
||||
private Fixtures() {
|
||||
}
|
||||
|
||||
private static final String[] RULES = {
|
||||
"Tickets that mention data loss, a security problem or a payment failure are priority high.",
|
||||
"Tickets about slow pages, wrong totals or broken exports are priority medium.",
|
||||
"Questions about how to use a feature, cosmetic problems and feature requests are priority low.",
|
||||
"Keep the title under eight words and write it in the present tense without the customer's name.",
|
||||
"If the ticket contains several problems, summarise the most severe one and ignore the rest.",
|
||||
"Never repeat an email address, phone number or order number from the ticket in the title.",
|
||||
};
|
||||
|
||||
/** A support policy of at least {@code chars} characters, built from numbered rules. */
|
||||
public static String policy(int chars) {
|
||||
StringBuilder sb = new StringBuilder("You triage support tickets for an online shop. Follow every rule below.\n");
|
||||
int n = 1;
|
||||
while (sb.length() < chars) {
|
||||
sb.append("Rule ").append(n).append(": ").append(RULES[(n - 1) % RULES.length]).append('\n');
|
||||
n++;
|
||||
}
|
||||
return sb.toString();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package com.ankurm.providers.support;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.PrintWriter;
|
||||
import java.io.StringWriter;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
|
||||
/**
|
||||
* Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
|
||||
* echoes it to the console. Every console block quoted in the article comes out of one of these
|
||||
* files verbatim.
|
||||
*/
|
||||
public final class Transcript implements AutoCloseable {
|
||||
|
||||
private final Path path;
|
||||
private final StringWriter buffer = new StringWriter();
|
||||
private final PrintWriter out = new PrintWriter(buffer);
|
||||
|
||||
public Transcript(String fileName, String title) {
|
||||
this.path = Path.of("output", fileName);
|
||||
out.println("# " + title);
|
||||
out.println();
|
||||
}
|
||||
|
||||
public Transcript line(String format, Object... args) {
|
||||
out.println(args.length == 0 ? format : String.format(format, args));
|
||||
return this;
|
||||
}
|
||||
|
||||
public Transcript blank() {
|
||||
out.println();
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
out.flush();
|
||||
try {
|
||||
Files.createDirectories(path.getParent());
|
||||
Files.writeString(path, buffer.toString());
|
||||
} catch (IOException e) {
|
||||
throw new IllegalStateException("could not write " + path, e);
|
||||
}
|
||||
System.out.print(buffer);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user