Add observability module: Spring AI built-in meters and spans, cost per endpoint from token usage, Prometheus and Grafana dashboard

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
Claude
2026-10-09 06:39:21 +00:00
parent 65f580c5bc
commit cc2a1b8bcb
45 changed files with 1731 additions and 0 deletions
@@ -0,0 +1,47 @@
package com.ankurm.observability;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import reactor.core.publisher.Flux;
/**
* Four endpoints, each tagging its ChatClient call with an {@code endpoint} value in the request
* context. {@link EndpointTagConvention} turns that into a metric tag, which is what lets the
* dashboard show cost per endpoint.
*/
@RestController
public class AiController {
private final ChatClient chat;
private final WeatherTools tools;
public AiController(ChatClient.Builder builder, WeatherTools tools) {
this.chat = builder.build();
this.tools = tools;
}
@GetMapping("/ask")
public String ask(@RequestParam String q) {
return chat.prompt().user(q).advisors(a -> a.param(EndpointTagConvention.KEY, "ask")).call().content();
}
@GetMapping("/summarize")
public String summarize(@RequestParam String text) {
return chat.prompt().system("Summarize the user's text in one sentence.").user(text)
.advisors(a -> a.param(EndpointTagConvention.KEY, "summarize")).call().content();
}
@GetMapping("/weather")
public String weather(@RequestParam String city) {
return chat.prompt().user("What is the weather in " + city + "?").tools(tools)
.advisors(a -> a.param(EndpointTagConvention.KEY, "weather")).call().content();
}
@GetMapping("/stream")
public Flux<String> stream(@RequestParam String q) {
return chat.prompt().user(q).advisors(a -> a.param(EndpointTagConvention.KEY, "stream")).stream().content();
}
}
@@ -0,0 +1,22 @@
package com.ankurm.observability;
import io.micrometer.common.KeyValue;
import io.micrometer.common.KeyValues;
import org.springframework.ai.chat.client.observation.ChatClientObservationContext;
import org.springframework.ai.chat.client.observation.DefaultChatClientObservationConvention;
/**
* Adds an {@code app.endpoint} tag to the ChatClient observation, taken from the request context.
* It is a LOW-cardinality key, so it becomes a metric tag: keep the set of values small and fixed
* (an endpoint name), never a user id or a prompt.
*/
public class EndpointTagConvention extends DefaultChatClientObservationConvention {
public static final String KEY = "endpoint";
@Override
public KeyValues getLowCardinalityKeyValues(ChatClientObservationContext context) {
Object endpoint = context.getRequest().context().get(KEY);
return super.getLowCardinalityKeyValues(context).and(KeyValue.of("app.endpoint", endpoint == null ? "none" : endpoint.toString()));
}
}
@@ -0,0 +1,14 @@
package com.ankurm.observability;
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
import org.springframework.boot.context.properties.ConfigurationPropertiesScan;
@SpringBootApplication
@ConfigurationPropertiesScan
public class ObservabilityApplication {
public static void main(String[] args) {
SpringApplication.run(ObservabilityApplication.class, args);
}
}
@@ -0,0 +1,20 @@
package com.ankurm.observability;
import io.micrometer.core.instrument.MeterRegistry;
import org.springframework.ai.chat.client.observation.ChatClientObservationConvention;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
@Configuration
public class ObservabilityConfig {
@Bean
ChatClientObservationConvention endpointTagConvention() {
return new EndpointTagConvention();
}
@Bean
UsageCostObservationHandler usageCostObservationHandler(MeterRegistry registry, Pricing pricing) {
return new UsageCostObservationHandler(registry, pricing);
}
}
@@ -0,0 +1,25 @@
package com.ankurm.observability;
import java.util.Map;
import org.springframework.boot.context.properties.ConfigurationProperties;
/** USD per one million tokens, by model name prefix. Prices are configuration, not code: they change. */
@ConfigurationProperties(prefix = "ai.pricing")
public record Pricing(Map<String, Price> models) {
public record Price(double inputPerMillion, double outputPerMillion) {
}
/** Longest configured prefix of the model name wins, so "gpt-4o-mini-2024-07-18" matches "gpt-4o-mini" and not "gpt-4o". */
public Price forModel(String model) {
if (model == null || models == null) {
return null;
}
return models.entrySet().stream()
.filter(e -> model.startsWith(e.getKey()))
.max(java.util.Comparator.comparingInt(e -> e.getKey().length()))
.map(Map.Entry::getValue)
.orElse(null);
}
}
@@ -0,0 +1,57 @@
package com.ankurm.observability;
import io.micrometer.core.instrument.Counter;
import io.micrometer.core.instrument.MeterRegistry;
import io.micrometer.observation.Observation;
import io.micrometer.observation.ObservationHandler;
import org.springframework.ai.chat.client.ChatClientResponse;
import org.springframework.ai.chat.client.observation.ChatClientObservationContext;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.model.ChatResponse;
/**
* When a ChatClient call finishes, reads the token usage the provider reported and records two
* counters tagged by endpoint and model: {@code app.ai.tokens} and {@code app.ai.cost.usd}. Cost is
* tokens times the configured price. If the provider reported no usage, or the model has no price,
* it counts the call under {@code app.ai.unpriced} instead of silently recording zero. A response
* with no usage block arrives as a usage of 0 tokens, not as null, and a real call always has at
* least one prompt token, so zero is treated as "no usage".
*/
public class UsageCostObservationHandler implements ObservationHandler<ChatClientObservationContext> {
private final MeterRegistry registry;
private final Pricing pricing;
public UsageCostObservationHandler(MeterRegistry registry, Pricing pricing) {
this.registry = registry;
this.pricing = pricing;
}
@Override
public boolean supportsContext(Observation.Context context) {
return context instanceof ChatClientObservationContext;
}
@Override
public void onStop(ChatClientObservationContext context) {
String endpoint = String.valueOf(context.getRequest().context().getOrDefault(EndpointTagConvention.KEY, "none"));
ChatClientResponse response = context.getResponse();
ChatResponse chat = response == null ? null : response.chatResponse();
Usage usage = chat == null || chat.getMetadata() == null ? null : chat.getMetadata().getUsage();
String model = chat == null || chat.getMetadata() == null ? null : chat.getMetadata().getModel();
Pricing.Price price = pricing.forModel(model);
boolean noUsage = usage == null || usage.getPromptTokens() == null || usage.getPromptTokens() == 0;
if (noUsage || price == null) {
Counter.builder("app.ai.unpriced").tag("endpoint", endpoint).tag("reason", noUsage ? "no-usage" : "no-price")
.register(registry).increment();
return;
}
long in = usage.getPromptTokens();
long out = usage.getCompletionTokens() == null ? 0 : usage.getCompletionTokens();
Counter.builder("app.ai.tokens").tag("endpoint", endpoint).tag("model", model).tag("type", "input").register(registry).increment(in);
Counter.builder("app.ai.tokens").tag("endpoint", endpoint).tag("model", model).tag("type", "output").register(registry).increment(out);
Counter.builder("app.ai.cost.usd").tag("endpoint", endpoint).tag("model", model).baseUnit("usd").register(registry)
.increment(in * price.inputPerMillion() / 1_000_000d + out * price.outputPerMillion() / 1_000_000d);
}
}
@@ -0,0 +1,14 @@
package com.ankurm.observability;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.stereotype.Component;
/** A tool the model can call, so the tool-calling observation has something to measure. */
@Component
public class WeatherTools {
@Tool(description = "Current temperature in Celsius for a city")
public String getWeather(String city) {
return city + ": 24 C, clear";
}
}
@@ -0,0 +1,158 @@
package com.ankurm.observability.fake;
import java.io.IOException;
import java.io.OutputStream;
import java.net.InetSocketAddress;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.CopyOnWriteArrayList;
import com.sun.net.httpserver.HttpExchange;
import com.sun.net.httpserver.HttpServer;
import tools.jackson.databind.JsonNode;
import tools.jackson.databind.json.JsonMapper;
/**
* A tiny stand-in for the OpenAI chat completions endpoint, so the REAL {@code OpenAiChatModel}
* (and with it Spring AI's real observation code) runs without an API key. It speaks just enough of
* the protocol: a plain answer, a tool call followed by a final answer, a streamed answer, and an
* HTTP 500. Token counts are a deterministic function of text length (about four characters per
* token) -- they are NOT real tokenizer output, so only their consistency, not their size, means
* anything. Words in the prompt steer it: "slow" adds 250 ms, "boom" returns a 500, "nousage" leaves out the usage block (as some OpenAI-compatible servers do).
*/
public class FakeOpenAiServer implements AutoCloseable {
private static final JsonMapper JSON = JsonMapper.builder().build();
public static final String MODEL = "gpt-4o-mini-2024-07-18";
private final HttpServer server;
private final List<String> requestBodies = new CopyOnWriteArrayList<>();
public FakeOpenAiServer(int port) throws IOException {
server = HttpServer.create(new InetSocketAddress("127.0.0.1", port), 0);
server.createContext("/chat/completions", this::handle);
server.createContext("/v1/chat/completions", this::handle);
server.start();
}
public int port() {
return server.getAddress().getPort();
}
public String url() {
return "http://127.0.0.1:" + port();
}
/** Raw request bodies received, oldest first. */
public List<String> requestBodies() {
return requestBodies;
}
@Override
public void close() {
server.stop(0);
}
private void handle(HttpExchange ex) throws IOException {
String body = new String(ex.getRequestBody().readAllBytes(), StandardCharsets.UTF_8);
requestBodies.add(body);
JsonNode req = JSON.readTree(body);
StringBuilder all = new StringBuilder();
String lastRole = "";
String lastUser = "";
for (JsonNode m : req.path("messages")) {
String content = m.path("content").isString() ? m.path("content").asString() : m.path("content").toString();
all.append(content);
lastRole = m.path("role").asString();
if (lastRole.equals("user")) {
lastUser = content;
}
}
String text = all.toString();
sleep(text.contains("slow") ? 250 : 30);
if (text.contains("boom")) {
send(ex, 500, "{\"error\":{\"message\":\"simulated upstream failure\",\"type\":\"server_error\"}}");
return;
}
int promptTokens = tokens(text);
boolean noUsage = text.contains("nousage");
boolean stream = req.path("stream").asBoolean(false);
boolean wantsTool = req.path("tools").size() > 0 && !lastRole.equals("tool") && lastUser.toLowerCase().contains("weather");
if (wantsTool) {
String city = lastUser.replaceAll("(?s).*weather in ([A-Za-z]+).*", "$1");
String message = "{\"role\":\"assistant\",\"content\":null,\"tool_calls\":[{\"id\":\"call_1\",\"type\":\"function\",\"function\":{\"name\":\"getWeather\",\"arguments\":\"{\\\"city\\\":\\\"" + city + "\\\"}\"}}]}";
send(ex, 200, completion(message, "tool_calls", promptTokens, 12, noUsage));
return;
}
String answer = lastRole.equals("tool") ? "It is 24 C and clear." : "Answer to: " + abbreviate(lastUser);
if (stream) {
streamAnswer(ex, answer, promptTokens, !noUsage && req.path("stream_options").path("include_usage").asBoolean(false));
return;
}
send(ex, 200, completion("{\"role\":\"assistant\",\"content\":" + JSON.writeValueAsString(answer) + "}", "stop", promptTokens, tokens(answer), noUsage));
}
private static String completion(String message, String finish, int in, int out, boolean noUsage) {
return "{\"id\":\"chatcmpl-fake\",\"object\":\"chat.completion\",\"created\":1700000000,\"model\":\"" + MODEL + "\","
+ "\"choices\":[{\"index\":0,\"message\":" + message + ",\"finish_reason\":\"" + finish + "\"}]"
+ (noUsage ? "" : ",\"usage\":{\"prompt_tokens\":" + in + ",\"completion_tokens\":" + out + ",\"total_tokens\":" + (in + out) + "}") + "}";
}
private void streamAnswer(HttpExchange ex, String answer, int promptTokens, boolean includeUsage) throws IOException {
ex.getResponseHeaders().add("Content-Type", "text/event-stream");
ex.sendResponseHeaders(200, 0);
List<String> parts = new ArrayList<>(List.of(answer.split("(?<= )")));
try (OutputStream os = ex.getResponseBody()) {
for (String part : parts) {
os.write(("data: " + chunk("{\"content\":" + JSON.writeValueAsString(part) + "}", "null", false, 0, 0) + "\n\n").getBytes(StandardCharsets.UTF_8));
}
os.write(("data: " + chunk("{}", "\"stop\"", false, 0, 0) + "\n\n").getBytes(StandardCharsets.UTF_8));
if (includeUsage) {
os.write(("data: " + chunk("{}", "null", true, promptTokens, tokens(answer)) + "\n\n").getBytes(StandardCharsets.UTF_8));
}
os.write("data: [DONE]\n\n".getBytes(StandardCharsets.UTF_8));
}
}
private static String chunk(String delta, String finish, boolean usage, int in, int out) {
String choices = usage ? "[]" : "[{\"index\":0,\"delta\":" + delta + ",\"finish_reason\":" + finish + "}]";
String u = usage ? ",\"usage\":{\"prompt_tokens\":" + in + ",\"completion_tokens\":" + out + ",\"total_tokens\":" + (in + out) + "}" : "";
return "{\"id\":\"chatcmpl-fake\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"" + MODEL + "\",\"choices\":" + choices + u + "}";
}
private static void send(HttpExchange ex, int status, String json) throws IOException {
byte[] bytes = json.getBytes(StandardCharsets.UTF_8);
ex.getResponseHeaders().add("Content-Type", "application/json");
ex.sendResponseHeaders(status, bytes.length);
try (OutputStream os = ex.getResponseBody()) {
os.write(bytes);
}
}
static int tokens(String s) {
return Math.max(1, (s.length() + 3) / 4);
}
private static String abbreviate(String s) {
return s.length() > 40 ? s.substring(0, 40) : s;
}
private static void sleep(long ms) {
try {
Thread.sleep(ms);
}
catch (InterruptedException e) {
Thread.currentThread().interrupt();
}
}
public static void main(String[] args) throws Exception {
int port = args.length > 0 ? Integer.parseInt(args[0]) : 8099;
new FakeOpenAiServer(port);
System.out.println("fake OpenAI server on " + port);
Thread.currentThread().join();
}
}
@@ -0,0 +1,11 @@
# Profile used by scripts/stack-up.sh: adds the latency histogram Grafana needs for percentiles.
management:
metrics:
distribution:
percentiles-histogram:
gen_ai.client.operation: true
spring.ai.chat.client: true
minimum-expected-value:
gen_ai.client.operation: 50ms
maximum-expected-value:
gen_ai.client.operation: 30s
@@ -0,0 +1,31 @@
spring:
application:
name: spring-ai-observability
ai:
openai:
api-key: test-key-not-a-secret
base-url: ${FAKE_OPENAI_URL:http://localhost:8099}
chat:
options:
model: gpt-4o-mini
management:
endpoints:
web:
exposure:
include: health,prometheus,metrics
tracing:
sampling:
probability: 1.0
export:
otlp:
enabled: false # set true (and management.opentelemetry.tracing.export.otlp.endpoint) to ship spans to a Collector
otlp:
metrics:
export:
enabled: false # Prometheus scrapes /actuator/prometheus instead
ai:
# Illustrative prices in USD per million tokens. They are example numbers for the demo, not a current price list.
pricing:
models:
gpt-4o-mini: { input-per-million: 0.15, output-per-million: 0.60 }
gpt-4o: { input-per-million: 2.50, output-per-million: 10.00 }