diff --git a/README.md b/README.md index 96c952f..78999ce 100644 --- a/README.md +++ b/README.md @@ -16,5 +16,6 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur | [`advisors/`](advisors) | Three custom advisors -- a logger, a PII redactor (with a stream-safe restore) and a per-request / per-user token budget -- and tests for how the chain is ordered, what `BaseAdvisor` does on a stream, where an advisor sits relative to memory and the tool loop, and what a refusal looks like on a call, a stream and over HTTP (429). A recording stub model, no live model. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Writing Custom Advisors in Spring AI 2.0: Logging, PII Redaction and Token Budgets](https://ankurm.com/spring-ai-2-0-custom-advisors-logging-pii-redaction-token-budgets/) | | [`vector-stores/`](vector-stores) | The same 30,000-document dataset behind `VectorStore` on pgvector, Redis, Qdrant and Elasticsearch: ingest time, recall@10, latency, metadata filtering and running cost, with the defaults that cost recall reproduced (Elasticsearch's quantised mapping, Redis `EF_RUNTIME`, pgvector post-filtering, Qdrant payload indexes). Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Choosing a Vector Store for Spring AI](https://ankurm.com/spring-ai-2-0-vector-store-comparison-pgvector-redis-qdrant-elasticsearch/) | | [`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/) | Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/). diff --git a/observability/.gitignore b/observability/.gitignore new file mode 100644 index 0000000..2f7896d --- /dev/null +++ b/observability/.gitignore @@ -0,0 +1 @@ +target/ diff --git a/observability/README.md b/observability/README.md new file mode 100644 index 0000000..3d90e1a --- /dev/null +++ b/observability/README.md @@ -0,0 +1,55 @@ +# observability + +Companion code for [Observability for Spring AI: Tokens, Latency and Cost with Micrometer and OpenTelemetry](https://ankurm.com/spring-ai-2-0-observability-micrometer-opentelemetry-tokens-cost/), part of the [Spring AI series](../README.md) on ankurm.com. + +A small Spring AI app (`/ask`, `/summarize`, `/weather` with a tool call, `/stream`), the meters and spans the framework records for it with no code of ours, a 60-line handler that turns token usage into cost per endpoint, and a Grafana dashboard fed by Prometheus. + +**The model class is real, the provider is not.** The app uses the real `OpenAiChatModel` and the real OpenAI Java SDK, pointed at `FakeOpenAiServer`, a local HTTP server that speaks the chat-completions protocol. So every meter name, tag, span and retry count is what Spring AI 2.0.1 really produces. But the token counts are about characters / 4, the latencies are the fake's, and the prices in `application.yml` are illustrative examples, not a current price list. No real provider was measured. The error test shows the SDK's retries against a server that always answers 500. + +## Versions + +| Component | Version | +|---|---| +| Spring Boot | 4.1.1 | +| Spring AI | 2.0.1 | +| Micrometer / micrometer-tracing | 1.17.1 / 1.7.1 | +| OpenTelemetry | 1.62.0 | +| Prometheus / Grafana (release tarballs, no Docker) | 3.15.0 / 13.2.3 | +| Java | 25 (LTS) | + +## Quickstart + +```bash +scripts/run-all.sh # 7 tests, then the stack, a load run and the dashboard check; regenerates output/01 .. 09 +scripts/stack-up.sh # or just the stack: app :8080, Prometheus :9090, Grafana :3000 (admin/admin) +scripts/load.sh 20 # mixed traffic so every panel has data +scripts/stack-down.sh +``` + +Two consecutive runs of `run-all.sh` produce byte-identical `output/` files. The Prometheus and Grafana tarballs must be unpacked under `/tmp/tools/obs/{prom,grafana}` (or set `PROM` and `GRAFANA`). + +## What's here + +| File | What it is | +|---|---| +| [`AiController.java`](src/main/java/com/ankurm/observability/AiController.java), [`WeatherTools.java`](src/main/java/com/ankurm/observability/WeatherTools.java) | The endpoints and the one tool | +| [`EndpointTagConvention.java`](src/main/java/com/ankurm/observability/EndpointTagConvention.java) | Adds a low-cardinality `app.endpoint` tag to the framework's chat client meter | +| [`UsageCostObservationHandler.java`](src/main/java/com/ankurm/observability/UsageCostObservationHandler.java), [`Pricing.java`](src/main/java/com/ankurm/observability/Pricing.java) | Tokens and cost per endpoint, plus a counter for calls that could not be priced | +| [`fake/FakeOpenAiServer.java`](src/main/java/com/ankurm/observability/fake/FakeOpenAiServer.java) | The local stand-in for the provider (keywords `slow`, `boom`, `nousage`, `weather`) | +| [`application.yml`](src/main/resources/application.yml), [`application-stack.yml`](src/main/resources/application-stack.yml) | Settings; the `stack` profile adds latency histograms | +| [`dashboards/spring-ai-observability.json`](dashboards/spring-ai-observability.json) | The Grafana dashboard, generated by [`scripts/make-dashboard.py`](scripts/make-dashboard.py); [`screenshot.png`](dashboards/screenshot.png) is how it rendered | +| [`stack/`](stack) | Prometheus scrape config and Grafana provisioning | + +## Output files + +| File | Written by | +|---|---| +| [`01-builtin-metrics.txt`](output/01-builtin-metrics.txt) | `BuiltInMetricsTest`: the meters the framework creates | +| [`02-tool-call-spans.txt`](output/02-tool-call-spans.txt) | `ToolCallSpansTest`: the span tree of a request with a tool call | +| [`03-cost-per-endpoint.txt`](output/03-cost-per-endpoint.txt) | `CostPerEndpointTest`: tokens and cost per endpoint | +| [`04-streaming-usage.txt`](output/04-streaming-usage.txt) | `StreamingUsageTest`: usage on a streamed response | +| [`05-error-and-retry.txt`](output/05-error-and-retry.txt) | `ErrorAndRetryTest`: a failed call and the SDK's retries | +| [`06-prompt-content.txt`](output/06-prompt-content.txt) | `PromptContentTest`: where prompt text appears, by default and when logging is on | +| [`07-histogram.txt`](output/07-histogram.txt) | `HistogramTest`: latency buckets are opt-in | +| [`08-dashboard-queries.txt`](output/08-dashboard-queries.txt) | `scripts/verify-dashboard.py`: every panel query through Prometheus and Grafana | +| [`09-prometheus-families.txt`](output/09-prometheus-families.txt) | The metric families as Prometheus sees them | diff --git a/observability/dashboards/screenshot.png b/observability/dashboards/screenshot.png new file mode 100644 index 0000000..8770f44 Binary files /dev/null and b/observability/dashboards/screenshot.png differ diff --git a/observability/dashboards/spring-ai-observability.json b/observability/dashboards/spring-ai-observability.json new file mode 100644 index 0000000..cbe4404 --- /dev/null +++ b/observability/dashboards/spring-ai-observability.json @@ -0,0 +1,308 @@ +{ + "uid": "spring-ai-observability", + "title": "Spring AI: tokens, latency and cost", + "schemaVersion": 39, + "version": 1, + "editable": true, + "refresh": "5s", + "time": { + "from": "now-15m", + "to": "now" + }, + "tags": [ + "spring-ai" + ], + "panels": [ + { + "id": 1, + "title": "Cost per hour by endpoint (USD, illustrative prices)", + "type": "timeseries", + "description": "Rate of app_ai_cost_usd_total scaled to an hourly figure. Prices come from ai.pricing in application.yml.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 0 + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (endpoint) (rate(app_ai_cost_usd_total[1m])) * 3600", + "legendFormat": "{{endpoint}}", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 2, + "title": "Tokens per minute by endpoint and direction", + "type": "timeseries", + "description": "Input and output tokens as the provider reported them.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 0 + }, + "fieldConfig": { + "defaults": { + "unit": "short" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (endpoint, type) (rate(app_ai_tokens_total[1m])) * 60", + "legendFormat": "{{endpoint}} {{type}}", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 3, + "title": "Cost per request by endpoint (USD)", + "type": "bargauge", + "description": "Cost divided by successful chat client calls. The two sides use different label names (endpoint vs app_endpoint), so label_replace renames one before the division.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 8 + }, + "fieldConfig": { + "defaults": { + "unit": "currencyUSD" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (endpoint) (increase(app_ai_cost_usd_total[5m])) / on (endpoint) label_replace(sum by (app_endpoint) (increase(spring_ai_chat_client_seconds_count{error=\"none\"}[5m])), \"endpoint\", \"$1\", \"app_endpoint\", \"(.*)\")", + "legendFormat": "{{endpoint}}", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 4, + "title": "Model call latency p50 / p95 (s)", + "type": "timeseries", + "description": "gen_ai.client.operation, one sample per HTTP call to the provider. Needs the percentiles-histogram property.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 8 + }, + "fieldConfig": { + "defaults": { + "unit": "s" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "histogram_quantile(0.50, sum by (le) (rate(gen_ai_client_operation_seconds_bucket{error=\"none\"}[1m])))", + "legendFormat": "p50", + "editorMode": "code", + "range": true + }, + { + "refId": "B", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_seconds_bucket{error=\"none\"}[1m])))", + "legendFormat": "p95", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 5, + "title": "Failed model calls (share)", + "type": "stat", + "description": "Calls whose error tag is not none, over all calls.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 16 + }, + "fieldConfig": { + "defaults": { + "unit": "percentunit" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum(rate(gen_ai_client_operation_seconds_count{error!=\"none\"}[5m])) / sum(rate(gen_ai_client_operation_seconds_count[5m]))", + "legendFormat": "failed", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 6, + "title": "Tool calls per minute", + "type": "timeseries", + "description": "spring.ai.tool, one series per tool name.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 16 + }, + "fieldConfig": { + "defaults": { + "unit": "short" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_count[1m])) * 60", + "legendFormat": "{{spring_ai_tool_definition_name}}", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 7, + "title": "Mean tool latency (s)", + "type": "timeseries", + "description": "Sum over count of the tool timer.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 0, + "y": 24 + }, + "fieldConfig": { + "defaults": { + "unit": "s" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_sum[1m])) / sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_count[1m]))", + "legendFormat": "{{spring_ai_tool_definition_name}}", + "editorMode": "code", + "range": true + } + ] + }, + { + "id": 8, + "title": "Calls with no cost recorded", + "type": "timeseries", + "description": "Calls that failed, reported no usage, or used a model with no configured price. Zero cost is never recorded silently.", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "gridPos": { + "h": 8, + "w": 12, + "x": 12, + "y": 24 + }, + "fieldConfig": { + "defaults": { + "unit": "short" + }, + "overrides": [] + }, + "targets": [ + { + "refId": "A", + "datasource": { + "type": "prometheus", + "uid": "prom" + }, + "expr": "sum by (endpoint, reason) (increase(app_ai_unpriced_total[5m]))", + "legendFormat": "{{endpoint}} {{reason}}", + "editorMode": "code", + "range": true + } + ] + } + ] +} diff --git a/observability/output/01-builtin-metrics.txt b/observability/output/01-builtin-metrics.txt new file mode 100644 index 0000000..f4d38d6 --- /dev/null +++ b/observability/output/01-builtin-metrics.txt @@ -0,0 +1,16 @@ +# Meters the framework creates for 3 requests (/ask, /weather with one tool call, /stream) + +gen_ai.client.operation TIMER [error, gen_ai.operation.name, gen_ai.request.model, gen_ai.response.model, gen_ai.system] +gen_ai.client.operation.active LONG_TASK_TIMER [gen_ai.operation.name, gen_ai.request.model, gen_ai.response.model, gen_ai.system] +gen_ai.client.token.usage COUNTER [gen_ai.operation.name, gen_ai.request.model, gen_ai.response.model, gen_ai.system, gen_ai.token.type] +spring.ai.advisor TIMER [error, gen_ai.operation.name, gen_ai.system, spring.ai.advisor.name, spring.ai.kind] +spring.ai.advisor.active LONG_TASK_TIMER [gen_ai.operation.name, gen_ai.system, spring.ai.advisor.name, spring.ai.kind] +spring.ai.chat.client TIMER [app.endpoint, error, gen_ai.operation.name, gen_ai.system, spring.ai.chat.client.stream, spring.ai.kind] +spring.ai.chat.client.active LONG_TASK_TIMER [app.endpoint, gen_ai.operation.name, gen_ai.system, spring.ai.chat.client.stream, spring.ai.kind] +spring.ai.tool TIMER [error, gen_ai.operation.name, gen_ai.system, spring.ai.kind, spring.ai.tool.definition.name, spring.ai.tool.type] +spring.ai.tool.active LONG_TASK_TIMER [gen_ai.operation.name, gen_ai.system, spring.ai.kind, spring.ai.tool.definition.name, spring.ai.tool.type] + +model calls recorded (gen_ai.client.operation, error=none): 4 +chat client calls recorded (spring.ai.chat.client): 3 +tool executions recorded (spring.ai.tool): 1 +tokens: input=24 output=29 total=53 diff --git a/observability/output/02-tool-call-spans.txt b/observability/output/02-tool-call-spans.txt new file mode 100644 index 0000000..7db744e --- /dev/null +++ b/observability/output/02-tool-call-spans.txt @@ -0,0 +1,24 @@ +# Span tree for GET /weather?city=Pune (one tool call, two model calls) + +http get /weather [SERVER] + spring_ai chat_client [INTERNAL] + tool _calling [INTERNAL] + call [INTERNAL] + chat gpt-4o-mini [INTERNAL] + POST [CLIENT] + execute_tool getWeather [INTERNAL] + call [INTERNAL] + chat gpt-4o-mini [INTERNAL] + POST [CLIENT] + +attribute keys on the model-call spans: + gen_ai.operation.name + gen_ai.request.model + gen_ai.response.finish_reasons + gen_ai.response.id + gen_ai.response.model + gen_ai.system + gen_ai.usage.input_tokens + gen_ai.usage.output_tokens + gen_ai.usage.total_tokens + spring.ai.model.request.tool.names diff --git a/observability/output/03-cost-per-endpoint.txt b/observability/output/03-cost-per-endpoint.txt new file mode 100644 index 0000000..3a75953 --- /dev/null +++ b/observability/output/03-cost-per-endpoint.txt @@ -0,0 +1,10 @@ +# Tokens and cost per endpoint (illustrative price: 0.15 USD in / 0.60 USD out per million tokens) + +endpoint input output cost USD +ask 1 4 0.00000255 +summarize 511 13 0.00008445 +weather 19 18 0.00001365 + +framework model-level tokens: input=531 output=35 +our per-endpoint tokens : input=531 output=35 +total cost USD: 0.00010065 diff --git a/observability/output/04-streaming-usage.txt b/observability/output/04-streaming-usage.txt new file mode 100644 index 0000000..9b68aaa --- /dev/null +++ b/observability/output/04-streaming-usage.txt @@ -0,0 +1,9 @@ +# Streaming asks for usage; a server that omits it is counted as unpriced + +stream request body asks for usage (stream_options.include_usage): true +request fragment: "stream":true + +two calls to a server that omits the usage block (/ask and /stream): + gen_ai.client.token.usage total, before 7 -> after 7 + app.ai.unpriced{reason=no-usage} for ask : 1 + app.ai.unpriced{reason=no-usage} for stream: 1 diff --git a/observability/output/05-error-and-retry.txt b/observability/output/05-error-and-retry.txt new file mode 100644 index 0000000..367cb75 --- /dev/null +++ b/observability/output/05-error-and-retry.txt @@ -0,0 +1,8 @@ +# One failed user request (the provider returns HTTP 500) + +HTTP requests the provider received for that one call: 4 +model-call timers recorded: + error=InternalServerException response.model=none count=1 +chat-client timers with an error: 1 +app.ai.unpriced{endpoint=ask,reason=no-usage}: 1 +app.ai.cost.usd recorded for the failed call: 0 diff --git a/observability/output/06-prompt-content.txt b/observability/output/06-prompt-content.txt new file mode 100644 index 0000000..7df4a93 --- /dev/null +++ b/observability/output/06-prompt-content.txt @@ -0,0 +1,7 @@ +# Where does the prompt text "ticket-4711-card-ending-0042" appear in telemetry? + +setting log lines with it / span attributes with it / span events with it +defaults 0 / 0 / 0 +spring.ai.chat.observations.log-prompt=true (+ client) 1 / 0 / 0 + +span attributes carrying the text when enabled: [] diff --git a/observability/output/07-histogram.txt b/observability/output/07-histogram.txt new file mode 100644 index 0000000..16246ce --- /dev/null +++ b/observability/output/07-histogram.txt @@ -0,0 +1,7 @@ +# Bucket series for gen_ai_client_operation_seconds after 2 calls + +setting _bucket series / _sum series +defaults 0 / 1 +percentiles-histogram.gen_ai.client.operation=true 69 / 1 + +first bucket series when enabled: gen_ai_client_operation_seconds_bucket{le="0.001"} diff --git a/observability/output/08-dashboard-queries.txt b/observability/output/08-dashboard-queries.txt new file mode 100644 index 0000000..4d11492 --- /dev/null +++ b/observability/output/08-dashboard-queries.txt @@ -0,0 +1,12 @@ +# Every dashboard query, run against Prometheus and through Grafana's query API after scripts/load.sh + +panel prom series grafana frames +Cost per hour by endpoint (USD, illustrative pri [A] 4 4 +Tokens per minute by endpoint and direction [A] 8 8 +Cost per request by endpoint (USD) [A] 4 4 +Model call latency p50 / p95 (s) [A] 1 1 +Model call latency p50 / p95 (s) [B] 1 1 +Failed model calls (share) [A] 1 1 +Tool calls per minute [A] 1 1 +Mean tool latency (s) [A] 1 1 +Calls with no cost recorded [A] 1 1 diff --git a/observability/output/09-prometheus-families.txt b/observability/output/09-prometheus-families.txt new file mode 100644 index 0000000..76734dc --- /dev/null +++ b/observability/output/09-prometheus-families.txt @@ -0,0 +1,26 @@ +# Prometheus metric families this app exposes at /actuator/prometheus (TYPE lines only) + +# TYPE app_ai_cost_usd_total counter +# TYPE app_ai_tokens_total counter +# TYPE app_ai_unpriced_total counter +# TYPE gen_ai_client_operation_active_seconds histogram +# TYPE gen_ai_client_operation_active_seconds_gcount gauge +# TYPE gen_ai_client_operation_active_seconds_gsum gauge +# TYPE gen_ai_client_operation_active_seconds_max gauge +# TYPE gen_ai_client_operation_seconds histogram +# TYPE gen_ai_client_operation_seconds_max gauge +# TYPE gen_ai_client_token_usage_total counter +# TYPE spring_ai_advisor_active_seconds summary +# TYPE spring_ai_advisor_active_seconds_max gauge +# TYPE spring_ai_advisor_seconds summary +# TYPE spring_ai_advisor_seconds_max gauge +# TYPE spring_ai_chat_client_active_seconds histogram +# TYPE spring_ai_chat_client_active_seconds_gcount gauge +# TYPE spring_ai_chat_client_active_seconds_gsum gauge +# TYPE spring_ai_chat_client_active_seconds_max gauge +# TYPE spring_ai_chat_client_seconds histogram +# TYPE spring_ai_chat_client_seconds_max gauge +# TYPE spring_ai_tool_active_seconds summary +# TYPE spring_ai_tool_active_seconds_max gauge +# TYPE spring_ai_tool_seconds summary +# TYPE spring_ai_tool_seconds_max gauge diff --git a/observability/pom.xml b/observability/pom.xml new file mode 100644 index 0000000..33c88fc --- /dev/null +++ b/observability/pom.xml @@ -0,0 +1,86 @@ + + + 4.0.0 + + + org.springframework.boot + spring-boot-starter-parent + 4.1.1 + + + + com.ankurm + observability + 1.0.0 + observability + Spring AI 2.0 observability: token usage, latency and cost with Micrometer, Prometheus and OpenTelemetry, against the real OpenAI model class and a local fake server. + + + 25 + 2.0.1 + + + + + + org.springframework.ai + spring-ai-bom + ${spring-ai.version} + pom + import + + + + + + + org.springframework.boot + spring-boot-starter-webmvc + + + org.springframework.boot + spring-boot-starter-actuator + + + io.micrometer + micrometer-registry-prometheus + + + org.springframework.boot + spring-boot-starter-opentelemetry + + + org.springframework.ai + spring-ai-starter-model-openai + + + + org.springframework.boot + spring-boot-starter-test + test + + + io.opentelemetry + opentelemetry-sdk-testing + test + + + + + + + org.springframework.boot + spring-boot-maven-plugin + + + org.apache.maven.plugins + maven-surefire-plugin + + -Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8 + + + + + diff --git a/observability/scripts/load.sh b/observability/scripts/load.sh new file mode 100755 index 0000000..c1ae752 --- /dev/null +++ b/observability/scripts/load.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash +# Sends a mixed workload so every dashboard panel has data. Usage: load.sh [rounds] +H=http://127.0.0.1:8080 +for i in $(seq 1 "${1:-10}"); do + curl -s -o /dev/null "$H/ask?q=what+is+observability+$i" + curl -s -o /dev/null "$H/summarize" --get --data-urlencode "text=$(head -c 2000 /dev/zero | tr '\0' 'x') round $i" + curl -s -o /dev/null "$H/weather?city=Pune" + curl -s -o /dev/null "$H/stream?q=tell+me+something+$i" + [ $((i % 4)) -eq 0 ] && curl -s -o /dev/null "$H/ask?q=slow+request+$i" + [ $((i % 5)) -eq 0 ] && curl -s -o /dev/null "$H/ask?q=boom+$i" + [ $((i % 6)) -eq 0 ] && curl -s -o /dev/null "$H/ask?q=nousage+$i" +done +echo "load done" diff --git a/observability/scripts/make-dashboard.py b/observability/scripts/make-dashboard.py new file mode 100755 index 0000000..610775c --- /dev/null +++ b/observability/scripts/make-dashboard.py @@ -0,0 +1,49 @@ +#!/usr/bin/env python3 +"""Generates dashboards/spring-ai-observability.json. Run it, commit the JSON; Grafana provisions the file as-is.""" +import json, pathlib + +DS = {"type": "prometheus", "uid": "prom"} +PANELS = [ + # (title, type, unit, queries[(expr, legend)], description) + ("Cost per hour by endpoint (USD, illustrative prices)", "timeseries", "currencyUSD", + [("sum by (endpoint) (rate(app_ai_cost_usd_total[1m])) * 3600", "{{endpoint}}")], + "Rate of app_ai_cost_usd_total scaled to an hourly figure. Prices come from ai.pricing in application.yml."), + ("Tokens per minute by endpoint and direction", "timeseries", "short", + [("sum by (endpoint, type) (rate(app_ai_tokens_total[1m])) * 60", "{{endpoint}} {{type}}")], + "Input and output tokens as the provider reported them."), + ("Cost per request by endpoint (USD)", "bargauge", "currencyUSD", + [("sum by (endpoint) (increase(app_ai_cost_usd_total[5m])) / on (endpoint) label_replace(sum by (app_endpoint) (increase(spring_ai_chat_client_seconds_count{error=\"none\"}[5m])), \"endpoint\", \"$1\", \"app_endpoint\", \"(.*)\")", "{{endpoint}}")], + "Cost divided by successful chat client calls. The two sides use different label names (endpoint vs app_endpoint), so label_replace renames one before the division."), + ("Model call latency p50 / p95 (s)", "timeseries", "s", + [("histogram_quantile(0.50, sum by (le) (rate(gen_ai_client_operation_seconds_bucket{error=\"none\"}[1m])))", "p50"), + ("histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_seconds_bucket{error=\"none\"}[1m])))", "p95")], + "gen_ai.client.operation, one sample per HTTP call to the provider. Needs the percentiles-histogram property."), + ("Failed model calls (share)", "stat", "percentunit", + [("sum(rate(gen_ai_client_operation_seconds_count{error!=\"none\"}[5m])) / sum(rate(gen_ai_client_operation_seconds_count[5m]))", "failed")], + "Calls whose error tag is not none, over all calls."), + ("Tool calls per minute", "timeseries", "short", + [("sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_count[1m])) * 60", "{{spring_ai_tool_definition_name}}")], + "spring.ai.tool, one series per tool name."), + ("Mean tool latency (s)", "timeseries", "s", + [("sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_sum[1m])) / sum by (spring_ai_tool_definition_name) (rate(spring_ai_tool_seconds_count[1m]))", "{{spring_ai_tool_definition_name}}")], + "Sum over count of the tool timer."), + ("Calls with no cost recorded", "timeseries", "short", + [("sum by (endpoint, reason) (increase(app_ai_unpriced_total[5m]))", "{{endpoint}} {{reason}}")], + "Calls that failed, reported no usage, or used a model with no configured price. Zero cost is never recorded silently."), +] + +panels = [] +for i, (title, kind, unit, queries, desc) in enumerate(PANELS): + panels.append({ + "id": i + 1, "title": title, "type": kind, "description": desc, "datasource": DS, + "gridPos": {"h": 8, "w": 12, "x": (i % 2) * 12, "y": (i // 2) * 8}, + "fieldConfig": {"defaults": {"unit": unit}, "overrides": []}, + "targets": [{"refId": chr(65 + j), "datasource": DS, "expr": e, "legendFormat": l, "editorMode": "code", "range": True} + for j, (e, l) in enumerate(queries)], + }) + +dash = {"uid": "spring-ai-observability", "title": "Spring AI: tokens, latency and cost", "schemaVersion": 39, "version": 1, + "editable": True, "refresh": "5s", "time": {"from": "now-15m", "to": "now"}, "tags": ["spring-ai"], "panels": panels} +out = pathlib.Path(__file__).resolve().parent.parent / "dashboards" / "spring-ai-observability.json" +out.write_text(json.dumps(dash, indent=2) + "\n") +print("wrote", out, len(panels), "panels") diff --git a/observability/scripts/run-all.sh b/observability/scripts/run-all.sh new file mode 100755 index 0000000..6b8ff8a --- /dev/null +++ b/observability/scripts/run-all.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash +# Regenerates every file in output/: 01-07 from the test suite, 08-09 from the live Prometheus + Grafana stack. +# Needs JDK 25 and Maven on PATH, plus the Prometheus and Grafana tarballs unpacked (see stack-up.sh). +set -euo pipefail +cd "$(dirname "$0")/.." +mvn -q -B test +scripts/stack-up.sh +trap 'scripts/stack-down.sh' EXIT +scripts/load.sh 12 +sleep 8 +{ echo "# Prometheus metric families this app exposes at /actuator/prometheus (TYPE lines only)"; echo + curl -s 127.0.0.1:8080/actuator/prometheus | grep -E '^# TYPE (gen_ai_|spring_ai_|app_ai_)' | sort; } > output/09-prometheus-families.txt +scripts/verify-dashboard.py output/08-dashboard-queries.txt diff --git a/observability/scripts/stack-down.sh b/observability/scripts/stack-down.sh new file mode 100755 index 0000000..e798090 --- /dev/null +++ b/observability/scripts/stack-down.sh @@ -0,0 +1,4 @@ +#!/usr/bin/env bash +# Stops what stack-up.sh started, by PID (never by pattern). +for f in /tmp/obs-run/*.pid; do [ -f "$f" ] && kill "$(cat "$f")" 2>/dev/null || true; done +echo stopped diff --git a/observability/scripts/stack-up.sh b/observability/scripts/stack-up.sh new file mode 100755 index 0000000..b3007d5 --- /dev/null +++ b/observability/scripts/stack-up.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +# Starts the whole demo stack without Docker, all on 127.0.0.1: +# fake OpenAI server :8099 -> the Spring app :8080 <-scraped every 2s- Prometheus :9090 <- Grafana :3000 +# PROM and GRAFANA default to where the article says to unpack the release tarballs. +set -euo pipefail +cd "$(dirname "$0")/.." +PROM="${PROM:-/tmp/tools/obs/prom}"; GRAFANA="${GRAFANA:-/tmp/tools/obs/grafana}" +RUN=/tmp/obs-run; rm -rf "$RUN"; mkdir -p "$RUN/dashboards" "$RUN/prom" "$RUN/grafana/data" "$RUN/grafana/logs" "$RUN/grafana/plugins" +cp dashboards/*.json "$RUN/dashboards/" +mvn -q -B -DskipTests package +JAR=$(ls target/*.jar | grep -v original | head -1) +start() { local name=$1; shift; setsid nohup "$@" > "$RUN/$name.log" 2>&1 < /dev/null & echo $! > "$RUN/$name.pid"; } +wait_for() { for i in $(seq 1 60); do curl -s -o /dev/null "$1" && return 0; sleep 1; done; echo "timeout waiting for $1" >&2; exit 1; } + +# the fake server ships inside the same Boot jar; PropertiesLauncher lets us pick its main class +start fake java -Dloader.main=com.ankurm.observability.fake.FakeOpenAiServer -cp "$JAR" org.springframework.boot.loader.launch.PropertiesLauncher 8099 +wait_for http://127.0.0.1:8099/ +start app java -jar "$JAR" --spring.profiles.active=stack +wait_for http://127.0.0.1:8080/actuator/health +start prometheus "$PROM/prometheus" --config.file=stack/prometheus.yml --storage.tsdb.path="$RUN/prom" --web.listen-address=127.0.0.1:9090 +wait_for http://127.0.0.1:9090/-/ready +GF_PATHS_DATA="$RUN/grafana/data" GF_PATHS_LOGS="$RUN/grafana/logs" GF_PATHS_PLUGINS="$RUN/grafana/plugins" \ +GF_PATHS_PROVISIONING="$PWD/stack/grafana/provisioning" GF_SERVER_HTTP_ADDR=127.0.0.1 GF_SERVER_HTTP_PORT=3000 \ +GF_SECURITY_ADMIN_PASSWORD=admin GF_ANALYTICS_REPORTING_ENABLED=false GF_ANALYTICS_CHECK_FOR_UPDATES=false \ + start grafana "$GRAFANA/bin/grafana" server --homepath "$GRAFANA" +wait_for http://127.0.0.1:3000/api/health +echo "up: app :8080, prometheus :9090, grafana :3000 (admin/admin). pids in $RUN/*.pid" diff --git a/observability/scripts/verify-dashboard.py b/observability/scripts/verify-dashboard.py new file mode 100755 index 0000000..9d92430 --- /dev/null +++ b/observability/scripts/verify-dashboard.py @@ -0,0 +1,38 @@ +#!/usr/bin/env python3 +"""Runs every dashboard query twice: straight against Prometheus, and through Grafana's /api/ds/query with the +provisioned datasource (what a panel really does). Prints one line per panel query and exits non-zero if any +query errors or returns no series. Usage: verify-dashboard.py [output-file]""" +import base64, json, pathlib, sys, time, urllib.parse, urllib.request + +root = pathlib.Path(__file__).resolve().parent.parent +dash = json.loads((root / "dashboards" / "spring-ai-observability.json").read_text()) +auth = "Basic " + base64.b64encode(b"admin:admin").decode() + +def get(url, data=None, headers=None): + req = urllib.request.Request(url, data=data, headers=headers or {}) + with urllib.request.urlopen(req, timeout=30) as r: + return json.loads(r.read()) + +lines, bad = [], 0 +now = int(time.time()) +lines.append(f"{'panel':<58} {'prom series':>11} {'grafana frames':>14}") +for p in dash["panels"]: + for t in p["targets"]: + q = urllib.parse.urlencode({"query": t["expr"]}) + prom = get("http://127.0.0.1:9090/api/v1/query?" + q) + n_prom = len(prom["data"]["result"]) if prom["status"] == "success" else -1 + body = json.dumps({"queries": [{"refId": t["refId"], "datasource": t["datasource"], "expr": t["expr"], "instant": False, + "range": True, "interval": "", "intervalMs": 5000, "maxDataPoints": 200}], + "from": str((now - 600) * 1000), "to": str(now * 1000)}).encode() + gf = get("http://127.0.0.1:3000/api/ds/query", body, {"Content-Type": "application/json", "Authorization": auth}) + res = gf["results"][t["refId"]] + n_gf = -1 if "error" in res else len(res.get("frames", [])) + name = f"{p['title'][:48]} [{t['refId']}]" + lines.append(f"{name:<58} {n_prom:>11} {n_gf:>14}") + if n_prom < 1 or n_gf < 1: + bad += 1 +out = "\n".join(lines) +print(out) +if len(sys.argv) > 1: + pathlib.Path(sys.argv[1]).write_text("# Every dashboard query, run against Prometheus and through Grafana's query API after scripts/load.sh\n\n" + out + "\n") +sys.exit(1 if bad else 0) diff --git a/observability/src/main/java/com/ankurm/observability/AiController.java b/observability/src/main/java/com/ankurm/observability/AiController.java new file mode 100644 index 0000000..8fade62 --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/AiController.java @@ -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 stream(@RequestParam String q) { + return chat.prompt().user(q).advisors(a -> a.param(EndpointTagConvention.KEY, "stream")).stream().content(); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/EndpointTagConvention.java b/observability/src/main/java/com/ankurm/observability/EndpointTagConvention.java new file mode 100644 index 0000000..f106b63 --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/EndpointTagConvention.java @@ -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())); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/ObservabilityApplication.java b/observability/src/main/java/com/ankurm/observability/ObservabilityApplication.java new file mode 100644 index 0000000..d8f881f --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/ObservabilityApplication.java @@ -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); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/ObservabilityConfig.java b/observability/src/main/java/com/ankurm/observability/ObservabilityConfig.java new file mode 100644 index 0000000..0c01eaa --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/ObservabilityConfig.java @@ -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); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/Pricing.java b/observability/src/main/java/com/ankurm/observability/Pricing.java new file mode 100644 index 0000000..c34f70c --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/Pricing.java @@ -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 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); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/UsageCostObservationHandler.java b/observability/src/main/java/com/ankurm/observability/UsageCostObservationHandler.java new file mode 100644 index 0000000..8ef48c2 --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/UsageCostObservationHandler.java @@ -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 { + + 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); + } +} diff --git a/observability/src/main/java/com/ankurm/observability/WeatherTools.java b/observability/src/main/java/com/ankurm/observability/WeatherTools.java new file mode 100644 index 0000000..3cc678b --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/WeatherTools.java @@ -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"; + } +} diff --git a/observability/src/main/java/com/ankurm/observability/fake/FakeOpenAiServer.java b/observability/src/main/java/com/ankurm/observability/fake/FakeOpenAiServer.java new file mode 100644 index 0000000..4ddfefd --- /dev/null +++ b/observability/src/main/java/com/ankurm/observability/fake/FakeOpenAiServer.java @@ -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 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 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 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(); + } +} diff --git a/observability/src/main/resources/application-stack.yml b/observability/src/main/resources/application-stack.yml new file mode 100644 index 0000000..843f480 --- /dev/null +++ b/observability/src/main/resources/application-stack.yml @@ -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 diff --git a/observability/src/main/resources/application.yml b/observability/src/main/resources/application.yml new file mode 100644 index 0000000..bf919af --- /dev/null +++ b/observability/src/main/resources/application.yml @@ -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 } diff --git a/observability/src/test/java/com/ankurm/observability/BuiltInMetricsTest.java b/observability/src/test/java/com/ankurm/observability/BuiltInMetricsTest.java new file mode 100644 index 0000000..1f42eb1 --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/BuiltInMetricsTest.java @@ -0,0 +1,63 @@ +package com.ankurm.observability; + +import java.util.Map; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.Obs; +import com.ankurm.observability.support.Transcript; +import io.micrometer.core.instrument.MeterRegistry; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * What Spring AI 2.0.1 records without any code of ours: the real OpenAiChatModel talks to a local + * fake server, and every meter below was created by the framework. Writes output/01. + */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class BuiltInMetricsTest { + + @DynamicPropertySource + static void fake(DynamicPropertyRegistry r) { + r.add("spring.ai.openai.base-url", Fake.SERVER::url); + } + + @LocalServerPort + int port; + + @Autowired + MeterRegistry registry; + + @Test + void whatTheFrameworkRecords() { + RestClient http = RestClient.create("http://localhost:" + port); + http.get().uri("/ask?q=hello world").retrieve().body(String.class); + http.get().uri("/weather?city=Pune").retrieve().body(String.class); + http.get().uri("/stream?q=tell me").retrieve().body(String.class); + + Map fam = Obs.families(registry, "gen_ai.", "spring.ai."); + double in = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "input"); + double out = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "output"); + double total = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "total"); + + try (Transcript t = new Transcript("01-builtin-metrics.txt", "Meters the framework creates for 3 requests (/ask, /weather with one tool call, /stream)")) { + fam.forEach((name, desc) -> t.line("%-34s %s", name, desc)); + t.blank(); + t.line("model calls recorded (gen_ai.client.operation, error=none): %d", Obs.timerCount(registry, "gen_ai.client.operation", "error", "none")); + t.line("chat client calls recorded (spring.ai.chat.client): %d", Obs.timerCount(registry, "spring.ai.chat.client")); + t.line("tool executions recorded (spring.ai.tool): %d", Obs.timerCount(registry, "spring.ai.tool")); + t.line("tokens: input=%.0f output=%.0f total=%.0f", in, out, total); + } + assertThat(fam).containsKeys("gen_ai.client.operation", "gen_ai.client.token.usage", "spring.ai.chat.client", "spring.ai.tool", "spring.ai.advisor"); + assertThat(Obs.timerCount(registry, "gen_ai.client.operation", "error", "none")).isEqualTo(4); + assertThat(Obs.timerCount(registry, "spring.ai.chat.client")).isEqualTo(3); + assertThat(Obs.timerCount(registry, "spring.ai.tool")).isEqualTo(1); + assertThat(total).isEqualTo(in + out); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/CostPerEndpointTest.java b/observability/src/test/java/com/ankurm/observability/CostPerEndpointTest.java new file mode 100644 index 0000000..967f19d --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/CostPerEndpointTest.java @@ -0,0 +1,67 @@ +package com.ankurm.observability; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.Obs; +import com.ankurm.observability.support.Transcript; +import io.micrometer.core.instrument.MeterRegistry; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.within; + +/** Tokens and cost per endpoint from our own handler, checked against what the fake server reported. Writes output/03. */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class CostPerEndpointTest { + + @DynamicPropertySource + static void fake(DynamicPropertyRegistry r) { + r.add("spring.ai.openai.base-url", Fake.SERVER::url); + } + + @LocalServerPort + int port; + + @Autowired + MeterRegistry registry; + + @Test + void costIsTokensTimesPriceAndSplitsByEndpoint() { + RestClient http = RestClient.create("http://localhost:" + port); + String shortText = "tiny"; + String longText = "x".repeat(2000); + http.get().uri("/ask?q={q}", shortText).retrieve().body(String.class); + http.get().uri("/summarize?text={t}", longText).retrieve().body(String.class); + http.get().uri("/weather?city=Pune").retrieve().body(String.class); + + try (Transcript t = new Transcript("03-cost-per-endpoint.txt", "Tokens and cost per endpoint (illustrative price: 0.15 USD in / 0.60 USD out per million tokens)")) { + t.line("%-10s %8s %8s %14s", "endpoint", "input", "output", "cost USD"); + double sum = 0; + for (String ep : new String[] { "ask", "summarize", "weather" }) { + double in = Obs.counter(registry, "app.ai.tokens", "endpoint", ep, "type", "input"); + double out = Obs.counter(registry, "app.ai.tokens", "endpoint", ep, "type", "output"); + double cost = Obs.counter(registry, "app.ai.cost.usd", "endpoint", ep); + sum += cost; + t.line("%-10s %8.0f %8.0f %14.8f", ep, in, out, cost); + assertThat(cost).isCloseTo(in * 0.15 / 1e6 + out * 0.60 / 1e6, within(1e-12)); + } + double modelIn = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "input"); + double modelOut = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "output"); + double ourIn = Obs.counter(registry, "app.ai.tokens", "type", "input"); + double ourOut = Obs.counter(registry, "app.ai.tokens", "type", "output"); + t.blank(); + t.line("framework model-level tokens: input=%.0f output=%.0f", modelIn, modelOut); + t.line("our per-endpoint tokens : input=%.0f output=%.0f", ourIn, ourOut); + t.line("total cost USD: %.8f", sum); + assertThat(ourIn).isEqualTo(modelIn); + assertThat(ourOut).isEqualTo(modelOut); + } + // The summarize endpoint sent ~2000 characters and must cost more than the one-word ask. + assertThat(Obs.counter(registry, "app.ai.cost.usd", "endpoint", "summarize")).isGreaterThan(Obs.counter(registry, "app.ai.cost.usd", "endpoint", "ask")); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/ErrorAndRetryTest.java b/observability/src/test/java/com/ankurm/observability/ErrorAndRetryTest.java new file mode 100644 index 0000000..a308d72 --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/ErrorAndRetryTest.java @@ -0,0 +1,54 @@ +package com.ankurm.observability; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.Obs; +import com.ankurm.observability.support.Transcript; +import io.micrometer.core.instrument.MeterRegistry; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.web.client.HttpServerErrorException; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +/** A failing provider: how many HTTP requests one call makes, what the error tag looks like, and what our cost handler does. Writes output/05. */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class ErrorAndRetryTest { + + @DynamicPropertySource + static void fake(DynamicPropertyRegistry r) { + r.add("spring.ai.openai.base-url", Fake.SERVER::url); + } + + @LocalServerPort + int port; + + @Autowired + MeterRegistry registry; + + @Test + void oneFailedCallIsSeveralHttpRequests() { + RestClient http = RestClient.create("http://localhost:" + port); + int before = Fake.SERVER.requestBodies().size(); + assertThatThrownBy(() -> http.get().uri("/ask?q=boom").retrieve().body(String.class)).isInstanceOf(HttpServerErrorException.class); + int attempts = Fake.SERVER.requestBodies().size() - before; + + try (Transcript t = new Transcript("05-error-and-retry.txt", "One failed user request (the provider returns HTTP 500)")) { + t.line("HTTP requests the provider received for that one call: %d", attempts); + t.line("model-call timers recorded:"); + registry.find("gen_ai.client.operation").timers().stream() + .sorted(java.util.Comparator.comparing(x -> x.getId().getTag("error"))) + .forEach(x -> t.line(" error=%-24s response.model=%-24s count=%d", x.getId().getTag("error"), x.getId().getTag("gen_ai.response.model"), x.count())); + t.line("chat-client timers with an error: %d", registry.find("spring.ai.chat.client").tag("error", "InternalServerException").timers().size()); + t.line("app.ai.unpriced{endpoint=ask,reason=no-usage}: %.0f", Obs.counter(registry, "app.ai.unpriced", "endpoint", "ask", "reason", "no-usage")); + t.line("app.ai.cost.usd recorded for the failed call: %.0f", Obs.counter(registry, "app.ai.cost.usd", "endpoint", "ask")); + } + assertThat(attempts).isGreaterThan(1); + assertThat(Obs.timerCount(registry, "gen_ai.client.operation", "error", "InternalServerException")).isEqualTo(1); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/HistogramTest.java b/observability/src/test/java/com/ankurm/observability/HistogramTest.java new file mode 100644 index 0000000..f260a4d --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/HistogramTest.java @@ -0,0 +1,61 @@ +package com.ankurm.observability; + +import java.util.Arrays; +import java.util.List; +import java.util.stream.Stream; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.Transcript; +import org.junit.jupiter.api.Test; +import org.springframework.boot.WebApplicationType; +import org.springframework.boot.builder.SpringApplicationBuilder; +import org.springframework.boot.web.server.context.WebServerApplicationContext; +import org.springframework.context.ConfigurableApplicationContext; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * Timers carry no histogram buckets unless asked, so histogram_quantile() has nothing to work with. + * Starts the app twice, scrapes /actuator/prometheus and counts the bucket series. Writes output/07. + */ +class HistogramTest { + + private static final String METER = "gen_ai_client_operation_seconds"; + + private record Scrape(long buckets, long sumLines, String firstBucket) { + } + + private static Scrape run(String... props) { + String[] args = Stream.concat(Stream.of("spring.ai.openai.base-url=" + Fake.SERVER.url(), "server.port=0"), Arrays.stream(props)) + .map(a -> "--" + a).toArray(String[]::new); + try (ConfigurableApplicationContext ctx = new SpringApplicationBuilder(ObservabilityApplication.class) + .web(WebApplicationType.SERVLET).run(args)) { + int port = ((WebServerApplicationContext) ctx).getWebServer().getPort(); + RestClient http = RestClient.create("http://localhost:" + port); + http.get().uri("/ask?q=hello").retrieve().body(String.class); + http.get().uri("/ask?q=slow one").retrieve().body(String.class); + String body = http.get().uri("/actuator/prometheus").retrieve().body(String.class); + List lines = body.lines().filter(l -> l.startsWith(METER)).toList(); + List buckets = lines.stream().filter(l -> l.startsWith(METER + "_bucket")).toList(); + return new Scrape(buckets.size(), lines.stream().filter(l -> l.startsWith(METER + "_sum")).count(), + buckets.isEmpty() ? "-" : buckets.get(0).replaceAll("\\{.*le=\"([^\"]+)\".*}", "{le=\"$1\"}").replaceAll(" .*", "")); + } + } + + @Test + void bucketsAreOptIn() { + Scrape off = run(); + Scrape on = run("management.metrics.distribution.percentiles-histogram.gen_ai.client.operation=true"); + try (Transcript t = new Transcript("07-histogram.txt", "Bucket series for " + METER + " after 2 calls")) { + t.line("%-70s %s", "setting", "_bucket series / _sum series"); + t.line("%-70s %d / %d", "defaults", off.buckets(), off.sumLines()); + t.line("%-70s %d / %d", "percentiles-histogram.gen_ai.client.operation=true", on.buckets(), on.sumLines()); + t.blank(); + t.line("first bucket series when enabled: %s", on.firstBucket()); + } + assertThat(off.buckets()).isZero(); + assertThat(off.sumLines()).isPositive(); + assertThat(on.buckets()).isGreaterThan(10); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/PromptContentTest.java b/observability/src/test/java/com/ankurm/observability/PromptContentTest.java new file mode 100644 index 0000000..9b5fe12 --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/PromptContentTest.java @@ -0,0 +1,80 @@ +package com.ankurm.observability; + +import java.util.ArrayList; +import java.util.List; + +import ch.qos.logback.classic.Logger; +import ch.qos.logback.classic.spi.ILoggingEvent; +import ch.qos.logback.core.read.ListAppender; +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.TestObs; +import com.ankurm.observability.support.Transcript; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.data.SpanData; +import org.junit.jupiter.api.Test; +import org.slf4j.LoggerFactory; +import org.springframework.boot.WebApplicationType; +import org.springframework.boot.builder.SpringApplicationBuilder; +import org.springframework.context.ConfigurableApplicationContext; + +import static org.assertj.core.api.Assertions.assertThat; + +/** Is the prompt text in your telemetry? By default no; two properties put it there. Writes output/06. */ +class PromptContentTest { + + private static final String SECRET = "ticket-4711-card-ending-0042"; + + private record Result(int logLines, List spanAttributesWithSecret, int spanEventsWithSecret) { + } + + private static Result run(String... props) { + ListAppender appender = new ListAppender<>(); + appender.start(); + Logger root = null; + List all = new ArrayList<>(List.of("spring.ai.openai.base-url=" + Fake.SERVER.url(), "server.port=0")); + all.addAll(List.of(props)); + try (ConfigurableApplicationContext ctx = new SpringApplicationBuilder(ObservabilityApplication.class, TestObs.class) + .web(WebApplicationType.SERVLET).run(all.stream().map(a -> "--" + a).toArray(String[]::new))) { + // Spring Boot re-initialises logging while the context starts and would drop an appender attached earlier. + root = (Logger) LoggerFactory.getLogger("org.springframework.ai"); + root.addAppender(appender); + ctx.getBean(AiController.class).ask(SECRET); + InMemorySpanExporter exporter = ctx.getBean(InMemorySpanExporter.class); + List attrs = new ArrayList<>(); + int events = 0; + for (SpanData s : exporter.getFinishedSpanItems()) { + s.getAttributes().forEach((k, v) -> { + if (String.valueOf(v).contains(SECRET)) { + attrs.add(s.getName() + " -> " + k.getKey()); + } + }); + events += (int) s.getEvents().stream().filter(e -> e.getAttributes().asMap().values().stream().anyMatch(v -> String.valueOf(v).contains(SECRET))).count(); + } + long logs = appender.list.stream().filter(e -> e.getFormattedMessage().contains(SECRET)).count(); + return new Result((int) logs, attrs.stream().sorted().toList(), events); + } + finally { + if (root != null) { + root.detachAppender(appender); + } + } + } + + @Test + void promptTextIsOffByDefault() { + Result off = run(); + Result on = run("spring.ai.chat.observations.log-prompt=true", "spring.ai.chat.client.observations.log-prompt=true"); + + try (Transcript t = new Transcript("06-prompt-content.txt", "Where does the prompt text \"" + SECRET + "\" appear in telemetry?")) { + t.line("%-62s %s", "setting", "log lines with it / span attributes with it / span events with it"); + t.line("%-62s %d / %d / %d", "defaults", off.logLines(), off.spanAttributesWithSecret().size(), off.spanEventsWithSecret()); + t.line("%-62s %d / %d / %d", "spring.ai.chat.observations.log-prompt=true (+ client)", on.logLines(), on.spanAttributesWithSecret().size(), on.spanEventsWithSecret()); + t.line(""); + t.line("span attributes carrying the text when enabled: %s", on.spanAttributesWithSecret()); + } + assertThat(off.logLines()).isZero(); + assertThat(off.spanAttributesWithSecret()).isEmpty(); + assertThat(off.spanEventsWithSecret()).isZero(); + assertThat(on.logLines() + on.spanAttributesWithSecret().size() + on.spanEventsWithSecret()).isPositive(); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/StreamingUsageTest.java b/observability/src/test/java/com/ankurm/observability/StreamingUsageTest.java new file mode 100644 index 0000000..93c2d8d --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/StreamingUsageTest.java @@ -0,0 +1,59 @@ +package com.ankurm.observability; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.Obs; +import com.ankurm.observability.support.Transcript; +import io.micrometer.core.instrument.MeterRegistry; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; + +/** Streaming responses and missing usage: what the request asks for, and what happens when a server leaves usage out. Writes output/04. */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +class StreamingUsageTest { + + @DynamicPropertySource + static void fake(DynamicPropertyRegistry r) { + r.add("spring.ai.openai.base-url", Fake.SERVER::url); + } + + @LocalServerPort + int port; + + @Autowired + MeterRegistry registry; + + @Test + void streamingRequestsUsageAndAMissingUsageBlockIsCountedNotIgnored() { + RestClient http = RestClient.create("http://localhost:" + port); + int before = Fake.SERVER.requestBodies().size(); + http.get().uri("/stream?q=tell me").retrieve().body(String.class); + String streamRequest = Fake.SERVER.requestBodies().get(before); + + double tokensBefore = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "total"); + http.get().uri("/ask?q=nousage please").retrieve().body(String.class); + http.get().uri("/stream?q=nousage please").retrieve().body(String.class); + double tokensAfter = Obs.counter(registry, "gen_ai.client.token.usage", "gen_ai.token.type", "total"); + + boolean asksForUsage = streamRequest.replace(" ", "").contains("\"stream_options\":{\"include_usage\":true"); + try (Transcript t = new Transcript("04-streaming-usage.txt", "Streaming asks for usage; a server that omits it is counted as unpriced")) { + t.line("stream request body asks for usage (stream_options.include_usage): %s", asksForUsage); + t.line("request fragment: %s", streamRequest.replaceAll(".*(\"stream\":true).*", "$1")); + t.line(""); + t.line("two calls to a server that omits the usage block (/ask and /stream):"); + t.line(" gen_ai.client.token.usage total, before %.0f -> after %.0f", tokensBefore, tokensAfter); + t.line(" app.ai.unpriced{reason=no-usage} for ask : %.0f", Obs.counter(registry, "app.ai.unpriced", "endpoint", "ask", "reason", "no-usage")); + t.line(" app.ai.unpriced{reason=no-usage} for stream: %.0f", Obs.counter(registry, "app.ai.unpriced", "endpoint", "stream", "reason", "no-usage")); + } + assertThat(asksForUsage).isTrue(); + assertThat(tokensAfter).isEqualTo(tokensBefore); + assertThat(Obs.counter(registry, "app.ai.unpriced", "endpoint", "ask", "reason", "no-usage")).isEqualTo(1); + assertThat(Obs.counter(registry, "app.ai.unpriced", "endpoint", "stream", "reason", "no-usage")).isEqualTo(1); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/ToolCallSpansTest.java b/observability/src/test/java/com/ankurm/observability/ToolCallSpansTest.java new file mode 100644 index 0000000..4aa302d --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/ToolCallSpansTest.java @@ -0,0 +1,74 @@ +package com.ankurm.observability; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; + +import com.ankurm.observability.support.Fake; +import com.ankurm.observability.support.TestObs; +import com.ankurm.observability.support.Transcript; +import io.opentelemetry.api.common.AttributeKey; +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.data.SpanData; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.boot.test.web.server.LocalServerPort; +import org.springframework.context.annotation.Import; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.web.client.RestClient; + +import static org.assertj.core.api.Assertions.assertThat; + +/** The trace of one tool-calling request, read from an in-memory OpenTelemetry exporter. Writes output/02. */ +@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT) +@Import(TestObs.class) +class ToolCallSpansTest { + + @DynamicPropertySource + static void fake(DynamicPropertyRegistry r) { + r.add("spring.ai.openai.base-url", Fake.SERVER::url); + } + + @LocalServerPort + int port; + + @Autowired + InMemorySpanExporter exporter; + + @Test + void oneToolCallingRequestIsOneTree() { + exporter.reset(); + RestClient.create("http://localhost:" + port).get().uri("/weather?city=Pune").retrieve().body(String.class); + + // The server span ends after the response has been written, so the client can get its answer first. + org.awaitility.Awaitility.await().atMost(java.time.Duration.ofSeconds(5)).until( + () -> exporter.getFinishedSpanItems().stream().anyMatch(s -> s.getName().equals("http get /weather"))); + List spans = new ArrayList<>(exporter.getFinishedSpanItems()); + // Drop the spans that are not part of this request's tree (there are none expected, but be explicit). + spans.sort(Comparator.comparingLong(SpanData::getStartEpochNanos)); + Map byId = spans.stream().collect(Collectors.toMap(s -> s.getSpanId(), s -> s)); + + try (Transcript t = new Transcript("02-tool-call-spans.txt", "Span tree for GET /weather?city=Pune (one tool call, two model calls)")) { + for (SpanData s : spans) { + int depth = 0; + SpanData p = byId.get(s.getParentSpanId()); + while (p != null) { + depth++; + p = byId.get(p.getParentSpanId()); + } + t.line("%s%s [%s]", " ".repeat(depth), s.getName(), s.getKind()); + } + t.blank(); + t.line("attribute keys on the model-call spans:"); + spans.stream().filter(s -> s.getName().startsWith("chat")).findFirst().ifPresent(s -> + s.getAttributes().asMap().keySet().stream().map(AttributeKey::getKey).sorted().forEach(k -> t.line(" %s", k))); + } + assertThat(spans).extracting(SpanData::getName).contains("spring_ai chat_client", "execute_tool getWeather"); + assertThat(spans.stream().filter(s -> s.getName().startsWith("chat ")).count()).isEqualTo(2); + assertThat(spans.stream().map(SpanData::getTraceId).distinct().count()).isEqualTo(1); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/support/Fake.java b/observability/src/test/java/com/ankurm/observability/support/Fake.java new file mode 100644 index 0000000..7acca4f --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/support/Fake.java @@ -0,0 +1,23 @@ +package com.ankurm.observability.support; + +import java.io.IOException; + +import com.ankurm.observability.fake.FakeOpenAiServer; + +/** One fake OpenAI server for the whole test JVM, so every Spring context can point at it. */ +public final class Fake { + + public static final FakeOpenAiServer SERVER = start(); + + private Fake() { + } + + private static FakeOpenAiServer start() { + try { + return new FakeOpenAiServer(0); + } + catch (IOException e) { + throw new IllegalStateException(e); + } + } +} diff --git a/observability/src/test/java/com/ankurm/observability/support/Obs.java b/observability/src/test/java/com/ankurm/observability/support/Obs.java new file mode 100644 index 0000000..acd1d9c --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/support/Obs.java @@ -0,0 +1,47 @@ +package com.ankurm.observability.support; + +import java.util.Comparator; +import java.util.List; +import java.util.Map; +import java.util.TreeMap; +import java.util.stream.Collectors; + +import io.micrometer.core.instrument.Meter; +import io.micrometer.core.instrument.MeterRegistry; +import io.micrometer.core.instrument.Tag; + +/** Small helpers for reading a MeterRegistry in tests. */ +public final class Obs { + + private Obs() { + } + + /** Sum of the counters with this name whose tags include all the given key/value pairs. */ + public static double counter(MeterRegistry registry, String name, String... tagPairs) { + return registry.find(name).tags(tagPairs).counters().stream().mapToDouble(c -> c.count()).sum(); + } + + /** Number of recordings of the timers with this name whose tags include all the given pairs. */ + public static long timerCount(MeterRegistry registry, String name, String... tagPairs) { + return registry.find(name).tags(tagPairs).timers().stream().mapToLong(t -> t.count()).sum(); + } + + /** One line per meter name starting with a prefix: name, type, and the sorted tag keys. */ + public static Map families(MeterRegistry registry, String... prefixes) { + Map out = new TreeMap<>(); + for (Meter m : registry.getMeters()) { + String name = m.getId().getName(); + for (String p : prefixes) { + if (name.startsWith(p)) { + String keys = m.getId().getTags().stream().map(Tag::getKey).sorted().collect(Collectors.joining(", ")); + out.merge(name, m.getId().getType() + " [" + keys + "]", (a, b) -> a.length() >= b.length() ? a : b); + } + } + } + return out; + } + + public static List names(MeterRegistry registry, String prefix) { + return registry.getMeters().stream().map(m -> m.getId().getName()).filter(n -> n.startsWith(prefix)).distinct().sorted(Comparator.naturalOrder()).toList(); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/support/TestObs.java b/observability/src/test/java/com/ankurm/observability/support/TestObs.java new file mode 100644 index 0000000..7e9385e --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/support/TestObs.java @@ -0,0 +1,22 @@ +package com.ankurm.observability.support; + +import io.opentelemetry.sdk.testing.exporter.InMemorySpanExporter; +import io.opentelemetry.sdk.trace.SpanProcessor; +import io.opentelemetry.sdk.trace.export.SimpleSpanProcessor; +import org.springframework.boot.test.context.TestConfiguration; +import org.springframework.context.annotation.Bean; + +/** Captures finished spans in memory (synchronously) so a test can read the span tree. */ +@TestConfiguration +public class TestObs { + + @Bean + InMemorySpanExporter inMemorySpanExporter() { + return InMemorySpanExporter.create(); + } + + @Bean + SpanProcessor simpleSpanProcessor(InMemorySpanExporter exporter) { + return SimpleSpanProcessor.create(exporter); + } +} diff --git a/observability/src/test/java/com/ankurm/observability/support/Transcript.java b/observability/src/test/java/com/ankurm/observability/support/Transcript.java new file mode 100644 index 0000000..c2350cd --- /dev/null +++ b/observability/src/test/java/com/ankurm/observability/support/Transcript.java @@ -0,0 +1,47 @@ +package com.ankurm.observability.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); + } +} diff --git a/observability/stack/grafana/provisioning/dashboards/ai.yml b/observability/stack/grafana/provisioning/dashboards/ai.yml new file mode 100644 index 0000000..ac62c52 --- /dev/null +++ b/observability/stack/grafana/provisioning/dashboards/ai.yml @@ -0,0 +1,6 @@ +apiVersion: 1 +providers: + - name: spring-ai + type: file + options: + path: /tmp/obs-run/dashboards diff --git a/observability/stack/grafana/provisioning/datasources/prom.yml b/observability/stack/grafana/provisioning/datasources/prom.yml new file mode 100644 index 0000000..c0404d8 --- /dev/null +++ b/observability/stack/grafana/provisioning/datasources/prom.yml @@ -0,0 +1,8 @@ +apiVersion: 1 +datasources: + - name: Prometheus + uid: prom + type: prometheus + access: proxy + url: http://127.0.0.1:9090 + isDefault: true diff --git a/observability/stack/prometheus.yml b/observability/stack/prometheus.yml new file mode 100644 index 0000000..28c94d4 --- /dev/null +++ b/observability/stack/prometheus.yml @@ -0,0 +1,7 @@ +global: + scrape_interval: 2s +scrape_configs: + - job_name: spring-ai-observability + metrics_path: /actuator/prometheus + static_configs: + - targets: ["127.0.0.1:8080"]