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:
@@ -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 }
|
||||
Reference in New Issue
Block a user