Add advisors module: custom logging, PII redaction and token-budget advisors
Tests pin down chain ordering (including ties), BaseAdvisor stream behaviour, redaction order versus memory and logging, the tool loop, and how a refusal surfaces on calls, streams and over HTTP. Co-Authored-By: Claude Sonnet 5 <[email protected]> Claude-Session: https://claude.ai/code/session_01Ja4jkzrbQ4LQZBNrb5mkZE
This commit is contained in:
@@ -13,5 +13,6 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur
|
||||
| [`structured-output/`](structured-output) | `ChatClient.entity()` mapping LLM responses to Java records, lists and maps; `StructuredOutputValidationAdvisor` retrying non-conforming JSON with a real enum-constrained schema, including a captured run that exhausts every retry without throwing. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Structured Output in Spring AI 2.0](https://ankurm.com/spring-ai-2-0-structured-output/) |
|
||||
| [`ollama-local/`](ollama-local) | Chat and embeddings against a real local `qwen2.5:0.5b`/`all-minilm`, no API key, driven by a Testcontainers-managed Ollama container started from a baked image; a confirmed model unload via `keep_alive: 0` and `/api/ps`, not a scripted model anywhere. Spring Boot 4.1.1, Spring AI 2.0.1, Testcontainers 2.0.5, Java 25. | [Run LLMs Locally with Spring AI and Ollama](https://ankurm.com/spring-ai-2-0-ollama-local/) |
|
||||
| [`chat-memory/`](chat-memory) | `MessageChatMemoryAdvisor`, `MessageWindowChatMemory`, the JDBC and Redis `ChatMemoryRepository`, per-user conversation IDs and a token-budget memory of our own, with the traps reproduced against a real PostgreSQL 16 and Redis Stack: a 36-character `conversation_id`, tool messages dropped on save, concurrent writers, a 1.x table under the 2.0 repository, and a Redis repository that silently steps aside for a custom `ChatMemory`. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Chat Memory in Spring AI 2.0: JDBC, Redis and Windowed Conversations](https://ankurm.com/spring-ai-2-0-chat-memory-jdbc-redis-windowed-conversations/) |
|
||||
| [`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/) |
|
||||
|
||||
Upgrading from Spring AI 1.x: [migration guide](https://ankurm.com/spring-ai-1-to-2-migration-guide/).
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
target/
|
||||
@@ -0,0 +1,65 @@
|
||||
# advisors
|
||||
|
||||
Companion code for [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/), part of the [Spring AI series](../README.md) on ankurm.com.
|
||||
|
||||
Three advisors of our own -- a logger, a PII redactor and a token budget -- plus the tests that pin down how the advisor chain orders them, what happens on a stream, and what a caller sees when one of them refuses a request.
|
||||
|
||||
There is no live model anywhere. [`RecordingModel`](src/test/java/com/ankurm/advisors/support/RecordingModel.java) records every prompt it is sent and answers with a fixed sentence, which is exactly what an advisor test needs: *what reached the model, and what came back*. No database, no Docker, no API key.
|
||||
|
||||
## Versions
|
||||
|
||||
| Component | Version |
|
||||
|---|---|
|
||||
| Spring Boot | 4.1.1 |
|
||||
| Spring AI | 2.0.1 (`spring-ai-client-chat` 2.0.1) |
|
||||
| Reactor Core | 3.8.7 |
|
||||
| JTokkit (token estimates) | 1.1.0 |
|
||||
| Java | 25 (LTS) |
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
scripts/run-all.sh # runs all 21 tests and regenerates output/01 .. 21
|
||||
```
|
||||
|
||||
Two consecutive runs produce byte-identical files. To run the app itself: `OPENAI_API_KEY=... mvn spring-boot:run`, then `POST /chat` with an `X-User` header and `{"text": "..."}`.
|
||||
|
||||
## What's here
|
||||
|
||||
| File | What it shows |
|
||||
|---|---|
|
||||
| [`advisor/Orders.java`](src/main/java/com/ankurm/advisors/advisor/Orders.java) | The order numbers, and the layers they sit between (memory +200, tool loop +300, model last) |
|
||||
| [`advisor/LoggingAdvisor.java`](src/main/java/com/ankurm/advisors/advisor/LoggingAdvisor.java) | Request / response / stream / failure lines; content off by default |
|
||||
| [`advisor/PiiRedactor.java`](src/main/java/com/ankurm/advisors/advisor/PiiRedactor.java) | Email, Luhn-checked card number and phone patterns; `<EMAIL_1>` placeholders |
|
||||
| [`advisor/PiiRedactionAdvisor.java`](src/main/java/com/ankurm/advisors/advisor/PiiRedactionAdvisor.java) | Redact on the way in, restore on the way out, on calls and on streams (buffers a split placeholder) |
|
||||
| [`advisor/TokenBudgetAdvisor.java`](src/main/java/com/ankurm/advisors/advisor/TokenBudgetAdvisor.java) | Per-request and per-user limits, checked before the model is called |
|
||||
| [`advisor/Texts.java`](src/main/java/com/ankurm/advisors/advisor/Texts.java) | Rewriting request and response text with the `mutate()` copies |
|
||||
| [`config/AdvisorConfig.java`](src/main/java/com/ankurm/advisors/config/AdvisorConfig.java) | The advisors as beans and one `ChatClient` that uses them with Spring AI's memory advisor |
|
||||
| [`web/ChatController.java`](src/main/java/com/ankurm/advisors/web/ChatController.java) | `POST /chat`; the user comes from a header only so the tests need no login |
|
||||
|
||||
## Output files
|
||||
|
||||
Every file is written by the test named in the right column, and every console block in the article is quoted from one of them.
|
||||
|
||||
| File | Written by |
|
||||
|---|---|
|
||||
| `01-chain-order.txt`, `02-chain-order-ties.txt`, `03-chain-contents.txt` | `ChainOrderTest` |
|
||||
| `04-base-advisor-stream.txt` | `BaseAdvisorStreamTest` |
|
||||
| `05-logging-advisor.txt`, `06-logging-failure.txt` | `LoggingAdvisorTest` |
|
||||
| `07-pii-redaction.txt`, `08-pii-limits.txt`, `09-pii-stream-boundary.txt` | `PiiRedactionTest` |
|
||||
| `10-pii-order.txt` | `PiiOrderTest` |
|
||||
| `11-pii-multi-turn.txt` | `PiiMultiTurnTest` |
|
||||
| `12-token-budget.txt`, `13-token-budget-stream.txt`, `14-token-budget-stream-refusal.txt` | `TokenBudgetTest` |
|
||||
| `15-tool-loop-order.txt` | `ToolLoopOrderTest` |
|
||||
| `16-error-propagation.txt` | `ErrorPropagationTest` |
|
||||
| `17-context.txt` | `ContextAndImmutabilityTest` |
|
||||
| `18-unit-test-no-model.txt` | `UnitTestingWithoutModelTest` |
|
||||
| `19-web-429.txt` | `AdvisorWebTest` |
|
||||
| `20-built-in-orders.txt` | `BuiltInOrdersTest` |
|
||||
| `21-tool-loop-budget.txt` | `ToolLoopBudgetTest` |
|
||||
|
||||
`TokenBudgetTest.aStreamIsRefusedAsAnErrorSignal` makes Spring AI's `MessageAggregator` log an `ERROR ... Aggregation Error` stack trace on the console. That is the refusal being reported, not a test failure.
|
||||
|
||||
## Requirements
|
||||
|
||||
JDK 25 and Maven.
|
||||
@@ -0,0 +1,8 @@
|
||||
# Advisor chain: what decides who runs first
|
||||
|
||||
registered as C300, A100, B200 (the number is HIGHEST_PRECEDENCE + n)
|
||||
call: A100> B200> C300> model C300< B200< A100<
|
||||
stream: A100> B200> C300> model C300< B200< A100<
|
||||
|
||||
plus one advisor added on the request with .advisors(...) at n=150:
|
||||
call: A100> R150> B200> C300> model C300< B200< R150< A100<
|
||||
@@ -0,0 +1,6 @@
|
||||
# Two advisors with the same order number
|
||||
|
||||
registered X, Y -> Y> X> model X< Y<
|
||||
registered Y, X -> X> Y> model Y< X<
|
||||
|
||||
For equal numbers the advisor registered LAST runs first (outermost).
|
||||
@@ -0,0 +1,16 @@
|
||||
# What is actually in the chain
|
||||
|
||||
one custom advisor at HIGHEST_PRECEDENCE + 100; the chain a call runs through:
|
||||
Dump HIGHEST_PRECEDENCE + 100
|
||||
Tool Calling Advisor HIGHEST_PRECEDENCE + 300
|
||||
call LOWEST_PRECEDENCE
|
||||
|
||||
and the chain a stream runs through:
|
||||
Dump HIGHEST_PRECEDENCE + 100
|
||||
Tool Calling Advisor HIGHEST_PRECEDENCE + 300
|
||||
stream LOWEST_PRECEDENCE
|
||||
|
||||
the same, with ToolCallingAdvisor.builder().build() added by hand:
|
||||
Dump HIGHEST_PRECEDENCE + 100
|
||||
Tool Calling Advisor HIGHEST_PRECEDENCE + 300
|
||||
call LOWEST_PRECEDENCE
|
||||
@@ -0,0 +1,11 @@
|
||||
# BaseAdvisor on a call and on a stream
|
||||
|
||||
call: before x1, after x1
|
||||
before ran on thread: the caller's thread
|
||||
|
||||
stream: the model streamed 5 chunks: one |two |thre|e fo|ur
|
||||
before x1, after x1
|
||||
after saw only: "ur"
|
||||
before ran on thread: a boundedElastic worker, not the caller's thread
|
||||
|
||||
stream whose last chunk has no finish reason: before x1, after x0
|
||||
@@ -0,0 +1,13 @@
|
||||
# LoggingAdvisor: a call and a stream
|
||||
|
||||
call, content logging off (the default):
|
||||
[LoggingAdvisor] request messages=2 roles=SU
|
||||
[LoggingAdvisor] response chars=14 tokens=12+3 took 5 ms
|
||||
|
||||
stream:
|
||||
[LoggingAdvisor] request messages=2 roles=SU
|
||||
[LoggingAdvisor] complete 3 chunks, 14 chars tokens=12+3 took 5 ms
|
||||
|
||||
call, content logging on:
|
||||
[LoggingAdvisor] request messages=1 roles=U last="Where is my refund?"
|
||||
[LoggingAdvisor] response chars=14 tokens=5+3 text="Refund issued." took 5 ms
|
||||
@@ -0,0 +1,5 @@
|
||||
# LoggingAdvisor when the model call fails
|
||||
|
||||
caller got: IllegalStateException: provider returned 503
|
||||
[LoggingAdvisor] request messages=1 roles=U
|
||||
[LoggingAdvisor] failed IllegalStateException after 5 ms
|
||||
@@ -0,0 +1,8 @@
|
||||
# PiiRedactionAdvisor on a call
|
||||
|
||||
caller sends: Hi, I'm Priya. Email [email protected] or call +91 98765 43210. Card 4111 1111 1111 1111, order 1234 5678 9012 3456, and again [email protected].
|
||||
model was sent: Hi, I'm Priya. Email <EMAIL_1> or call <PHONE_1>. Card <CARD_1>, order 1234 5678 9012 3456, and again <EMAIL_1>.
|
||||
caller receives: You said: Hi, I'm Priya. Email [email protected] or call +91 98765 43210. Card 4111 1111 1111 1111, order 1234 5678 9012 3456, and again [email protected].
|
||||
|
||||
with restore switched off, the caller receives:
|
||||
You said: Hi, I'm Priya. Email <EMAIL_1> or call <PHONE_1>. Card <CARD_1>, order 1234 5678 9012 3456, and again <EMAIL_1>.
|
||||
@@ -0,0 +1,14 @@
|
||||
# What pattern-based redaction misses
|
||||
|
||||
in: My name is Priya Sharma and I live at 14 Hill Road, Bandra, Mumbai 400050.
|
||||
out: My name is Priya Sharma and I live at 14 Hill Road, Bandra, Mumbai 400050.
|
||||
|
||||
in: Passport N1234567, PAN ABCDE1234F.
|
||||
out: Passport N1234567, PAN ABCDE1234F.
|
||||
|
||||
in: Write to priya (at) example (dot) com
|
||||
out: Write to priya (at) example (dot) com
|
||||
|
||||
in: Card 4111-1111-1111-1112 and order 1234 5678 9012 3456
|
||||
out: Card 4111-1111-1111-1112 and order 1234 5678 9012 3456
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
# Restoring placeholders in a stream
|
||||
|
||||
the model streams 5-character chunks: Sure,| I wi|ll wr|ite t|o <EM|AIL_1|> now|.
|
||||
|
||||
restore each chunk on its own:
|
||||
chunks: Sure,| I wi|ll wr|ite t|o <EM|AIL_1|> now|.
|
||||
joined: Sure, I will write to <EMAIL_1> now.
|
||||
|
||||
PiiRedactionAdvisor (holds back from an unfinished "<"):
|
||||
chunks: Sure,| I wi|ll wr|ite t|o |[email protected] now|.
|
||||
joined: Sure, I will write to [email protected] now.
|
||||
@@ -0,0 +1,19 @@
|
||||
# Redaction order versus logging and memory
|
||||
|
||||
memory advisor is fixed at HIGHEST_PRECEDENCE + 200
|
||||
user says: My email is [email protected]
|
||||
|
||||
A redaction +100, logging +400 (redaction outside both):
|
||||
log line saw: "My email is <EMAIL_1>"
|
||||
memory stored: "My email is <EMAIL_1>"
|
||||
model was sent: "My email is <EMAIL_1>"
|
||||
|
||||
B redaction +100, logging +50 (logging outside redaction):
|
||||
log line saw: "My email is [email protected]"
|
||||
memory stored: "My email is <EMAIL_1>"
|
||||
model was sent: "My email is <EMAIL_1>"
|
||||
|
||||
C redaction +300 (inside the memory advisor), logging +400:
|
||||
log line saw: "My email is <EMAIL_1>"
|
||||
memory stored: "My email is [email protected]"
|
||||
model was sent: "My email is <EMAIL_1>"
|
||||
@@ -0,0 +1,18 @@
|
||||
# Placeholders across two turns with memory
|
||||
|
||||
turn 1: My email is [email protected]
|
||||
turn 2: Also cc [email protected]
|
||||
|
||||
numbering restarts on every request:
|
||||
model was sent on turn 2: U:My email is <EMAIL_1> | A:You said: My email is <EMAIL_1> | U:Also cc <EMAIL_1>
|
||||
caller receives: You said: Also cc [email protected]
|
||||
|
||||
numbering kept per conversation (the default):
|
||||
model was sent on turn 2: U:My email is <EMAIL_1> | A:You said: My email is <EMAIL_1> | U:Also cc <EMAIL_2>
|
||||
caller receives: You said: Also cc [email protected]
|
||||
|
||||
what the memory stores (placeholders, never the addresses):
|
||||
USER My email is <EMAIL_1>
|
||||
ASSISTANT You said: My email is <EMAIL_1>
|
||||
USER Also cc <EMAIL_2>
|
||||
ASSISTANT You said: Also cc <EMAIL_2>
|
||||
@@ -0,0 +1,15 @@
|
||||
# TokenBudgetAdvisor: 40 tokens per request, 60 per user
|
||||
|
||||
alice asks short questions; "spent" is the usage the model reported:
|
||||
call 1: answered, spent=13, model calls=1
|
||||
call 2: answered, spent=26, model calls=2
|
||||
call 3: answered, spent=39, model calls=3
|
||||
call 4: answered, spent=52, model calls=4
|
||||
call 5: answered, spent=65, model calls=5
|
||||
call 6: refused (user alice has used 65 of 60 tokens), spent=65, model calls=5
|
||||
|
||||
bob pastes a stack trace of 103 estimated tokens:
|
||||
refused: prompt is about 103 tokens, the limit per request is 40
|
||||
model calls: 5 (was 5), bob's spent: 0
|
||||
|
||||
bob then asks a short question: answered, bob spent=13, alice spent=65
|
||||
@@ -0,0 +1,9 @@
|
||||
# Token accounting on a stream
|
||||
|
||||
same question, same answer:
|
||||
call, usage from the response: 9 tokens
|
||||
stream, usage on the last chunk: 9 tokens
|
||||
stream, no usage reported (estimated): 9 tokens
|
||||
|
||||
A stream that reports no usage is billed by the provider all the same. The estimate above
|
||||
matches only because this scripted model and the advisor use the same tokenizer.
|
||||
@@ -0,0 +1,5 @@
|
||||
# A refusal on the stream path
|
||||
|
||||
subscriber got: TokenBudgetExceededException
|
||||
message: prompt is about 13 tokens, the limit per request is 5
|
||||
model calls: 0
|
||||
@@ -0,0 +1,13 @@
|
||||
# Where an advisor sits relative to the tool loop
|
||||
|
||||
one question, one tool call, so the model is called twice (ToolCallingAdvisor is at +300)
|
||||
|
||||
logging advisor at +250 (outside the tool loop):
|
||||
[LoggingAdvisor] request messages=1 roles=U
|
||||
[LoggingAdvisor] response chars=20 took 5 ms
|
||||
|
||||
logging advisor at +400 (inside the tool loop):
|
||||
[LoggingAdvisor] request messages=1 roles=U
|
||||
[LoggingAdvisor] response chars=0 took 5 ms
|
||||
[LoggingAdvisor] request messages=3 roles=UAT
|
||||
[LoggingAdvisor] response chars=20 took 5 ms
|
||||
@@ -0,0 +1,11 @@
|
||||
# Where a refusal surfaces
|
||||
|
||||
an advisor that throws IllegalStateException("refused"), model calls counted afterwards
|
||||
|
||||
plain advisor, call failed while running it with IllegalStateException: refused
|
||||
plain advisor, stream, throws eagerly failed while running it with IllegalStateException: refused
|
||||
plain advisor, stream, error inside Flux.defer failed while running it with IllegalStateException: refused
|
||||
BaseAdvisor.before, call failed while running it with IllegalStateException: refused
|
||||
BaseAdvisor.before, stream failed while running it with IllegalStateException: Stream processing failed (cause: IllegalStateException: refused)
|
||||
|
||||
model calls: 0
|
||||
@@ -0,0 +1,7 @@
|
||||
# Advisor context: request down, response up
|
||||
|
||||
outer: original request has tenant=null, the copy it forwards has tenant=acme
|
||||
inner: request context has tenant=acme
|
||||
outer: response context has verdict=clean, tenant=acme
|
||||
|
||||
req.context().put(...) on the incoming request: allowed
|
||||
@@ -0,0 +1,6 @@
|
||||
# Testing an advisor with a stub chain
|
||||
|
||||
what the caller sent: Mail [email protected] please
|
||||
what reached the rest: Mail <EMAIL_1> please
|
||||
what the caller got back: Noted: Mail [email protected] please
|
||||
caller's request unchanged: true
|
||||
@@ -0,0 +1,9 @@
|
||||
# Over HTTP: 200, then 429
|
||||
|
||||
limits for this run: 60 tokens per request, 30 per user
|
||||
dana request 1 -> HTTP 200, model calls so far: 1
|
||||
dana request 2 -> HTTP 200, model calls so far: 2
|
||||
dana request 3 -> HTTP 429, model calls so far: 2
|
||||
dana request 4 -> HTTP 429, model calls so far: 2
|
||||
|
||||
erin pastes a long stack trace -> HTTP 429, model calls: 2 (was 2)
|
||||
@@ -0,0 +1,8 @@
|
||||
# Built-in advisor orders in 2.0.1
|
||||
|
||||
SimpleLoggerAdvisor getOrder() = 0
|
||||
SafeGuardAdvisor getOrder() = 0
|
||||
MessageChatMemoryAdvisor getOrder() = -2147483448 (HIGHEST_PRECEDENCE + 200)
|
||||
Tool Calling Advisor getOrder() = -2147483348 (HIGHEST_PRECEDENCE + 300)
|
||||
|
||||
Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER = HIGHEST_PRECEDENCE + 200
|
||||
@@ -0,0 +1,7 @@
|
||||
# Token budget and the tool loop
|
||||
|
||||
the model reports 100+10 tokens for round 1 (asks for the tool) and 130+20 for round 2 (answers)
|
||||
so the provider would bill 260 tokens for this one question
|
||||
|
||||
budget advisor at +250 (outside the tool loop): recorded 260
|
||||
budget advisor at +400 (inside the tool loop): recorded 260
|
||||
@@ -0,0 +1,69 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project xmlns="http://maven.apache.org/POM/4.0.0"
|
||||
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
|
||||
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
|
||||
<modelVersion>4.0.0</modelVersion>
|
||||
|
||||
<parent>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-parent</artifactId>
|
||||
<version>4.1.1</version>
|
||||
<relativePath/>
|
||||
</parent>
|
||||
|
||||
<groupId>com.ankurm</groupId>
|
||||
<artifactId>advisors</artifactId>
|
||||
<version>1.0.0</version>
|
||||
<name>advisors</name>
|
||||
<description>Custom Spring AI 2.0 advisors: audit logging, PII redaction and token budgets, with the ordering and streaming traps reproduced offline.</description>
|
||||
|
||||
<properties>
|
||||
<java.version>25</java.version>
|
||||
<spring-ai.version>2.0.1</spring-ai.version>
|
||||
</properties>
|
||||
|
||||
<dependencyManagement>
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-bom</artifactId>
|
||||
<version>${spring-ai.version}</version>
|
||||
<type>pom</type>
|
||||
<scope>import</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
</dependencyManagement>
|
||||
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-webmvc</artifactId>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-starter-model-openai</artifactId>
|
||||
</dependency>
|
||||
|
||||
<dependency>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-test</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
|
||||
<build>
|
||||
<plugins>
|
||||
<plugin>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-maven-plugin</artifactId>
|
||||
</plugin>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-surefire-plugin</artifactId>
|
||||
<configuration>
|
||||
<argLine>-Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8</argLine>
|
||||
</configuration>
|
||||
</plugin>
|
||||
</plugins>
|
||||
</build>
|
||||
</project>
|
||||
Executable
+9
@@ -0,0 +1,9 @@
|
||||
#!/usr/bin/env bash
|
||||
# Regenerates every file under output/ (01-21). The test suite writes all of them itself through
|
||||
# the Transcript helper. No database, no Docker and no API key is needed: the "model" is a
|
||||
# recording stub, so what the tests prove is what reached it and what came back.
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.."
|
||||
rm -rf target
|
||||
mvn -q -B test 2>&1 | grep -E "Tests run:|BUILD|FAIL" || true
|
||||
ls output
|
||||
@@ -0,0 +1,12 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import org.springframework.boot.SpringApplication;
|
||||
import org.springframework.boot.autoconfigure.SpringBootApplication;
|
||||
|
||||
@SpringBootApplication
|
||||
public class AdvisorsApplication {
|
||||
|
||||
public static void main(String[] args) {
|
||||
SpringApplication.run(AdvisorsApplication.class, args);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.function.Consumer;
|
||||
import java.util.function.LongSupplier;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.metadata.Usage;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* One log line going in, one coming out, with latency and token usage. It wraps {@code
|
||||
* chain.nextCall} / {@code nextStream}, so "took" is the time of everything <em>inside</em> this
|
||||
* advisor in the chain: the model, plus any advisor with a higher order number.
|
||||
*
|
||||
* <p>Message text is off by default. When it is on, what you read is whatever reaches <em>this</em>
|
||||
* position in the chain, which is the point of the ordering experiments in the article.
|
||||
*/
|
||||
public final class LoggingAdvisor implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
private static final Logger LOG = LoggerFactory.getLogger(LoggingAdvisor.class);
|
||||
|
||||
private final String name;
|
||||
|
||||
private final int order;
|
||||
|
||||
private final boolean logContent;
|
||||
|
||||
private final Consumer<String> sink;
|
||||
|
||||
private final LongSupplier nanoClock;
|
||||
|
||||
public LoggingAdvisor(String name, int order, boolean logContent, Consumer<String> sink, LongSupplier nanoClock) {
|
||||
this.name = name;
|
||||
this.order = order;
|
||||
this.logContent = logContent;
|
||||
this.sink = sink;
|
||||
this.nanoClock = nanoClock;
|
||||
}
|
||||
|
||||
public static LoggingAdvisor toSlf4j(int order, boolean logContent) {
|
||||
return new LoggingAdvisor("LoggingAdvisor", order, logContent, LOG::info, System::nanoTime);
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return name;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return order;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
long start = nanoClock.getAsLong();
|
||||
sink.accept("[" + name + "] request " + describe(request));
|
||||
try {
|
||||
ChatClientResponse response = chain.nextCall(request);
|
||||
sink.accept("[" + name + "] response " + describe(response) + " took " + millis(start) + " ms");
|
||||
return response;
|
||||
}
|
||||
catch (RuntimeException e) {
|
||||
sink.accept("[" + name + "] failed " + e.getClass().getSimpleName() + " after " + millis(start) + " ms");
|
||||
throw e;
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
return Flux.defer(() -> {
|
||||
long start = nanoClock.getAsLong();
|
||||
sink.accept("[" + name + "] request " + describe(request));
|
||||
AtomicInteger chunks = new AtomicInteger();
|
||||
AtomicInteger chars = new AtomicInteger();
|
||||
AtomicReference<Usage> usage = new AtomicReference<>();
|
||||
return chain.nextStream(request).doOnNext(chunk -> {
|
||||
chunks.incrementAndGet();
|
||||
chars.addAndGet(Texts.text(chunk).length());
|
||||
Usage u = usageOf(chunk);
|
||||
if (u != null) {
|
||||
usage.set(u);
|
||||
}
|
||||
}).doOnComplete(() -> sink.accept("[" + name + "] complete " + chunks + " chunks, " + chars + " chars"
|
||||
+ usageText(usage.get()) + " took " + millis(start) + " ms")).doOnError(e -> sink
|
||||
.accept("[" + name + "] failed " + e.getClass().getSimpleName() + " after " + millis(start) + " ms"));
|
||||
});
|
||||
}
|
||||
|
||||
private String describe(ChatClientRequest request) {
|
||||
List<Message> messages = request.prompt().getInstructions();
|
||||
StringBuilder roles = new StringBuilder();
|
||||
for (Message m : messages) {
|
||||
roles.append(m.getMessageType().name().charAt(0));
|
||||
}
|
||||
String text = "messages=" + messages.size() + " roles=" + roles;
|
||||
if (logContent) {
|
||||
text += " last=\"" + messages.getLast().getText() + "\"";
|
||||
}
|
||||
return text;
|
||||
}
|
||||
|
||||
private String describe(ChatClientResponse response) {
|
||||
String text = "chars=" + Texts.text(response).length() + usageText(usageOf(response));
|
||||
if (logContent) {
|
||||
text += " text=\"" + Texts.text(response) + "\"";
|
||||
}
|
||||
return text;
|
||||
}
|
||||
|
||||
private static Usage usageOf(ChatClientResponse response) {
|
||||
if (response.chatResponse() == null || response.chatResponse().getMetadata() == null) {
|
||||
return null;
|
||||
}
|
||||
Usage u = response.chatResponse().getMetadata().getUsage();
|
||||
return u != null && u.getTotalTokens() != null && u.getTotalTokens() > 0 ? u : null;
|
||||
}
|
||||
|
||||
private static String usageText(Usage u) {
|
||||
return u == null ? "" : " tokens=" + u.getPromptTokens() + "+" + u.getCompletionTokens();
|
||||
}
|
||||
|
||||
private long millis(long startNanos) {
|
||||
return (nanoClock.getAsLong() - startNanos) / 1_000_000;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/**
|
||||
* Where the three custom advisors sit, written down once. A <em>lower</em> number runs earlier
|
||||
* on the way in (and later on the way out), so it sits further out.
|
||||
*
|
||||
* <pre>
|
||||
* HIGHEST + 100 PiiRedactionAdvisor (this module) nothing outside it sees raw text
|
||||
* HIGHEST + 200 MessageChatMemoryAdvisor (Spring AI) memory stores the redacted text
|
||||
* HIGHEST + 250 TokenBudgetAdvisor (this module) counts history + question
|
||||
* HIGHEST + 300 ToolCallingAdvisor (Spring AI) the tool loop
|
||||
* HIGHEST + 400 LoggingAdvisor (this module) logs every model round trip
|
||||
* LOWEST ChatModelCallAdvisor / ChatModelStreamAdvisor (Spring AI) the model itself
|
||||
* </pre>
|
||||
*/
|
||||
public final class Orders {
|
||||
|
||||
public static final int PII_REDACTION = Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
|
||||
public static final int TOKEN_BUDGET = Ordered.HIGHEST_PRECEDENCE + 250;
|
||||
|
||||
public static final int LOGGING = Ordered.HIGHEST_PRECEDENCE + 400;
|
||||
|
||||
private Orders() {
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import reactor.core.publisher.Flux;
|
||||
import reactor.core.publisher.Mono;
|
||||
|
||||
/**
|
||||
* Redacts PII from everything sent to the model and, optionally, puts the originals back into the
|
||||
* answer the caller sees. It implements {@link CallAdvisor} and {@link StreamAdvisor} directly
|
||||
* rather than {@code BaseAdvisor}: {@code BaseAdvisor.after} is only invoked for the last chunk of
|
||||
* a stream, which is no use for restoring text that arrives in pieces.
|
||||
*/
|
||||
public final class PiiRedactionAdvisor implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
private final int order;
|
||||
|
||||
private final boolean restore;
|
||||
|
||||
private final boolean conversationScoped;
|
||||
|
||||
private final ConcurrentHashMap<String, PiiRedactor.Session> vault = new ConcurrentHashMap<>();
|
||||
|
||||
/** Conversation-scoped placeholders (the default): the same value keeps the same placeholder across turns. */
|
||||
public PiiRedactionAdvisor(int order, boolean restore) {
|
||||
this(order, restore, true);
|
||||
}
|
||||
|
||||
/**
|
||||
* @param conversationScoped {@code false} starts a fresh numbering on every request, which is
|
||||
* simpler and wrong as soon as memory replays an earlier turn ("<EMAIL_1>" then means two people)
|
||||
*/
|
||||
public PiiRedactionAdvisor(int order, boolean restore, boolean conversationScoped) {
|
||||
this.order = order;
|
||||
this.restore = restore;
|
||||
this.conversationScoped = conversationScoped;
|
||||
}
|
||||
|
||||
/** Drops the stored originals for a conversation. The vault holds real PII in memory: bound it. */
|
||||
public void forget(String conversationId) {
|
||||
vault.remove(conversationId);
|
||||
}
|
||||
|
||||
private PiiRedactor.Session sessionFor(ChatClientRequest request) {
|
||||
Object id = request.context().get(ChatMemory.CONVERSATION_ID);
|
||||
if (!conversationScoped || id == null) {
|
||||
return new PiiRedactor.Session();
|
||||
}
|
||||
return vault.computeIfAbsent(id.toString(), k -> new PiiRedactor.Session());
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "PiiRedactionAdvisor";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return order;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
PiiRedactor.Session session = sessionFor(request);
|
||||
ChatClientResponse response = chain.nextCall(Texts.mapRequestText(request, session::redact));
|
||||
return restore ? Texts.mapResponseText(response, session::restore) : response;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
return Flux.defer(() -> {
|
||||
PiiRedactor.Session session = sessionFor(request);
|
||||
Flux<ChatClientResponse> upstream = chain.nextStream(Texts.mapRequestText(request, session::redact));
|
||||
if (!restore) {
|
||||
return upstream;
|
||||
}
|
||||
StreamRestorer restorer = new StreamRestorer(session);
|
||||
return upstream.map(restorer::onChunk).concatWith(Mono.defer(restorer::flush).flux());
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* A placeholder such as {@code <EMAIL_1>} can be split across chunks ("<EMA" then "IL_1>"),
|
||||
* and restoring each chunk on its own would miss it. So text from the last {@code <} that has
|
||||
* no {@code >} yet is held back until the next chunk (or the end of the stream) completes it.
|
||||
*/
|
||||
static final class StreamRestorer {
|
||||
|
||||
private static final int LONGEST_PLACEHOLDER = 16;
|
||||
|
||||
private final PiiRedactor.Session session;
|
||||
|
||||
private final StringBuilder pending = new StringBuilder();
|
||||
|
||||
private ChatClientResponse last;
|
||||
|
||||
StreamRestorer(PiiRedactor.Session session) {
|
||||
this.session = session;
|
||||
}
|
||||
|
||||
ChatClientResponse onChunk(ChatClientResponse chunk) {
|
||||
String text = Texts.text(chunk);
|
||||
last = chunk;
|
||||
if (text.isEmpty()) {
|
||||
return chunk;
|
||||
}
|
||||
pending.append(text);
|
||||
int cut = safeCut();
|
||||
String emit = session.restore(pending.substring(0, cut));
|
||||
pending.delete(0, cut);
|
||||
return Texts.mapResponseText(chunk, t -> emit);
|
||||
}
|
||||
|
||||
Mono<ChatClientResponse> flush() {
|
||||
if (pending.isEmpty() || last == null) {
|
||||
return Mono.empty();
|
||||
}
|
||||
String tail = session.restore(pending.toString());
|
||||
pending.setLength(0);
|
||||
ChatResponse tailResponse = ChatResponse.builder()
|
||||
.generations(java.util.List.of(new Generation(new AssistantMessage(tail))))
|
||||
.build();
|
||||
return Mono.just(ChatClientResponse.builder().chatResponse(tailResponse).context(last.context()).build());
|
||||
}
|
||||
|
||||
private int safeCut() {
|
||||
int open = pending.lastIndexOf("<");
|
||||
if (open >= 0 && pending.indexOf(">", open) < 0 && pending.length() - open <= LONGEST_PLACEHOLDER) {
|
||||
return open;
|
||||
}
|
||||
return pending.length();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import java.util.HashMap;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.Map;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/**
|
||||
* Finds emails, card numbers and phone numbers in free text and swaps each for a numbered
|
||||
* placeholder such as {@code <EMAIL_1>}. The same value always gets the same placeholder within
|
||||
* one {@link Session}, so the model can still tell "the same address twice" from "two addresses",
|
||||
* and {@link Session#restore} can put the originals back into the answer.
|
||||
*
|
||||
* <p>This is pattern matching, not understanding: it finds what looks like an email, a Luhn-valid
|
||||
* card number and a 10-digit or NANP-style phone number. It does not find names, street addresses
|
||||
* or an ID number in a shape it does not know. Treat it as a floor, not as a guarantee.
|
||||
*/
|
||||
public final class PiiRedactor {
|
||||
|
||||
private static final Pattern EMAIL = Pattern.compile("[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\\.[A-Za-z]{2,}");
|
||||
|
||||
private static final Pattern CARD = Pattern.compile("(?<!\\d)(?:\\d[ -]?){12,18}\\d(?!\\d)");
|
||||
|
||||
private static final Pattern PHONE = Pattern
|
||||
.compile("(?<![\\w+])(?:\\+\\d{1,3}[ -]?)?(?:\\d{5}[ -]?\\d{5}|\\(?\\d{3}\\)?[ -]\\d{3}-\\d{4})(?!\\w)");
|
||||
|
||||
private PiiRedactor() {
|
||||
}
|
||||
|
||||
/**
|
||||
* Which original value became which placeholder. Thread-safe. One session can span a whole
|
||||
* conversation, so "<EMAIL_1>" means the same address in turn 1 and in turn 5.
|
||||
*/
|
||||
public static final class Session {
|
||||
|
||||
private final Map<String, String> placeholderByValue = new LinkedHashMap<>();
|
||||
|
||||
private final Map<String, String> valueByPlaceholder = new LinkedHashMap<>();
|
||||
|
||||
private final Map<String, Integer> counters = new HashMap<>();
|
||||
|
||||
public synchronized String redact(String text) {
|
||||
if (text == null || text.isEmpty()) {
|
||||
return text;
|
||||
}
|
||||
String out = replace(EMAIL, text, "EMAIL", false);
|
||||
out = replace(CARD, out, "CARD", true);
|
||||
return replace(PHONE, out, "PHONE", false);
|
||||
}
|
||||
|
||||
public synchronized String restore(String text) {
|
||||
if (text == null || valueByPlaceholder.isEmpty()) {
|
||||
return text;
|
||||
}
|
||||
String out = text;
|
||||
for (Map.Entry<String, String> e : valueByPlaceholder.entrySet()) {
|
||||
out = out.replace(e.getKey(), e.getValue());
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
public synchronized boolean isEmpty() {
|
||||
return valueByPlaceholder.isEmpty();
|
||||
}
|
||||
|
||||
/** Placeholder to original value. Never log this: it is the PII. */
|
||||
public synchronized Map<String, String> mapping() {
|
||||
return Map.copyOf(valueByPlaceholder);
|
||||
}
|
||||
|
||||
private String replace(Pattern pattern, String text, String kind, boolean luhn) {
|
||||
Matcher m = pattern.matcher(text);
|
||||
StringBuilder sb = new StringBuilder();
|
||||
while (m.find()) {
|
||||
String value = m.group();
|
||||
if (luhn && !luhnValid(value)) {
|
||||
m.appendReplacement(sb, Matcher.quoteReplacement(value));
|
||||
continue;
|
||||
}
|
||||
String placeholder = placeholderByValue.computeIfAbsent(kind + ":" + value, k -> {
|
||||
String p = "<" + kind + "_" + counters.merge(kind, 1, Integer::sum) + ">";
|
||||
valueByPlaceholder.put(p, value);
|
||||
return p;
|
||||
});
|
||||
m.appendReplacement(sb, Matcher.quoteReplacement(placeholder));
|
||||
}
|
||||
m.appendTail(sb);
|
||||
return sb.toString();
|
||||
}
|
||||
}
|
||||
|
||||
static boolean luhnValid(String candidate) {
|
||||
String digits = candidate.replaceAll("\\D", "");
|
||||
if (digits.length() < 13 || digits.length() > 19) {
|
||||
return false;
|
||||
}
|
||||
int sum = 0;
|
||||
boolean dbl = false;
|
||||
for (int i = digits.length() - 1; i >= 0; i--) {
|
||||
int d = digits.charAt(i) - '0';
|
||||
if (dbl) {
|
||||
d *= 2;
|
||||
if (d > 9) {
|
||||
d -= 9;
|
||||
}
|
||||
}
|
||||
sum += d;
|
||||
dbl = !dbl;
|
||||
}
|
||||
return sum % 10 == 0;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.function.UnaryOperator;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.SystemMessage;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
|
||||
/**
|
||||
* Two rewriting helpers shared by the advisors. Everything in a {@code ChatClientRequest} and a
|
||||
* {@code ChatClientResponse} is immutable, so "changing the prompt" always means building a new one;
|
||||
* these keep the parts an advisor should not touch (options, metadata, media, tool calls).
|
||||
*/
|
||||
public final class Texts {
|
||||
|
||||
private Texts() {
|
||||
}
|
||||
|
||||
/** Applies {@code f} to the text of every system, user and assistant message in the prompt. */
|
||||
public static ChatClientRequest mapRequestText(ChatClientRequest request, UnaryOperator<String> f) {
|
||||
List<Message> rewritten = new ArrayList<>();
|
||||
for (Message m : request.prompt().getInstructions()) {
|
||||
rewritten.add(switch (m) {
|
||||
case UserMessage u -> u.mutate().text(f.apply(u.getText())).build();
|
||||
case SystemMessage s -> s.mutate().text(f.apply(s.getText())).build();
|
||||
case AssistantMessage a when !a.hasToolCalls() -> a.mutate().content(f.apply(a.getText())).build();
|
||||
default -> m;
|
||||
});
|
||||
}
|
||||
return request.mutate().prompt(request.prompt().mutate().messages(rewritten).build()).build();
|
||||
}
|
||||
|
||||
/** Applies {@code f} to the text of every generation, keeping metadata and usage. */
|
||||
public static ChatClientResponse mapResponseText(ChatClientResponse response, UnaryOperator<String> f) {
|
||||
ChatResponse original = response.chatResponse();
|
||||
if (original == null) {
|
||||
return response;
|
||||
}
|
||||
List<Generation> generations = new ArrayList<>();
|
||||
for (Generation g : original.getResults()) {
|
||||
AssistantMessage out = g.getOutput();
|
||||
String text = out.getText();
|
||||
if (text == null || out.hasToolCalls()) {
|
||||
generations.add(g);
|
||||
}
|
||||
else {
|
||||
generations.add(new Generation(out.mutate().content(f.apply(text)).build(), g.getMetadata()));
|
||||
}
|
||||
}
|
||||
ChatResponse rewritten = ChatResponse.builder().from(original).generations(generations).build();
|
||||
return response.mutate().chatResponse(rewritten).build();
|
||||
}
|
||||
|
||||
/** The text of the first generation, or empty. */
|
||||
public static String text(ChatClientResponse response) {
|
||||
ChatResponse r = response.chatResponse();
|
||||
if (r == null || r.getResult() == null || r.getResult().getOutput().getText() == null) {
|
||||
return "";
|
||||
}
|
||||
return r.getResult().getOutput().getText();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
import java.util.concurrent.atomic.LongAdder;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.metadata.Usage;
|
||||
import org.springframework.ai.tokenizer.TokenCountEstimator;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* Two limits, both checked before the model is called:
|
||||
* <ul>
|
||||
* <li>a per-request cap on the <em>estimated</em> prompt size, so one pasted log file cannot cost
|
||||
* a fortune;</li>
|
||||
* <li>a running per-user total built from the token usage the provider reports on each response,
|
||||
* so a user who has used their allowance is refused without a model call.</li>
|
||||
* </ul>
|
||||
* The user comes from the request context (parameter {@link #USER_KEY}); the state is in memory,
|
||||
* so it is per instance and resets on restart.
|
||||
*/
|
||||
public final class TokenBudgetAdvisor implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
public static final String USER_KEY = "budget_user";
|
||||
|
||||
private final TokenCountEstimator estimator;
|
||||
|
||||
private final int maxPromptTokens;
|
||||
|
||||
private final long maxTokensPerUser;
|
||||
|
||||
private final int order;
|
||||
|
||||
private final ConcurrentHashMap<String, LongAdder> spent = new ConcurrentHashMap<>();
|
||||
|
||||
public TokenBudgetAdvisor(TokenCountEstimator estimator, int maxPromptTokens, long maxTokensPerUser, int order) {
|
||||
this.estimator = estimator;
|
||||
this.maxPromptTokens = maxPromptTokens;
|
||||
this.maxTokensPerUser = maxTokensPerUser;
|
||||
this.order = order;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "TokenBudgetAdvisor";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return order;
|
||||
}
|
||||
|
||||
public long spent(String user) {
|
||||
LongAdder a = spent.get(user);
|
||||
return a == null ? 0 : a.sum();
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
String user = userOf(request);
|
||||
check(user, request);
|
||||
ChatClientResponse response = chain.nextCall(request);
|
||||
Usage usage = usageOf(response);
|
||||
record(user, usage != null ? usage.getTotalTokens() : estimate(request) + estimator.estimate(Texts.text(response)));
|
||||
return response;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
return Flux.defer(() -> {
|
||||
String user = userOf(request);
|
||||
check(user, request);
|
||||
AtomicLong reported = new AtomicLong(-1);
|
||||
StringBuilder streamed = new StringBuilder();
|
||||
return chain.nextStream(request).doOnNext(chunk -> {
|
||||
Usage usage = usageOf(chunk);
|
||||
if (usage != null) {
|
||||
reported.set(usage.getTotalTokens());
|
||||
}
|
||||
streamed.append(Texts.text(chunk));
|
||||
}).doOnComplete(() -> record(user,
|
||||
reported.get() >= 0 ? reported.get() : estimate(request) + estimator.estimate(streamed.toString())));
|
||||
});
|
||||
}
|
||||
|
||||
private void check(String user, ChatClientRequest request) {
|
||||
int prompt = estimate(request);
|
||||
if (prompt > maxPromptTokens) {
|
||||
throw new TokenBudgetExceededException(
|
||||
"prompt is about " + prompt + " tokens, the limit per request is " + maxPromptTokens);
|
||||
}
|
||||
if (spent(user) >= maxTokensPerUser) {
|
||||
throw new TokenBudgetExceededException(
|
||||
"user " + user + " has used " + spent(user) + " of " + maxTokensPerUser + " tokens");
|
||||
}
|
||||
}
|
||||
|
||||
private int estimate(ChatClientRequest request) {
|
||||
int total = 0;
|
||||
for (Message m : request.prompt().getInstructions()) {
|
||||
total += estimator.estimate(m.getText());
|
||||
}
|
||||
return total;
|
||||
}
|
||||
|
||||
private void record(String user, long tokens) {
|
||||
spent.computeIfAbsent(user, k -> new LongAdder()).add(tokens);
|
||||
}
|
||||
|
||||
private static String userOf(ChatClientRequest request) {
|
||||
Object u = request.context().get(USER_KEY);
|
||||
return u == null ? "anonymous" : u.toString();
|
||||
}
|
||||
|
||||
private static Usage usageOf(ChatClientResponse response) {
|
||||
if (response.chatResponse() == null || response.chatResponse().getMetadata() == null) {
|
||||
return null;
|
||||
}
|
||||
Usage u = response.chatResponse().getMetadata().getUsage();
|
||||
return u != null && u.getTotalTokens() != null && u.getTotalTokens() > 0 ? u : null;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
package com.ankurm.advisors.advisor;
|
||||
|
||||
import org.springframework.http.HttpStatus;
|
||||
import org.springframework.web.bind.annotation.ResponseStatus;
|
||||
|
||||
/** Thrown before the model is called, so the request costs nothing. Maps to HTTP 429 in the web layer. */
|
||||
@ResponseStatus(HttpStatus.TOO_MANY_REQUESTS)
|
||||
public class TokenBudgetExceededException extends RuntimeException {
|
||||
|
||||
public TokenBudgetExceededException(String message) {
|
||||
super(message);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.advisors.config;
|
||||
|
||||
import com.ankurm.advisors.advisor.LoggingAdvisor;
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.advisor.PiiRedactionAdvisor;
|
||||
import com.ankurm.advisors.advisor.TokenBudgetAdvisor;
|
||||
import com.knuddels.jtokkit.api.EncodingType;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository;
|
||||
import org.springframework.ai.chat.memory.MessageWindowChatMemory;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator;
|
||||
import org.springframework.ai.tokenizer.TokenCountEstimator;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.context.annotation.Configuration;
|
||||
|
||||
/**
|
||||
* The advisors as beans, and one {@link ChatClient} with all of them plus Spring AI's own memory
|
||||
* advisor. The order numbers are what decide the chain, not the order they are listed in
|
||||
* {@code defaultAdvisors(...)}; see {@link Orders}.
|
||||
*/
|
||||
@Configuration
|
||||
public class AdvisorConfig {
|
||||
|
||||
@Bean
|
||||
TokenCountEstimator tokenCountEstimator() {
|
||||
return new JTokkitTokenCountEstimator(EncodingType.O200K_BASE);
|
||||
}
|
||||
|
||||
@Bean
|
||||
PiiRedactionAdvisor piiRedactionAdvisor() {
|
||||
return new PiiRedactionAdvisor(Orders.PII_REDACTION, true);
|
||||
}
|
||||
|
||||
@Bean
|
||||
TokenBudgetAdvisor tokenBudgetAdvisor(TokenCountEstimator estimator,
|
||||
@Value("${app.budget.max-prompt-tokens:2000}") int maxPromptTokens,
|
||||
@Value("${app.budget.max-tokens-per-user:20000}") long maxTokensPerUser) {
|
||||
return new TokenBudgetAdvisor(estimator, maxPromptTokens, maxTokensPerUser, Orders.TOKEN_BUDGET);
|
||||
}
|
||||
|
||||
@Bean
|
||||
LoggingAdvisor loggingAdvisor(@Value("${app.logging.log-content:false}") boolean logContent) {
|
||||
return LoggingAdvisor.toSlf4j(Orders.LOGGING, logContent);
|
||||
}
|
||||
|
||||
@Bean
|
||||
ChatMemory chatMemory() {
|
||||
return MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).build();
|
||||
}
|
||||
|
||||
@Bean
|
||||
ChatClient chatClient(ChatModel model, PiiRedactionAdvisor pii, TokenBudgetAdvisor budget, LoggingAdvisor logging,
|
||||
ChatMemory memory) {
|
||||
return ChatClient.builder(model)
|
||||
.defaultSystem("You are a terse support assistant.")
|
||||
.defaultAdvisors(logging, budget, MessageChatMemoryAdvisor.builder(memory).build(), pii)
|
||||
.build();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package com.ankurm.advisors.web;
|
||||
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.advisors.advisor.TokenBudgetAdvisor;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.web.bind.annotation.PostMapping;
|
||||
import org.springframework.web.bind.annotation.RequestBody;
|
||||
import org.springframework.web.bind.annotation.RequestHeader;
|
||||
import org.springframework.web.bind.annotation.RestController;
|
||||
|
||||
/**
|
||||
* One conversation per user. The user comes from an {@code X-User} header only so the tests need
|
||||
* no login; in a real service it is the authenticated principal, never a header the client picks.
|
||||
*/
|
||||
@RestController
|
||||
public class ChatController {
|
||||
|
||||
public record Ask(String text) {
|
||||
}
|
||||
|
||||
private final ChatClient chat;
|
||||
|
||||
public ChatController(ChatClient chat) {
|
||||
this.chat = chat;
|
||||
}
|
||||
|
||||
@PostMapping("/chat")
|
||||
Map<String, String> ask(@RequestHeader("X-User") String user, @RequestBody Ask ask) {
|
||||
String reply = chat.prompt()
|
||||
.user(ask.text())
|
||||
.advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "conv-" + user).param(TokenBudgetAdvisor.USER_KEY, user))
|
||||
.call()
|
||||
.content();
|
||||
return Map.of("reply", reply);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
spring:
|
||||
application:
|
||||
name: advisors
|
||||
ai:
|
||||
model:
|
||||
chat: openai
|
||||
openai:
|
||||
api-key: ${OPENAI_API_KEY:sk-not-set}
|
||||
chat:
|
||||
model: gpt-5-mini
|
||||
|
||||
app:
|
||||
budget:
|
||||
max-prompt-tokens: 2000
|
||||
max-tokens-per-user: 20000
|
||||
logging:
|
||||
log-content: false
|
||||
@@ -0,0 +1,51 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import com.ankurm.advisors.support.Http;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.TestModelConfig;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
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;
|
||||
|
||||
/** The whole app, over HTTP, with a scripted model: a refused request is a 429, and the model is never asked. */
|
||||
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT,
|
||||
properties = { "spring.ai.model.chat=none", "app.budget.max-prompt-tokens=60", "app.budget.max-tokens-per-user=30" })
|
||||
@Import(TestModelConfig.class)
|
||||
class AdvisorWebTest {
|
||||
|
||||
@LocalServerPort
|
||||
int port;
|
||||
|
||||
@Autowired
|
||||
RecordingModel model;
|
||||
|
||||
@Test
|
||||
void aRefusedRequestIsA429AndTheModelIsNeverCalled() {
|
||||
try (Transcript t = new Transcript("19-web-429.txt", "Over HTTP: 200, then 429")) {
|
||||
Http http = new Http(port);
|
||||
t.line("limits for this run: 60 tokens per request, 30 per user");
|
||||
int first = 0;
|
||||
for (int i = 1; i <= 4; i++) {
|
||||
Http.Reply r = http.post("dana", "/chat", "{\"text\":\"Where is my refund?\"}");
|
||||
t.line("dana request %d -> HTTP %d, model calls so far: %d", i, r.status(), model.callCount());
|
||||
if (i == 1) {
|
||||
first = r.status();
|
||||
}
|
||||
}
|
||||
int callsBefore = model.callCount();
|
||||
String big = "java.lang.NullPointerException at com.example.Checkout.pay ".repeat(10);
|
||||
Http.Reply refused = http.post("erin", "/chat", "{\"text\":\"" + big + "\"}");
|
||||
t.blank().line("erin pastes a long stack trace -> HTTP %d, model calls: %d (was %d)", refused.status(),
|
||||
model.callCount(), callsBefore);
|
||||
|
||||
assertThat(first).isEqualTo(200);
|
||||
assertThat(refused.status()).isEqualTo(429);
|
||||
assertThat(model.callCount()).isEqualTo(callsBefore);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.AdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.BaseAdvisor;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/**
|
||||
* {@code BaseAdvisor} is the friendly base class: implement {@code before} and {@code after} and
|
||||
* it wires both the call and the stream path. What it does with a stream is not what the name
|
||||
* "after" suggests.
|
||||
*/
|
||||
class BaseAdvisorStreamTest {
|
||||
|
||||
/** Counts what BaseAdvisor calls, and on which thread. */
|
||||
static final class Counting implements BaseAdvisor {
|
||||
|
||||
final AtomicInteger befores = new AtomicInteger();
|
||||
|
||||
final List<String> afterTexts = new CopyOnWriteArrayList<>();
|
||||
|
||||
volatile String beforeThread;
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "Counting";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientRequest before(ChatClientRequest request, AdvisorChain chain) {
|
||||
befores.incrementAndGet();
|
||||
beforeThread = Thread.currentThread().getName();
|
||||
return request;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse after(ChatClientResponse response, AdvisorChain chain) {
|
||||
afterTexts.add(response.chatResponse().getResult().getOutput().getText());
|
||||
return response;
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void afterRunsOncePerCallButOnlyOnTheLastChunkOfAStream() {
|
||||
try (Transcript t = new Transcript("04-base-advisor-stream.txt", "BaseAdvisor on a call and on a stream")) {
|
||||
Counting advisor = new Counting();
|
||||
RecordingModel model = new RecordingModel().replier(p -> "one two three four").chunkSize(4);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
|
||||
client.prompt().user("hi").call().content();
|
||||
t.line("call: before x%d, after x%d", advisor.befores.get(), advisor.afterTexts.size());
|
||||
t.line(" before ran on thread: %s", advisor.beforeThread.equals(Thread.currentThread().getName())
|
||||
? "the caller's thread" : advisor.beforeThread);
|
||||
String callThread = advisor.beforeThread;
|
||||
|
||||
Counting streamed = new Counting();
|
||||
ChatClient streamClient = ChatClient.builder(model).defaultAdvisors(streamed).build();
|
||||
List<String> chunks = streamClient.prompt().user("hi").stream().content().collectList().block();
|
||||
t.blank().line("stream: the model streamed %d chunks: %s", chunks.size(), String.join("|", chunks));
|
||||
t.line(" before x%d, after x%d", streamed.befores.get(), streamed.afterTexts.size());
|
||||
t.line(" after saw only: \"%s\"", streamed.afterTexts.getFirst());
|
||||
t.line(" before ran on thread: %s", streamed.beforeThread.startsWith("boundedElastic")
|
||||
? "a boundedElastic worker, not the caller's thread" : streamed.beforeThread);
|
||||
|
||||
Counting noFinish = new Counting();
|
||||
ChatClient noFinishClient = ChatClient.builder(new RecordingModel().replier(p -> "one two three four").chunkSize(4)
|
||||
.finishReason(false)).defaultAdvisors(noFinish).build();
|
||||
noFinishClient.prompt().user("hi").stream().content().blockLast();
|
||||
t.blank().line("stream whose last chunk has no finish reason: before x%d, after x%d", noFinish.befores.get(),
|
||||
noFinish.afterTexts.size());
|
||||
assertThat(noFinish.afterTexts).isEmpty();
|
||||
assertThat(advisor.afterTexts).hasSize(1);
|
||||
assertThat(streamed.afterTexts).hasSize(1);
|
||||
assertThat(chunks.size()).isGreaterThan(1);
|
||||
assertThat(callThread).isEqualTo(Thread.currentThread().getName());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.SafeGuardAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.SimpleLoggerAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.ToolCallingAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.Advisor;
|
||||
import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository;
|
||||
import org.springframework.ai.chat.memory.MessageWindowChatMemory;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/** Where the advisors that ship with Spring AI sit, read from the objects themselves. */
|
||||
class BuiltInOrdersTest {
|
||||
|
||||
private static String row(Advisor a) {
|
||||
long offset = (long) a.getOrder() - Ordered.HIGHEST_PRECEDENCE;
|
||||
return "%-28s getOrder() = %d%s".formatted(a.getName(), a.getOrder(),
|
||||
offset < 1000 ? " (HIGHEST_PRECEDENCE + " + offset + ")" : "");
|
||||
}
|
||||
|
||||
@Test
|
||||
void theShippedAdvisorsAndTheirOrders() {
|
||||
try (Transcript t = new Transcript("20-built-in-orders.txt", "Built-in advisor orders in 2.0.1")) {
|
||||
var memory = MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).build();
|
||||
List<Advisor> shipped = List.of(new SimpleLoggerAdvisor(), new SafeGuardAdvisor(List.of("secret")),
|
||||
MessageChatMemoryAdvisor.builder(memory).build(), ToolCallingAdvisor.builder().build());
|
||||
shipped.forEach(a -> t.line(row(a)));
|
||||
t.blank().line("Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER = HIGHEST_PRECEDENCE + %d",
|
||||
(long) Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER - Ordered.HIGHEST_PRECEDENCE);
|
||||
assertThat(shipped).isNotEmpty();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.advisors.support.Probe;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.ToolCallingAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.core.Ordered;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* How Spring AI turns a list of advisors into a chain. The chain is sorted by
|
||||
* {@code getOrder()}; the order you pass them to {@code defaultAdvisors(...)} only matters for ties.
|
||||
*/
|
||||
class ChainOrderTest {
|
||||
|
||||
|
||||
private static String joined(List<String> events) {
|
||||
return String.join(" ", events);
|
||||
}
|
||||
|
||||
@Test
|
||||
void theChainIsSortedByOrderNotByRegistrationOrder() {
|
||||
try (Transcript t = new Transcript("01-chain-order.txt", "Advisor chain: what decides who runs first")) {
|
||||
List<String> events = new ArrayList<>();
|
||||
RecordingModel model = new RecordingModel().trace(events::add);
|
||||
ChatClient client = ChatClient.builder(model)
|
||||
.defaultAdvisors(new Probe("C300", 300, events), new Probe("A100", 100, events), new Probe("B200", 200, events))
|
||||
.build();
|
||||
|
||||
client.prompt().user("hi").call().content();
|
||||
String call = joined(events);
|
||||
t.line("registered as C300, A100, B200 (the number is HIGHEST_PRECEDENCE + n)");
|
||||
t.line("call: %s", call);
|
||||
|
||||
events.clear();
|
||||
client.prompt().user("hi").stream().content().blockLast();
|
||||
String stream = joined(events);
|
||||
t.line("stream: %s", stream);
|
||||
assertThat(call).isEqualTo("A100> B200> C300> model C300< B200< A100<");
|
||||
assertThat(stream).isEqualTo(call);
|
||||
|
||||
events.clear();
|
||||
ChatClient withRequestLevel = client.mutate().build();
|
||||
withRequestLevel.prompt().user("hi").advisors(new Probe("R150", 150, events)).call().content();
|
||||
t.blank().line("plus one advisor added on the request with .advisors(...) at n=150:");
|
||||
t.line("call: %s", joined(events));
|
||||
assertThat(joined(events)).isEqualTo("A100> R150> B200> C300> model C300< B200< R150< A100<");
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void twoAdvisorsWithTheSameOrderKeepRegistrationOrder() {
|
||||
try (Transcript t = new Transcript("02-chain-order-ties.txt", "Two advisors with the same order number")) {
|
||||
List<String> events = new ArrayList<>();
|
||||
RecordingModel model = new RecordingModel().trace(events::add);
|
||||
|
||||
ChatClient xy = ChatClient.builder(model)
|
||||
.defaultAdvisors(new Probe("X", 100, events), new Probe("Y", 100, events)).build();
|
||||
xy.prompt().user("hi").call().content();
|
||||
t.line("registered X, Y -> %s", joined(events));
|
||||
String first = joined(events);
|
||||
|
||||
events.clear();
|
||||
ChatClient yx = ChatClient.builder(model)
|
||||
.defaultAdvisors(new Probe("Y", 100, events), new Probe("X", 100, events)).build();
|
||||
yx.prompt().user("hi").call().content();
|
||||
t.line("registered Y, X -> %s", joined(events));
|
||||
|
||||
t.blank().line("For equal numbers the advisor registered LAST runs first (outermost).");
|
||||
assertThat(first).isEqualTo("Y> X> model X< Y<");
|
||||
assertThat(joined(events)).isEqualTo("X> Y> model Y< X<");
|
||||
}
|
||||
}
|
||||
|
||||
private static String describe(int order) {
|
||||
return order == Integer.MAX_VALUE ? "LOWEST_PRECEDENCE"
|
||||
: "HIGHEST_PRECEDENCE + " + (order - (long) Integer.MIN_VALUE);
|
||||
}
|
||||
|
||||
/** Lists the chain the way an advisor inside it sees it. */
|
||||
private static final class Dump implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
final List<String> callChain = new ArrayList<>();
|
||||
|
||||
final List<String> streamChain = new ArrayList<>();
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "Dump";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
chain.getCallAdvisors().forEach(a -> callChain.add(String.format("%-24s %s", a.getName(), describe(a.getOrder()))));
|
||||
return chain.nextCall(request);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
chain.getStreamAdvisors()
|
||||
.forEach(a -> streamChain.add(String.format("%-24s %s", a.getName(), describe(a.getOrder()))));
|
||||
return chain.nextStream(request);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void chainCanBeInspectedFromInsideAnAdvisor() {
|
||||
try (Transcript t = new Transcript("03-chain-contents.txt", "What is actually in the chain")) {
|
||||
Dump dump = new Dump();
|
||||
ChatClient client = ChatClient.builder(new RecordingModel()).defaultAdvisors(dump).build();
|
||||
client.prompt().user("hi").call().content();
|
||||
t.line("one custom advisor at HIGHEST_PRECEDENCE + 100; the chain a call runs through:");
|
||||
dump.callChain.forEach(s -> t.line(" %s", s));
|
||||
client.prompt().user("hi").stream().content().blockLast();
|
||||
t.blank().line("and the chain a stream runs through:");
|
||||
dump.streamChain.forEach(s -> t.line(" %s", s));
|
||||
|
||||
Dump withExplicit = new Dump();
|
||||
ChatClient explicit = ChatClient.builder(new RecordingModel())
|
||||
.defaultAdvisors(withExplicit, ToolCallingAdvisor.builder().build()).build();
|
||||
explicit.prompt().user("hi").call().content();
|
||||
t.blank().line("the same, with ToolCallingAdvisor.builder().build() added by hand:");
|
||||
withExplicit.callChain.forEach(s -> t.line(" %s", s));
|
||||
|
||||
assertThat(dump.callChain).hasSize(3);
|
||||
assertThat(dump.callChain.get(1)).contains("Tool Calling Advisor");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/** Requests and responses are records. An advisor changes them by making a copy, and context is how advisors talk. */
|
||||
class ContextAndImmutabilityTest {
|
||||
|
||||
private static CallAdvisor advisor(String name, int offset, java.util.function.BiFunction<ChatClientRequest, CallAdvisorChain, ChatClientResponse> body) {
|
||||
return new CallAdvisor() {
|
||||
@Override
|
||||
public String getName() {
|
||||
return name;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + offset;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
return body.apply(request, chain);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
@Test
|
||||
void contextTravelsDownOnTheRequestAndUpOnTheResponse() {
|
||||
try (Transcript t = new Transcript("17-context.txt", "Advisor context: request down, response up")) {
|
||||
List<String> seen = new ArrayList<>();
|
||||
CallAdvisor outer = advisor("outer", 100, (req, chain) -> {
|
||||
ChatClientRequest copy = req.mutate().context("tenant", "acme").build();
|
||||
seen.add("outer: original request has tenant=" + req.context().get("tenant")
|
||||
+ ", the copy it forwards has tenant=" + copy.context().get("tenant"));
|
||||
ChatClientResponse resp = chain.nextCall(copy);
|
||||
seen.add("outer: response context has verdict=" + resp.context().get("verdict")
|
||||
+ ", tenant=" + resp.context().get("tenant"));
|
||||
return resp;
|
||||
});
|
||||
CallAdvisor inner = advisor("inner", 200, (req, chain) -> {
|
||||
seen.add("inner: request context has tenant=" + req.context().get("tenant"));
|
||||
ChatClientResponse resp = chain.nextCall(req);
|
||||
return resp.mutate().context("verdict", "clean").build();
|
||||
});
|
||||
|
||||
RecordingModel model = new RecordingModel();
|
||||
ChatClient.builder(model).defaultAdvisors(outer, inner).build().prompt().user("hi").call().content();
|
||||
seen.forEach(t::line);
|
||||
|
||||
boolean mutable;
|
||||
try {
|
||||
ChatClient.builder(model).defaultAdvisors(advisor("poke", 100, (req, chain) -> {
|
||||
req.context().put("sneaky", "yes");
|
||||
return chain.nextCall(req);
|
||||
})).build().prompt().user("hi").call().content();
|
||||
mutable = true;
|
||||
}
|
||||
catch (UnsupportedOperationException e) {
|
||||
mutable = false;
|
||||
}
|
||||
t.blank().line("req.context().put(...) on the incoming request: %s",
|
||||
mutable ? "allowed" : "UnsupportedOperationException (the map is unmodifiable)");
|
||||
|
||||
assertThat(seen).anyMatch(s -> s.startsWith("inner:") && s.endsWith("tenant=acme"));
|
||||
assertThat(seen).anyMatch(s -> s.contains("original request has tenant=null"));
|
||||
assertThat(seen).anyMatch(s -> s.contains("verdict=clean"));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.function.Supplier;
|
||||
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.BaseAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.AdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.core.Ordered;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* What the caller sees when an advisor refuses a request: on a call, and on a stream. An advisor that throws
|
||||
* from {@code adviseStream} before it returns a Flux and one that throws inside the Flux are different things.
|
||||
*/
|
||||
class ErrorPropagationTest {
|
||||
|
||||
private static final class Refuse implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
private final boolean insideFlux;
|
||||
|
||||
Refuse(boolean insideFlux) {
|
||||
this.insideFlux = insideFlux;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return insideFlux ? "refuse-inside-flux" : "refuse-eagerly";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
throw new IllegalStateException("refused");
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
if (insideFlux) {
|
||||
return Flux.defer(() -> Flux.error(new IllegalStateException("refused")));
|
||||
}
|
||||
throw new IllegalStateException("refused");
|
||||
}
|
||||
}
|
||||
|
||||
private static final class RefusingBase implements BaseAdvisor {
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "refusing-base";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientRequest before(ChatClientRequest request, AdvisorChain chain) {
|
||||
throw new IllegalStateException("refused");
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse after(ChatClientResponse response, AdvisorChain chain) {
|
||||
return response;
|
||||
}
|
||||
}
|
||||
|
||||
/** Runs the action and reports at which point it failed and with what. */
|
||||
private static String outcome(Supplier<Object> build, Supplier<Object> consume) {
|
||||
String phase = "building the Flux";
|
||||
try {
|
||||
build.get();
|
||||
phase = "running it";
|
||||
consume.get();
|
||||
return "no error";
|
||||
}
|
||||
catch (Throwable e) {
|
||||
String cause = e.getCause() == null ? ""
|
||||
: " (cause: " + e.getCause().getClass().getSimpleName() + ": " + e.getCause().getMessage() + ")";
|
||||
return "failed while " + phase + " with " + e.getClass().getSimpleName() + ": " + e.getMessage() + cause;
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void whereARefusalSurfacesOnCallAndOnStream() {
|
||||
try (Transcript t = new Transcript("16-error-propagation.txt", "Where a refusal surfaces")) {
|
||||
RecordingModel model = new RecordingModel();
|
||||
t.line("an advisor that throws IllegalStateException(\"refused\"), model calls counted afterwards");
|
||||
t.blank();
|
||||
|
||||
String[] labels = { "plain advisor, call", "plain advisor, stream, throws eagerly",
|
||||
"plain advisor, stream, error inside Flux.defer", "BaseAdvisor.before, call",
|
||||
"BaseAdvisor.before, stream" };
|
||||
String[] results = new String[5];
|
||||
|
||||
ChatClient eager = ChatClient.builder(model).defaultAdvisors(new Refuse(false)).build();
|
||||
ChatClient deferred = ChatClient.builder(model).defaultAdvisors(new Refuse(true)).build();
|
||||
ChatClient base = ChatClient.builder(model).defaultAdvisors(new RefusingBase()).build();
|
||||
|
||||
results[0] = outcome(() -> null, () -> eager.prompt().user("hi").call().content());
|
||||
results[1] = outcome(() -> eager.prompt().user("hi").stream().content(),
|
||||
() -> eager.prompt().user("hi").stream().content().blockLast());
|
||||
results[2] = outcome(() -> deferred.prompt().user("hi").stream().content(),
|
||||
() -> deferred.prompt().user("hi").stream().content().blockLast());
|
||||
results[3] = outcome(() -> null, () -> base.prompt().user("hi").call().content());
|
||||
results[4] = outcome(() -> base.prompt().user("hi").stream().content(),
|
||||
() -> base.prompt().user("hi").stream().content().blockLast());
|
||||
|
||||
for (int i = 0; i < labels.length; i++) {
|
||||
t.line("%-48s %s", labels[i], results[i]);
|
||||
}
|
||||
t.blank().line("model calls: %d", model.callCount());
|
||||
|
||||
assertThat(model.callCount()).isZero();
|
||||
for (String r : results) {
|
||||
assertThat(r).contains("IllegalStateException").contains("refused");
|
||||
}
|
||||
assertThat(results[4]).contains("Stream processing failed");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
|
||||
import com.ankurm.advisors.advisor.LoggingAdvisor;
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
|
||||
/** The audit-logging advisor on a call and on a stream. A fake clock makes "took N ms" deterministic. */
|
||||
class LoggingAdvisorTest {
|
||||
|
||||
/** Advances 5 ms every time it is read. */
|
||||
private static AtomicLong clock() {
|
||||
return new AtomicLong();
|
||||
}
|
||||
|
||||
private static LoggingAdvisor logger(List<String> sink, boolean content) {
|
||||
AtomicLong nanos = clock();
|
||||
return new LoggingAdvisor("LoggingAdvisor", Orders.LOGGING, content, sink::add, () -> nanos.addAndGet(5_000_000));
|
||||
}
|
||||
|
||||
@Test
|
||||
void oneLineInOneLineOutWithLatencyAndUsage() {
|
||||
try (Transcript t = new Transcript("05-logging-advisor.txt", "LoggingAdvisor: a call and a stream")) {
|
||||
List<String> log = new ArrayList<>();
|
||||
RecordingModel model = new RecordingModel().replier(p -> "Refund issued.").chunkSize(6);
|
||||
ChatClient client = ChatClient.builder(model).defaultSystem("You are a terse support assistant.")
|
||||
.defaultAdvisors(logger(log, false)).build();
|
||||
|
||||
client.prompt().user("Where is my refund?").call().content();
|
||||
t.line("call, content logging off (the default):");
|
||||
log.forEach(l -> t.line(" %s", l));
|
||||
assertThat(log).hasSize(2);
|
||||
|
||||
log.clear();
|
||||
client.prompt().user("Where is my refund?").stream().content().blockLast();
|
||||
t.blank().line("stream:");
|
||||
log.forEach(l -> t.line(" %s", l));
|
||||
|
||||
log.clear();
|
||||
ChatClient withText = ChatClient.builder(model).defaultAdvisors(logger(log, true)).build();
|
||||
withText.prompt().user("Where is my refund?").call().content();
|
||||
t.blank().line("call, content logging on:");
|
||||
log.forEach(l -> t.line(" %s", l));
|
||||
assertThat(log.getFirst()).contains("Where is my refund?");
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void aFailureIsLoggedAndRethrown() {
|
||||
try (Transcript t = new Transcript("06-logging-failure.txt", "LoggingAdvisor when the model call fails")) {
|
||||
List<String> log = new ArrayList<>();
|
||||
RecordingModel model = new RecordingModel().replier(p -> {
|
||||
throw new IllegalStateException("provider returned 503");
|
||||
});
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(logger(log, false)).build();
|
||||
try {
|
||||
client.prompt().user("hi").call().content();
|
||||
}
|
||||
catch (IllegalStateException e) {
|
||||
t.line("caller got: %s: %s", e.getClass().getSimpleName(), e.getMessage());
|
||||
}
|
||||
log.forEach(l -> t.line(" %s", l));
|
||||
assertThat(log.getLast()).contains("failed").contains("IllegalStateException");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.advisor.PiiRedactionAdvisor;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Show;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository;
|
||||
import org.springframework.ai.chat.memory.MessageWindowChatMemory;
|
||||
|
||||
/** Placeholders and memory: numbering that restarts every request gives one label to two people. */
|
||||
class PiiMultiTurnTest {
|
||||
|
||||
private record Turns(String modelSawTurn2, String callerGotTurn2, List<String> stored) {
|
||||
}
|
||||
|
||||
private static Turns twoTurns(boolean conversationScoped) {
|
||||
ChatMemory memory = MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).build();
|
||||
RecordingModel model = new RecordingModel();
|
||||
ChatClient client = ChatClient.builder(model)
|
||||
.defaultAdvisors(new PiiRedactionAdvisor(Orders.PII_REDACTION, true, conversationScoped),
|
||||
MessageChatMemoryAdvisor.builder(memory).build())
|
||||
.build();
|
||||
client.prompt().user("My email is [email protected]").advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "c")).call()
|
||||
.content();
|
||||
String reply = client.prompt().user("Also cc [email protected]").advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "c"))
|
||||
.call().content();
|
||||
return new Turns(Show.inline(model.lastPrompt().getInstructions()), reply,
|
||||
memory.get("c").stream().map(Show::one).toList());
|
||||
}
|
||||
|
||||
@Test
|
||||
void placeholdersMustSurviveAcrossTurnsWhenMemoryReplaysThem() {
|
||||
try (Transcript t = new Transcript("11-pii-multi-turn.txt", "Placeholders across two turns with memory")) {
|
||||
t.line("turn 1: My email is [email protected]");
|
||||
t.line("turn 2: Also cc [email protected]");
|
||||
|
||||
Turns perRequest = twoTurns(false);
|
||||
t.blank().line("numbering restarts on every request:");
|
||||
t.line(" model was sent on turn 2: %s", perRequest.modelSawTurn2());
|
||||
t.line(" caller receives: %s", perRequest.callerGotTurn2());
|
||||
|
||||
Turns scoped = twoTurns(true);
|
||||
t.blank().line("numbering kept per conversation (the default):");
|
||||
t.line(" model was sent on turn 2: %s", scoped.modelSawTurn2());
|
||||
t.line(" caller receives: %s", scoped.callerGotTurn2());
|
||||
t.blank().line("what the memory stores (placeholders, never the addresses):");
|
||||
scoped.stored().forEach(s -> t.line(" %s", s));
|
||||
|
||||
assertThat(perRequest.modelSawTurn2()).contains("U:My email is <EMAIL_1>").contains("U:Also cc <EMAIL_1>");
|
||||
assertThat(scoped.modelSawTurn2()).contains("U:My email is <EMAIL_1>").contains("U:Also cc <EMAIL_2>");
|
||||
assertThat(scoped.stored()).noneMatch(s -> s.contains("@"));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
|
||||
import com.ankurm.advisors.advisor.LoggingAdvisor;
|
||||
import com.ankurm.advisors.advisor.PiiRedactionAdvisor;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.Advisor;
|
||||
import org.springframework.ai.chat.memory.ChatMemory;
|
||||
import org.springframework.ai.chat.memory.InMemoryChatMemoryRepository;
|
||||
import org.springframework.ai.chat.memory.MessageWindowChatMemory;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/**
|
||||
* Three places to put the redaction advisor relative to the logging advisor and Spring AI's memory
|
||||
* advisor (HIGHEST_PRECEDENCE + 200), and what each of the three ends up holding.
|
||||
*/
|
||||
class PiiOrderTest {
|
||||
|
||||
private static final String INPUT = "My email is [email protected]";
|
||||
|
||||
private record Outcome(String logged, String stored, String modelSaw) {
|
||||
}
|
||||
|
||||
private static Outcome run(int piiOffset, int loggingOffset) {
|
||||
List<String> log = new ArrayList<>();
|
||||
AtomicLong nanos = new AtomicLong();
|
||||
LoggingAdvisor logging = new LoggingAdvisor("LoggingAdvisor", Ordered.HIGHEST_PRECEDENCE + loggingOffset, true,
|
||||
log::add, () -> nanos.addAndGet(5_000_000));
|
||||
ChatMemory memory = MessageWindowChatMemory.builder().chatMemoryRepository(new InMemoryChatMemoryRepository()).build();
|
||||
RecordingModel model = new RecordingModel();
|
||||
ChatClient client = ChatClient.builder(model)
|
||||
.defaultAdvisors(logging, new PiiRedactionAdvisor(Ordered.HIGHEST_PRECEDENCE + piiOffset, true),
|
||||
MessageChatMemoryAdvisor.builder(memory).build())
|
||||
.build();
|
||||
client.prompt().user(INPUT).advisors(a -> a.param(ChatMemory.CONVERSATION_ID, "c")).call().content();
|
||||
String logged = log.getFirst().replaceFirst(".*last=", "");
|
||||
return new Outcome(logged, "\"" + memory.get("c").getFirst().getText() + "\"",
|
||||
"\"" + model.lastPrompt().getInstructions().getLast().getText() + "\"");
|
||||
}
|
||||
|
||||
@Test
|
||||
void whereTheRedactionAdvisorSitsDecidesWhatIsLoggedAndWhatIsStored() {
|
||||
try (Transcript t = new Transcript("10-pii-order.txt", "Redaction order versus logging and memory")) {
|
||||
t.line("memory advisor is fixed at HIGHEST_PRECEDENCE + %d",
|
||||
Advisor.DEFAULT_CHAT_MEMORY_PRECEDENCE_ORDER - (long) Ordered.HIGHEST_PRECEDENCE);
|
||||
t.line("user says: %s", INPUT);
|
||||
|
||||
Outcome a = run(100, 400);
|
||||
t.blank().line("A redaction +100, logging +400 (redaction outside both):");
|
||||
t.line(" log line saw: %s", a.logged());
|
||||
t.line(" memory stored: %s", a.stored());
|
||||
t.line(" model was sent: %s", a.modelSaw());
|
||||
|
||||
Outcome b = run(100, 50);
|
||||
t.blank().line("B redaction +100, logging +50 (logging outside redaction):");
|
||||
t.line(" log line saw: %s", b.logged());
|
||||
t.line(" memory stored: %s", b.stored());
|
||||
t.line(" model was sent: %s", b.modelSaw());
|
||||
|
||||
Outcome c = run(300, 400);
|
||||
t.blank().line("C redaction +300 (inside the memory advisor), logging +400:");
|
||||
t.line(" log line saw: %s", c.logged());
|
||||
t.line(" memory stored: %s", c.stored());
|
||||
t.line(" model was sent: %s", c.modelSaw());
|
||||
|
||||
assertThat(a.logged()).doesNotContain("[email protected]");
|
||||
assertThat(a.stored()).doesNotContain("[email protected]");
|
||||
assertThat(b.logged()).contains("[email protected]");
|
||||
assertThat(b.stored()).doesNotContain("[email protected]");
|
||||
assertThat(c.stored()).contains("[email protected]");
|
||||
assertThat(c.modelSaw()).doesNotContain("[email protected]");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.advisor.PiiRedactionAdvisor;
|
||||
import com.ankurm.advisors.support.NaiveStreamingPii;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Show;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
|
||||
/** PII redaction before the model, and restoring the originals afterwards. */
|
||||
class PiiRedactionTest {
|
||||
|
||||
private static final String INPUT = "Hi, I'm Priya. Email [email protected] or call +91 98765 43210. "
|
||||
+ "Card 4111 1111 1111 1111, order 1234 5678 9012 3456, and again [email protected].";
|
||||
|
||||
@Test
|
||||
void theModelSeesPlaceholdersAndTheCallerSeesTheOriginals() {
|
||||
try (Transcript t = new Transcript("07-pii-redaction.txt", "PiiRedactionAdvisor on a call")) {
|
||||
RecordingModel model = new RecordingModel();
|
||||
ChatClient restoring = ChatClient.builder(model)
|
||||
.defaultAdvisors(new PiiRedactionAdvisor(Orders.PII_REDACTION, true)).build();
|
||||
String reply = restoring.prompt().user(INPUT).call().content();
|
||||
|
||||
t.line("caller sends: %s", INPUT);
|
||||
t.line("model was sent: %s", model.lastPrompt().getInstructions().getLast().getText());
|
||||
t.line("caller receives: %s", reply);
|
||||
|
||||
ChatClient plain = ChatClient.builder(model)
|
||||
.defaultAdvisors(new PiiRedactionAdvisor(Orders.PII_REDACTION, false)).build();
|
||||
t.blank().line("with restore switched off, the caller receives:");
|
||||
t.line(" %s", plain.prompt().user(INPUT).call().content());
|
||||
|
||||
String sent = model.prompts().getFirst().getInstructions().getLast().getText();
|
||||
assertThat(sent).contains("<EMAIL_1>", "<PHONE_1>", "<CARD_1>").doesNotContain("[email protected]", "98765", "4111");
|
||||
assertThat(sent).contains("1234 5678 9012 3456").contains("Priya");
|
||||
assertThat(sent.split("<EMAIL_1>", -1)).hasSize(3);
|
||||
assertThat(reply).isEqualTo("You said: " + INPUT);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void whatItDoesNotCatch() {
|
||||
try (Transcript t = new Transcript("08-pii-limits.txt", "What pattern-based redaction misses")) {
|
||||
RecordingModel model = new RecordingModel();
|
||||
ChatClient client = ChatClient.builder(model)
|
||||
.defaultAdvisors(new PiiRedactionAdvisor(Orders.PII_REDACTION, true)).build();
|
||||
List<String> inputs = List.of(
|
||||
"My name is Priya Sharma and I live at 14 Hill Road, Bandra, Mumbai 400050.",
|
||||
"Passport N1234567, PAN ABCDE1234F.",
|
||||
"Write to priya (at) example (dot) com",
|
||||
"Card 4111-1111-1111-1112 and order 1234 5678 9012 3456");
|
||||
for (String in : inputs) {
|
||||
client.prompt().user(in).call().content();
|
||||
String sent = model.lastPrompt().getInstructions().getLast().getText();
|
||||
t.line("in: %s", in);
|
||||
t.line("out: %s", sent);
|
||||
t.blank();
|
||||
}
|
||||
assertThat(model.prompts().get(0).getInstructions().getLast().getText()).contains("Priya Sharma", "14 Hill Road");
|
||||
assertThat(model.prompts().get(1).getInstructions().getLast().getText()).contains("N1234567", "ABCDE1234F");
|
||||
assertThat(model.prompts().get(3).getInstructions().getLast().getText()).contains("4111-1111-1111-1112");
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void aPlaceholderSplitAcrossChunksSurvivesOnlyWithTheBufferedRestore() {
|
||||
try (Transcript t = new Transcript("09-pii-stream-boundary.txt", "Restoring placeholders in a stream")) {
|
||||
Pattern placeholder = Pattern.compile("<[A-Z]+_\\d+>");
|
||||
RecordingModel model = new RecordingModel().chunkSize(5).replier(p -> {
|
||||
Matcher m = placeholder.matcher(p.getInstructions().getLast().getText());
|
||||
return "Sure, I will write to " + (m.find() ? m.group() : "nobody") + " now.";
|
||||
});
|
||||
String input = "Please email [email protected]";
|
||||
|
||||
List<String> naive = ChatClient.builder(model).defaultAdvisors(new NaiveStreamingPii()).build().prompt()
|
||||
.user(input).stream().content().collectList().block();
|
||||
List<String> raw = model.stream(new org.springframework.ai.chat.prompt.Prompt("<EMAIL_1>"))
|
||||
.map(r -> r.getResult().getOutput().getText()).collectList().block();
|
||||
t.line("the model streams 5-character chunks: %s", String.join("|", raw));
|
||||
t.blank().line("restore each chunk on its own:");
|
||||
t.line(" chunks: %s", String.join("|", naive));
|
||||
t.line(" joined: %s", String.join("", naive));
|
||||
|
||||
List<String> buffered = ChatClient.builder(model)
|
||||
.defaultAdvisors(new PiiRedactionAdvisor(Orders.PII_REDACTION, true)).build().prompt().user(input)
|
||||
.stream().content().collectList().block();
|
||||
t.blank().line("PiiRedactionAdvisor (holds back from an unfinished \"<\"):");
|
||||
t.line(" chunks: %s", String.join("|", buffered));
|
||||
t.line(" joined: %s", String.join("", buffered));
|
||||
|
||||
assertThat(raw).contains("o <EM", "AIL_1");
|
||||
assertThat(String.join("", naive)).contains("<EMA").doesNotContain("[email protected]");
|
||||
assertThat(String.join("", buffered)).isEqualTo("Sure, I will write to [email protected] now.");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
import static org.assertj.core.api.Assertions.assertThatThrownBy;
|
||||
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.advisor.TokenBudgetAdvisor;
|
||||
import com.ankurm.advisors.advisor.TokenBudgetExceededException;
|
||||
import com.ankurm.advisors.support.RecordingModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import com.knuddels.jtokkit.api.EncodingType;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator;
|
||||
import org.springframework.ai.tokenizer.TokenCountEstimator;
|
||||
|
||||
/** A per-request cap and a per-user running total, both enforced before the model is called. */
|
||||
class TokenBudgetTest {
|
||||
|
||||
private static final TokenCountEstimator ESTIMATOR = new JTokkitTokenCountEstimator(EncodingType.O200K_BASE);
|
||||
|
||||
private static String ask(ChatClient client, String user, String text) {
|
||||
return client.prompt().user(text).advisors(a -> a.param(TokenBudgetAdvisor.USER_KEY, user)).call().content();
|
||||
}
|
||||
|
||||
@Test
|
||||
void refusesBeforeTheModelIsCalled() {
|
||||
try (Transcript t = new Transcript("12-token-budget.txt", "TokenBudgetAdvisor: 40 tokens per request, 60 per user")) {
|
||||
RecordingModel model = new RecordingModel();
|
||||
TokenBudgetAdvisor budget = new TokenBudgetAdvisor(ESTIMATOR, 40, 60, Orders.TOKEN_BUDGET);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(budget).build();
|
||||
|
||||
t.line("alice asks short questions; \"spent\" is the usage the model reported:");
|
||||
for (int i = 1; i <= 6; i++) {
|
||||
try {
|
||||
ask(client, "alice", "Where is my refund?");
|
||||
t.line(" call %d: answered, spent=%d, model calls=%d", i, budget.spent("alice"), model.callCount());
|
||||
}
|
||||
catch (TokenBudgetExceededException e) {
|
||||
t.line(" call %d: refused (%s), spent=%d, model calls=%d", i, e.getMessage(), budget.spent("alice"),
|
||||
model.callCount());
|
||||
}
|
||||
}
|
||||
|
||||
int before = model.callCount();
|
||||
String paste = "java.lang.NullPointerException at com.example.Checkout.pay(Checkout.java:88) ".repeat(6);
|
||||
t.blank().line("bob pastes a stack trace of %d estimated tokens:", ESTIMATOR.estimate(paste));
|
||||
assertThatThrownBy(() -> ask(client, "bob", paste)).isInstanceOf(TokenBudgetExceededException.class)
|
||||
.satisfies(e -> t.line(" refused: %s", e.getMessage()));
|
||||
t.line(" model calls: %d (was %d), bob's spent: %d", model.callCount(), before, budget.spent("bob"));
|
||||
|
||||
ask(client, "bob", "Where is my refund?");
|
||||
t.blank().line("bob then asks a short question: answered, bob spent=%d, alice spent=%d", budget.spent("bob"),
|
||||
budget.spent("alice"));
|
||||
|
||||
assertThat(model.callCount()).isEqualTo(before + 1);
|
||||
assertThat(budget.spent("alice")).isGreaterThanOrEqualTo(60);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void streamsAreCountedFromTheLastChunkOrEstimatedWhenThereIsNoUsage() {
|
||||
try (Transcript t = new Transcript("13-token-budget-stream.txt", "Token accounting on a stream")) {
|
||||
TokenBudgetAdvisor budget = new TokenBudgetAdvisor(ESTIMATOR, 1000, 100_000, Orders.TOKEN_BUDGET);
|
||||
RecordingModel withUsage = new RecordingModel().replier(p -> "Refund issued today.").chunkSize(6);
|
||||
ChatClient client = ChatClient.builder(withUsage).defaultAdvisors(budget).build();
|
||||
|
||||
client.prompt().user("Where is my refund?").advisors(a -> a.param(TokenBudgetAdvisor.USER_KEY, "u1")).call()
|
||||
.content();
|
||||
long callSpent = budget.spent("u1");
|
||||
client.prompt().user("Where is my refund?").advisors(a -> a.param(TokenBudgetAdvisor.USER_KEY, "u2")).stream()
|
||||
.content().blockLast();
|
||||
long streamSpent = budget.spent("u2");
|
||||
|
||||
RecordingModel noUsage = new RecordingModel().replier(p -> "Refund issued today.").chunkSize(6)
|
||||
.streamUsage(false);
|
||||
ChatClient.builder(noUsage).defaultAdvisors(budget).build().prompt().user("Where is my refund?")
|
||||
.advisors(a -> a.param(TokenBudgetAdvisor.USER_KEY, "u3")).stream().content().blockLast();
|
||||
long estimated = budget.spent("u3");
|
||||
|
||||
t.line("same question, same answer:");
|
||||
t.line(" call, usage from the response: %d tokens", callSpent);
|
||||
t.line(" stream, usage on the last chunk: %d tokens", streamSpent);
|
||||
t.line(" stream, no usage reported (estimated): %d tokens", estimated);
|
||||
t.blank().line("A stream that reports no usage is billed by the provider all the same. The estimate above");
|
||||
t.line("matches only because this scripted model and the advisor use the same tokenizer.");
|
||||
|
||||
assertThat(streamSpent).isEqualTo(callSpent);
|
||||
assertThat(estimated).isEqualTo(callSpent);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void aStreamIsRefusedAsAnErrorSignal() {
|
||||
try (Transcript t = new Transcript("14-token-budget-stream-refusal.txt", "A refusal on the stream path")) {
|
||||
RecordingModel model = new RecordingModel();
|
||||
TokenBudgetAdvisor budget = new TokenBudgetAdvisor(ESTIMATOR, 5, 100, Orders.TOKEN_BUDGET);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(budget).build();
|
||||
try {
|
||||
client.prompt().user("This question is longer than five tokens, so it is refused.").stream().content()
|
||||
.blockLast();
|
||||
}
|
||||
catch (RuntimeException e) {
|
||||
t.line("subscriber got: %s", e.getClass().getSimpleName());
|
||||
t.line("message: %s", e.getMessage());
|
||||
assertThat(e).isInstanceOf(TokenBudgetExceededException.class);
|
||||
}
|
||||
t.line("model calls: %d", model.callCount());
|
||||
assertThat(model.callCount()).isZero();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import com.ankurm.advisors.advisor.TokenBudgetAdvisor;
|
||||
import com.ankurm.advisors.support.ScriptedToolModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import com.knuddels.jtokkit.api.EncodingType;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.support.ToolCallbacks;
|
||||
import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator;
|
||||
import org.springframework.ai.tool.annotation.Tool;
|
||||
import org.springframework.ai.tool.annotation.ToolParam;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/** Does a budget advisor count every model round trip of a tool loop, or only what comes out of it? */
|
||||
class ToolLoopBudgetTest {
|
||||
|
||||
static class Weather {
|
||||
|
||||
@Tool(description = "Current temperature for a city")
|
||||
String currentWeather(@ToolParam(description = "city name") String city) {
|
||||
return "31C in " + city;
|
||||
}
|
||||
}
|
||||
|
||||
private static long spentWith(int offset) {
|
||||
TokenBudgetAdvisor budget = new TokenBudgetAdvisor(new JTokkitTokenCountEstimator(EncodingType.O200K_BASE), 10_000, 100_000,
|
||||
Ordered.HIGHEST_PRECEDENCE + offset);
|
||||
ScriptedToolModel model = ScriptedToolModel.builder().reportingUsage(100, 10)
|
||||
.thenCallTools(new AssistantMessage.ToolCall("call-1", "function", "currentWeather", "{\"city\":\"Mumbai\"}"))
|
||||
.reportingUsage(130, 20).thenRespond("It is 31C in Mumbai.").build();
|
||||
ChatClient.builder(model).defaultToolCallbacks(ToolCallbacks.from(new Weather())).defaultAdvisors(budget).build().prompt()
|
||||
.user("Weather in Mumbai?").advisors(a -> a.param(TokenBudgetAdvisor.USER_KEY, "u")).call().content();
|
||||
return budget.spent("u");
|
||||
}
|
||||
|
||||
@Test
|
||||
void everyRoundIsCountedWhereverTheBudgetSits() {
|
||||
try (Transcript t = new Transcript("21-tool-loop-budget.txt", "Token budget and the tool loop")) {
|
||||
long outside = spentWith(250);
|
||||
long inside = spentWith(400);
|
||||
t.line("the model reports 100+10 tokens for round 1 (asks for the tool) and 130+20 for round 2 (answers)");
|
||||
t.line("so the provider would bill 260 tokens for this one question");
|
||||
t.blank().line("budget advisor at +250 (outside the tool loop): recorded %d", outside);
|
||||
t.line("budget advisor at +400 (inside the tool loop): recorded %d", inside);
|
||||
assertThat(inside).isEqualTo(260);
|
||||
assertThat(outside).isEqualTo(260);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
|
||||
import com.ankurm.advisors.advisor.LoggingAdvisor;
|
||||
import com.ankurm.advisors.support.ScriptedToolModel;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.support.ToolCallbacks;
|
||||
import org.springframework.ai.tool.annotation.Tool;
|
||||
import org.springframework.ai.tool.annotation.ToolParam;
|
||||
import org.springframework.core.Ordered;
|
||||
|
||||
/**
|
||||
* The tool loop runs at HIGHEST_PRECEDENCE + 300. An advisor outside it (a lower number) is entered once
|
||||
* per question; an advisor inside it (a higher number) is entered once per model round trip.
|
||||
*/
|
||||
class ToolLoopOrderTest {
|
||||
|
||||
static class Weather {
|
||||
|
||||
@Tool(description = "Current temperature for a city")
|
||||
String currentWeather(@ToolParam(description = "city name") String city) {
|
||||
return "31C in " + city;
|
||||
}
|
||||
}
|
||||
|
||||
private static List<String> run(int loggingOffset) {
|
||||
List<String> log = new ArrayList<>();
|
||||
AtomicLong nanos = new AtomicLong();
|
||||
LoggingAdvisor logging = new LoggingAdvisor("LoggingAdvisor", Ordered.HIGHEST_PRECEDENCE + loggingOffset, false,
|
||||
log::add, () -> nanos.addAndGet(5_000_000));
|
||||
ScriptedToolModel model = ScriptedToolModel.builder()
|
||||
.thenCallTools(new AssistantMessage.ToolCall("call-1", "function", "currentWeather", "{\"city\":\"Mumbai\"}"))
|
||||
.thenRespond("It is 31C in Mumbai.")
|
||||
.build();
|
||||
ChatClient.builder(model).defaultToolCallbacks(ToolCallbacks.from(new Weather())).defaultAdvisors(logging).build()
|
||||
.prompt().user("Weather in Mumbai?").call().content();
|
||||
return log;
|
||||
}
|
||||
|
||||
@Test
|
||||
void anAdvisorInsideTheToolLoopRunsOncePerModelRoundTrip() {
|
||||
try (Transcript t = new Transcript("15-tool-loop-order.txt", "Where an advisor sits relative to the tool loop")) {
|
||||
List<String> outside = run(250);
|
||||
List<String> inside = run(400);
|
||||
t.line("one question, one tool call, so the model is called twice (ToolCallingAdvisor is at +300)");
|
||||
t.blank().line("logging advisor at +250 (outside the tool loop):");
|
||||
outside.forEach(l -> t.line(" %s", l));
|
||||
t.blank().line("logging advisor at +400 (inside the tool loop):");
|
||||
inside.forEach(l -> t.line(" %s", l));
|
||||
|
||||
assertThat(outside).hasSize(2);
|
||||
assertThat(inside).hasSize(4);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package com.ankurm.advisors;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
|
||||
import com.ankurm.advisors.advisor.Orders;
|
||||
import com.ankurm.advisors.advisor.PiiRedactionAdvisor;
|
||||
import com.ankurm.advisors.advisor.Texts;
|
||||
import com.ankurm.advisors.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
|
||||
/** The smallest advisor test: no ChatClient, no ChatModel, no Spring. A stub chain stands in for "the rest". */
|
||||
class UnitTestingWithoutModelTest {
|
||||
|
||||
/** The rest of the chain: records what reached it and answers with a fixed sentence built from the request. */
|
||||
private static final class StubChain implements CallAdvisorChain {
|
||||
|
||||
final AtomicReference<ChatClientRequest> received = new AtomicReference<>();
|
||||
|
||||
@Override
|
||||
public ChatClientResponse nextCall(ChatClientRequest request) {
|
||||
received.set(request);
|
||||
String last = request.prompt().getUserMessage().getText();
|
||||
return ChatClientResponse.builder()
|
||||
.chatResponse(ChatResponse.builder()
|
||||
.generations(List.of(new Generation(new AssistantMessage("Noted: " + last))))
|
||||
.build())
|
||||
.build();
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<CallAdvisor> getCallAdvisors() {
|
||||
return new ArrayList<>();
|
||||
}
|
||||
|
||||
@Override
|
||||
public CallAdvisorChain copy(CallAdvisor after) {
|
||||
return this;
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void anAdvisorIsAnOrdinaryObjectYouCanCallDirectly() {
|
||||
try (Transcript t = new Transcript("18-unit-test-no-model.txt", "Testing an advisor with a stub chain")) {
|
||||
PiiRedactionAdvisor advisor = new PiiRedactionAdvisor(Orders.PII_REDACTION, true, false);
|
||||
StubChain chain = new StubChain();
|
||||
ChatClientRequest request = ChatClientRequest.builder().prompt(new Prompt("Mail [email protected] please")).build();
|
||||
|
||||
ChatClientResponse response = advisor.adviseCall(request, chain);
|
||||
|
||||
String forwarded = chain.received.get().prompt().getUserMessage().getText();
|
||||
String answer = Texts.text(response);
|
||||
t.line("what the caller sent: %s", request.prompt().getUserMessage().getText());
|
||||
t.line("what reached the rest: %s", forwarded);
|
||||
t.line("what the caller got back: %s", answer);
|
||||
t.line("caller's request unchanged: %s", request.prompt().getUserMessage().getText().contains("[email protected]"));
|
||||
|
||||
assertThat(forwarded).isEqualTo("Mail <EMAIL_1> please");
|
||||
assertThat(answer).isEqualTo("Noted: Mail [email protected] please");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import java.net.URI;
|
||||
import java.net.http.HttpClient;
|
||||
import java.net.http.HttpRequest;
|
||||
import java.net.http.HttpResponse;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/** A synchronous JDK HttpClient -- one send() per request, so a transcript line is one request. */
|
||||
public final class Http {
|
||||
|
||||
public record Reply(int status, String body) {
|
||||
}
|
||||
|
||||
private final HttpClient client = HttpClient.newHttpClient();
|
||||
private final String base;
|
||||
|
||||
public Http(int port) {
|
||||
this.base = "http://127.0.0.1:" + port;
|
||||
}
|
||||
|
||||
public Reply post(String user, String path, String json) {
|
||||
return send(HttpRequest.newBuilder(URI.create(base + path)).header("X-User", user)
|
||||
.header("Content-Type", "application/json").POST(HttpRequest.BodyPublishers.ofString(json)).build());
|
||||
}
|
||||
|
||||
public Reply get(String user, String path) {
|
||||
return send(HttpRequest.newBuilder(URI.create(base + path)).header("X-User", user).GET().build());
|
||||
}
|
||||
|
||||
private Reply send(HttpRequest request) {
|
||||
try {
|
||||
HttpResponse<String> r = client.send(request, HttpResponse.BodyHandlers.ofString());
|
||||
return new Reply(r.statusCode(), r.body());
|
||||
}
|
||||
catch (Exception e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
|
||||
public static String field(String json, String name) {
|
||||
Matcher m = Pattern.compile("\"" + name + "\":\"([^\"]*)\"").matcher(json);
|
||||
return m.find() ? m.group(1) : null;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import com.ankurm.advisors.advisor.PiiRedactor;
|
||||
import com.ankurm.advisors.advisor.Texts;
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.core.Ordered;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* The obvious way to restore placeholders in a stream, and the wrong one: restore each chunk on its
|
||||
* own. It works whenever a placeholder happens to arrive whole and fails whenever a chunk boundary
|
||||
* falls inside one. Kept only so a test can show the failure next to the fix.
|
||||
*/
|
||||
public final class NaiveStreamingPii implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return "NaiveStreamingPii";
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return Ordered.HIGHEST_PRECEDENCE + 100;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
PiiRedactor.Session session = new PiiRedactor.Session();
|
||||
return Texts.mapResponseText(chain.nextCall(Texts.mapRequestText(request, session::redact)), session::restore);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
return Flux.defer(() -> {
|
||||
PiiRedactor.Session session = new PiiRedactor.Session();
|
||||
return chain.nextStream(Texts.mapRequestText(request, session::redact))
|
||||
.map(chunk -> Texts.mapResponseText(chunk, session::restore));
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClientRequest;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.CallAdvisorChain;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisor;
|
||||
import org.springframework.ai.chat.client.advisor.api.StreamAdvisorChain;
|
||||
import org.springframework.core.Ordered;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/** An advisor that only records when it is entered ("X>") and left ("X<"). */
|
||||
public final class Probe implements CallAdvisor, StreamAdvisor {
|
||||
|
||||
private final String name;
|
||||
|
||||
private final int order;
|
||||
|
||||
private final List<String> events;
|
||||
|
||||
public Probe(String name, int offsetFromHighest, List<String> events) {
|
||||
this.name = name;
|
||||
this.order = Ordered.HIGHEST_PRECEDENCE + offsetFromHighest;
|
||||
this.events = events;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return name;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getOrder() {
|
||||
return order;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
events.add(name + ">");
|
||||
ChatClientResponse response = chain.nextCall(request);
|
||||
events.add(name + "<");
|
||||
return response;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatClientResponse> adviseStream(ChatClientRequest request, StreamAdvisorChain chain) {
|
||||
return Flux.defer(() -> {
|
||||
events.add(name + ">");
|
||||
return chain.nextStream(request).doOnComplete(() -> events.add(name + "<"));
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.function.Consumer;
|
||||
import java.util.function.Function;
|
||||
|
||||
import com.knuddels.jtokkit.api.EncodingType;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
|
||||
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
|
||||
import org.springframework.ai.chat.metadata.DefaultUsage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator;
|
||||
import org.springframework.ai.tokenizer.TokenCountEstimator;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
/**
|
||||
* A scripted {@link ChatModel}: no network, no API key, no randomness. It records every {@link
|
||||
* Prompt} it is sent, answers with whatever the {@code replier} function returns (by default it
|
||||
* echoes the last message, which makes redaction visible: the model "repeats" what it was given),
|
||||
* and reports token usage counted with the same JTokkit estimator the budget advisor uses, so the
|
||||
* numbers are self-consistent. It says nothing about what a real model would reply.
|
||||
*
|
||||
* <p>{@link #stream} splits the reply into {@code chunkSize}-character pieces. The last piece
|
||||
* carries a finish reason and, when {@code streamUsage} is on, the usage -- the shape OpenAI's
|
||||
* streaming API has when usage reporting is switched on.
|
||||
*/
|
||||
public class RecordingModel implements ChatModel {
|
||||
|
||||
private static final TokenCountEstimator ESTIMATOR = new JTokkitTokenCountEstimator(EncodingType.O200K_BASE);
|
||||
|
||||
private final List<Prompt> prompts = new CopyOnWriteArrayList<>();
|
||||
|
||||
private Function<Prompt, String> replier = p -> "You said: " + p.getInstructions().getLast().getText();
|
||||
|
||||
private Consumer<String> trace = s -> {
|
||||
};
|
||||
|
||||
private int chunkSize = 8;
|
||||
|
||||
private boolean streamUsage = true;
|
||||
|
||||
private boolean finishReason = true;
|
||||
|
||||
public RecordingModel replier(Function<Prompt, String> replier) {
|
||||
this.replier = replier;
|
||||
return this;
|
||||
}
|
||||
|
||||
public RecordingModel trace(Consumer<String> trace) {
|
||||
this.trace = trace;
|
||||
return this;
|
||||
}
|
||||
|
||||
public RecordingModel chunkSize(int chunkSize) {
|
||||
this.chunkSize = chunkSize;
|
||||
return this;
|
||||
}
|
||||
|
||||
public RecordingModel streamUsage(boolean streamUsage) {
|
||||
this.streamUsage = streamUsage;
|
||||
return this;
|
||||
}
|
||||
|
||||
/** Whether the last streamed chunk carries a finish reason (real providers do; this lets a test see what happens if not). */
|
||||
public RecordingModel finishReason(boolean finishReason) {
|
||||
this.finishReason = finishReason;
|
||||
return this;
|
||||
}
|
||||
|
||||
public List<Prompt> prompts() {
|
||||
return Collections.unmodifiableList(prompts);
|
||||
}
|
||||
|
||||
public Prompt lastPrompt() {
|
||||
return prompts.getLast();
|
||||
}
|
||||
|
||||
public int callCount() {
|
||||
return prompts.size();
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
prompts.add(prompt);
|
||||
trace.accept("model");
|
||||
String reply = replier.apply(prompt);
|
||||
return new ChatResponse(List.of(new Generation(new AssistantMessage(reply), finish())), metadata(prompt, reply));
|
||||
}
|
||||
|
||||
@Override
|
||||
public Flux<ChatResponse> stream(Prompt prompt) {
|
||||
return Flux.defer(() -> {
|
||||
prompts.add(prompt);
|
||||
trace.accept("model");
|
||||
String reply = replier.apply(prompt);
|
||||
List<ChatResponse> chunks = new ArrayList<>();
|
||||
for (int i = 0; i < reply.length(); i += chunkSize) {
|
||||
String piece = reply.substring(i, Math.min(reply.length(), i + chunkSize));
|
||||
boolean last = i + chunkSize >= reply.length();
|
||||
chunks.add(last
|
||||
? new ChatResponse(List.of(finishReason ? new Generation(new AssistantMessage(piece), finish()) : new Generation(new AssistantMessage(piece))),
|
||||
streamUsage ? metadata(prompt, reply) : ChatResponseMetadata.builder().build())
|
||||
: new ChatResponse(List.of(new Generation(new AssistantMessage(piece)))));
|
||||
}
|
||||
return Flux.fromIterable(chunks);
|
||||
});
|
||||
}
|
||||
|
||||
private static ChatGenerationMetadata finish() {
|
||||
return ChatGenerationMetadata.builder().finishReason("STOP").build();
|
||||
}
|
||||
|
||||
private static ChatResponseMetadata metadata(Prompt prompt, String reply) {
|
||||
int in = 0;
|
||||
for (Message m : prompt.getInstructions()) {
|
||||
in += ESTIMATOR.estimate(m.getText());
|
||||
}
|
||||
return ChatResponseMetadata.builder().usage(new DefaultUsage(in, ESTIMATOR.estimate(reply))).build();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.model.Generation;
|
||||
import org.springframework.ai.chat.prompt.ChatOptions;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.model.tool.ToolCallingChatOptions;
|
||||
|
||||
import java.util.ArrayDeque;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Deque;
|
||||
import java.util.List;
|
||||
import java.util.Queue;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
|
||||
/**
|
||||
* A hand-written {@link ChatModel} that returns a pre-programmed queue of {@link ChatResponse}s,
|
||||
* one per {@link #call(Prompt)} invocation, instead of calling a real LLM API. {@code ChatModel}
|
||||
* has exactly one abstract method -- confirmed with {@code javap} against
|
||||
* {@code spring-ai-model-2.0.1.jar}, everything else on the interface has a default implementation
|
||||
* -- so this is the entire surface a deterministic test double needs to implement.
|
||||
*
|
||||
* <p>{@link org.springframework.ai.chat.client.advisor.ToolCallingAdvisor ToolCallingAdvisor} (and
|
||||
* its subclass, {@code ToolSearchToolCallingAdvisor}) drive the tool-calling loop by calling the
|
||||
* underlying {@link ChatModel} once per round: once to get the model's first response (which may
|
||||
* contain tool calls), then once more per round of tool results fed back in, until a response
|
||||
* comes back with no tool calls. Queuing responses here lets a test assert the exact shape of that
|
||||
* loop -- how many rounds it took, what tool calls appeared, what the final answer was -- without
|
||||
* an API key, network access, or the nondeterminism of an actual model.
|
||||
*
|
||||
* <p>Every {@link Prompt} the advisor sends is recorded in {@link #capturedPrompts()} so a test can
|
||||
* also assert on what the advisor sent back on the next round -- in particular, that a
|
||||
* {@code ToolResponseMessage} was appended after tool execution.
|
||||
*/
|
||||
public class ScriptedToolModel implements ChatModel {
|
||||
|
||||
private final Queue<ChatResponse> script;
|
||||
private final List<Prompt> capturedPrompts = new CopyOnWriteArrayList<>();
|
||||
|
||||
private ScriptedToolModel(Deque<ChatResponse> script) {
|
||||
this.script = script;
|
||||
}
|
||||
|
||||
public static Builder builder() {
|
||||
return new Builder();
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
this.capturedPrompts.add(prompt);
|
||||
ChatResponse next = this.script.poll();
|
||||
if (next == null) {
|
||||
throw new IllegalStateException(
|
||||
"ScriptedChatModel ran out of queued responses after " + this.capturedPrompts.size()
|
||||
+ " calls. Prompts so far: " + this.capturedPrompts);
|
||||
}
|
||||
return next;
|
||||
}
|
||||
|
||||
public List<Prompt> capturedPrompts() {
|
||||
return List.copyOf(this.capturedPrompts);
|
||||
}
|
||||
|
||||
/**
|
||||
* {@code DefaultChatClientUtils} builds every outgoing {@link Prompt}'s options from
|
||||
* {@code chatModel.getOptions().mutate()} -- not {@code getDefaultOptions()}, which is a
|
||||
* separate default method nobody in the request-building path actually calls. Confirmed by
|
||||
* disassembling both: {@link ChatModel#getOptions()}'s default body is a bare
|
||||
* {@code ChatOptions.builder().build()}, a plain {@link ChatOptions} that is not a
|
||||
* {@link ToolCallingChatOptions}. Since {@code ToolCallingAdvisor.adviseCall} starts with an
|
||||
* {@code instanceof ToolCallingChatOptions} check on that exact object and falls straight
|
||||
* through to the underlying model with no tool loop at all when it fails, leaving this method's
|
||||
* default in place silently turns every tool call in this repository into a no-op -- confirmed
|
||||
* the hard way, by a first version of this class that overrode {@code getDefaultOptions()}
|
||||
* instead and watched every test below get back an empty answer after exactly one model call.
|
||||
* Real providers (OpenAI, Anthropic) return their own {@code ToolCallingChatOptions}
|
||||
* implementation from {@code getOptions()} for the same reason.
|
||||
*/
|
||||
@Override
|
||||
public ChatOptions getOptions() {
|
||||
return ToolCallingChatOptions.builder().build();
|
||||
}
|
||||
|
||||
public int callCount() {
|
||||
return this.capturedPrompts.size();
|
||||
}
|
||||
|
||||
public static final class Builder {
|
||||
|
||||
private final Deque<ChatResponse> script = new ArrayDeque<>();
|
||||
|
||||
private int promptTokens = -1;
|
||||
|
||||
private int completionTokens;
|
||||
|
||||
private Builder() {
|
||||
}
|
||||
|
||||
/** Responses queued after this call report this usage, the way a provider does. */
|
||||
public Builder reportingUsage(int prompt, int completion) {
|
||||
this.promptTokens = prompt;
|
||||
this.completionTokens = completion;
|
||||
return this;
|
||||
}
|
||||
|
||||
private ChatResponse response(Generation generation) {
|
||||
if (promptTokens < 0) {
|
||||
return new ChatResponse(List.of(generation));
|
||||
}
|
||||
return ChatResponse.builder().generations(List.of(generation)).metadata(
|
||||
org.springframework.ai.chat.metadata.ChatResponseMetadata.builder()
|
||||
.usage(new org.springframework.ai.chat.metadata.DefaultUsage(promptTokens, completionTokens)).build())
|
||||
.build();
|
||||
}
|
||||
|
||||
/** Queues a plain-text final answer with no tool calls -- ends the tool-calling loop. */
|
||||
public Builder thenRespond(String text) {
|
||||
this.script.add(response(new Generation(new AssistantMessage(text))));
|
||||
return this;
|
||||
}
|
||||
|
||||
/** Queues an assistant turn that calls one or more tools, continuing the loop. */
|
||||
public Builder thenCallTools(AssistantMessage.ToolCall... toolCalls) {
|
||||
AssistantMessage message = AssistantMessage.builder()
|
||||
.content("")
|
||||
.toolCalls(List.of(toolCalls))
|
||||
.build();
|
||||
this.script.add(response(new Generation(message)));
|
||||
return this;
|
||||
}
|
||||
|
||||
public ScriptedToolModel build() {
|
||||
return new ScriptedToolModel(new ArrayDeque<>(this.script));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
|
||||
/** One-line renderings of messages so transcripts stay short and diff cleanly. */
|
||||
public final class Show {
|
||||
|
||||
private Show() {
|
||||
}
|
||||
|
||||
public static String one(Message m) {
|
||||
return String.format("%-9s %s", m.getMessageType(), m.getText());
|
||||
}
|
||||
|
||||
public static String inline(List<Message> messages) {
|
||||
return messages.stream().map(Show::brief).reduce((a, b) -> a + " | " + b).orElse("(empty)");
|
||||
}
|
||||
|
||||
private static String brief(Message m) {
|
||||
char kind = m.getMessageType().name().charAt(0);
|
||||
if (m instanceof org.springframework.ai.chat.messages.AssistantMessage a && a.hasToolCalls()) {
|
||||
return kind + ":[tool call " + a.getToolCalls().getFirst().name() + "]";
|
||||
}
|
||||
if (m instanceof org.springframework.ai.chat.messages.ToolResponseMessage r) {
|
||||
return kind + ":[tool result " + r.getResponses().getFirst().responseData() + "]";
|
||||
}
|
||||
return kind + ":" + m.getText();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import org.springframework.boot.test.context.TestConfiguration;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
|
||||
/** Puts one shared {@link RecordingModel} in the context so a test can inspect what it was sent. */
|
||||
@TestConfiguration(proxyBeanMethods = false)
|
||||
public class TestModelConfig {
|
||||
|
||||
@Bean
|
||||
RecordingModel recordingModel() {
|
||||
return new RecordingModel();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package com.ankurm.advisors.support;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.PrintWriter;
|
||||
import java.io.StringWriter;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
|
||||
/**
|
||||
* Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
|
||||
* echoes it to the console. Every console block quoted in the article comes out of one of these
|
||||
* files verbatim.
|
||||
*/
|
||||
public final class Transcript implements AutoCloseable {
|
||||
|
||||
private final Path path;
|
||||
private final StringWriter buffer = new StringWriter();
|
||||
private final PrintWriter out = new PrintWriter(buffer);
|
||||
|
||||
public Transcript(String fileName, String title) {
|
||||
this.path = Path.of("output", fileName);
|
||||
out.println("# " + title);
|
||||
out.println();
|
||||
}
|
||||
|
||||
public Transcript line(String format, Object... args) {
|
||||
out.println(args.length == 0 ? format : String.format(format, args));
|
||||
return this;
|
||||
}
|
||||
|
||||
public Transcript blank() {
|
||||
out.println();
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
out.flush();
|
||||
try {
|
||||
Files.createDirectories(path.getParent());
|
||||
Files.writeString(path, buffer.toString());
|
||||
} catch (IOException e) {
|
||||
throw new IllegalStateException("could not write " + path, e);
|
||||
}
|
||||
System.out.print(buffer);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user