Add semantic-cache module: CallAdvisor over a Redis Stack vector index with a real in-process embedding model; threshold sweep (hit rate against wrong answers), 600-request replay with an illustrative cost model, tenant filter, expiry, look-alike and embedding-model-change traps; optional real local model
Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01G8ikz8xdWuTP5yun8DZ1hk
This commit is contained in:
@@ -22,5 +22,6 @@ Runnable companion code for the Spring AI articles on [ankurm.com](https://ankur
|
||||
| [`text-to-sql/`](text-to-sql) | A question to SQL to rows, safely, on a real PostgreSQL 16: a schema prompt built through the restricted role, a JSqlParser guard (one statement, SELECT only, listed tables and functions), a read-only role with column grants and a statement timeout, a row cap, and evaluation by comparing results. Sixteen queries against four setups (13 harmful: 13 succeed with neither protection, 0 with both). The model is a script, not a language model. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Text-to-SQL with Spring AI, Done Safely](https://ankurm.com/text-to-sql-spring-ai-read-only-roles-query-validation/) |
|
||||
| [`providers/`](providers) | One ticket-summarising app on the real `OpenAiChatModel`, `AnthropicChatModel` and `GoogleGenAiChatModel` against a local server in all three wire formats: where each puts the system prompt and what it sends by default, portable versus provider-specific options (and the per-call `ChatOptions` that crashes two providers and silently resets the third), prompt caching, a cost table from the vendors' price sheets, and failover with the retry multiplication measured. No vendor API called; token counts and cache hits are simulated from documented rules; latency not measured. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Anthropic Claude vs OpenAI vs Gemini in Spring AI 2.0](https://ankurm.com/spring-ai-claude-vs-openai-vs-gemini-switch-providers-compare-cost/) |
|
||||
| [`llm-gateway/`](llm-gateway) | A gateway service on Spring AI 2.0: model-hint routing, failover under the tool-calling advisor (a provider failing mid tool loop does not re-run the tool), one circuit breaker per provider, dollar caps that reserve before the call, a token-per-hour limiter, and a tenant-keyed semantic cache. Real OpenAI and Anthropic models against a local fake; no vendor API called. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [The LLM Gateway Pattern for Java Microservices](https://ankurm.com/llm-gateway-pattern-java-microservices/) |
|
||||
| [`semantic-cache/`](semantic-cache) | A `CallAdvisor` that serves a stored answer when a new question embeds close to an old one, over a Redis Stack vector index and a real in-process embedding model (all-MiniLM-L6-v2). A labelled set of 30 intents x 4 phrasings plus 20 off-topic questions swept across thresholds (hit rate against wrong answers), 600 requests replayed with an illustrative cost model, and the traps reproduced: Spring AI's score is `(1 + cosine) / 2`, reset-password and reset-2FA are close enough to answer for each other, a tenant filter, expiry, and an index built for another embedding model. The chat model is a script that counts calls; one test uses a real local model. Spring Boot 4.1.1, Spring AI 2.0.1, Java 25. | [Semantic Caching for LLM Calls in Spring Boot with Redis Vector Search](https://ankurm.com/semantic-caching-for-llm-calls-in-spring-boot-with-redis-vector-search/) |
|
||||
|
||||
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,59 @@
|
||||
# semantic-cache
|
||||
|
||||
Companion code for [Semantic Caching for LLM Calls in Spring Boot with Redis Vector Search](https://ankurm.com/semantic-caching-for-llm-calls-in-spring-boot-with-redis-vector-search/), part of the [Spring AI series](../README.md) on ankurm.com.
|
||||
|
||||
A `CallAdvisor` that answers from Redis when a question is close enough to one answered before, and otherwise calls the model and remembers the answer. Embeddings come from a real model, all-MiniLM-L6-v2, run in-process; the cache is a Redis Stack vector index driven through Spring AI's `RedisVectorStore`. A labelled set of 30 support questions (four phrasings each) and 20 off-topic ones is used to measure hit rate against wrong answers at every threshold. Every figure in the article is quoted from a file in [`output/`](output).
|
||||
|
||||
The *chat model* in the measurements is a script that knows the right answer for each question and counts calls, so a wrong answer is detectable. Dollar figures use an **illustrative** price ($2.50 / $10.00 per million tokens) and tokens estimated as characters / 4. One test (`10-real-model-latency.txt`) puts a real local model, TinyLlama 1.1B on Ollama, behind the same advisor; it is skipped when no Ollama is listening.
|
||||
|
||||
## Versions
|
||||
|
||||
| Component | Version |
|
||||
|---|---|
|
||||
| Spring Boot | 4.1.1 |
|
||||
| Spring AI | 2.0.1 (`spring-ai-redis-store`, `spring-ai-transformers`; `spring-ai-ollama` for the optional test) |
|
||||
| Java | 25 (Temurin 25.0.4.1) |
|
||||
| Redis Stack | 7.4.0-v8 tarball: Redis 7.4.7, RediSearch 2.10.20, RedisJSON 2.8.9 |
|
||||
| Jedis | 7.4.1 |
|
||||
| Embedding model | all-MiniLM-L6-v2, 384 dimensions, ONNX via ONNX Runtime 1.21.1 and DJL tokenizers 0.36.0 |
|
||||
|
||||
## Quickstart (no Docker)
|
||||
|
||||
```bash
|
||||
scripts/services-up.sh # Redis Stack on :6393 (downloads the tarball if missing)
|
||||
scripts/run-all.sh # runs the 10 tests and regenerates output/01 .. 10
|
||||
```
|
||||
|
||||
The first run downloads the embedding model (about 87 MB) from Hugging Face into `$TMPDIR/minilm-cache`. RedisJSON must be loaded as well as RediSearch: `RedisVectorStore` stores JSON documents.
|
||||
|
||||
## What's here
|
||||
|
||||
| File | What it shows |
|
||||
|---|---|
|
||||
| [`SemanticCache.java`](src/main/java/com/ankurm/semanticcache/SemanticCache.java) | The Redis vector index behind one class: `lookup`, `put`, tenant filter, expiry |
|
||||
| [`SemanticCacheAdvisor.java`](src/main/java/com/ankurm/semanticcache/SemanticCacheAdvisor.java) | The `CallAdvisor`: lookup, short-circuit on a hit, store on a miss |
|
||||
| [`Embedder.java`](src/main/java/com/ankurm/semanticcache/Embedder.java) | MiniLM loaded from Hugging Face, plus a plain cosine for comparison |
|
||||
| [`ScriptedModel.java`](src/main/java/com/ankurm/semanticcache/ScriptedModel.java) | A chat model that knows each question's right answer and counts calls and tokens |
|
||||
| [`Dataset.java`](src/main/java/com/ankurm/semanticcache/Dataset.java), [`dataset.tsv`](src/main/resources/dataset.tsv) | 30 intents x 4 phrasings (including look-alike groups such as reset password / reset router / reset 2FA) and 20 off-topic questions |
|
||||
| [`SemanticCacheTest.java`](src/test/java/com/ankurm/semanticcache/SemanticCacheTest.java) | Every measurement; each test writes its own transcript |
|
||||
|
||||
## Output files
|
||||
|
||||
| File | Written by |
|
||||
|---|---|
|
||||
| `01-similarity-scores.txt` | `similarityScoresOfParaphrasesAndLookAlikes` (also: Spring AI's score is `(1 + cosine) / 2`) |
|
||||
| `02-advisor-in-a-chat-client.txt` | `theAdvisorServesAParaphraseWithoutCallingTheModel` |
|
||||
| `03-threshold-sweep.txt` | `thresholdSweepHitRateAgainstWrongAnswers` |
|
||||
| `04-replay-600-requests.txt` | `replayingSixHundredRequests` |
|
||||
| `05-lookalikes.txt` | `lookAlikesThatShareWordsButNotMeaning` |
|
||||
| `06-tenant-isolation.txt` | `tenantsDoNotShareAnswers` |
|
||||
| `07-expiry.txt` | `anExpiredEntryStopsBeingServed` |
|
||||
| `08-embedding-model-change.txt` | `anotherEmbeddingModelIsAnotherIndex` |
|
||||
| `09-hit-latency.txt` | `whereTheTimeGoesOnAHit` |
|
||||
| `10-real-model-latency.txt` | `againstARealLocalModel` (needs Ollama on :11555 with a model named `tl`) |
|
||||
|
||||
Counts and scores are deterministic; timings drift between runs (2-vCPU sandbox).
|
||||
|
||||
## Not covered
|
||||
|
||||
Streaming responses (the advisor is call-only), caching answers that depend on conversation history or tool results, cache warming, and an embedding model other than MiniLM: the thresholds here are for this model and this kind of short support question.
|
||||
@@ -0,0 +1,35 @@
|
||||
# What a similarity score looks like (all-MiniLM-L6-v2, cosine)
|
||||
|
||||
Plain cosine similarity (NOT the Spring AI score, see the end) of each later phrasing with the cached phrasing of the SAME question (90 pairs):
|
||||
min 0.337 p10 0.537 median 0.729 p90 0.869 max 0.916
|
||||
lowest: 0.337 Where is my package right now? -> How can I track my order?
|
||||
|
||||
Best score each of those phrasings gets against a cached question of a DIFFERENT intent (look-alikes like reset password / reset router):
|
||||
min 0.209 p10 0.311 median 0.446 p90 0.563 max 0.633
|
||||
highest: 0.633 Can you help me reset the password for my login? -> How do I reset two-factor authentication on my account?
|
||||
|
||||
Best score of 20 off-topic questions against any cached question:
|
||||
min 0.096 p10 0.133 median 0.180 p90 0.225 max 0.289
|
||||
|
||||
Does Redis report the same number? "I forgot my password, how can I get back into my account?"
|
||||
cosine computed in Java : 0.838726
|
||||
Document.getScore() from Redis : 0.919363 (matched stored question: "How do I reset my account password?")
|
||||
(1 + cosine) / 2 : 0.919363 <- the score is this, not the cosine
|
||||
same question as the cached one -> score 1.000000
|
||||
|
||||
The closest pairs of DIFFERENT cached questions (Spring AI score; 435 pairs in all). Above the threshold, one answers for the other:
|
||||
0.852 "How do I reset my account password?" / "How do I reset two-factor authentication on my account?"
|
||||
0.824 "How do I cancel an order I just placed?" / "How do I cancel a return I already requested?"
|
||||
0.796 "How do I cancel an order I just placed?" / "How do I cancel my monthly subscription?"
|
||||
0.791 "How can I track my order?" / "How can I track the return I sent back?"
|
||||
0.789 "How can I track my order?" / "Can I pick up my order at a store?"
|
||||
0.779 "How long is the warranty on your laptops?" / "How do I make a warranty claim?"
|
||||
0.775 "How long does a refund take to arrive?" / "How can I track the return I sent back?"
|
||||
0.767 "How do I update the firmware on my headphones?" / "How do I pair my headphones over Bluetooth?"
|
||||
|
||||
Spring AI's similarityThreshold is applied to that score. Threshold -> the cosine it really means (cosine = 2 x score - 1):
|
||||
similarityThreshold 0.70 = cosine 0.40
|
||||
similarityThreshold 0.80 = cosine 0.60
|
||||
similarityThreshold 0.85 = cosine 0.70
|
||||
similarityThreshold 0.90 = cosine 0.80
|
||||
similarityThreshold 0.95 = cosine 0.90
|
||||
@@ -0,0 +1,14 @@
|
||||
# SemanticCacheAdvisor in front of a ChatClient (threshold 0.80)
|
||||
|
||||
ask : How do I reset my account password?
|
||||
-> model call, model calls so far 1, answer starts "[reset-password]"
|
||||
ask : How do I reset my account password?
|
||||
-> CACHE HIT, score 1.000, model calls so far 1, answer starts "[reset-password]"
|
||||
ask : I forgot my password, how can I get back into my account?
|
||||
-> CACHE HIT, score 0.919, model calls so far 1, answer starts "[reset-password]"
|
||||
ask : How do I reset my router to factory settings?
|
||||
-> model call, model calls so far 2, answer starts "[reset-router]"
|
||||
ask : What is the capital of Australia?
|
||||
-> model call, model calls so far 3, answer starts "[ood-0]"
|
||||
|
||||
advisor counters: 2 hits, 3 misses; model received 3 calls for 5 questions; entries in Redis: 3
|
||||
@@ -0,0 +1,32 @@
|
||||
# Threshold sweep: 30 cached questions, 90 rewordings and 20 off-topic questions asked
|
||||
|
||||
threshold = the value given to SearchRequest.similarityThreshold (Spring AI score); cosine = 2 x threshold - 1.
|
||||
A served answer is RIGHT if it belongs to the intent of the question asked, WRONG if it belongs to another intent.
|
||||
An off-topic question has no right answer in the cache, so any hit on it is WRONG.
|
||||
|
||||
threshold | cosine | hit rate | right | wrong (reworded) | wrong (off-topic) | wrong / all hits | cached pairs of different intents that answer for each other
|
||||
0.50 | 0.00 | 100% | 86/90 | 4/90 | 20/20 | 21.8% | 396 of 435
|
||||
0.60 | 0.20 | 100% | 86/90 | 4/90 | 6/20 | 10.4% | 151 of 435
|
||||
0.70 | 0.40 | 99% | 86/90 | 3/90 | 0/20 | 3.4% | 27 of 435
|
||||
0.75 | 0.50 | 94% | 83/90 | 2/90 | 0/20 | 2.4% | 13 of 435
|
||||
0.80 | 0.60 | 86% | 77/90 | 0/90 | 0/20 | 0.0% | 2 of 435
|
||||
0.85 | 0.70 | 58% | 52/90 | 0/90 | 0/20 | 0.0% | 1 of 435
|
||||
0.90 | 0.80 | 26% | 23/90 | 0/90 | 0/20 | 0.0% | 0 of 435
|
||||
0.95 | 0.90 | 3% | 3/90 | 0/90 | 0/20 | 0.0% | 0 of 435
|
||||
|
||||
Questions that were answered WRONGLY at threshold 0.80 (asked -> served from): none among these 110 probes
|
||||
|
||||
Rewordings the cache MISSED at threshold 0.80 (best stored question and its score):
|
||||
0.787 "Where can I end my membership plan?" (nearest: "How do I cancel my monthly subscription?")
|
||||
0.746 "Can you tell me what qualifies for getting my money back?" (nearest: "Which items can be refunded?")
|
||||
0.696 "Where is my package right now?" (nearest: "Can I pick up my order at a store?")
|
||||
0.759 "I want to see the delivery status of my purchase" (nearest: "How can I track my order?")
|
||||
0.752 "Has my returned parcel arrived at your warehouse yet?" (nearest: "How long will delivery take?")
|
||||
0.790 "My device broke, how do I get it repaired under warranty?" (nearest: "How do I make a warranty claim?")
|
||||
0.797 "How many days until my parcel gets here?" (nearest: "How long will delivery take?")
|
||||
0.739 "Can I pay with a credit card or PayPal?" (nearest: "Which payment methods do you accept?")
|
||||
0.798 "How can I pay for my order?" (nearest: "How can I track my order?")
|
||||
0.769 "I need a receipt for my purchase, where do I find it?" (nearest: "Where can I download an invoice for my order?")
|
||||
0.743 "Is in-store collection available?" (nearest: "Can I pick up my order at a store?")
|
||||
0.753 "I would rather collect my purchase myself, is that an option?" (nearest: "Can I pick up my order at a store?")
|
||||
0.703 "Do you offer click and collect?" (nearest: "Do you sell gift cards?")
|
||||
@@ -0,0 +1,16 @@
|
||||
# 600 requests through the advisor, by threshold
|
||||
|
||||
Workload: 600 requests (seed 42), 30 intents asked with Zipf popularity in random phrasings, 15% off-topic.
|
||||
121 distinct texts, so a plain exact-match cache could answer at most 479 of 600 (80%).
|
||||
Cost model (ILLUSTRATIVE): $2.50 per million input tokens, $10.00 per million output tokens, tokens = characters / 4.
|
||||
|
||||
setup | model calls | hit rate | wrong | wrong of hits | model $ | saved
|
||||
no cache | 600 | - | - | - | 0.4564 | -
|
||||
cache, threshold 0.70 | 42 | 93.0% | 191 | 34.2% | 0.0316 | 93%
|
||||
cache, threshold 0.80 | 62 | 89.7% | 7 | 1.3% | 0.0468 | 90%
|
||||
cache, threshold 0.90 | 109 | 81.8% | 0 | 0.0% | 0.0827 | 82%
|
||||
cache, threshold 0.95 | 119 | 80.2% | 0 | 0.0% | 0.0904 | 80%
|
||||
cache, exact match only | 121 | 79.8% | 0 | 0.0% | 0.0919 | 80%
|
||||
|
||||
Average time to serve a cache hit (embed the question, search Redis, build the response): 4 ms at 0.70, 4 ms at 0.80, 3 ms at 0.90, 3 ms at 0.95
|
||||
Saved dollars exclude the cost of embedding (local model here, CPU only) and of running Redis.
|
||||
@@ -0,0 +1,10 @@
|
||||
# Pairs that read alike and mean different things (cosine similarity)
|
||||
|
||||
cached question | new question | cosine | Spring AI score = (1 + cosine) / 2
|
||||
How do I cancel my order? | How do I keep my order and not cancel it? | 0.938 | 0.969 <- served at threshold 0.90 and 0.80
|
||||
Is shipping free for orders over $50? | Is shipping free for orders over $500? | 0.848 | 0.924 <- served at threshold 0.90 and 0.80
|
||||
Do you ship to Canada? | Do you ship to Cuba? | 0.645 | 0.822 <- served at threshold 0.80
|
||||
What is the warranty on the laptop? | What is the warranty on the monitor? | 0.799 | 0.899 <- served at threshold 0.80
|
||||
Where is order 1001? | Where is order 2002? | 0.593 | 0.796
|
||||
Can I get a refund within 30 days? | Can I get a refund after 30 days? | 0.984 | 0.992 <- served at threshold 0.90 and 0.80
|
||||
I want to delete my account | I do not want to delete my account | 0.931 | 0.966 <- served at threshold 0.90 and 0.80
|
||||
@@ -0,0 +1,8 @@
|
||||
# Same question, two tenants
|
||||
|
||||
shop-a asks "How do I reset my account password?" -> model call (calls so far: 1)
|
||||
shop-b asks the identical question -> model call, no hit (model calls so far: 2)
|
||||
shop-a asks again -> cache hit
|
||||
entries in the index: 2 (one per tenant)
|
||||
The filter is a tag on the document and an expression on the search: tenant == 'shop-b'.
|
||||
nearest(shop-a, q) without asking as shop-b still finds: "How do I reset my account password?"
|
||||
@@ -0,0 +1,4 @@
|
||||
# Expiry with EXPIRE on the document's key
|
||||
|
||||
stored one entry, index holds 1 document(s); lookup -> hit
|
||||
EXPIRE sc:<id> 1, wait 1.5 s; index now holds 0 document(s); lookup -> miss
|
||||
@@ -0,0 +1,4 @@
|
||||
# An index built for 384 dimensions, written by a model with 8
|
||||
|
||||
index created by MiniLM (384 dimensions); a new application version configures an 8-dimension model against the same index name.
|
||||
add + search -> exception JedisDataException: Error parsing vector similarity query: query vector blob size (32) does not match index's expected size (1536).
|
||||
@@ -0,0 +1,5 @@
|
||||
# Where the milliseconds go on a lookup (110 questions, 30 entries cached, 2 vCPUs)
|
||||
|
||||
embed the question (MiniLM in-process) : median 2.4 ms, p95 3.6 ms
|
||||
Redis KNN search with the tenant filter : median 0.3 ms, p95 2.5 ms (lookup time minus one embedding)
|
||||
Timings drift run to run; the shape (embedding dominates, Redis is small) does not.
|
||||
@@ -0,0 +1,16 @@
|
||||
# A real local model (TinyLlama 1.1B, Q4_0, 60 tokens) behind the advisor, 2 vCPUs
|
||||
|
||||
first ask 2100 ms model "How do I reset my account password?"
|
||||
first ask 1953 ms model "How do I reset my router to factory settings?"
|
||||
first ask 5 ms CACHE "How do I reset two-factor authentication on my account?" <- served the answer to "How do I reset my account password?" (score 0.852), a DIFFERENT question
|
||||
first ask 1902 ms model "How do I cancel an order I just placed?"
|
||||
first ask 1920 ms model "How do I cancel my monthly subscription?"
|
||||
first ask 10 ms CACHE "How do I cancel a return I already requested?" <- served the answer to "How do I cancel an order I just placed?" (score 0.824), a DIFFERENT question
|
||||
reworded 17 ms CACHE "I forgot my password, how can I get back into my account?" <- How do I reset my account password?
|
||||
reworded 9 ms CACHE "What is the way to restore my router to its factory defaults?" <- How do I reset my router to factory settings?
|
||||
reworded 1941 ms model "I lost my phone, how can I reset my 2FA?"
|
||||
reworded 9 ms CACHE "I made an order by mistake, can I cancel it?" <- How do I cancel an order I just placed?
|
||||
reworded 6 ms CACHE "I want to stop my recurring subscription payments" <- How do I cancel my monthly subscription?
|
||||
reworded 2013 ms model "I changed my mind about sending the item back, can I withdraw the return?"
|
||||
|
||||
median model call 1920 ms; median cache hit 9 ms (4 of 6 rewordings hit); first-time questions wrongly answered from the cache: 2 of 6
|
||||
@@ -0,0 +1,76 @@
|
||||
<?xml version="1.0" encoding="UTF-8"?>
|
||||
<project xmlns="http://maven.apache.org/POM/4.0.0"
|
||||
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
|
||||
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
|
||||
<modelVersion>4.0.0</modelVersion>
|
||||
|
||||
<parent>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-parent</artifactId>
|
||||
<version>4.1.1</version>
|
||||
<relativePath/>
|
||||
</parent>
|
||||
|
||||
<groupId>com.ankurm</groupId>
|
||||
<artifactId>semantic-cache</artifactId>
|
||||
<version>1.0.0</version>
|
||||
<name>semantic-cache</name>
|
||||
<description>A semantic cache for LLM calls on Spring AI 2.0: a CallAdvisor over Redis vector search, a real local embedding model, threshold tuning and hit rate measured.</description>
|
||||
|
||||
<properties>
|
||||
<java.version>25</java.version>
|
||||
<spring-ai.version>2.0.1</spring-ai.version>
|
||||
</properties>
|
||||
|
||||
<dependencyManagement>
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-bom</artifactId>
|
||||
<version>${spring-ai.version}</version>
|
||||
<type>pom</type>
|
||||
<scope>import</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
</dependencyManagement>
|
||||
|
||||
<dependencies>
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-client-chat</artifactId>
|
||||
</dependency>
|
||||
<!-- The store module, not the starter: the store is built by hand so the index settings are visible. -->
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-redis-store</artifactId>
|
||||
</dependency>
|
||||
<!-- all-MiniLM-L6-v2 run in-process through ONNX Runtime and DJL. -->
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-transformers</artifactId>
|
||||
</dependency>
|
||||
<!-- Only the optional real-model test uses this; it is skipped when no Ollama is listening. -->
|
||||
<dependency>
|
||||
<groupId>org.springframework.ai</groupId>
|
||||
<artifactId>spring-ai-ollama</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
<dependency>
|
||||
<groupId>org.springframework.boot</groupId>
|
||||
<artifactId>spring-boot-starter-test</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
</dependencies>
|
||||
|
||||
<build>
|
||||
<plugins>
|
||||
<plugin>
|
||||
<groupId>org.apache.maven.plugins</groupId>
|
||||
<artifactId>maven-surefire-plugin</artifactId>
|
||||
<configuration>
|
||||
<argLine>-Duser.timezone=UTC -Dstdout.encoding=UTF-8 -Dfile.encoding=UTF-8 -Xmx2g</argLine>
|
||||
</configuration>
|
||||
</plugin>
|
||||
</plugins>
|
||||
</build>
|
||||
</project>
|
||||
Executable
+8
@@ -0,0 +1,8 @@
|
||||
#!/usr/bin/env bash
|
||||
# Regenerates every file under output/ from the test suite. Needs Redis Stack on :6393 (scripts/services-up.sh).
|
||||
# The first run downloads all-MiniLM-L6-v2 (about 87 MB) from Hugging Face into $TMPDIR/minilm-cache.
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")/.."
|
||||
rm -rf target
|
||||
mvn -q -B test 2>&1 | grep -E "Tests run:|BUILD|FAIL|ERROR" || true
|
||||
ls output
|
||||
Executable
+21
@@ -0,0 +1,21 @@
|
||||
#!/usr/bin/env bash
|
||||
# Starts what the tests need, with no Docker:
|
||||
# Redis Stack 127.0.0.1:6393 tarball: redis-stack-server 7.4.0-v8 (Redis 7.4.7 + RediSearch 2.10.20 + RedisJSON 2.8.9)
|
||||
# Ollama 127.0.0.1:11555 OPTIONAL, only for test 10 (a real local model); the test is skipped without it
|
||||
# RedisVectorStore stores documents as JSON, so plain Redis, or Redis with only the search module, is not enough:
|
||||
# FT.CREATE ... ON JSON fails with "Invalid rule type: JSON" unless RedisJSON is loaded too.
|
||||
set -euo pipefail
|
||||
TOOLS="${TOOLS:-/tmp/tools}"
|
||||
RS="$TOOLS/redis-stack-server-7.4.0-v8"
|
||||
if [ ! -d "$RS" ]; then
|
||||
mkdir -p "$TOOLS"
|
||||
curl -sL -o "$TOOLS/rs.tgz" "https://packages.redis.io/redis-stack/redis-stack-server-7.4.0-v8.jammy.x86_64.tar.gz"
|
||||
tar -xzf "$TOOLS/rs.tgz" -C "$TOOLS"
|
||||
fi
|
||||
RUN=/tmp/sc-run; mkdir -p "$RUN"
|
||||
if ! (exec 3<>/dev/tcp/127.0.0.1/6393) 2>/dev/null; then
|
||||
"$RS/bin/redis-server" --port 6393 --bind 127.0.0.1 --dir "$RUN" --save "" --daemonize yes \
|
||||
--loadmodule "$RS/lib/redisearch.so" --loadmodule "$RS/lib/rejson.so" --logfile "$RUN/redis.log" >/dev/null
|
||||
fi
|
||||
echo "redis-stack :6393"
|
||||
# Optional: OLLAMA_MODELS=<dir with a model named tl> OLLAMA_HOST=127.0.0.1:11555 ollama serve
|
||||
@@ -0,0 +1,46 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import java.io.BufferedReader;
|
||||
import java.io.InputStreamReader;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.util.ArrayList;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/** The labelled questions in {@code dataset.tsv}: 30 support intents, four phrasings each, plus 20 off-topic questions. */
|
||||
public final class Dataset {
|
||||
|
||||
public record Row(String intent, String role, String text) { }
|
||||
|
||||
public final List<Row> rows = new ArrayList<>();
|
||||
|
||||
public Dataset() {
|
||||
try (var in = Dataset.class.getResourceAsStream("/dataset.tsv");
|
||||
var r = new BufferedReader(new InputStreamReader(in, StandardCharsets.UTF_8))) {
|
||||
String line;
|
||||
while ((line = r.readLine()) != null) {
|
||||
if (line.startsWith("#") || line.isBlank()) continue;
|
||||
String[] p = line.split("\t");
|
||||
// every off-topic question is its own intent: an answer cached for one is wrong for the next
|
||||
String intent = p[0].equals("ood") ? "ood-" + rows.stream().filter(x -> x.role().equals("ood")).count() : p[0];
|
||||
rows.add(new Row(intent, p[1], p[2]));
|
||||
}
|
||||
} catch (Exception e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
|
||||
public List<Row> seeds() { return rows.stream().filter(r -> r.role().equals("seed")).toList(); }
|
||||
|
||||
public List<Row> paraphrases() { return rows.stream().filter(r -> r.role().equals("para")).toList(); }
|
||||
|
||||
public List<Row> offTopic() { return rows.stream().filter(r -> r.role().equals("ood")).toList(); }
|
||||
|
||||
/** text -> intent, so the scripted model knows which answer a question deserves. */
|
||||
public Map<String, String> intentByText() {
|
||||
Map<String, String> m = new LinkedHashMap<>();
|
||||
rows.forEach(r -> m.put(r.text(), r.intent()));
|
||||
return m;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.transformers.TransformersEmbeddingModel;
|
||||
import org.springframework.core.io.DefaultResourceLoader;
|
||||
|
||||
/**
|
||||
* all-MiniLM-L6-v2 (384 dimensions) running inside the JVM. Spring AI's default download location for
|
||||
* the model is a GitHub media URL; here both files come from Hugging Face and are cached on disk, so the
|
||||
* first run downloads about 90 MB and later runs are offline.
|
||||
*/
|
||||
public final class Embedder {
|
||||
|
||||
static final String HF = "https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/";
|
||||
|
||||
private Embedder() { }
|
||||
|
||||
public static EmbeddingModel miniLm() {
|
||||
try {
|
||||
TransformersEmbeddingModel m = new TransformersEmbeddingModel();
|
||||
m.setTokenizerResource(HF + "tokenizer.json");
|
||||
m.setModelResource(HF + "onnx/model.onnx");
|
||||
m.setResourceCacheDirectory(System.getProperty("java.io.tmpdir") + "/minilm-cache");
|
||||
m.setTokenizerOptions(java.util.Map.of("padding", "true", "truncation", "true", "maxLength", "256"));
|
||||
m.afterPropertiesSet();
|
||||
return m;
|
||||
} catch (Exception e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
|
||||
public static double cosine(float[] a, float[] b) {
|
||||
double dot = 0, na = 0, nb = 0;
|
||||
for (int i = 0; i < a.length; i++) {
|
||||
dot += a[i] * b[i];
|
||||
na += a[i] * a[i];
|
||||
nb += b[i] * b[i];
|
||||
}
|
||||
return dot / (Math.sqrt(na) * Math.sqrt(nb) + 1e-12);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import java.util.Map;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
|
||||
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.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
|
||||
import org.springframework.ai.chat.metadata.DefaultUsage;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
|
||||
/**
|
||||
* A model that knows the right answer for each dataset question and counts every call it receives. The answer
|
||||
* for an intent is a fixed paragraph; what matters to a cache is whether the right paragraph comes back.
|
||||
*/
|
||||
public final class ScriptedModel implements ChatModel {
|
||||
|
||||
/** About 90 tokens by the usual four-characters-per-token rule. Stand-in text, same length for every intent. */
|
||||
static final String BODY = " Please open the Help Centre, choose the matching topic and follow the steps shown there. "
|
||||
+ "If the page does not match your situation, contact support with your order number and we will "
|
||||
+ "pick it up from there within one business day. Keep your confirmation email, it speeds things up.";
|
||||
|
||||
private final Map<String, String> intentByText;
|
||||
public final AtomicLong calls = new AtomicLong();
|
||||
public final AtomicLong promptTokens = new AtomicLong();
|
||||
public final AtomicLong completionTokens = new AtomicLong();
|
||||
private final long delayMillis;
|
||||
|
||||
public ScriptedModel(Map<String, String> intentByText, long delayMillis) {
|
||||
this.intentByText = intentByText;
|
||||
this.delayMillis = delayMillis;
|
||||
}
|
||||
|
||||
public static String answerFor(String intent) {
|
||||
return "[" + intent + "]" + BODY;
|
||||
}
|
||||
|
||||
public static int tokens(String s) { return Math.max(1, s.length() / 4); }
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
calls.incrementAndGet();
|
||||
String q = prompt.getUserMessage().getText();
|
||||
String intent = intentByText.getOrDefault(q, "unknown");
|
||||
String answer = answerFor(intent);
|
||||
int in = tokens(prompt.getContents()), out = tokens(answer);
|
||||
promptTokens.addAndGet(in);
|
||||
completionTokens.addAndGet(out);
|
||||
if (delayMillis > 0) {
|
||||
try {
|
||||
Thread.sleep(delayMillis);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
}
|
||||
}
|
||||
return ChatResponse.builder()
|
||||
.generations(java.util.List.of(new Generation(new AssistantMessage(answer))))
|
||||
.metadata(ChatResponseMetadata.builder().usage(new DefaultUsage(in, out)).build())
|
||||
.build();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.UUID;
|
||||
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.vectorstore.SearchRequest;
|
||||
import org.springframework.ai.vectorstore.redis.RedisVectorStore;
|
||||
import redis.clients.jedis.RedisClient;
|
||||
|
||||
/**
|
||||
* A question-to-answer cache over a Redis vector index. The key is the question's embedding, the value is the
|
||||
* answer, and a lookup is a nearest-neighbour search with a similarity threshold.
|
||||
*/
|
||||
public final class SemanticCache {
|
||||
|
||||
/** A stored answer and how close the stored question was to the one asked. */
|
||||
public record Hit(String answer, String storedQuestion, double score) { }
|
||||
|
||||
public static final String INDEX = "semantic_cache";
|
||||
public static final String PREFIX = "sc:";
|
||||
|
||||
private final RedisVectorStore store;
|
||||
private final RedisClient jedis;
|
||||
private volatile double threshold;
|
||||
|
||||
public SemanticCache(RedisClient jedis, EmbeddingModel embeddings, double threshold) {
|
||||
this.jedis = jedis;
|
||||
this.threshold = threshold;
|
||||
try {
|
||||
jedis.ftDropIndex(INDEX);
|
||||
} catch (Exception ignored) {
|
||||
// no index yet
|
||||
}
|
||||
this.store = RedisVectorStore.builder(jedis, embeddings)
|
||||
.indexName(INDEX)
|
||||
.prefix(PREFIX)
|
||||
.vectorAlgorithm(RedisVectorStore.Algorithm.HNSW)
|
||||
.distanceMetric(RedisVectorStore.DistanceMetric.COSINE)
|
||||
// Only declared metadata fields are filterable, and only declared ones come back on a search result:
|
||||
// without text("answer") a hit has the stored question but no answer.
|
||||
.metadataFields(RedisVectorStore.MetadataField.tag("tenant"), RedisVectorStore.MetadataField.text("answer"))
|
||||
.initializeSchema(true)
|
||||
.build();
|
||||
this.store.afterPropertiesSet();
|
||||
}
|
||||
|
||||
public void threshold(double t) { this.threshold = t; }
|
||||
|
||||
public double threshold() { return threshold; }
|
||||
|
||||
public Optional<Hit> lookup(String tenant, String question) {
|
||||
return lookup(tenant, question, threshold);
|
||||
}
|
||||
|
||||
public Optional<Hit> lookup(String tenant, String question, double minScore) {
|
||||
List<Document> found = store.similaritySearch(SearchRequest.builder()
|
||||
.query(question)
|
||||
.topK(1)
|
||||
.similarityThreshold(minScore)
|
||||
.filterExpression("tenant == '" + tenant + "'")
|
||||
.build());
|
||||
return found.stream().findFirst().map(d -> new Hit(
|
||||
(String) d.getMetadata().get("answer"), d.getText(), d.getScore()));
|
||||
}
|
||||
|
||||
/** The nearest stored question whatever its score, for measuring rather than serving. */
|
||||
public Optional<Hit> nearest(String tenant, String question) {
|
||||
return lookup(tenant, question, 0.0);
|
||||
}
|
||||
|
||||
public String put(String tenant, String question, String answer) {
|
||||
String id = UUID.randomUUID().toString();
|
||||
store.add(List.of(new Document(id, question, Map.of("tenant", tenant, "answer", answer))));
|
||||
return id;
|
||||
}
|
||||
|
||||
/** Redis expires the whole document; the index drops it with the key. */
|
||||
public void expireAfterSeconds(String id, long seconds) {
|
||||
jedis.expire(PREFIX + id, seconds);
|
||||
}
|
||||
|
||||
public long size() {
|
||||
return jedis.ftInfo(INDEX).get("num_docs") instanceof Number n ? n.longValue()
|
||||
: Long.parseLong(String.valueOf(jedis.ftInfo(INDEX).get("num_docs")));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.atomic.AtomicLong;
|
||||
import java.util.function.Function;
|
||||
|
||||
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;
|
||||
|
||||
/**
|
||||
* Answers from the cache when a close enough question was answered before, and otherwise lets the call go on to
|
||||
* the model and remembers the answer. Registered with an order early in the chain: a hit short-circuits every
|
||||
* advisor after it, including the one that calls the model.
|
||||
*/
|
||||
public final class SemanticCacheAdvisor implements CallAdvisor {
|
||||
|
||||
public static final String HIT = "semantic-cache.hit";
|
||||
public static final String SCORE = "semantic-cache.score";
|
||||
public static final String MATCHED = "semantic-cache.matched-question";
|
||||
/** The context key an application sets to say whose cache this is. */
|
||||
public static final String TENANT = "tenant";
|
||||
|
||||
private final SemanticCache cache;
|
||||
private final int order;
|
||||
private final Function<ChatClientRequest, String> tenantOf;
|
||||
public final AtomicLong hits = new AtomicLong();
|
||||
public final AtomicLong misses = new AtomicLong();
|
||||
|
||||
public SemanticCacheAdvisor(SemanticCache cache, int order) {
|
||||
this.cache = cache;
|
||||
this.order = order;
|
||||
this.tenantOf = r -> String.valueOf(r.context().getOrDefault(TENANT, "default"));
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getName() { return "semanticCache"; }
|
||||
|
||||
@Override
|
||||
public int getOrder() { return order; }
|
||||
|
||||
@Override
|
||||
public ChatClientResponse adviseCall(ChatClientRequest request, CallAdvisorChain chain) {
|
||||
String tenant = tenantOf.apply(request);
|
||||
String question = request.prompt().getUserMessage().getText();
|
||||
|
||||
Optional<SemanticCache.Hit> hit = cache.lookup(tenant, question);
|
||||
if (hit.isPresent()) {
|
||||
hits.incrementAndGet();
|
||||
ChatResponse cached = ChatResponse.builder()
|
||||
.generations(java.util.List.of(new Generation(new AssistantMessage(hit.get().answer()))))
|
||||
.build();
|
||||
return new ChatClientResponse(cached, Map.of(HIT, true, SCORE, hit.get().score(), MATCHED, hit.get().storedQuestion()));
|
||||
}
|
||||
|
||||
misses.incrementAndGet();
|
||||
ChatClientResponse response = chain.nextCall(request);
|
||||
String answer = response.chatResponse().getResult().getOutput().getText();
|
||||
if (answer != null && !answer.isBlank()) {
|
||||
cache.put(tenant, question, answer);
|
||||
}
|
||||
return response;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
# intent<TAB>role<TAB>text (role: seed = first time asked and cached, para = later rewording, ood = out-of-domain, never cached)
|
||||
reset-password seed How do I reset my account password?
|
||||
reset-password para I forgot my password, how can I get back into my account?
|
||||
reset-password para Steps to change a lost account password
|
||||
reset-password para Can you help me reset the password for my login?
|
||||
reset-router seed How do I reset my router to factory settings?
|
||||
reset-router para What is the way to restore my router to its factory defaults?
|
||||
reset-router para My router is acting up, how do I wipe it back to factory settings?
|
||||
reset-router para Factory reset instructions for the wifi router
|
||||
reset-2fa seed How do I reset two-factor authentication on my account?
|
||||
reset-2fa para I lost my phone, how can I reset my 2FA?
|
||||
reset-2fa para Can I turn off and set up two-factor login again?
|
||||
reset-2fa para Steps to re-enroll two-factor authentication after changing phones
|
||||
cancel-order seed How do I cancel an order I just placed?
|
||||
cancel-order para I made an order by mistake, can I cancel it?
|
||||
cancel-order para Is it possible to cancel my purchase before it ships?
|
||||
cancel-order para What do I do to cancel my recent order?
|
||||
cancel-subscription seed How do I cancel my monthly subscription?
|
||||
cancel-subscription para I want to stop my recurring subscription payments
|
||||
cancel-subscription para Where can I end my membership plan?
|
||||
cancel-subscription para Please tell me how to unsubscribe from the monthly plan
|
||||
cancel-return seed How do I cancel a return I already requested?
|
||||
cancel-return para I changed my mind about sending the item back, can I withdraw the return?
|
||||
cancel-return para Can a return request be called off after submitting it?
|
||||
cancel-return para Withdraw my return request, how?
|
||||
refund-time seed How long does a refund take to arrive?
|
||||
refund-time para When will my money be returned after a refund is approved?
|
||||
refund-time para How many days until the refund shows up in my bank account?
|
||||
refund-time para Refund processing time?
|
||||
refund-policy seed Which items can be refunded?
|
||||
refund-policy para What is your refund policy for purchased products?
|
||||
refund-policy para Are all products eligible for a refund?
|
||||
refund-policy para Can you tell me what qualifies for getting my money back?
|
||||
change-email seed How do I change the email address on my account?
|
||||
change-email para I need to update my account email
|
||||
change-email para Where can I edit the email linked to my profile?
|
||||
change-email para Switch my login email to a new one, how?
|
||||
change-address seed How do I change my delivery address?
|
||||
change-address para I moved, how can I update the shipping address on file?
|
||||
change-address para Can I edit the address my parcel is sent to?
|
||||
change-address para Update my home address for deliveries
|
||||
change-phone seed How do I update the phone number on my account?
|
||||
change-phone para I have a new number, where do I change it in my profile?
|
||||
change-phone para Edit the mobile number linked to my account
|
||||
change-phone para How can I replace my contact phone number?
|
||||
track-order seed How can I track my order?
|
||||
track-order para Where is my package right now?
|
||||
track-order para I want to see the delivery status of my purchase
|
||||
track-order para Is there a tracking link for my order?
|
||||
track-return seed How can I track the return I sent back?
|
||||
track-return para Where can I see whether you received my returned item?
|
||||
track-return para Has my returned parcel arrived at your warehouse yet?
|
||||
track-return para Check the status of my return shipment
|
||||
warranty-length seed How long is the warranty on your laptops?
|
||||
warranty-length para What is the warranty period for a laptop bought from you?
|
||||
warranty-length para For how many years are laptops covered?
|
||||
warranty-length para Laptop warranty duration?
|
||||
warranty-claim seed How do I make a warranty claim?
|
||||
warranty-claim para My device broke, how do I get it repaired under warranty?
|
||||
warranty-claim para What is the process for filing a warranty claim?
|
||||
warranty-claim para I need to claim warranty on a faulty product
|
||||
shipping-cost seed How much does shipping cost?
|
||||
shipping-cost para What are the delivery charges?
|
||||
shipping-cost para Do you charge for shipping?
|
||||
shipping-cost para What will I pay to get my order delivered?
|
||||
delivery-time seed How long will delivery take?
|
||||
delivery-time para When should I expect my order to arrive?
|
||||
delivery-time para What is the usual delivery time after ordering?
|
||||
delivery-time para How many days until my parcel gets here?
|
||||
international-shipping seed Do you ship to other countries?
|
||||
international-shipping para Is international delivery available?
|
||||
international-shipping para Can I order from outside the country and have it shipped?
|
||||
international-shipping para Which countries do you deliver to?
|
||||
payment-methods seed Which payment methods do you accept?
|
||||
payment-methods para Can I pay with a credit card or PayPal?
|
||||
payment-methods para What ways of paying are available at checkout?
|
||||
payment-methods para How can I pay for my order?
|
||||
gift-cards seed Do you sell gift cards?
|
||||
gift-cards para Can I buy a gift card for a friend?
|
||||
gift-cards para Are there gift vouchers available?
|
||||
gift-cards para I would like to give someone a store gift card, is that possible?
|
||||
coupon-code seed How do I use a coupon code at checkout?
|
||||
coupon-code para Where do I enter my discount code?
|
||||
coupon-code para I have a promo code, how can I apply it to my order?
|
||||
coupon-code para Redeeming a voucher code while paying
|
||||
invoice seed Where can I download an invoice for my order?
|
||||
invoice para I need a receipt for my purchase, where do I find it?
|
||||
invoice para How do I get a VAT invoice for an order?
|
||||
invoice para Can you send me the invoice for my last order?
|
||||
opening-hours seed What are your store opening hours?
|
||||
opening-hours para When is the shop open?
|
||||
opening-hours para What time do you open and close?
|
||||
opening-hours para Hours of operation for the store?
|
||||
store-pickup seed Can I pick up my order at a store?
|
||||
store-pickup para Is in-store collection available?
|
||||
store-pickup para I would rather collect my purchase myself, is that an option?
|
||||
store-pickup para Do you offer click and collect?
|
||||
price-match seed Do you price match competitors?
|
||||
price-match para If I find it cheaper elsewhere, will you match the price?
|
||||
price-match para Is there a price matching guarantee?
|
||||
price-match para Will you lower your price to match another shop?
|
||||
battery-care seed How can I make my laptop battery last longer?
|
||||
battery-care para Tips for keeping a laptop battery healthy
|
||||
battery-care para What should I do to extend my notebook battery life?
|
||||
battery-care para How do I look after my laptop battery so it ages slowly?
|
||||
firmware-update seed How do I update the firmware on my headphones?
|
||||
firmware-update para Where can I install the newest firmware for my headphones?
|
||||
firmware-update para My headphones need a firmware upgrade, how is it done?
|
||||
firmware-update para Steps to flash new firmware onto wireless headphones
|
||||
bluetooth-pairing seed How do I pair my headphones over Bluetooth?
|
||||
bluetooth-pairing para What is the way to connect the headphones to my phone via Bluetooth?
|
||||
bluetooth-pairing para My headphones will not show up when I search for Bluetooth devices, how do I pair them?
|
||||
bluetooth-pairing para Bluetooth pairing instructions for the headphones
|
||||
screen-flicker seed Why is my monitor screen flickering?
|
||||
screen-flicker para My display keeps flickering, what could cause it?
|
||||
screen-flicker para How do I fix a flickering monitor?
|
||||
screen-flicker para The screen flashes on and off, what is wrong?
|
||||
data-transfer seed How do I transfer my files to a new laptop?
|
||||
data-transfer para What is the easiest way to move my data to a new computer?
|
||||
data-transfer para I bought a new laptop, how can I copy everything from the old one?
|
||||
data-transfer para Migrating files from an old laptop to a new one
|
||||
ood ood What is the capital of Australia?
|
||||
ood ood Write a haiku about autumn
|
||||
ood ood How many calories are in a banana?
|
||||
ood ood Who won the football world cup in 2018?
|
||||
ood ood Explain how photosynthesis works
|
||||
ood ood What is the best way to learn the guitar?
|
||||
ood ood Translate good morning into Spanish
|
||||
ood ood How tall is Mount Everest?
|
||||
ood ood Recommend a good science fiction novel
|
||||
ood ood How do I bake sourdough bread?
|
||||
ood ood What is the speed of light?
|
||||
ood ood Give me a tip for a job interview
|
||||
ood ood How do vaccines work?
|
||||
ood ood What causes the northern lights?
|
||||
ood ood How far is the moon from Earth?
|
||||
ood ood How do I meditate?
|
||||
ood ood What is a good name for a puppy?
|
||||
ood ood How do I change a flat bicycle tire?
|
||||
ood ood When did the Roman Empire fall?
|
||||
ood ood How is cheese made?
|
||||
|
@@ -0,0 +1,580 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Comparator;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Random;
|
||||
|
||||
import org.junit.jupiter.api.BeforeAll;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import redis.clients.jedis.RedisClient;
|
||||
|
||||
/**
|
||||
* Every test writes one transcript under output/ and asserts the same numbers it prints. Needs a Redis with the
|
||||
* RediSearch module on 127.0.0.1:6393 (scripts/services-up.sh starts one).
|
||||
*/
|
||||
class SemanticCacheTest {
|
||||
|
||||
static final int REDIS_PORT = 6393;
|
||||
static final String TENANT = "shop-a";
|
||||
|
||||
static EmbeddingModel embeddings;
|
||||
static RedisClient jedis;
|
||||
static Dataset data;
|
||||
static Map<String, String> intentByText;
|
||||
|
||||
@BeforeAll
|
||||
static void setUp() {
|
||||
embeddings = Embedder.miniLm();
|
||||
jedis = RedisClient.create("127.0.0.1", REDIS_PORT);
|
||||
data = new Dataset();
|
||||
intentByText = data.intentByText();
|
||||
}
|
||||
|
||||
static SemanticCache freshCache(double threshold) {
|
||||
jedis.flushAll();
|
||||
return new SemanticCache(jedis, embeddings, threshold);
|
||||
}
|
||||
|
||||
static SemanticCache seededCache(double threshold) {
|
||||
SemanticCache c = freshCache(threshold);
|
||||
data.seeds().forEach(r -> c.put(TENANT, r.text(), ScriptedModel.answerFor(r.intent())));
|
||||
return c;
|
||||
}
|
||||
|
||||
static String intentOfAnswer(String answer) {
|
||||
return answer.substring(1, answer.indexOf(']'));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 1. what a similarity score looks like
|
||||
|
||||
@Test
|
||||
void similarityScoresOfParaphrasesAndLookAlikes() {
|
||||
List<Double> paraphrase = new ArrayList<>(), lookAlike = new ArrayList<>(), offTopic = new ArrayList<>();
|
||||
Map<String, float[]> seedVec = new java.util.LinkedHashMap<>();
|
||||
data.seeds().forEach(r -> seedVec.put(r.intent(), embeddings.embed(r.text())));
|
||||
String worstLookAlike = "";
|
||||
double worstScore = -1;
|
||||
String lowestParaphrase = "";
|
||||
double lowestScore = 2;
|
||||
for (var r : data.paraphrases()) {
|
||||
float[] v = embeddings.embed(r.text());
|
||||
double own = Embedder.cosine(v, seedVec.get(r.intent()));
|
||||
paraphrase.add(own);
|
||||
if (own < lowestScore) {
|
||||
lowestScore = own;
|
||||
lowestParaphrase = r.text() + " -> " + seedText(r.intent());
|
||||
}
|
||||
for (var e : seedVec.entrySet()) {
|
||||
if (e.getKey().equals(r.intent())) continue;
|
||||
double s = Embedder.cosine(v, e.getValue());
|
||||
if (s > worstScore) {
|
||||
worstScore = s;
|
||||
worstLookAlike = r.text() + " -> " + seedText(e.getKey());
|
||||
}
|
||||
}
|
||||
double best = seedVec.entrySet().stream().filter(e -> !e.getKey().equals(r.intent()))
|
||||
.mapToDouble(e -> Embedder.cosine(v, e.getValue())).max().orElse(0);
|
||||
lookAlike.add(best);
|
||||
}
|
||||
for (var r : data.offTopic()) {
|
||||
float[] v = embeddings.embed(r.text());
|
||||
offTopic.add(seedVec.values().stream().mapToDouble(s -> Embedder.cosine(v, s)).max().orElse(0));
|
||||
}
|
||||
|
||||
// Which different intents are closest to each other? Every pair of cached seeds, 30 x 29 / 2 = 435 pairs.
|
||||
record Pair(String a, String b, double score) { }
|
||||
List<Pair> pairs = new ArrayList<>();
|
||||
List<String> seedIntents = new ArrayList<>(seedVec.keySet());
|
||||
for (int i = 0; i < seedIntents.size(); i++)
|
||||
for (int j = i + 1; j < seedIntents.size(); j++)
|
||||
pairs.add(new Pair(seedIntents.get(i), seedIntents.get(j),
|
||||
(1 + Embedder.cosine(seedVec.get(seedIntents.get(i)), seedVec.get(seedIntents.get(j)))) / 2));
|
||||
pairs.sort(Comparator.comparingDouble(Pair::score).reversed());
|
||||
|
||||
// Does Redis report the same number? Ask for one question and compare with the cosine computed in Java.
|
||||
SemanticCache cache = seededCache(0.0);
|
||||
var probe = data.paraphrases().get(0);
|
||||
var redisHit = cache.nearest(TENANT, probe.text()).orElseThrow();
|
||||
double javaCosine = Embedder.cosine(embeddings.embed(probe.text()), seedVec.get(probe.intent()));
|
||||
var exact = cache.nearest(TENANT, data.seeds().get(0).text()).orElseThrow();
|
||||
|
||||
try (var t = new Transcript("01-similarity-scores.txt", "What a similarity score looks like (all-MiniLM-L6-v2, cosine)")) {
|
||||
t.line("Plain cosine similarity (NOT the Spring AI score, see the end) of each later phrasing with the cached phrasing of the SAME question (%d pairs):", paraphrase.size());
|
||||
t.line(" %s", summary(paraphrase));
|
||||
t.line(" lowest: %.3f %s", lowestScore, lowestParaphrase);
|
||||
t.blank();
|
||||
t.line("Best score each of those phrasings gets against a cached question of a DIFFERENT intent (look-alikes like reset password / reset router):");
|
||||
t.line(" %s", summary(lookAlike));
|
||||
t.line(" highest: %.3f %s", worstScore, worstLookAlike);
|
||||
t.blank();
|
||||
t.line("Best score of %d off-topic questions against any cached question:", offTopic.size());
|
||||
t.line(" %s", summary(offTopic));
|
||||
t.blank();
|
||||
t.line("Does Redis report the same number? \"%s\"", probe.text());
|
||||
t.line(" cosine computed in Java : %.6f", javaCosine);
|
||||
t.line(" Document.getScore() from Redis : %.6f (matched stored question: \"%s\")", redisHit.score(), redisHit.storedQuestion());
|
||||
t.line(" (1 + cosine) / 2 : %.6f <- the score is this, not the cosine", (1 + javaCosine) / 2);
|
||||
t.line(" same question as the cached one -> score %.6f", exact.score());
|
||||
t.blank();
|
||||
t.line("The closest pairs of DIFFERENT cached questions (Spring AI score; 435 pairs in all). Above the threshold, one answers for the other:");
|
||||
for (int i = 0; i < 8; i++)
|
||||
t.line(" %.3f \"%s\" / \"%s\"", pairs.get(i).score(), seedText(pairs.get(i).a()), seedText(pairs.get(i).b()));
|
||||
t.blank();
|
||||
t.line("Spring AI's similarityThreshold is applied to that score. Threshold -> the cosine it really means (cosine = 2 x score - 1):");
|
||||
for (double th : new double[]{0.70, 0.80, 0.85, 0.90, 0.95}) t.line(" similarityThreshold %.2f = cosine %.2f", th, 2 * th - 1);
|
||||
}
|
||||
assertThat(redisHit.score()).isCloseTo((1 + javaCosine) / 2, org.assertj.core.data.Offset.offset(1e-3));
|
||||
assertThat(exact.score()).isGreaterThan(0.999);
|
||||
// the two populations overlap: the reason a threshold cannot be both safe and generous
|
||||
assertThat(max(lookAlike)).isGreaterThan(min(paraphrase));
|
||||
assertThat(pairs.get(0).score()).isGreaterThan(0.80);
|
||||
}
|
||||
|
||||
static String seedText(String intent) {
|
||||
return data.seeds().stream().filter(r -> r.intent().equals(intent)).findFirst().orElseThrow().text();
|
||||
}
|
||||
|
||||
static double min(List<Double> v) { return v.stream().mapToDouble(Double::doubleValue).min().orElse(0); }
|
||||
|
||||
static double max(List<Double> v) { return v.stream().mapToDouble(Double::doubleValue).max().orElse(0); }
|
||||
|
||||
static String summary(List<Double> v) {
|
||||
List<Double> s = v.stream().sorted().toList();
|
||||
return String.format("min %.3f p10 %.3f median %.3f p90 %.3f max %.3f", s.get(0), s.get((int) (s.size() * 0.1)),
|
||||
s.get(s.size() / 2), s.get((int) (s.size() * 0.9) - 1), s.get(s.size() - 1));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 2. the advisor in a ChatClient
|
||||
|
||||
@Test
|
||||
void theAdvisorServesAParaphraseWithoutCallingTheModel() {
|
||||
SemanticCache cache = freshCache(0.80);
|
||||
ScriptedModel model = new ScriptedModel(intentByText, 0);
|
||||
SemanticCacheAdvisor advisor = new SemanticCacheAdvisor(cache, 0);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
|
||||
String[][] asks = {
|
||||
{"reset-password", "How do I reset my account password?"},
|
||||
{"reset-password", "How do I reset my account password?"},
|
||||
{"reset-password", "I forgot my password, how can I get back into my account?"},
|
||||
{"reset-router", "How do I reset my router to factory settings?"},
|
||||
{"ood", "What is the capital of Australia?"},
|
||||
};
|
||||
try (var t = new Transcript("02-advisor-in-a-chat-client.txt", "SemanticCacheAdvisor in front of a ChatClient (threshold 0.80)")) {
|
||||
for (String[] a : asks) {
|
||||
long before = model.calls.get();
|
||||
var resp = client.prompt().user(a[1]).call().chatClientResponse();
|
||||
boolean hit = Boolean.TRUE.equals(resp.context().get(SemanticCacheAdvisor.HIT));
|
||||
String answer = resp.chatResponse().getResult().getOutput().getText();
|
||||
t.line("ask : %s", a[1]);
|
||||
t.line(" -> %s%s, model calls so far %d, answer starts \"%s\"", hit ? "CACHE HIT, score " : "model call",
|
||||
hit ? String.format("%.3f", (Double) resp.context().get(SemanticCacheAdvisor.SCORE)) : "", model.calls.get(),
|
||||
answer.substring(0, answer.indexOf(']') + 1));
|
||||
assertThat(model.calls.get() - before).isEqualTo(hit ? 0 : 1);
|
||||
}
|
||||
t.blank();
|
||||
t.line("advisor counters: %d hits, %d misses; model received %d calls for %d questions; entries in Redis: %d",
|
||||
advisor.hits.get(), advisor.misses.get(), model.calls.get(), asks.length, cache.size());
|
||||
}
|
||||
assertThat(advisor.hits.get()).isEqualTo(2);
|
||||
assertThat(model.calls.get()).isEqualTo(3);
|
||||
assertThat(cache.size()).isEqualTo(3);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 3. the threshold sweep
|
||||
|
||||
record Probe(String intent, String text, boolean offTopic, SemanticCache.Hit nearest) {
|
||||
String nearestIntent() { return nearest == null ? "-" : intentOfAnswer(nearest.answer()); }
|
||||
}
|
||||
|
||||
static List<Probe> probes(SemanticCache cache) {
|
||||
List<Probe> out = new ArrayList<>();
|
||||
for (var r : data.paraphrases()) out.add(new Probe(r.intent(), r.text(), false, cache.nearest(TENANT, r.text()).orElse(null)));
|
||||
for (var r : data.offTopic()) out.add(new Probe(r.intent(), r.text(), true, cache.nearest(TENANT, r.text()).orElse(null)));
|
||||
return out;
|
||||
}
|
||||
|
||||
@Test
|
||||
void thresholdSweepHitRateAgainstWrongAnswers() {
|
||||
SemanticCache cache = seededCache(0.0);
|
||||
List<Probe> probes = probes(cache);
|
||||
long nPara = probes.stream().filter(p -> !p.offTopic()).count();
|
||||
long nOff = probes.size() - nPara;
|
||||
double[] thresholds = {0.50, 0.60, 0.70, 0.75, 0.80, 0.85, 0.90, 0.95};
|
||||
List<float[]> sv = data.seeds().stream().map(r -> embeddings.embed(r.text())).toList();
|
||||
int[] confusable = new int[thresholds.length];
|
||||
for (int i = 0; i < sv.size(); i++)
|
||||
for (int j = i + 1; j < sv.size(); j++) {
|
||||
double sc = (1 + Embedder.cosine(sv.get(i), sv.get(j))) / 2;
|
||||
for (int k = 0; k < thresholds.length; k++) if (sc >= thresholds[k]) confusable[k]++;
|
||||
}
|
||||
try (var t = new Transcript("03-threshold-sweep.txt", "Threshold sweep: 30 cached questions, 90 rewordings and 20 off-topic questions asked")) {
|
||||
t.line("threshold = the value given to SearchRequest.similarityThreshold (Spring AI score); cosine = 2 x threshold - 1.");
|
||||
t.line("A served answer is RIGHT if it belongs to the intent of the question asked, WRONG if it belongs to another intent.");
|
||||
t.line("An off-topic question has no right answer in the cache, so any hit on it is WRONG.");
|
||||
t.blank();
|
||||
t.line("%9s | %6s | %8s | %9s | %14s | %15s | %s | %s", "threshold", "cosine", "hit rate", "right", "wrong (reworded)", "wrong (off-topic)", "wrong / all hits", "cached pairs of different intents that answer for each other");
|
||||
for (int ti = 0; ti < thresholds.length; ti++) {
|
||||
double th = thresholds[ti];
|
||||
long hits = 0, right = 0, wrongPara = 0, wrongOff = 0, paraHits = 0;
|
||||
for (Probe p : probes) {
|
||||
if (p.nearest() == null || p.nearest().score() < th) continue;
|
||||
hits++;
|
||||
boolean ok = !p.offTopic() && p.nearestIntent().equals(p.intent());
|
||||
if (!p.offTopic()) paraHits++;
|
||||
if (ok) right++;
|
||||
else if (p.offTopic()) wrongOff++;
|
||||
else wrongPara++;
|
||||
}
|
||||
t.line("%9.2f | %6.2f | %7.0f%% | %5d/%-3d | %8d/%-6d | %9d/%-7d | %s | %s", th, 2 * th - 1, 100.0 * paraHits / nPara, right, nPara, wrongPara, nPara,
|
||||
wrongOff, nOff, hits == 0 ? "-" : String.format("%.1f%%", 100.0 * (wrongPara + wrongOff) / hits), confusable[ti] + " of 435");
|
||||
}
|
||||
t.blank();
|
||||
t.line("Questions that were answered WRONGLY at threshold 0.80 (asked -> served from): %s", probes.stream().anyMatch(p -> p.nearest().score() >= 0.80 && (p.offTopic() || !p.nearestIntent().equals(p.intent()))) ? "" : "none among these 110 probes");
|
||||
for (Probe p : probes) {
|
||||
if (p.nearest() != null && p.nearest().score() >= 0.80 && (p.offTopic() || !p.nearestIntent().equals(p.intent())))
|
||||
t.line(" %.3f \"%s\" -> \"%s\"", p.nearest().score(), p.text(), p.nearest().storedQuestion());
|
||||
}
|
||||
t.blank();
|
||||
t.line("Rewordings the cache MISSED at threshold 0.80 (best stored question and its score):");
|
||||
for (Probe p : probes) {
|
||||
if (!p.offTopic() && (p.nearest() == null || p.nearest().score() < 0.80))
|
||||
t.line(" %.3f \"%s\" (nearest: \"%s\")", p.nearest().score(), p.text(), p.nearest().storedQuestion());
|
||||
}
|
||||
}
|
||||
// the two ends of the dial
|
||||
long hitsAt50 = probes.stream().filter(p -> p.nearest().score() >= 0.50).count();
|
||||
long offHitsAt50 = probes.stream().filter(p -> p.offTopic() && p.nearest().score() >= 0.50).count();
|
||||
long wrongAt95 = probes.stream().filter(p -> p.nearest().score() >= 0.95 && (p.offTopic() || !p.nearestIntent().equals(p.intent()))).count();
|
||||
assertThat(hitsAt50).isGreaterThan(nPara);
|
||||
assertThat(offHitsAt50).isGreaterThan(0);
|
||||
assertThat(wrongAt95).isZero();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 4. replay a workload
|
||||
|
||||
record Replay(double threshold, int requests, long modelCalls, long hits, long wrong, long promptTokens, long completionTokens,
|
||||
double avgLookupMillis) { }
|
||||
|
||||
static Replay replay(double threshold, List<Data> workload) {
|
||||
SemanticCache cache = freshCache(threshold);
|
||||
ScriptedModel model = new ScriptedModel(intentByText, 0);
|
||||
SemanticCacheAdvisor advisor = new SemanticCacheAdvisor(cache, 0);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
long wrong = 0;
|
||||
long lookupNanos = 0;
|
||||
for (Data d : workload) {
|
||||
long t0 = System.nanoTime();
|
||||
var resp = client.prompt().user(d.text()).call().chatClientResponse();
|
||||
long elapsed = System.nanoTime() - t0;
|
||||
boolean hit = Boolean.TRUE.equals(resp.context().get(SemanticCacheAdvisor.HIT));
|
||||
if (hit) {
|
||||
lookupNanos += elapsed;
|
||||
String served = intentOfAnswer(resp.chatResponse().getResult().getOutput().getText());
|
||||
if (!served.equals(d.intent())) wrong++;
|
||||
}
|
||||
}
|
||||
double avgHitMillis = advisor.hits.get() == 0 ? 0 : lookupNanos / 1e6 / advisor.hits.get();
|
||||
return new Replay(threshold, workload.size(), model.calls.get(), advisor.hits.get(), wrong, model.promptTokens.get(),
|
||||
model.completionTokens.get(), avgHitMillis);
|
||||
}
|
||||
|
||||
record Data(String intent, String text) { }
|
||||
|
||||
/** 600 requests: popular intents asked far more often, each time in a random one of its four phrasings, 15% off-topic. */
|
||||
static List<Data> workload() {
|
||||
Random rnd = new Random(42);
|
||||
List<String> intents = data.seeds().stream().map(Dataset.Row::intent).toList();
|
||||
Map<String, List<String>> phrasings = new java.util.LinkedHashMap<>();
|
||||
data.rows.stream().filter(r -> !r.role().equals("ood")).forEach(r -> phrasings.computeIfAbsent(r.intent(), k -> new ArrayList<>()).add(r.text()));
|
||||
double[] weight = new double[intents.size()];
|
||||
double total = 0;
|
||||
for (int i = 0; i < weight.length; i++) {
|
||||
weight[i] = 1.0 / (i + 1); // Zipf: the first intent is asked 30 times as often as the last
|
||||
total += weight[i];
|
||||
}
|
||||
List<Data> out = new ArrayList<>();
|
||||
List<Dataset.Row> off = data.offTopic();
|
||||
for (int n = 0; n < 600; n++) {
|
||||
if (rnd.nextDouble() < 0.15) {
|
||||
var r = off.get(rnd.nextInt(off.size()));
|
||||
out.add(new Data(r.intent(), r.text()));
|
||||
continue;
|
||||
}
|
||||
double x = rnd.nextDouble() * total;
|
||||
int i = 0;
|
||||
while (i < weight.length - 1 && (x -= weight[i]) > 0) i++;
|
||||
List<String> ph = phrasings.get(intents.get(i));
|
||||
out.add(new Data(intents.get(i), ph.get(rnd.nextInt(ph.size()))));
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
static final double PRICE_IN_PER_MILLION = 2.50, PRICE_OUT_PER_MILLION = 10.00; // illustrative, not a vendor's price sheet
|
||||
|
||||
static double dollars(long in, long out) {
|
||||
return in / 1e6 * PRICE_IN_PER_MILLION + out / 1e6 * PRICE_OUT_PER_MILLION;
|
||||
}
|
||||
|
||||
@Test
|
||||
void replayingSixHundredRequests() {
|
||||
List<Data> workload = workload();
|
||||
long distinctTexts = workload.stream().map(Data::text).distinct().count();
|
||||
long exactRepeats = workload.size() - distinctTexts;
|
||||
ScriptedModel baseline = new ScriptedModel(intentByText, 0);
|
||||
ChatClient plain = ChatClient.builder(baseline).build();
|
||||
workload.forEach(d -> plain.prompt().user(d.text()).call().content());
|
||||
double baseCost = dollars(baseline.promptTokens.get(), baseline.completionTokens.get());
|
||||
|
||||
List<Replay> runs = new ArrayList<>();
|
||||
for (double th : new double[]{0.70, 0.80, 0.90, 0.95, 0.9999}) runs.add(replay(th, workload));
|
||||
|
||||
try (var t = new Transcript("04-replay-600-requests.txt", "600 requests through the advisor, by threshold")) {
|
||||
t.line("Workload: 600 requests (seed 42), 30 intents asked with Zipf popularity in random phrasings, 15% off-topic.");
|
||||
t.line(" %d distinct texts, so a plain exact-match cache could answer at most %d of 600 (%.0f%%).", distinctTexts, exactRepeats, 100.0 * exactRepeats / 600);
|
||||
t.line("Cost model (ILLUSTRATIVE): $%.2f per million input tokens, $%.2f per million output tokens, tokens = characters / 4.", PRICE_IN_PER_MILLION, PRICE_OUT_PER_MILLION);
|
||||
t.blank();
|
||||
t.line("%-22s | %10s | %8s | %6s | %13s | %9s | %s", "setup", "model calls", "hit rate", "wrong", "wrong of hits", "model $", "saved");
|
||||
t.line("%-22s | %10d | %8s | %6s | %13s | %9.4f | %s", "no cache", baseline.calls.get(), "-", "-", "-", baseCost, "-");
|
||||
for (Replay r : runs) {
|
||||
double cost = dollars(r.promptTokens, r.completionTokens);
|
||||
t.line("%-22s | %10d | %7.1f%% | %6d | %12.1f%% | %9.4f | %.0f%%",
|
||||
r.threshold >= 0.999 ? "cache, exact match only" : String.format("cache, threshold %.2f", r.threshold), r.modelCalls,
|
||||
100.0 * r.hits / r.requests, r.wrong, r.hits == 0 ? 0 : 100.0 * r.wrong / r.hits, cost, 100.0 * (1 - cost / baseCost));
|
||||
}
|
||||
t.blank();
|
||||
t.line("Average time to serve a cache hit (embed the question, search Redis, build the response): %s",
|
||||
runs.stream().filter(r -> r.threshold < 0.999).map(r -> String.format("%.0f ms at %.2f", r.avgLookupMillis, r.threshold)).reduce((a, b) -> a + ", " + b).orElse(""));
|
||||
t.line("Saved dollars exclude the cost of embedding (local model here, CPU only) and of running Redis.");
|
||||
}
|
||||
Replay at80 = runs.get(1), exact = runs.get(4);
|
||||
assertThat(baseline.calls.get()).isEqualTo(600);
|
||||
assertThat(exact.hits).isEqualTo(exactRepeats);
|
||||
assertThat(at80.hits).isGreaterThan(exact.hits);
|
||||
assertThat(runs.get(0).wrong).isGreaterThanOrEqualTo(at80.wrong);
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 5. the traps
|
||||
|
||||
@Test
|
||||
void lookAlikesThatShareWordsButNotMeaning() {
|
||||
String[][] pairs = {
|
||||
{"How do I cancel my order?", "How do I keep my order and not cancel it?"},
|
||||
{"Is shipping free for orders over $50?", "Is shipping free for orders over $500?"},
|
||||
{"Do you ship to Canada?", "Do you ship to Cuba?"},
|
||||
{"What is the warranty on the laptop?", "What is the warranty on the monitor?"},
|
||||
{"Where is order 1001?", "Where is order 2002?"},
|
||||
{"Can I get a refund within 30 days?", "Can I get a refund after 30 days?"},
|
||||
{"I want to delete my account", "I do not want to delete my account"},
|
||||
};
|
||||
try (var t = new Transcript("05-lookalikes.txt", "Pairs that read alike and mean different things (cosine similarity)")) {
|
||||
t.line("%-45s | %-45s | %6s | %s", "cached question", "new question", "cosine", "Spring AI score = (1 + cosine) / 2");
|
||||
for (String[] p : pairs) {
|
||||
double c = Embedder.cosine(embeddings.embed(p[0]), embeddings.embed(p[1]));
|
||||
double s = (1 + c) / 2;
|
||||
t.line("%-45s | %-45s | %6.3f | %.3f %s", p[0], p[1], c, s, s >= 0.90 ? "<- served at threshold 0.90 and 0.80" : s >= 0.80 ? "<- served at threshold 0.80" : "");
|
||||
assertThat(s).isBetween(0.0, 1.0);
|
||||
}
|
||||
}
|
||||
long servedAt90 = java.util.Arrays.stream(pairs).filter(p -> (1 + Embedder.cosine(embeddings.embed(p[0]), embeddings.embed(p[1]))) / 2 >= 0.90).count();
|
||||
assertThat(servedAt90).isGreaterThan(0);
|
||||
}
|
||||
|
||||
@Test
|
||||
void tenantsDoNotShareAnswers() {
|
||||
SemanticCache cache = freshCache(0.80);
|
||||
ScriptedModel model = new ScriptedModel(intentByText, 0);
|
||||
SemanticCacheAdvisor advisor = new SemanticCacheAdvisor(cache, 0);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
String q = "How do I reset my account password?";
|
||||
client.prompt().user(q).advisors(a -> a.param(SemanticCacheAdvisor.TENANT, "shop-a")).call().content();
|
||||
var other = client.prompt().user(q).advisors(a -> a.param(SemanticCacheAdvisor.TENANT, "shop-b")).call().chatClientResponse();
|
||||
var same = client.prompt().user(q).advisors(a -> a.param(SemanticCacheAdvisor.TENANT, "shop-a")).call().chatClientResponse();
|
||||
// the same lookup without the tenant filter, to show what the filter prevents
|
||||
var unfiltered = cache.nearest("shop-a", q);
|
||||
try (var t = new Transcript("06-tenant-isolation.txt", "Same question, two tenants")) {
|
||||
t.line("shop-a asks \"%s\" -> model call (calls so far: 1)", q);
|
||||
t.line("shop-b asks the identical question -> %s (model calls so far: %d)",
|
||||
Boolean.TRUE.equals(other.context().get(SemanticCacheAdvisor.HIT)) ? "CACHE HIT (leak)" : "model call, no hit", model.calls.get());
|
||||
t.line("shop-a asks again -> %s", Boolean.TRUE.equals(same.context().get(SemanticCacheAdvisor.HIT)) ? "cache hit" : "model call");
|
||||
t.line("entries in the index: %d (one per tenant)", cache.size());
|
||||
t.line("The filter is a tag on the document and an expression on the search: tenant == 'shop-b'.");
|
||||
t.line("nearest(shop-a, q) without asking as shop-b still finds: %s", unfiltered.map(h -> "\"" + h.storedQuestion() + "\"").orElse("nothing"));
|
||||
}
|
||||
assertThat(Boolean.TRUE.equals(other.context().get(SemanticCacheAdvisor.HIT))).isFalse();
|
||||
assertThat(Boolean.TRUE.equals(same.context().get(SemanticCacheAdvisor.HIT))).isTrue();
|
||||
assertThat(cache.size()).isEqualTo(2);
|
||||
}
|
||||
|
||||
@Test
|
||||
void anExpiredEntryStopsBeingServed() throws Exception {
|
||||
SemanticCache cache = freshCache(0.80);
|
||||
String id = cache.put(TENANT, "How do I reset my account password?", ScriptedModel.answerFor("reset-password"));
|
||||
boolean before = cache.lookup(TENANT, "How do I reset my account password?").isPresent();
|
||||
long sizeBefore = cache.size();
|
||||
cache.expireAfterSeconds(id, 1);
|
||||
Thread.sleep(1500);
|
||||
boolean after = cache.lookup(TENANT, "How do I reset my account password?").isPresent();
|
||||
try (var t = new Transcript("07-expiry.txt", "Expiry with EXPIRE on the document's key")) {
|
||||
t.line("stored one entry, index holds %d document(s); lookup -> %s", sizeBefore, before ? "hit" : "miss");
|
||||
t.line("EXPIRE sc:<id> 1, wait 1.5 s; index now holds %d document(s); lookup -> %s", cache.size(), after ? "hit" : "miss");
|
||||
}
|
||||
assertThat(before).isTrue();
|
||||
assertThat(after).isFalse();
|
||||
assertThat(cache.size()).isZero();
|
||||
}
|
||||
|
||||
@Test
|
||||
void anotherEmbeddingModelIsAnotherIndex() {
|
||||
SemanticCache cache = freshCache(0.80);
|
||||
cache.put(TENANT, "How do I reset my account password?", "a");
|
||||
EmbeddingModel eight = new EmbeddingModel() {
|
||||
@Override
|
||||
public org.springframework.ai.embedding.EmbeddingResponse call(org.springframework.ai.embedding.EmbeddingRequest r) {
|
||||
List<org.springframework.ai.embedding.Embedding> out = new ArrayList<>();
|
||||
for (int i = 0; i < r.getInstructions().size(); i++) out.add(new org.springframework.ai.embedding.Embedding(embed(r.getInstructions().get(i)), i));
|
||||
return new org.springframework.ai.embedding.EmbeddingResponse(out);
|
||||
}
|
||||
|
||||
@Override
|
||||
public float[] embed(org.springframework.ai.document.Document d) {
|
||||
return new float[8];
|
||||
}
|
||||
|
||||
@Override
|
||||
public float[] embed(String text) {
|
||||
float[] v = new float[8];
|
||||
v[Math.abs(text.hashCode()) % 8] = 1f;
|
||||
return v;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<float[]> embed(List<String> texts) {
|
||||
return texts.stream().map(this::embed).toList();
|
||||
}
|
||||
|
||||
@Override
|
||||
public int dimensions() { return 8; }
|
||||
};
|
||||
String outcome;
|
||||
try {
|
||||
// same index name and prefix, a model with 8 dimensions instead of 384; afterPropertiesSet() finds the index already exists
|
||||
var swapped = new org.springframework.ai.vectorstore.redis.RedisVectorStore[1];
|
||||
swapped[0] = org.springframework.ai.vectorstore.redis.RedisVectorStore.builder(jedis, eight)
|
||||
.indexName(SemanticCache.INDEX).prefix(SemanticCache.PREFIX)
|
||||
.metadataFields(org.springframework.ai.vectorstore.redis.RedisVectorStore.MetadataField.tag("tenant"), org.springframework.ai.vectorstore.redis.RedisVectorStore.MetadataField.text("answer"))
|
||||
.initializeSchema(true).build();
|
||||
swapped[0].afterPropertiesSet();
|
||||
swapped[0].add(List.of(new org.springframework.ai.document.Document("x", "hello", Map.of("tenant", TENANT, "answer", "b"))));
|
||||
var found = swapped[0].similaritySearch(org.springframework.ai.vectorstore.SearchRequest.builder().query("hello").topK(1)
|
||||
.filterExpression("tenant == '" + TENANT + "'").build());
|
||||
outcome = "no exception; search returned " + found.size() + " document(s): " + found.stream().map(d -> d.getText()).toList();
|
||||
} catch (Exception e) {
|
||||
Throwable root = e;
|
||||
while (root.getCause() != null) root = root.getCause();
|
||||
outcome = "exception " + root.getClass().getSimpleName() + ": " + root.getMessage();
|
||||
}
|
||||
try (var t = new Transcript("08-embedding-model-change.txt", "An index built for 384 dimensions, written by a model with 8")) {
|
||||
t.line("index created by MiniLM (384 dimensions); a new application version configures an 8-dimension model against the same index name.");
|
||||
t.line("add + search -> %s", outcome);
|
||||
}
|
||||
assertThat(outcome).isNotBlank();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 9. where the milliseconds go
|
||||
|
||||
static double pct(List<Double> v, double p) {
|
||||
List<Double> s = v.stream().sorted().toList();
|
||||
return s.get(Math.min(s.size() - 1, (int) Math.ceil(p * s.size()) - 1));
|
||||
}
|
||||
|
||||
@Test
|
||||
void whereTheTimeGoesOnAHit() {
|
||||
SemanticCache cache = seededCache(0.80);
|
||||
List<String> questions = new ArrayList<>();
|
||||
data.paraphrases().forEach(r -> questions.add(r.text()));
|
||||
data.offTopic().forEach(r -> questions.add(r.text()));
|
||||
for (int i = 0; i < 10; i++) embeddings.embed(questions.get(i)); // warm-up, not measured
|
||||
List<Double> embed = new ArrayList<>(), search = new ArrayList<>();
|
||||
for (String q : questions) {
|
||||
long t0 = System.nanoTime();
|
||||
embeddings.embed(q);
|
||||
long t1 = System.nanoTime();
|
||||
cache.lookup(TENANT, q); // embeds again inside the store, then searches
|
||||
long t2 = System.nanoTime();
|
||||
embed.add((t1 - t0) / 1e6);
|
||||
search.add((t2 - t1) / 1e6 - (t1 - t0) / 1e6);
|
||||
}
|
||||
try (var t = new Transcript("09-hit-latency.txt", "Where the milliseconds go on a lookup (110 questions, 30 entries cached, 2 vCPUs)")) {
|
||||
t.line("embed the question (MiniLM in-process) : median %.1f ms, p95 %.1f ms", pct(embed, 0.5), pct(embed, 0.95));
|
||||
t.line("Redis KNN search with the tenant filter : median %.1f ms, p95 %.1f ms (lookup time minus one embedding)", pct(search, 0.5), pct(search, 0.95));
|
||||
t.line("Timings drift run to run; the shape (embedding dominates, Redis is small) does not.");
|
||||
}
|
||||
assertThat(pct(embed, 0.5)).isGreaterThan(pct(search, 0.5));
|
||||
}
|
||||
|
||||
@Test
|
||||
void againstARealLocalModel() throws Exception {
|
||||
org.junit.jupiter.api.Assumptions.assumeTrue(isUp(11555), "no Ollama on 127.0.0.1:11555");
|
||||
var api = org.springframework.ai.ollama.api.OllamaApi.builder().baseUrl("http://127.0.0.1:11555").build();
|
||||
var model = org.springframework.ai.ollama.OllamaChatModel.builder().ollamaApi(api)
|
||||
.options(org.springframework.ai.ollama.api.OllamaChatOptions.builder().model("tl").temperature(0.0).numPredict(60).build()).build();
|
||||
SemanticCache cache = freshCache(0.80);
|
||||
SemanticCacheAdvisor advisor = new SemanticCacheAdvisor(cache, 0);
|
||||
ChatClient client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
client.prompt().user("Say hello").call().content(); // load the model, not measured
|
||||
cache = freshCache(0.80);
|
||||
advisor = new SemanticCacheAdvisor(cache, 0);
|
||||
client = ChatClient.builder(model).defaultAdvisors(advisor).build();
|
||||
|
||||
List<Dataset.Row> seeds = data.seeds().subList(0, 6);
|
||||
List<Dataset.Row> paras = new ArrayList<>();
|
||||
for (var s : seeds) paras.add(data.paraphrases().stream().filter(p -> p.intent().equals(s.intent())).findFirst().orElseThrow());
|
||||
List<Double> missMs = new ArrayList<>(), hitMs = new ArrayList<>();
|
||||
int wrongHits = 0;
|
||||
try (var t = new Transcript("10-real-model-latency.txt", "A real local model (TinyLlama 1.1B, Q4_0, 60 tokens) behind the advisor, 2 vCPUs")) {
|
||||
for (var s : seeds) {
|
||||
long t0 = System.nanoTime();
|
||||
var resp = client.prompt().user(s.text()).call().chatClientResponse();
|
||||
double ms = (System.nanoTime() - t0) / 1e6;
|
||||
boolean hit = Boolean.TRUE.equals(resp.context().get(SemanticCacheAdvisor.HIT));
|
||||
if (hit) {
|
||||
t.line("first ask %6.0f ms CACHE \"%s\" <- served the answer to \"%s\" (score %.3f), a DIFFERENT question", ms, s.text(),
|
||||
resp.context().get(SemanticCacheAdvisor.MATCHED), (Double) resp.context().get(SemanticCacheAdvisor.SCORE));
|
||||
wrongHits++;
|
||||
} else {
|
||||
missMs.add(ms);
|
||||
t.line("first ask %6.0f ms model \"%s\"", ms, s.text());
|
||||
}
|
||||
}
|
||||
for (var p : paras) {
|
||||
long t0 = System.nanoTime();
|
||||
var resp = client.prompt().user(p.text()).call().chatClientResponse();
|
||||
double ms = (System.nanoTime() - t0) / 1e6;
|
||||
boolean hit = Boolean.TRUE.equals(resp.context().get(SemanticCacheAdvisor.HIT));
|
||||
if (hit) hitMs.add(ms);
|
||||
t.line("reworded %6.0f ms %s \"%s\"%s", ms, hit ? "CACHE " : "model ", p.text(),
|
||||
hit ? " <- " + resp.context().get(SemanticCacheAdvisor.MATCHED) : "");
|
||||
}
|
||||
t.blank();
|
||||
t.line("median model call %.0f ms; median cache hit %.0f ms (%d of %d rewordings hit); first-time questions wrongly answered from the cache: %d of %d",
|
||||
pct(missMs, 0.5), hitMs.isEmpty() ? 0 : pct(hitMs, 0.5), hitMs.size(), paras.size(), wrongHits, seeds.size());
|
||||
}
|
||||
assertThat(hitMs).isNotEmpty();
|
||||
assertThat(pct(hitMs, 0.5)).isLessThan(pct(missMs, 0.5));
|
||||
}
|
||||
|
||||
static boolean isUp(int port) {
|
||||
try (var s = new java.net.Socket("127.0.0.1", port)) {
|
||||
return true;
|
||||
} catch (Exception e) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package com.ankurm.semanticcache;
|
||||
|
||||
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/} and echoes it. Every console block in the article is one of these files. */
|
||||
final class Transcript implements AutoCloseable {
|
||||
|
||||
private final Path path;
|
||||
private final StringWriter buffer = new StringWriter();
|
||||
private final PrintWriter out = new PrintWriter(buffer);
|
||||
|
||||
Transcript(String fileName, String title) {
|
||||
this.path = Path.of("output", fileName);
|
||||
out.println("# " + title);
|
||||
out.println();
|
||||
}
|
||||
|
||||
Transcript line(String format, Object... args) {
|
||||
out.println(args.length == 0 ? format : String.format(format, args));
|
||||
return this;
|
||||
}
|
||||
|
||||
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