Add rag module: Spring AI 2.0 RAG with pgvector, chunking, reranking and a faithfulness check
Co-Authored-By: Claude Sonnet 5 <[email protected]> Claude-Session: https://claude.ai/code/session_01B38FGKKam5SCGgwgduVAh3
This commit is contained in:
@@ -0,0 +1,18 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import org.springframework.boot.SpringApplication;
|
||||
import org.springframework.boot.autoconfigure.SpringBootApplication;
|
||||
import org.springframework.boot.context.properties.ConfigurationPropertiesScan;
|
||||
|
||||
/**
|
||||
* Entry point. The pipeline is described in
|
||||
* <a href="../../../../../../docs/01-the-shape-of-a-rag-pipeline.md">chapter 1</a>.
|
||||
*/
|
||||
@SpringBootApplication
|
||||
@ConfigurationPropertiesScan
|
||||
public class RagApplication {
|
||||
|
||||
public static void main(String[] args) {
|
||||
SpringApplication.run(RagApplication.class, args);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package com.ankurm.rag.chunk;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
|
||||
/**
|
||||
* Cuts text at the coarsest boundary that keeps a piece under {@code maxChars}: a blank line first,
|
||||
* then a line break, then a sentence end, then a space, and only as a last resort in the middle of a
|
||||
* word. Neighbouring chunks share up to {@code overlapChars} characters, which
|
||||
* {@link org.springframework.ai.transformer.splitter.TokenTextSplitter} does not offer.
|
||||
*
|
||||
* <p>Sizes are in characters, not tokens. See
|
||||
* <a href="../../../../../../../docs/02-chunking.md">chapter 2</a>.
|
||||
*/
|
||||
public class RecursiveChunker implements DocumentTransformer {
|
||||
|
||||
/** Coarsest first. Each separator stays attached to the end of the piece it closes. */
|
||||
private static final List<String> SEPARATORS = List.of("\n\n", "\n", ". ", " ");
|
||||
|
||||
private final int maxChars;
|
||||
private final int overlapChars;
|
||||
|
||||
public RecursiveChunker(int maxChars, int overlapChars) {
|
||||
if (overlapChars >= maxChars) {
|
||||
throw new IllegalArgumentException("overlap must be smaller than the chunk size");
|
||||
}
|
||||
this.maxChars = maxChars;
|
||||
this.overlapChars = overlapChars;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<Document> apply(List<Document> documents) {
|
||||
List<Document> chunks = new ArrayList<>();
|
||||
for (Document document : documents) {
|
||||
List<String> texts = merge(split(document.getText(), 0));
|
||||
for (int i = 0; i < texts.size(); i++) {
|
||||
Map<String, Object> metadata = new HashMap<>(document.getMetadata());
|
||||
metadata.put("chunk_index", i);
|
||||
metadata.put("chunk_total", texts.size());
|
||||
chunks.add(new Document(texts.get(i), metadata));
|
||||
}
|
||||
}
|
||||
return chunks;
|
||||
}
|
||||
|
||||
/** Breaks {@code text} into pieces that are each at most {@code maxChars} long. */
|
||||
private List<String> split(String text, int level) {
|
||||
if (text.length() <= maxChars) {
|
||||
return List.of(text);
|
||||
}
|
||||
List<String> pieces = new ArrayList<>();
|
||||
if (level >= SEPARATORS.size()) {
|
||||
for (int start = 0; start < text.length(); start += maxChars) {
|
||||
pieces.add(text.substring(start, Math.min(text.length(), start + maxChars)));
|
||||
}
|
||||
return pieces;
|
||||
}
|
||||
String separator = SEPARATORS.get(level);
|
||||
for (String part : text.split("(?<=" + Pattern.quote(separator) + ")")) {
|
||||
pieces.addAll(split(part, level + 1));
|
||||
}
|
||||
return pieces;
|
||||
}
|
||||
|
||||
/** Packs consecutive pieces into chunks, starting each new chunk with the tail of the last one. */
|
||||
private List<String> merge(List<String> pieces) {
|
||||
List<String> chunks = new ArrayList<>();
|
||||
StringBuilder current = new StringBuilder();
|
||||
for (String piece : pieces) {
|
||||
if (current.length() + piece.length() > maxChars && !current.isEmpty()) {
|
||||
chunks.add(current.toString().strip());
|
||||
String tail = overlapTail(current, Math.min(overlapChars, maxChars - piece.length()));
|
||||
current = new StringBuilder(tail);
|
||||
}
|
||||
current.append(piece);
|
||||
}
|
||||
if (!current.toString().isBlank()) {
|
||||
chunks.add(current.toString().strip());
|
||||
}
|
||||
return chunks;
|
||||
}
|
||||
|
||||
/** The last {@code n} characters, moved forward to a word start so a chunk never begins mid-word. */
|
||||
private static String overlapTail(CharSequence text, int n) {
|
||||
if (n <= 0) {
|
||||
return "";
|
||||
}
|
||||
String tail = text.subSequence(Math.max(0, text.length() - n), text.length()).toString();
|
||||
int space = tail.indexOf(' ');
|
||||
return space < 0 ? "" : tail.substring(space + 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.rag.chunk;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
|
||||
/**
|
||||
* Cuts where the topic changes: every sentence is embedded, and a new chunk starts wherever two
|
||||
* neighbouring sentences are further apart than {@code maxDistance} (cosine distance, 0 = identical).
|
||||
*
|
||||
* <p>All sentences of a document go to the embedding model in one batched call. Whether the cuts are
|
||||
* good depends on the embedding model; see
|
||||
* <a href="../../../../../../../docs/02-chunking.md">chapter 2</a>.
|
||||
*/
|
||||
public class SemanticChunker implements DocumentTransformer {
|
||||
|
||||
private final EmbeddingModel embeddingModel;
|
||||
private final double maxDistance;
|
||||
|
||||
public SemanticChunker(EmbeddingModel embeddingModel, double maxDistance) {
|
||||
this.embeddingModel = embeddingModel;
|
||||
this.maxDistance = maxDistance;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<Document> apply(List<Document> documents) {
|
||||
List<Document> chunks = new ArrayList<>();
|
||||
for (Document document : documents) {
|
||||
String[] sentences = document.getText().strip().split("(?<=[.!?])\\s+");
|
||||
if (sentences.length == 0) {
|
||||
continue;
|
||||
}
|
||||
List<float[]> vectors = embeddingModel.embed(List.of(sentences));
|
||||
|
||||
int start = 0;
|
||||
for (int i = 1; i <= sentences.length; i++) {
|
||||
boolean end = i == sentences.length;
|
||||
if (end || 1.0 - cosine(vectors.get(i - 1), vectors.get(i)) > maxDistance) {
|
||||
Map<String, Object> metadata = new HashMap<>(document.getMetadata());
|
||||
metadata.put("chunk_index", chunks.size());
|
||||
chunks.add(new Document(String.join(" ", List.of(sentences).subList(start, i)), metadata));
|
||||
start = i;
|
||||
}
|
||||
}
|
||||
}
|
||||
return chunks;
|
||||
}
|
||||
|
||||
static double cosine(float[] a, float[] b) {
|
||||
double dot = 0, normA = 0, normB = 0;
|
||||
for (int i = 0; i < a.length; i++) {
|
||||
dot += a[i] * b[i];
|
||||
normA += a[i] * a[i];
|
||||
normB += b[i] * b[i];
|
||||
}
|
||||
return dot / (Math.sqrt(normA) * Math.sqrt(normB) + 1e-12);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package com.ankurm.rag.config;
|
||||
|
||||
import com.ankurm.rag.query.LlmReranker;
|
||||
import io.micrometer.core.instrument.MeterRegistry;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.evaluation.FactCheckingEvaluator;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
import org.springframework.ai.rag.advisor.RetrievalAugmentationAdvisor;
|
||||
import org.springframework.ai.rag.generation.augmentation.ContextualQueryAugmenter;
|
||||
import org.springframework.ai.rag.retrieval.search.VectorStoreDocumentRetriever;
|
||||
import org.springframework.ai.transformer.splitter.TokenTextSplitter;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.context.annotation.Configuration;
|
||||
|
||||
/** Wires the pipeline. Each bean is one stage; replace one and the others do not notice. */
|
||||
@Configuration
|
||||
public class RagConfig {
|
||||
|
||||
/** Stage 1, chunking. Swap in {@code RecursiveChunker} or {@code SemanticChunker} here. */
|
||||
@Bean
|
||||
DocumentTransformer chunker(RagProperties properties) {
|
||||
return TokenTextSplitter.builder()
|
||||
.withChunkSize(properties.chunking().chunkSize())
|
||||
.withMinChunkSizeChars(properties.chunking().minChunkSizeChars())
|
||||
.build();
|
||||
}
|
||||
|
||||
/** Stage 3, reranking. Built from the builder so it gets its own client, not the RAG one. */
|
||||
@Bean
|
||||
LlmReranker reranker(ChatClient.Builder builder, RagProperties properties, MeterRegistry metrics) {
|
||||
return new LlmReranker(builder.build(), properties.rerank().topN(), metrics);
|
||||
}
|
||||
|
||||
/** Stages 2 to 4: retrieve, rerank, and put the survivors into the prompt. */
|
||||
@Bean
|
||||
RetrievalAugmentationAdvisor retrievalAdvisor(VectorStore vectorStore, LlmReranker reranker,
|
||||
RagProperties properties) {
|
||||
return RetrievalAugmentationAdvisor.builder()
|
||||
.documentRetriever(VectorStoreDocumentRetriever.builder()
|
||||
.vectorStore(vectorStore)
|
||||
.topK(properties.retrieval().topK())
|
||||
.similarityThreshold(properties.retrieval().similarityThreshold())
|
||||
.build())
|
||||
.documentPostProcessors(reranker)
|
||||
.queryAugmenter(ContextualQueryAugmenter.builder().allowEmptyContext(false).build())
|
||||
.build();
|
||||
}
|
||||
|
||||
/** The chat client that answers questions. The advisor runs on every call it makes. */
|
||||
@Bean
|
||||
ChatClient ragChatClient(ChatClient.Builder builder, RetrievalAugmentationAdvisor advisor) {
|
||||
return builder.defaultAdvisors(advisor).build();
|
||||
}
|
||||
|
||||
/** The judge for the faithfulness check. */
|
||||
@Bean
|
||||
FactCheckingEvaluator factChecker(ChatClient.Builder builder) {
|
||||
return FactCheckingEvaluator.builder(builder).build();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package com.ankurm.rag.config;
|
||||
|
||||
import org.springframework.boot.context.properties.ConfigurationProperties;
|
||||
import org.springframework.boot.context.properties.bind.DefaultValue;
|
||||
|
||||
/**
|
||||
* The numbers you tune. Every default is a starting point, not a measured optimum: this repository
|
||||
* has no relevance benchmark, so nothing here claims to be the best value for your documents.
|
||||
*/
|
||||
@ConfigurationProperties("rag")
|
||||
public record RagProperties(
|
||||
@DefaultValue Retrieval retrieval,
|
||||
@DefaultValue Rerank rerank,
|
||||
@DefaultValue Chunking chunking) {
|
||||
|
||||
/** How wide the first search is, and how similar a chunk must be to be a candidate at all. */
|
||||
public record Retrieval(@DefaultValue("20") int topK, @DefaultValue("0.0") double similarityThreshold) {
|
||||
}
|
||||
|
||||
/** How many chunks survive reranking and reach the prompt. */
|
||||
public record Rerank(@DefaultValue("5") int topN) {
|
||||
}
|
||||
|
||||
/** Settings for {@link org.springframework.ai.transformer.splitter.TokenTextSplitter}. */
|
||||
public record Chunking(@DefaultValue("512") int chunkSize, @DefaultValue("350") int minChunkSizeChars) {
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package com.ankurm.rag.ingest;
|
||||
|
||||
/** What {@link IngestionService} did with one file. */
|
||||
public record IngestionResult(String filename, String status, int chunksWritten, int chunksReplaced) {
|
||||
|
||||
public static IngestionResult ingested(String filename, int chunks) {
|
||||
return new IngestionResult(filename, "ingested", chunks, 0);
|
||||
}
|
||||
|
||||
public static IngestionResult updated(String filename, int chunks, int replaced) {
|
||||
return new IngestionResult(filename, "updated", chunks, replaced);
|
||||
}
|
||||
|
||||
public static IngestionResult skipped(String filename) {
|
||||
return new IngestionResult(filename, "skipped", 0, 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package com.ankurm.rag.ingest;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.security.MessageDigest;
|
||||
import java.security.NoSuchAlgorithmException;
|
||||
import java.util.HashMap;
|
||||
import java.util.HexFormat;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import io.micrometer.core.instrument.MeterRegistry;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
import org.springframework.ai.reader.pdf.PagePdfDocumentReader;
|
||||
import org.springframework.ai.reader.pdf.config.PdfDocumentReaderConfig;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
/**
|
||||
* Reads a PDF page by page, tags every page with where it came from, cuts pages into chunks and
|
||||
* writes the chunks to the vector store.
|
||||
*
|
||||
* <p>Ingesting the same bytes twice does nothing. Ingesting a changed file deletes the chunks the
|
||||
* old version produced before adding the new ones. See
|
||||
* <a href="../../../../../../../docs/03-ingestion.md">chapter 3</a>.
|
||||
*/
|
||||
@Service
|
||||
public class IngestionService {
|
||||
|
||||
/** Metadata key carrying the file a chunk came from; the delete-by-source filter depends on it. */
|
||||
public static final String SOURCE_FILE = "source_file";
|
||||
|
||||
private final VectorStore vectorStore;
|
||||
private final DocumentTransformer chunker;
|
||||
private final IngestionTracker tracker;
|
||||
private final MeterRegistry metrics;
|
||||
|
||||
public IngestionService(VectorStore vectorStore, DocumentTransformer chunker,
|
||||
IngestionTracker tracker, MeterRegistry metrics) {
|
||||
this.vectorStore = vectorStore;
|
||||
this.chunker = chunker;
|
||||
this.tracker = tracker;
|
||||
this.metrics = metrics;
|
||||
}
|
||||
|
||||
public IngestionResult ingestPdf(Resource pdf, String filename, Map<String, Object> callerMetadata) {
|
||||
String hash = sha256(pdf);
|
||||
var known = tracker.find(filename);
|
||||
if (known.isPresent() && known.get().hash().equals(hash)) {
|
||||
metrics.counter("rag.ingestion.skipped").increment();
|
||||
return IngestionResult.skipped(filename);
|
||||
}
|
||||
|
||||
List<Document> pages = new PagePdfDocumentReader(pdf, PdfDocumentReaderConfig.builder()
|
||||
.withPagesPerDocument(1)
|
||||
.build()).get();
|
||||
|
||||
List<Document> tagged = pages.stream().map(page -> {
|
||||
Map<String, Object> metadata = new HashMap<>(page.getMetadata());
|
||||
metadata.put(SOURCE_FILE, filename);
|
||||
metadata.put("source_hash", hash);
|
||||
metadata.putAll(callerMetadata);
|
||||
return new Document(tidy(page.getText()), metadata);
|
||||
}).toList();
|
||||
|
||||
List<Document> chunks = chunker.apply(tagged);
|
||||
|
||||
int replaced = 0;
|
||||
if (known.isPresent()) {
|
||||
// The file changed: remove what the old version wrote, by id, before adding the new chunks.
|
||||
replaced = known.get().chunkIds().size();
|
||||
vectorStore.delete(known.get().chunkIds());
|
||||
}
|
||||
vectorStore.add(chunks);
|
||||
tracker.record(filename, hash, chunks.stream().map(Document::getId).toList());
|
||||
|
||||
metrics.counter("rag.chunks.ingested").increment(chunks.size());
|
||||
return known.isPresent()
|
||||
? IngestionResult.updated(filename, chunks.size(), replaced)
|
||||
: IngestionResult.ingested(filename, chunks.size());
|
||||
}
|
||||
|
||||
/**
|
||||
* The PDF reader pads text with runs of spaces to reproduce the page layout. Those runs cost tokens
|
||||
* and change what gets embedded, so collapse them and the blank-line stacks around them.
|
||||
*/
|
||||
public static String tidy(String text) {
|
||||
return text.replaceAll("[ \\t]+", " ")
|
||||
.replaceAll(" ?\\n ?", "\n")
|
||||
.replaceAll("\\n{3,}", "\n\n")
|
||||
.strip();
|
||||
}
|
||||
|
||||
private static String sha256(Resource resource) {
|
||||
try {
|
||||
return HexFormat.of().formatHex(MessageDigest.getInstance("SHA-256")
|
||||
.digest(resource.getInputStream().readAllBytes()));
|
||||
} catch (IOException | NoSuchAlgorithmException e) {
|
||||
throw new IllegalStateException("cannot hash " + resource, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package com.ankurm.rag.ingest;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
import org.springframework.stereotype.Component;
|
||||
|
||||
/**
|
||||
* Remembers, per file name, the SHA-256 of the bytes that were ingested and the ids of the chunks
|
||||
* they produced.
|
||||
*
|
||||
* <p>It lives in memory, so a restart forgets everything and the next start re-embeds every file.
|
||||
* Back it with a table if that matters to you; the chunk ids are what make a targeted delete
|
||||
* possible when a file changes.
|
||||
*/
|
||||
@Component
|
||||
public class IngestionTracker {
|
||||
|
||||
public record Entry(String hash, List<String> chunkIds) {
|
||||
}
|
||||
|
||||
private final Map<String, Entry> entries = new ConcurrentHashMap<>();
|
||||
|
||||
public Optional<Entry> find(String filename) {
|
||||
return Optional.ofNullable(entries.get(filename));
|
||||
}
|
||||
|
||||
public void record(String filename, String hash, List<String> chunkIds) {
|
||||
entries.put(filename, new Entry(hash, List.copyOf(chunkIds)));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package com.ankurm.rag.query;
|
||||
|
||||
import java.util.Comparator;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.ExecutionException;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.Future;
|
||||
|
||||
import io.micrometer.core.instrument.MeterRegistry;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.rag.Query;
|
||||
import org.springframework.ai.rag.postretrieval.document.DocumentPostProcessor;
|
||||
|
||||
/**
|
||||
* Asks a chat model to rate each candidate chunk from 0 to 10 for the question, keeps the best
|
||||
* {@code topN}, and returns them best first.
|
||||
*
|
||||
* <p>Spring AI ships the hook ({@link DocumentPostProcessor}) but no reranker, so this is ours. It
|
||||
* costs one model call per candidate. The calls run on virtual threads so they wait together rather
|
||||
* than one after another. See <a href="../../../../../../../docs/04-retrieval-and-reranking.md">chapter 4</a>.
|
||||
*/
|
||||
public class LlmReranker implements DocumentPostProcessor {
|
||||
|
||||
/** Metadata key holding the score the model gave. */
|
||||
public static final String SCORE = "rerank_score";
|
||||
|
||||
private static final String PROMPT = """
|
||||
Rate how well the PASSAGE helps answer the QUESTION, from 0 (irrelevant) to 10 (answers it).
|
||||
Reply with a single integer and nothing else.
|
||||
QUESTION: {question}
|
||||
PASSAGE: {passage}
|
||||
SCORE:""";
|
||||
|
||||
private final ChatClient client;
|
||||
private final int topN;
|
||||
private final MeterRegistry metrics;
|
||||
|
||||
public LlmReranker(ChatClient client, int topN, MeterRegistry metrics) {
|
||||
this.client = client;
|
||||
this.topN = topN;
|
||||
this.metrics = metrics;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<Document> process(Query query, List<Document> candidates) {
|
||||
try (var executor = Executors.newVirtualThreadPerTaskExecutor()) {
|
||||
List<Future<Document>> futures = candidates.stream()
|
||||
.map(candidate -> executor.submit(() -> withScore(query.text(), candidate)))
|
||||
.toList();
|
||||
return futures.stream()
|
||||
.map(LlmReranker::join)
|
||||
.sorted(Comparator.comparingInt((Document d) -> (int) d.getMetadata().get(SCORE)).reversed())
|
||||
.limit(topN)
|
||||
.toList();
|
||||
}
|
||||
}
|
||||
|
||||
private Document withScore(String question, Document candidate) {
|
||||
int score = 0;
|
||||
try {
|
||||
String reply = client.prompt()
|
||||
.user(u -> u.text(PROMPT).param("question", question).param("passage", candidate.getText()))
|
||||
.call().content();
|
||||
score = Math.max(0, Math.min(10, Integer.parseInt(reply.strip())));
|
||||
} catch (RuntimeException e) {
|
||||
// An unparseable or failed rating counts as 0 rather than failing the whole question.
|
||||
metrics.counter("rag.rerank.failures").increment();
|
||||
}
|
||||
return candidate.mutate().metadata(SCORE, score).build();
|
||||
}
|
||||
|
||||
private static Document join(Future<Document> future) {
|
||||
try {
|
||||
return future.get();
|
||||
} catch (InterruptedException | ExecutionException e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package com.ankurm.rag.query;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import io.micrometer.core.instrument.MeterRegistry;
|
||||
import io.micrometer.core.instrument.Timer;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.ChatClientResponse;
|
||||
import org.springframework.ai.chat.evaluation.FactCheckingEvaluator;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.evaluation.EvaluationRequest;
|
||||
import org.springframework.ai.rag.advisor.RetrievalAugmentationAdvisor;
|
||||
import org.springframework.ai.rag.retrieval.search.VectorStoreDocumentRetriever;
|
||||
import org.springframework.ai.vectorstore.filter.Filter;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
/**
|
||||
* Asks one question: retrieve, rerank and augment (all inside the {@link RetrievalAugmentationAdvisor}
|
||||
* the chat client was built with), generate, then check the answer against the chunks it was given.
|
||||
*
|
||||
* <p>See <a href="../../../../../../../docs/05-generation-and-the-faithfulness-check.md">chapter 5</a>.
|
||||
*/
|
||||
@Service
|
||||
public class RagQueryService {
|
||||
|
||||
private final ChatClient chatClient;
|
||||
private final FactCheckingEvaluator factChecker;
|
||||
private final MeterRegistry metrics;
|
||||
|
||||
public RagQueryService(ChatClient ragChatClient, FactCheckingEvaluator factChecker, MeterRegistry metrics) {
|
||||
this.chatClient = ragChatClient;
|
||||
this.factChecker = factChecker;
|
||||
this.metrics = metrics;
|
||||
}
|
||||
|
||||
/** @param filter restricts the search, for example to one tenant; {@code null} searches everything */
|
||||
public RagResponse ask(String question, Filter.Expression filter) {
|
||||
Timer.Sample sample = Timer.start(metrics);
|
||||
String status = "error";
|
||||
try {
|
||||
ChatClientResponse response = chatClient.prompt()
|
||||
.user(question)
|
||||
.advisors(a -> {
|
||||
if (filter != null) {
|
||||
a.param(VectorStoreDocumentRetriever.FILTER_EXPRESSION, filter);
|
||||
}
|
||||
})
|
||||
.call()
|
||||
.chatClientResponse();
|
||||
|
||||
String answer = response.chatResponse().getResult().getOutput().getText();
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Document> context = (List<Document>) response.context()
|
||||
.getOrDefault(RetrievalAugmentationAdvisor.DOCUMENT_CONTEXT, List.of());
|
||||
metrics.summary("rag.context.chunks").record(context.size());
|
||||
|
||||
List<RagResponse.Source> sources = context.stream().map(RagQueryService::toSource).toList();
|
||||
if (context.isEmpty()) {
|
||||
status = "no_context";
|
||||
return new RagResponse(answer, sources, false, status);
|
||||
}
|
||||
boolean grounded = isGrounded(question, answer, context);
|
||||
status = grounded ? "answered" : "ungrounded";
|
||||
return new RagResponse(answer, sources, grounded, status);
|
||||
} finally {
|
||||
sample.stop(metrics.timer("rag.query.duration", "status", status));
|
||||
metrics.counter("rag.queries", "status", status).increment();
|
||||
}
|
||||
}
|
||||
|
||||
private boolean isGrounded(String question, String answer, List<Document> context) {
|
||||
try {
|
||||
return factChecker.evaluate(new EvaluationRequest(question, context, answer)).isPass();
|
||||
} catch (RuntimeException e) {
|
||||
// If the judge itself fails, the answer is unchecked, which is not the same as supported.
|
||||
metrics.counter("rag.faithfulness.judge_failures").increment();
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
private static RagResponse.Source toSource(Document d) {
|
||||
String text = d.getText();
|
||||
Object score = d.getMetadata().get(LlmReranker.SCORE);
|
||||
return new RagResponse.Source(
|
||||
String.valueOf(d.getMetadata().getOrDefault("source_file", "unknown")),
|
||||
d.getMetadata().get("page_number"),
|
||||
score instanceof Integer i ? i : -1,
|
||||
text.substring(0, Math.min(80, text.length())));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
package com.ankurm.rag.query;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* The answer, where it came from, and whether the guard found it supported by that context.
|
||||
* {@code status} is {@code answered}, {@code ungrounded} or {@code no_context}.
|
||||
*/
|
||||
public record RagResponse(String answer, List<Source> sources, boolean grounded, String status) {
|
||||
|
||||
/** One chunk that reached the prompt. {@code rerankScore} is the 0-10 rating, or -1 if not reranked. */
|
||||
public record Source(String file, Object page, int rerankScore, String preview) {
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package com.ankurm.rag.web;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.ingest.IngestionResult;
|
||||
import com.ankurm.rag.ingest.IngestionService;
|
||||
import com.ankurm.rag.query.RagQueryService;
|
||||
import com.ankurm.rag.query.RagResponse;
|
||||
import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder;
|
||||
import org.springframework.core.io.ByteArrayResource;
|
||||
import org.springframework.web.bind.annotation.PostMapping;
|
||||
import org.springframework.web.bind.annotation.RequestBody;
|
||||
import org.springframework.web.bind.annotation.RequestMapping;
|
||||
import org.springframework.web.bind.annotation.RequestParam;
|
||||
import org.springframework.web.bind.annotation.RestController;
|
||||
import org.springframework.web.multipart.MultipartFile;
|
||||
|
||||
/** Two endpoints: upload a PDF, ask a question. */
|
||||
@RestController
|
||||
@RequestMapping("/api")
|
||||
public class RagController {
|
||||
|
||||
public record QueryRequest(String question, String tenantId) {
|
||||
}
|
||||
|
||||
private final IngestionService ingestion;
|
||||
private final RagQueryService queries;
|
||||
|
||||
public RagController(IngestionService ingestion, RagQueryService queries) {
|
||||
this.ingestion = ingestion;
|
||||
this.queries = queries;
|
||||
}
|
||||
|
||||
@PostMapping("/ingest")
|
||||
public IngestionResult ingest(@RequestParam("file") MultipartFile file,
|
||||
@RequestParam(defaultValue = "default") String tenantId,
|
||||
@RequestParam(defaultValue = "general") String docType) throws IOException {
|
||||
return ingestion.ingestPdf(new ByteArrayResource(file.getBytes()), file.getOriginalFilename(),
|
||||
Map.of("tenant_id", tenantId, "doc_type", docType));
|
||||
}
|
||||
|
||||
@PostMapping("/query")
|
||||
public RagResponse query(@RequestBody QueryRequest request) {
|
||||
// The tenant goes in as a value of an expression, never as text glued into a filter string.
|
||||
var filter = request.tenantId() == null ? null : new FilterExpressionBuilder()
|
||||
.eq("tenant_id", request.tenantId()).build();
|
||||
return queries.ask(request.question(), filter);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
spring:
|
||||
application:
|
||||
name: spring-ai-rag
|
||||
threads:
|
||||
virtual:
|
||||
enabled: true
|
||||
datasource:
|
||||
url: ${RAG_PG_URL:jdbc:postgresql://localhost:5432/ragdb}
|
||||
username: ${RAG_PG_USER:raguser}
|
||||
password: ${RAG_PG_PASSWORD:ragpass}
|
||||
ai:
|
||||
openai:
|
||||
api-key: ${OPENAI_API_KEY}
|
||||
# In 2.0.1 the "options" level (chat.options.model) is marked deprecated in the jars' metadata; the
|
||||
# settings now sit directly under chat: and embedding:. Spring Boot ignores keys it does not know,
|
||||
# so a misspelt key fails nothing: ConfigKeysTest checks every spring.ai key against the jars.
|
||||
chat:
|
||||
model: gpt-4o
|
||||
temperature: 0.1
|
||||
embedding:
|
||||
model: text-embedding-3-small
|
||||
vectorstore:
|
||||
pgvector:
|
||||
schema-name: rag
|
||||
table-name: document_chunks
|
||||
index-type: HNSW
|
||||
distance-type: COSINE_DISTANCE
|
||||
dimensions: 1536
|
||||
# false: the schema comes from init.sql, not from the application.
|
||||
initialize-schema: false
|
||||
|
||||
rag:
|
||||
retrieval:
|
||||
top-k: 20
|
||||
similarity-threshold: 0.0
|
||||
rerank:
|
||||
top-n: 5
|
||||
chunking:
|
||||
chunk-size: 512
|
||||
min-chunk-size-chars: 350
|
||||
|
||||
management:
|
||||
endpoints:
|
||||
web:
|
||||
exposure:
|
||||
include: health, prometheus
|
||||
@@ -0,0 +1,158 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.chunk.RecursiveChunker;
|
||||
import com.ankurm.rag.chunk.SemanticChunker;
|
||||
import com.ankurm.rag.support.HashingEmbeddingModel;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import com.knuddels.jtokkit.Encodings;
|
||||
import com.knuddels.jtokkit.api.Encoding;
|
||||
import com.knuddels.jtokkit.api.EncodingType;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.transformer.splitter.TokenTextSplitter;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/** Chapter 2: what the chunkers do to a document, measured. Transcripts 02 and 03. */
|
||||
class ChunkingTest {
|
||||
|
||||
private static final Encoding CL100K = Encodings.newDefaultEncodingRegistry().getEncoding(EncodingType.CL100K_BASE);
|
||||
|
||||
/** Forty numbered sentences, about 880 tokens in all, so a document has a known size. */
|
||||
private static String longText() {
|
||||
StringBuilder sb = new StringBuilder();
|
||||
for (int i = 1; i <= 40; i++) {
|
||||
sb.append("Sentence ").append(i).append(" describes the policy for item ").append(i)
|
||||
.append(" and states that the rule applies to every full-time employee.");
|
||||
sb.append(i % 5 == 0 ? "\n\n" : " ");
|
||||
}
|
||||
return sb.toString().strip();
|
||||
}
|
||||
|
||||
private static Document doc(String text) {
|
||||
return new Document(text, Map.of("source_file", "long.txt", "page_number", 1));
|
||||
}
|
||||
|
||||
private static int tokens(String text) {
|
||||
return CL100K.countTokens(text);
|
||||
}
|
||||
|
||||
@Test
|
||||
void tokenTextSplitter() {
|
||||
String text = longText();
|
||||
try (Transcript t = new Transcript("02-token-text-splitter.txt", "What TokenTextSplitter does to a 40-sentence document")) {
|
||||
t.line("document: %d characters, %d cl100k_base tokens", text.length(), tokens(text));
|
||||
|
||||
t.section("new TokenTextSplitter() (defaults: 800 tokens, 350 min chars, 5 min length to embed)");
|
||||
List<Document> byDefault = new TokenTextSplitter().apply(List.of(doc(text)));
|
||||
t.line("chunks: %d, sizes in tokens: %s", byDefault.size(),
|
||||
byDefault.stream().map(c -> tokens(c.getText())).toList());
|
||||
assertThat(byDefault).hasSize(2);
|
||||
assertThat(tokens(byDefault.get(0).getText())).isLessThanOrEqualTo(800);
|
||||
|
||||
t.section("chunk size 100 tokens");
|
||||
TokenTextSplitter small = TokenTextSplitter.builder().withChunkSize(100).build();
|
||||
List<Document> chunks = small.apply(List.of(doc(text)));
|
||||
t.line("chunks: %d", chunks.size());
|
||||
for (int i = 0; i < chunks.size(); i++) {
|
||||
String c = chunks.get(i).getText();
|
||||
t.line("chunk %d: %3d tokens, %4d chars, starts \"%s\", ends \"%s\"", i, tokens(c), c.length(),
|
||||
c.substring(0, Math.min(22, c.length())).replace("\n", "\\n"),
|
||||
c.substring(Math.max(0, c.length() - 22)).replace("\n", "\\n"));
|
||||
}
|
||||
var chunkMetadata = new java.util.TreeMap<>(chunks.get(1).getMetadata());
|
||||
chunkMetadata.remove("parent_document_id"); // a random UUID, different on every run
|
||||
t.line("metadata of chunk 1 (parent_document_id, a random UUID, left out): %s", chunkMetadata);
|
||||
|
||||
t.section("is there any overlap between neighbouring chunks?");
|
||||
int shared = 0;
|
||||
for (int i = 1; i < chunks.size(); i++) {
|
||||
String opening = chunks.get(i).getText().substring(0, 30);
|
||||
if (chunks.get(i - 1).getText().contains(opening)) {
|
||||
shared++;
|
||||
}
|
||||
}
|
||||
t.line("boundaries where the first 30 characters of a chunk already appear in the chunk before it: %d of %d", shared, chunks.size() - 1);
|
||||
assertThat(shared).isZero();
|
||||
assertThat(chunks.stream().mapToInt(c -> tokens(c.getText())).max().orElseThrow()).isLessThanOrEqualTo(100);
|
||||
|
||||
t.section("minChunkSizeChars: where the cut lands");
|
||||
for (int minChars : new int[] {350, 50}) {
|
||||
List<Document> cut = TokenTextSplitter.builder().withChunkSize(100).withMinChunkSizeChars(minChars).build()
|
||||
.apply(List.of(doc(text)));
|
||||
long atSentenceEnd = cut.stream().limit(cut.size() - 1L)
|
||||
.filter(c -> c.getText().stripTrailing().endsWith(".")).count();
|
||||
t.line("minChunkSizeChars=%3d: %d chunks, %d of the first %d end on a full stop",
|
||||
minChars, cut.size(), atSentenceEnd, cut.size() - 1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void recursiveChunker() {
|
||||
String text = longText();
|
||||
try (Transcript t = new Transcript("03-recursive-and-semantic-chunkers.txt", "RecursiveChunker and SemanticChunker")) {
|
||||
t.section("RecursiveChunker(maxChars=400, overlapChars=80)");
|
||||
List<Document> chunks = new RecursiveChunker(400, 80).apply(List.of(doc(text)));
|
||||
t.line("chunks: %d, longest: %d characters", chunks.size(), chunks.stream().mapToInt(c -> c.getText().length()).max().orElse(0));
|
||||
for (int i = 0; i < chunks.size(); i++) {
|
||||
String c = chunks.get(i).getText();
|
||||
t.line("chunk %d: %3d chars, starts \"%s\", ends \"%s\"", i, c.length(),
|
||||
c.substring(0, 24).replace("\n", "\\n"), c.substring(c.length() - 24).replace("\n", "\\n"));
|
||||
}
|
||||
t.line("metadata of chunk 2: %s", new java.util.TreeMap<>(chunks.get(2).getMetadata()));
|
||||
|
||||
assertThat(chunks).allSatisfy(c -> assertThat(c.getText().length()).isLessThanOrEqualTo(400));
|
||||
int overlapping = 0;
|
||||
for (int i = 1; i < chunks.size(); i++) {
|
||||
String[] words = chunks.get(i).getText().split(" ");
|
||||
String firstWords = String.join(" ", List.of(words).subList(0, 3));
|
||||
if (chunks.get(i - 1).getText().contains(firstWords)) {
|
||||
overlapping++;
|
||||
}
|
||||
}
|
||||
t.line("boundaries where the next chunk opens with words the previous one ended with: %d of %d", overlapping, chunks.size() - 1);
|
||||
assertThat(overlapping).isGreaterThan(0);
|
||||
|
||||
t.section("one 900-character word-salad with no separators falls back to a hard cut");
|
||||
List<Document> hard = new RecursiveChunker(400, 0).apply(List.of(doc("x".repeat(900))));
|
||||
t.line("chunks: %s", hard.stream().map(c -> c.getText().length()).toList());
|
||||
assertThat(hard).hasSize(3);
|
||||
|
||||
t.section("SemanticChunker(distance 0.9) on three topics, four sentences each");
|
||||
String topics = String.join(" ",
|
||||
"Annual leave is twenty days per year.", "Unused annual leave carries over until March.",
|
||||
"Leave requests go through the HR portal.", "Annual leave accrues monthly.",
|
||||
"Expenses need a receipt above fifty euros.", "Expense claims must be filed within thirty days.",
|
||||
"Receipts for expenses are uploaded as photos.", "Approved expenses are paid with salary.",
|
||||
"Remote work is allowed two days per week.", "Managers agree the remote work days.",
|
||||
"Remote work needs a quiet workspace.", "Remote work days are recorded in the calendar.");
|
||||
HashingEmbeddingModel embeddings = new HashingEmbeddingModel(256);
|
||||
List<Document> semantic = new SemanticChunker(embeddings, 0.9).apply(List.of(doc(topics)));
|
||||
t.line("chunks: %d", semantic.size());
|
||||
for (Document c : semantic) {
|
||||
t.line(" [%s] %s", c.getMetadata().get("chunk_index"), c.getText());
|
||||
}
|
||||
t.line("texts sent to the embedding model: %d (12 sentences, one batched call)", embeddings.textsEmbedded());
|
||||
assertThat(semantic).hasSize(3);
|
||||
assertThat(semantic.get(0).getText()).startsWith("Annual leave is twenty days");
|
||||
assertThat(semantic.get(1).getText()).startsWith("Expenses need a receipt");
|
||||
assertThat(semantic.get(2).getText()).startsWith("Remote work is allowed");
|
||||
}
|
||||
}
|
||||
|
||||
/** The same fixture, kept so the numbers in transcript 03 do not depend on test order. */
|
||||
@Test
|
||||
void chunkersKeepMetadata() {
|
||||
List<Document> chunks = new RecursiveChunker(400, 80).apply(List.of(doc(longText())));
|
||||
assertThat(chunks).allSatisfy(c -> {
|
||||
assertThat(c.getMetadata()).containsEntry("source_file", "long.txt").containsEntry("page_number", 1);
|
||||
assertThat(c.getMetadata()).containsKeys("chunk_index", "chunk_total");
|
||||
});
|
||||
assertThat(new ArrayList<>(chunks).size()).isEqualTo((int) chunks.get(0).getMetadata().get("chunk_total"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.TreeMap;
|
||||
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.beans.factory.config.YamlPropertiesFactoryBean;
|
||||
import org.springframework.core.io.ClassPathResource;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.core.io.support.PathMatchingResourcePatternResolver;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
import tools.jackson.databind.json.JsonMapper;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Spring Boot ignores configuration keys it does not know, so a wrong key fails neither the build
|
||||
* nor the start-up: the setting is just never applied. This test checks every {@code spring.ai.*}
|
||||
* key in the shipped {@code application.yml} against the configuration metadata inside the jars on
|
||||
* the classpath, and fails on a key that is unknown or deprecated. Transcript 15.
|
||||
*/
|
||||
class ConfigKeysTest {
|
||||
|
||||
/** name to the replacement it is deprecated in favour of, or "" when the name is current. */
|
||||
private static Map<String, String> metadata() throws IOException {
|
||||
JsonMapper mapper = JsonMapper.builder().build();
|
||||
Map<String, String> names = new TreeMap<>();
|
||||
for (Resource resource : new PathMatchingResourcePatternResolver()
|
||||
.getResources("classpath*:META-INF/spring-configuration-metadata.json")) {
|
||||
JsonNode root = mapper.readTree(resource.getInputStream().readAllBytes());
|
||||
for (JsonNode property : root.get("properties")) {
|
||||
JsonNode deprecation = property.get("deprecation");
|
||||
String replacement = deprecation == null || deprecation.get("replacement") == null
|
||||
? "(no replacement named)" : deprecation.get("replacement").stringValue();
|
||||
boolean deprecated = property.get("deprecated") != null && property.get("deprecated").asBoolean();
|
||||
names.merge(property.get("name").stringValue(), deprecated ? replacement : "", (a, b) -> a.isEmpty() ? a : b);
|
||||
}
|
||||
}
|
||||
return names;
|
||||
}
|
||||
|
||||
private static String status(Map<String, String> known, String key) {
|
||||
if (known.containsKey(key)) {
|
||||
String replacement = known.get(key);
|
||||
return replacement.isEmpty() ? "current" : "DEPRECATED, use " + replacement;
|
||||
}
|
||||
boolean underMap = known.keySet().stream()
|
||||
.anyMatch(name -> (name.endsWith("custom-headers") || name.endsWith("extra-body")) && key.startsWith(name + "."));
|
||||
return underMap ? "current" : "UNKNOWN";
|
||||
}
|
||||
|
||||
@Test
|
||||
void everySpringAiKeyIsCurrent() throws IOException {
|
||||
Map<String, String> known = metadata();
|
||||
YamlPropertiesFactoryBean yaml = new YamlPropertiesFactoryBean();
|
||||
yaml.setResources(new ClassPathResource("application.yml"));
|
||||
|
||||
List<String> problems = new ArrayList<>();
|
||||
try (Transcript t = new Transcript("15-config-keys.txt", "spring.ai.* keys in application.yml against the jars' configuration metadata")) {
|
||||
t.line("property names in the jars' metadata: %d", known.size());
|
||||
t.section("keys in the shipped application.yml");
|
||||
for (String key : yaml.getObject().keySet().stream().map(Object::toString).sorted().toList()) {
|
||||
if (key.startsWith("spring.ai.")) {
|
||||
String status = status(known, key);
|
||||
t.line("%-52s %s", key, status);
|
||||
if (!status.equals("current")) {
|
||||
problems.add(key + " " + status);
|
||||
}
|
||||
}
|
||||
}
|
||||
t.section("keys written the 1.x way");
|
||||
for (String key : List.of("spring.ai.openai.chat.options.model", "spring.ai.openai.chat.options.temperature",
|
||||
"spring.ai.openai.embedding.options.model", "spring.ai.openai.chat.optoins.model")) {
|
||||
t.line("%-52s %s", key, status(known, key));
|
||||
}
|
||||
assertThat(status(known, "spring.ai.openai.chat.options.model")).startsWith("DEPRECATED");
|
||||
assertThat(status(known, "spring.ai.openai.chat.optoins.model")).isEqualTo("UNKNOWN");
|
||||
}
|
||||
assertThat(problems).isEmpty();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.support.FakeChatModel;
|
||||
import com.ankurm.rag.support.HashingEmbeddingModel;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import com.ankurm.rag.ingest.IngestionService;
|
||||
import com.ankurm.rag.query.LlmReranker;
|
||||
import com.ankurm.rag.query.RagQueryService;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.evaluation.FactCheckingEvaluator;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.rag.advisor.RetrievalAugmentationAdvisor;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
import org.springframework.ai.vectorstore.pgvector.PgVectorStore;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.context.ApplicationContext;
|
||||
import org.springframework.boot.test.context.SpringBootTest;
|
||||
import org.springframework.boot.test.context.TestConfiguration;
|
||||
import org.springframework.boot.test.web.server.LocalServerPort;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.core.io.ByteArrayResource;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.http.MediaType;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
import org.springframework.test.context.DynamicPropertyRegistry;
|
||||
import org.springframework.test.context.DynamicPropertySource;
|
||||
import org.springframework.util.LinkedMultiValueMap;
|
||||
import org.springframework.web.client.RestClient;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* The whole application, started for real: the shipped {@code application.yml}, the {@code init.sql}
|
||||
* schema in a real PostgreSQL with pgvector, real HTTP requests to a real port, real PDFs. Only the
|
||||
* two model beans are replaced, by the scripted fakes, so it runs with no API key. Transcript 10.
|
||||
*/
|
||||
@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT, properties = {
|
||||
"spring.ai.model.chat=none",
|
||||
"spring.ai.model.embedding=none",
|
||||
"spring.ai.openai.api-key=not-used-in-tests"
|
||||
})
|
||||
class EndToEndTest {
|
||||
|
||||
@TestConfiguration
|
||||
static class Fakes {
|
||||
@Bean
|
||||
FakeChatModel chatModel() {
|
||||
return new FakeChatModel();
|
||||
}
|
||||
|
||||
@Bean
|
||||
HashingEmbeddingModel embeddingModel() {
|
||||
return new HashingEmbeddingModel(Stores.DIMS);
|
||||
}
|
||||
}
|
||||
|
||||
@DynamicPropertySource
|
||||
static void database(DynamicPropertyRegistry registry) throws Exception {
|
||||
JdbcTemplate jdbc = Stores.jdbc();
|
||||
jdbc.execute("DROP SCHEMA IF EXISTS rag CASCADE");
|
||||
jdbc.execute(Files.readString(Path.of("init.sql")));
|
||||
registry.add("spring.datasource.url", () -> Stores.PG_URL);
|
||||
registry.add("spring.datasource.username", () -> Stores.PG_USER);
|
||||
registry.add("spring.datasource.password", () -> Stores.PG_PASSWORD);
|
||||
}
|
||||
|
||||
@LocalServerPort
|
||||
int port;
|
||||
|
||||
@Autowired
|
||||
JdbcTemplate jdbc;
|
||||
|
||||
@Autowired
|
||||
ApplicationContext context;
|
||||
|
||||
@Test
|
||||
void pipelineBeans() {
|
||||
try (Transcript t = new Transcript("01-pipeline-beans.txt", "The pipeline as Spring beans (the shipped configuration, fake models)")) {
|
||||
t.line("%-28s %-28s %s", "stage / role", "bean name", "actual class");
|
||||
List<Class<?>> types = List.of(DocumentTransformer.class, EmbeddingModel.class, VectorStore.class,
|
||||
RetrievalAugmentationAdvisor.class, LlmReranker.class, ChatModel.class, ChatClient.class,
|
||||
FactCheckingEvaluator.class, IngestionService.class, RagQueryService.class);
|
||||
for (Class<?> type : types) {
|
||||
for (String name : context.getBeanNamesForType(type)) {
|
||||
t.line("%-28s %-28s %s", type.getSimpleName(), name, context.getType(name).getSimpleName());
|
||||
}
|
||||
}
|
||||
assertThat(context.getBeanNamesForType(RetrievalAugmentationAdvisor.class)).hasSize(1);
|
||||
assertThat(context.getBean(VectorStore.class)).isInstanceOf(PgVectorStore.class);
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void uploadAskAndScrape() {
|
||||
RestClient http = RestClient.create("http://localhost:" + port);
|
||||
|
||||
try (Transcript t = new Transcript("10-end-to-end.txt", "End to end: HTTP, real PostgreSQL + pgvector, the shipped application.yml")) {
|
||||
t.line("schema from init.sql: %s", jdbc.queryForList(
|
||||
"select indexname from pg_indexes where schemaname = 'rag' order by indexname", String.class));
|
||||
|
||||
t.section("POST /api/ingest (acme, then globex, then acme again)");
|
||||
Resource acmePdf = new ByteArrayResource(bytes(SamplePdf.of(SamplePdf.ACME)));
|
||||
t.line("%s", upload(http, "acme-handbook.pdf", acmePdf, "acme"));
|
||||
t.line("%s", upload(http, "globex-manual.pdf", SamplePdf.of(SamplePdf.GLOBEX), "globex"));
|
||||
String again = upload(http, "acme-handbook.pdf", acmePdf, "acme");
|
||||
t.line("%s", again);
|
||||
|
||||
t.section("what is in the table");
|
||||
t.line("rows: %d", jdbc.queryForObject("select count(*) from rag.document_chunks", Integer.class));
|
||||
t.line("rows per tenant: %s", jdbc.queryForList(
|
||||
"select metadata::jsonb ->> 'tenant_id' as tenant, count(*) as n from rag.document_chunks group by 1 order by 1")
|
||||
.stream().map(r -> r.get("tenant") + "=" + r.get("n")).toList());
|
||||
t.line("metadata of one row: %s", jdbc.queryForObject(
|
||||
"select metadata::jsonb - 'parent_document_id' - 'source_hash' from rag.document_chunks where content like 'Globex%' limit 1",
|
||||
String.class));
|
||||
|
||||
t.section("POST /api/query");
|
||||
String acme = query(http, "How many days of annual leave do employees get?", "acme");
|
||||
String globex = query(http, "How many days of annual leave do employees get?", "globex");
|
||||
t.line("tenantId acme: %s", acme);
|
||||
t.line("tenantId globex: %s", globex);
|
||||
|
||||
t.section("GET /actuator/prometheus (only the rag_ series; the timer's sum and max are left out because they change every run)");
|
||||
String scrape = http.get().uri("/actuator/prometheus").retrieve().body(String.class);
|
||||
List<String> lines = scrape.lines().filter(l -> l.startsWith("rag_") && !l.startsWith("rag_query_duration_seconds_bucket")
|
||||
&& !l.startsWith("rag_query_duration_seconds_sum") && !l.startsWith("rag_query_duration_seconds_max")).toList();
|
||||
lines.forEach(l -> t.line("%s", l));
|
||||
|
||||
assertThat(again).contains("\"status\":\"skipped\"");
|
||||
assertThat(acme).contains("20 working days").contains("\"grounded\":true").doesNotContain("25 working days");
|
||||
assertThat(globex).contains("25 working days").contains("\"grounded\":true").doesNotContain("20 working days");
|
||||
assertThat(scrape).contains("rag_chunks_ingested_total").contains("rag_ingestion_skipped_total");
|
||||
}
|
||||
}
|
||||
|
||||
private static String upload(RestClient http, String filename, Resource pdf, String tenant) {
|
||||
var body = new LinkedMultiValueMap<String, Object>();
|
||||
body.add("file", new ByteArrayResource(bytes(pdf)) {
|
||||
@Override
|
||||
public String getFilename() {
|
||||
return filename;
|
||||
}
|
||||
});
|
||||
body.add("tenantId", tenant);
|
||||
return http.post().uri("/api/ingest").contentType(MediaType.MULTIPART_FORM_DATA).body(body)
|
||||
.retrieve().body(String.class);
|
||||
}
|
||||
|
||||
private static String query(RestClient http, String question, String tenant) {
|
||||
return http.post().uri("/api/query").contentType(MediaType.APPLICATION_JSON)
|
||||
.body(Map.of("question", question, "tenantId", tenant)).retrieve().body(String.class);
|
||||
}
|
||||
|
||||
private static byte[] bytes(Resource resource) {
|
||||
try {
|
||||
return resource.getInputStream().readAllBytes();
|
||||
} catch (java.io.IOException e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.config.RagProperties;
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.query.RagResponse;
|
||||
import com.ankurm.rag.support.FakeChatModel;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Chapter 5: the check that runs after generation, and the four outcomes {@code RagQueryService}
|
||||
* reports. Transcript 09. The scripted model decides what the "answer" and the "verdict" are, so
|
||||
* this tests Spring AI's plumbing and this repository's handling of each outcome. It does not
|
||||
* test how good a real judge model is at spotting a hallucination.
|
||||
*/
|
||||
class FaithfulnessTest {
|
||||
|
||||
private static final String QUESTION = "How many days of annual leave do employees get?";
|
||||
private static final String OFF_TOPIC = "What is the capital of Mongolia?";
|
||||
|
||||
@Test
|
||||
void fourOutcomes() {
|
||||
TestPipeline p = TestPipeline.on(Stores::simple);
|
||||
p.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
TestPipeline strict = TestPipeline.on(new RagProperties(new RagProperties.Retrieval(20, 0.3),
|
||||
new RagProperties.Rerank(5), new RagProperties.Chunking(512, 350)), Stores::simple);
|
||||
strict.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
|
||||
try (Transcript t = new Transcript("09-faithfulness-check.txt", "Generation and the faithfulness check")) {
|
||||
t.section("1. the answer is in the chunks");
|
||||
p.chat.reset();
|
||||
RagResponse grounded = p.queries.ask(QUESTION, null);
|
||||
show(t, grounded);
|
||||
|
||||
t.section("2. the model answers with something the chunks do not say");
|
||||
p.chat.reset();
|
||||
p.chat.forcedAnswer("Employees get 30 days of annual leave.");
|
||||
RagResponse invented = p.queries.ask(QUESTION, null);
|
||||
show(t, invented);
|
||||
t.line("the check the judge model was given:");
|
||||
t.line(p.chat.promptsContaining("Evaluate whether or not the following claim is supported").getFirst()
|
||||
.lines().limit(3).reduce((a, b) -> a + "\n" + b).orElse(""));
|
||||
|
||||
t.section("3. the judge says \"Yes.\" instead of \"yes\"");
|
||||
p.chat.reset();
|
||||
p.chat.forcedVerdict("Yes.");
|
||||
RagResponse punctuated = p.queries.ask(QUESTION, null);
|
||||
show(t, punctuated);
|
||||
p.chat.reset();
|
||||
p.chat.forcedVerdict("YES");
|
||||
RagResponse upper = p.queries.ask(QUESTION, null);
|
||||
t.line("with the reply \"YES\": grounded=%s", upper.grounded());
|
||||
|
||||
t.section("4. the judge call itself fails");
|
||||
p.chat.reset();
|
||||
p.metrics.clear();
|
||||
p.chat.failFactChecks(true);
|
||||
RagResponse judgeDown = p.queries.ask(QUESTION, null);
|
||||
show(t, judgeDown);
|
||||
t.line("rag.faithfulness.judge_failures = %.0f", p.metrics.counter("rag.faithfulness.judge_failures").count());
|
||||
|
||||
t.section("5. nothing is retrieved (threshold 0.3, off-topic question)");
|
||||
strict.chat.reset();
|
||||
RagResponse none = strict.queries.ask(OFF_TOPIC, null);
|
||||
show(t, none);
|
||||
t.line("fact-check calls made: %d", strict.chat.promptsContaining("Evaluate whether or not").size());
|
||||
|
||||
assertThat(grounded.status()).isEqualTo("answered");
|
||||
assertThat(grounded.grounded()).isTrue();
|
||||
assertThat(invented.status()).isEqualTo("ungrounded");
|
||||
assertThat(punctuated.grounded()).isFalse();
|
||||
assertThat(upper.grounded()).isTrue();
|
||||
assertThat(judgeDown.status()).isEqualTo("ungrounded");
|
||||
assertThat(p.metrics.counter("rag.faithfulness.judge_failures").count()).isEqualTo(1.0);
|
||||
assertThat(none.status()).isEqualTo("no_context");
|
||||
assertThat(none.sources()).isEmpty();
|
||||
assertThat(strict.chat.promptsContaining("Evaluate whether or not")).isEmpty();
|
||||
assertThat(none.answer()).isEqualTo(FakeChatModel.REFUSAL);
|
||||
}
|
||||
}
|
||||
|
||||
private static void show(Transcript t, RagResponse r) {
|
||||
t.line("status=%s grounded=%s sources=%d", r.status(), r.grounded(), r.sources().size());
|
||||
t.line("answer: %s", r.answer());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.ingest.IngestionResult;
|
||||
import com.ankurm.rag.ingest.IngestionService;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.reader.pdf.PagePdfDocumentReader;
|
||||
import org.springframework.ai.reader.pdf.config.PdfDocumentReaderConfig;
|
||||
import org.springframework.core.io.Resource;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/** Chapter 3: what a PDF becomes on the way in, and what re-ingesting does. Transcripts 04 and 05. */
|
||||
class IngestionTest {
|
||||
|
||||
private static final Map<String, Object> ACME_TAGS = Map.of("tenant_id", "acme", "doc_type", "handbook");
|
||||
|
||||
@Test
|
||||
void pdfPagesAndMetadata() {
|
||||
Resource pdf = SamplePdf.of(SamplePdf.ACME);
|
||||
List<Document> pages = new PagePdfDocumentReader(pdf, PdfDocumentReaderConfig.builder()
|
||||
.withPagesPerDocument(1).build()).get();
|
||||
|
||||
try (Transcript t = new Transcript("04-pdf-pages-and-metadata.txt", "PagePdfDocumentReader: what one PDF page becomes")) {
|
||||
t.line("pages in the PDF: 4, documents read: %d", pages.size());
|
||||
for (Document page : pages) {
|
||||
t.line("page document metadata: %s", page.getMetadata());
|
||||
}
|
||||
|
||||
Document page2 = pages.get(1);
|
||||
String raw = page2.getText();
|
||||
String tidy = IngestionService.tidy(raw);
|
||||
t.section("page 2 as read (single spaces as dots, runs of 4+ as [n spaces], line ends as a pilcrow)");
|
||||
t.line(visible(raw));
|
||||
t.section("page 2 after IngestionService.tidy");
|
||||
t.line(visible(tidy));
|
||||
t.section("size of page 2");
|
||||
t.line("as read : %d characters, longest run of spaces %d", raw.length(), longestSpaceRun(raw));
|
||||
t.line("tidied : %d characters, longest run of spaces %d", tidy.length(), longestSpaceRun(tidy));
|
||||
|
||||
assertThat(pages).hasSize(4);
|
||||
assertThat(pages.get(1).getMetadata()).containsKey("page_number");
|
||||
assertThat(tidy.length()).isLessThanOrEqualTo(raw.length());
|
||||
assertThat(tidy).contains("20 working days").doesNotContain(" ");
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
void reingestingTheSameFile() {
|
||||
JdbcTemplate jdbc = Stores.jdbc();
|
||||
TestPipeline p = TestPipeline.on(embeddings -> Stores.pg(jdbc, embeddings));
|
||||
Resource v1 = SamplePdf.of(SamplePdf.ACME);
|
||||
List<List<String>> edited = new ArrayList<>(SamplePdf.ACME);
|
||||
edited.set(1, List.of("4.1 Annual Leave Entitlement. Full-time employees are entitled to 22 working days",
|
||||
"of annual leave per calendar year. Part-time employees receive leave pro rata.",
|
||||
"",
|
||||
"4.2 Leave Carryover. Unused annual leave may be carried over for a maximum of 5 days",
|
||||
"into the next calendar year and must be used by 31 March."));
|
||||
Resource v2 = SamplePdf.of(edited);
|
||||
|
||||
try (Transcript t = new Transcript("05-ingestion-idempotency.txt", "Ingesting the same handbook more than once (real PostgreSQL + pgvector)")) {
|
||||
t.line("rows are counted with: select count(*) from vector_store");
|
||||
IngestionResult first = p.ingestion.ingestPdf(v1, "acme-handbook.pdf", ACME_TAGS);
|
||||
int rowsAfterFirst = rows(jdbc);
|
||||
int embeddedAfterFirst = p.embeddings.textsEmbedded();
|
||||
report(t, "1. first upload", first, rowsAfterFirst, embeddedAfterFirst);
|
||||
|
||||
IngestionResult second = p.ingestion.ingestPdf(v1, "acme-handbook.pdf", ACME_TAGS);
|
||||
report(t, "2. same bytes again", second, rows(jdbc), p.embeddings.textsEmbedded());
|
||||
|
||||
IngestionResult third = p.ingestion.ingestPdf(v2, "acme-handbook.pdf", ACME_TAGS);
|
||||
report(t, "3. page 2 edited (20 -> 22 days)", third, rows(jdbc), p.embeddings.textsEmbedded());
|
||||
int stale = jdbc.queryForObject("select count(*) from vector_store where content like '%20 working days%'", Integer.class);
|
||||
int fresh = jdbc.queryForObject("select count(*) from vector_store where content like '%22 working days%'", Integer.class);
|
||||
t.line("rows still saying \"20 working days\": %d, rows saying \"22 working days\": %d", stale, fresh);
|
||||
|
||||
Resource v2Again = SamplePdf.of(edited);
|
||||
assertThat(bytes(v2Again)).isNotEqualTo(bytes(v2));
|
||||
IngestionResult reexport = p.ingestion.ingestPdf(v2Again, "acme-handbook.pdf", ACME_TAGS);
|
||||
report(t, "4. same text, exported again", reexport, rows(jdbc), p.embeddings.textsEmbedded());
|
||||
t.line("the two PDFs have identical text and different bytes: %s", !java.util.Arrays.equals(bytes(v2Again), bytes(v2)));
|
||||
t.line("the file hash is a hash of bytes, so a re-export counts as a change and is re-embedded");
|
||||
|
||||
t.section("the naive version: vectorStore.add() on every upload, nothing remembered");
|
||||
List<Document> pages = new PagePdfDocumentReader(v2, PdfDocumentReaderConfig.builder()
|
||||
.withPagesPerDocument(1).build()).get();
|
||||
List<Document> chunks = p.chunker.apply(pages.stream()
|
||||
.map(pg -> new Document(IngestionService.tidy(pg.getText()), pg.getMetadata())).toList());
|
||||
int before = rows(jdbc);
|
||||
p.store.add(chunks);
|
||||
p.store.add(chunks.stream().map(c -> new Document(c.getText(), c.getMetadata())).toList());
|
||||
int after = rows(jdbc);
|
||||
t.line("the same %d chunks added on two more uploads: rows in table %d -> %d", chunks.size(), before, after);
|
||||
t.line("rows saying \"22 working days\" now: %d",
|
||||
jdbc.queryForObject("select count(*) from vector_store where content like '%22 working days%'", Integer.class));
|
||||
assertThat(after).isEqualTo(before + 2 * chunks.size());
|
||||
assertThat(first.status()).isEqualTo("ingested");
|
||||
assertThat(first.chunksWritten()).isEqualTo(rowsAfterFirst);
|
||||
assertThat(second.status()).isEqualTo("skipped");
|
||||
assertThat(third.status()).isEqualTo("updated");
|
||||
assertThat(third.chunksReplaced()).isEqualTo(first.chunksWritten());
|
||||
assertThat(reexport.status()).isEqualTo("updated");
|
||||
assertThat(stale).isZero();
|
||||
assertThat(fresh).isEqualTo(1);
|
||||
assertThat(p.embeddings.textsEmbedded()).isGreaterThan(embeddedAfterFirst);
|
||||
}
|
||||
}
|
||||
|
||||
private static void report(Transcript t, String label, IngestionResult r, int rows, int embedded) {
|
||||
t.line("%-34s -> status=%-8s chunksWritten=%d chunksReplaced=%d | rows in table=%d, texts embedded so far=%d",
|
||||
label, r.status(), r.chunksWritten(), r.chunksReplaced(), rows, embedded);
|
||||
}
|
||||
|
||||
private static byte[] bytes(Resource resource) {
|
||||
try {
|
||||
return resource.getInputStream().readAllBytes();
|
||||
} catch (java.io.IOException e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
|
||||
private static int rows(JdbcTemplate jdbc) {
|
||||
return jdbc.queryForObject("select count(*) from vector_store", Integer.class);
|
||||
}
|
||||
|
||||
/** Short runs of spaces become dots, runs of four or more become "[n spaces]", line ends a pilcrow. */
|
||||
private static String visible(String text) {
|
||||
Matcher m = Pattern.compile(" +").matcher(text);
|
||||
StringBuilder sb = new StringBuilder();
|
||||
while (m.find()) {
|
||||
int n = m.end() - m.start();
|
||||
m.appendReplacement(sb, n >= 4 ? "[" + n + " spaces]" : "\u00b7".repeat(n));
|
||||
}
|
||||
m.appendTail(sb);
|
||||
return sb.toString().replace("\n", "\u00b6\n");
|
||||
}
|
||||
|
||||
private static int longestSpaceRun(String text) {
|
||||
int longest = 0;
|
||||
Matcher m = Pattern.compile(" +").matcher(text);
|
||||
while (m.find()) {
|
||||
longest = Math.max(longest, m.end() - m.start());
|
||||
}
|
||||
return longest;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.config.RagProperties;
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.query.LlmReranker;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.rag.Query;
|
||||
import org.springframework.ai.rag.retrieval.search.VectorStoreDocumentRetriever;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Chapter 4, second half: what an LLM reranker costs, how it fails, and why it runs on virtual
|
||||
* threads. Transcript 08. The chat model here is the scripted fake, so the timings are its
|
||||
* {@code Thread.sleep}, standing in for network latency, and the ratings are word overlap. What
|
||||
* this shows is the shape of the cost and the failure handling; it does not show that reranking
|
||||
* improves answers, and it cannot.
|
||||
*/
|
||||
class RerankTest {
|
||||
|
||||
private static final String QUESTION = "How many days of annual leave do employees get?";
|
||||
|
||||
@Test
|
||||
void costOrderingAndFailure() {
|
||||
TestPipeline p = TestPipeline.on(withTopN(2), Stores::simple);
|
||||
p.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
List<Document> candidates = VectorStoreDocumentRetriever.builder().vectorStore(p.store).topK(20)
|
||||
.similarityThreshold(0.0).build().retrieve(new Query(QUESTION));
|
||||
|
||||
try (Transcript t = new Transcript("08-reranking.txt", "LLM reranking: calls, order, latency and failure")) {
|
||||
t.section("one model call per candidate; only topN survive (topN = 2)");
|
||||
p.chat.reset();
|
||||
List<Document> kept = p.reranker.process(new Query(QUESTION), candidates);
|
||||
t.line("candidates in: %d, model calls made: %d, chunks out: %d", candidates.size(), p.chat.calls(), kept.size());
|
||||
t.line("order from the vector search, best first:");
|
||||
for (Document d : candidates) {
|
||||
t.line(" similarity %.4f page %s", d.getScore(), d.getMetadata().get("page_number"));
|
||||
}
|
||||
t.line("order after reranking, best first:");
|
||||
for (Document d : kept) {
|
||||
t.line(" rerank_score %s page %s", d.getMetadata().get(LlmReranker.SCORE), d.getMetadata().get("page_number"));
|
||||
}
|
||||
t.line("the rating prompt for the page 2 candidate (calls run concurrently, so pick it by content):");
|
||||
t.line(p.chat.promptsContaining("4.1 Annual Leave").getFirst().lines().limit(4).reduce((a, b) -> a + "\n" + b).orElse(""));
|
||||
|
||||
t.section("latency: 20 candidates, each rating call takes 200 ms (a Thread.sleep in the fake model)");
|
||||
List<Document> twenty = new ArrayList<>();
|
||||
for (int i = 1; i <= 20; i++) {
|
||||
twenty.add(new Document("Passage " + i + " says employees get " + i + " days of annual leave."));
|
||||
}
|
||||
p.chat.reset();
|
||||
p.chat.latencyMillis(200);
|
||||
long start = System.nanoTime();
|
||||
p.reranker.process(new Query(QUESTION), twenty);
|
||||
long concurrent = (System.nanoTime() - start) / 1_000_000;
|
||||
|
||||
p.chat.reset();
|
||||
p.chat.latencyMillis(200);
|
||||
start = System.nanoTime();
|
||||
for (Document d : twenty) {
|
||||
p.chat.call(new Prompt("Rate how well the PASSAGE helps answer the QUESTION\nQUESTION: " + QUESTION
|
||||
+ "\nPASSAGE: " + d.getText()));
|
||||
}
|
||||
long sequential = (System.nanoTime() - start) / 1_000_000;
|
||||
t.line("the 20 calls one after another take 4000 ms or more: %s", sequential >= 4000);
|
||||
t.line("LlmReranker, one virtual thread per candidate, takes under 1000 ms: %s", concurrent < 1000);
|
||||
t.line("a real API adds its own rate limits, which this test cannot show");
|
||||
|
||||
t.section("failure: the model does not reply with a bare integer");
|
||||
p.chat.reset();
|
||||
p.metrics.clear();
|
||||
p.chat.forcedRating("Score: 8");
|
||||
List<Document> degraded = p.reranker.process(new Query(QUESTION), candidates);
|
||||
double failures = p.metrics.counter("rag.rerank.failures").count();
|
||||
t.line("reply \"Score: 8\" for every candidate -> failures counted: %.0f of %d", failures, candidates.size());
|
||||
t.line("scores assigned: %s", degraded.stream().map(d -> d.getMetadata().get(LlmReranker.SCORE)).toList());
|
||||
t.line("pages kept, in order: %s (the vector-search order, because every score is 0)",
|
||||
degraded.stream().map(d -> d.getMetadata().get("page_number")).toList());
|
||||
p.chat.forcedRating(" 9\n");
|
||||
List<Document> padded = p.reranker.process(new Query(QUESTION), candidates.subList(0, 1));
|
||||
t.line("reply \" 9\\n\" (padded) -> score %s", padded.getFirst().getMetadata().get(LlmReranker.SCORE));
|
||||
|
||||
assertThat(p.chat.calls()).isPositive();
|
||||
assertThat(kept).hasSize(2);
|
||||
assertThat(concurrent).isLessThan(1000);
|
||||
assertThat(sequential).isGreaterThanOrEqualTo(4000);
|
||||
assertThat(failures).isEqualTo(candidates.size());
|
||||
assertThat(degraded.stream().map(d -> d.getMetadata().get(LlmReranker.SCORE))).containsOnly(0);
|
||||
assertThat(degraded.stream().map(d -> d.getMetadata().get("page_number"))).containsExactly(2, 3);
|
||||
assertThat(padded.getFirst().getMetadata().get(LlmReranker.SCORE)).isEqualTo(9);
|
||||
}
|
||||
}
|
||||
|
||||
private static RagProperties withTopN(int topN) {
|
||||
return new RagProperties(new RagProperties.Retrieval(20, 0.0), new RagProperties.Rerank(topN),
|
||||
new RagProperties.Chunking(512, 350));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.BeforeAll;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.rag.Query;
|
||||
import org.springframework.ai.rag.advisor.RetrievalAugmentationAdvisor;
|
||||
import org.springframework.ai.rag.generation.augmentation.ContextualQueryAugmenter;
|
||||
import org.springframework.ai.rag.retrieval.search.VectorStoreDocumentRetriever;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Chapter 4, first half: what the retriever returns, what the similarity threshold does, what the
|
||||
* prompt looks like once the chunks are in it, and what happens when nothing is retrieved.
|
||||
* Transcript 06. Uses the hashing embedding model, so the scores are this repository's, not a real
|
||||
* model's: the point is the behaviour of the Spring AI components, not the numbers.
|
||||
*/
|
||||
class RetrievalTest {
|
||||
|
||||
private static final String LEAVE = "How many days of annual leave do employees get?";
|
||||
private static final String OFF_TOPIC = "What is the capital of Mongolia?";
|
||||
|
||||
private static TestPipeline pipeline;
|
||||
|
||||
@BeforeAll
|
||||
static void ingest() {
|
||||
pipeline = TestPipeline.on(Stores::simple);
|
||||
pipeline.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void thresholdPromptShapeAndEmptyContext() {
|
||||
try (Transcript t = new Transcript("06-retrieval-threshold-and-prompt.txt",
|
||||
"Retrieval: threshold, prompt shape and the empty-context path")) {
|
||||
t.line("the store holds the 4 pages of one handbook as 4 chunks; topK = 20");
|
||||
|
||||
t.section("question: \"" + LEAVE + "\"");
|
||||
List<Document> allAccepted = retrieve(0.0, LEAVE);
|
||||
List<Document> filtered = retrieve(0.3, LEAVE);
|
||||
scores(t, "similarityThreshold 0.0 (the default)", allAccepted);
|
||||
scores(t, "similarityThreshold 0.3", filtered);
|
||||
|
||||
t.section("question: \"" + OFF_TOPIC + "\"");
|
||||
List<Document> offAll = retrieve(0.0, OFF_TOPIC);
|
||||
List<Document> offFiltered = retrieve(0.3, OFF_TOPIC);
|
||||
scores(t, "similarityThreshold 0.0 (the default)", offAll);
|
||||
scores(t, "similarityThreshold 0.3", offFiltered);
|
||||
|
||||
t.section("the prompt the model receives (threshold 0.3, question about leave)");
|
||||
String augmented = promptFor(0.3, false, LEAVE);
|
||||
t.line(augmented);
|
||||
|
||||
t.section("nothing retrieved, allowEmptyContext(false): the prompt the model receives");
|
||||
String refusing = promptFor(0.3, false, OFF_TOPIC);
|
||||
t.line(refusing);
|
||||
|
||||
t.section("nothing retrieved, allowEmptyContext(true): the prompt the model receives");
|
||||
String passthrough = promptFor(0.3, true, OFF_TOPIC);
|
||||
t.line(passthrough);
|
||||
|
||||
t.section("off-topic question, default threshold 0.0, allowEmptyContext(false)");
|
||||
String offDefault = promptFor(0.0, false, OFF_TOPIC);
|
||||
t.line("chunks placed in the prompt: %d of 4", count(offDefault, "Full-time") + count(offDefault, "Expenses")
|
||||
+ count(offDefault, "Requesting Leave") + count(offDefault, "Probationary"));
|
||||
t.line("the empty-context safety net fired: %s", offDefault.contains("outside your knowledge base"));
|
||||
|
||||
assertThat(allAccepted).hasSize(4);
|
||||
assertThat(filtered).hasSize(2);
|
||||
assertThat(filtered.getFirst().getText()).contains("20 working days");
|
||||
assertThat(offAll).hasSize(4);
|
||||
assertThat(offFiltered).isEmpty();
|
||||
assertThat(augmented).contains("Context information is below").contains("20 working days");
|
||||
assertThat(refusing).contains("outside your knowledge base").doesNotContain("Context information is below");
|
||||
assertThat(passthrough).isEqualTo(OFF_TOPIC);
|
||||
assertThat(offDefault).doesNotContain("outside your knowledge base").contains("Context information is below");
|
||||
}
|
||||
}
|
||||
|
||||
private static List<Document> retrieve(double threshold, String question) {
|
||||
return VectorStoreDocumentRetriever.builder().vectorStore(pipeline.store).topK(20)
|
||||
.similarityThreshold(threshold).build().retrieve(new Query(question));
|
||||
}
|
||||
|
||||
private static String promptFor(double threshold, boolean allowEmptyContext, String question) {
|
||||
pipeline.chat.reset();
|
||||
var advisor = RetrievalAugmentationAdvisor.builder()
|
||||
.documentRetriever(VectorStoreDocumentRetriever.builder().vectorStore(pipeline.store).topK(20)
|
||||
.similarityThreshold(threshold).build())
|
||||
.queryAugmenter(ContextualQueryAugmenter.builder().allowEmptyContext(allowEmptyContext).build())
|
||||
.build();
|
||||
ChatClient.builder(pipeline.chat).defaultAdvisors(advisor).build().prompt().user(question).call().content();
|
||||
return pipeline.chat.prompts().getFirst().getContents();
|
||||
}
|
||||
|
||||
private static void scores(Transcript t, String label, List<Document> docs) {
|
||||
t.line("%s: %d chunk(s)", label, docs.size());
|
||||
// Equal scores come back in no promised order, so ties are listed by page to keep this file stable.
|
||||
List<Document> ordered = docs.stream().sorted(java.util.Comparator
|
||||
.comparing(Document::getScore, java.util.Comparator.reverseOrder())
|
||||
.thenComparing(d -> (Integer) d.getMetadata().get("page_number"))).toList();
|
||||
for (Document d : ordered) {
|
||||
String text = d.getText();
|
||||
t.line(" score %.4f page %s \"%s...\"", d.getScore(), d.getMetadata().get("page_number"),
|
||||
text.substring(0, Math.min(44, text.length())).replace('\n', ' '));
|
||||
}
|
||||
}
|
||||
|
||||
private static int count(String text, String needle) {
|
||||
return text.contains(needle) ? 1 : 0;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.client.advisor.vectorstore.QuestionAnswerAdvisor;
|
||||
import org.springframework.ai.vectorstore.SearchRequest;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Chapter 1: the smallest RAG that works, one advisor and no other class. Transcript 11.
|
||||
* It is where the article starts, and the comparison at the end shows what it does not do.
|
||||
*/
|
||||
class SimpleAdvisorTest {
|
||||
|
||||
private static final String QUESTION = "How many days of annual leave do employees get?";
|
||||
private static final String OFF_TOPIC = "What is the capital of Mongolia?";
|
||||
|
||||
@Test
|
||||
void smallestWorkingRag() {
|
||||
TestPipeline p = TestPipeline.on(Stores::simple);
|
||||
p.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
p.chat.reset();
|
||||
|
||||
ChatClient client = ChatClient.builder(p.chat)
|
||||
.defaultAdvisors(QuestionAnswerAdvisor.builder(p.store)
|
||||
.searchRequest(SearchRequest.builder().topK(2).similarityThreshold(0.3).build())
|
||||
.build())
|
||||
.build();
|
||||
|
||||
try (Transcript t = new Transcript("11-question-answer-advisor.txt", "QuestionAnswerAdvisor: the smallest RAG")) {
|
||||
client.prompt().user(QUESTION).call().content();
|
||||
String prompt = p.chat.prompts().getFirst().getContents();
|
||||
t.line("only the prompts are recorded: what a real model would reply is not something this repository tests");
|
||||
t.section("the prompt the model received");
|
||||
t.line(prompt);
|
||||
|
||||
p.chat.reset();
|
||||
client.prompt().user(OFF_TOPIC).call().content();
|
||||
String offTopicPrompt = p.chat.prompts().getFirst().getContents();
|
||||
t.section("the prompt for an off-topic question (nothing passes the threshold)");
|
||||
t.line(offTopicPrompt);
|
||||
|
||||
assertThat(prompt).contains("20 working days").contains(QUESTION);
|
||||
assertThat(offTopicPrompt).contains(OFF_TOPIC);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package com.ankurm.rag;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import com.ankurm.rag.config.TestPipeline;
|
||||
import com.ankurm.rag.support.SamplePdf;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import com.ankurm.rag.support.Transcript;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.vectorstore.SearchRequest;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
/**
|
||||
* Chapter 3 and 4: keeping one tenant's chunks away from another's, and how a filter can be built
|
||||
* so that user input cannot rewrite it. Runs the same checks on the in-memory store and on real
|
||||
* PostgreSQL, because the two translate a filter differently. Transcript 07.
|
||||
*/
|
||||
class TenantFilterTest {
|
||||
|
||||
private static final String QUESTION = "How many days of annual leave do employees get?";
|
||||
/** A caller who belongs to globex sends this as their tenant id. */
|
||||
private static final String HOSTILE_SINGLE = "globex' || tenant_id == 'acme";
|
||||
private static final String HOSTILE_DOUBLE = "globex\" || @.tenant_id == \"acme";
|
||||
|
||||
@Test
|
||||
void filtersOnBothStores() {
|
||||
try (Transcript t = new Transcript("07-tenant-filter-and-injection.txt",
|
||||
"Tenant isolation: metadata filters, and a filter built from user input")) {
|
||||
t.line("two tenants each upload a handbook with a section 4.1 on annual leave");
|
||||
t.line("acme says 20 working days, globex says 25");
|
||||
|
||||
var simple = TestPipeline.on(Stores::simple);
|
||||
load(simple);
|
||||
check(t, "SimpleVectorStore (in memory)", simple.store);
|
||||
|
||||
var jdbc = Stores.jdbc();
|
||||
var pg = TestPipeline.on(e -> Stores.pg(jdbc, e));
|
||||
load(pg);
|
||||
check(t, "PgVectorStore (PostgreSQL + pgvector)", pg.store);
|
||||
}
|
||||
}
|
||||
|
||||
private static void load(TestPipeline p) {
|
||||
p.ingestion.ingestPdf(SamplePdf.of(SamplePdf.ACME), "acme-handbook.pdf", Map.of("tenant_id", "acme"));
|
||||
p.ingestion.ingestPdf(SamplePdf.of(SamplePdf.GLOBEX), "globex-manual.pdf", Map.of("tenant_id", "globex"));
|
||||
}
|
||||
|
||||
private static void check(Transcript t, String storeName, VectorStore store) {
|
||||
t.section(storeName);
|
||||
|
||||
List<Document> unfiltered = store.similaritySearch(SearchRequest.builder().query(QUESTION).topK(3).build());
|
||||
describe(t, "no filter, top 3", unfiltered);
|
||||
|
||||
var acmeOnly = new FilterExpressionBuilder().eq("tenant_id", "acme").build();
|
||||
List<Document> scoped = store.similaritySearch(
|
||||
SearchRequest.builder().query(QUESTION).topK(3).filterExpression(acmeOnly).build());
|
||||
describe(t, "eq(\"tenant_id\", \"acme\") built with FilterExpressionBuilder, top 3", scoped);
|
||||
|
||||
// The mistake: build the filter as text from a value the caller controls.
|
||||
String glued = "tenant_id == '" + HOSTILE_SINGLE + "'";
|
||||
t.line("filter string built by concatenation: %s", glued);
|
||||
List<Document> injected = store.similaritySearch(
|
||||
SearchRequest.builder().query(QUESTION).topK(5).filterExpression(glued).build());
|
||||
describe(t, " result, top 5", injected);
|
||||
|
||||
// The same hostile text as a value handed to the builder is only ever a value.
|
||||
List<Document> asValue = store.similaritySearch(SearchRequest.builder().query(QUESTION).topK(5)
|
||||
.filterExpression(new FilterExpressionBuilder().eq("tenant_id", HOSTILE_SINGLE).build()).build());
|
||||
describe(t, "same text passed to FilterExpressionBuilder.eq(), top 5", asValue);
|
||||
List<Document> asValueDouble = store.similaritySearch(SearchRequest.builder().query(QUESTION).topK(5)
|
||||
.filterExpression(new FilterExpressionBuilder().eq("tenant_id", HOSTILE_DOUBLE).build()).build());
|
||||
describe(t, "double-quote variant passed to FilterExpressionBuilder.eq(), top 5", asValueDouble);
|
||||
|
||||
assertThat(tenants(unfiltered)).contains("globex");
|
||||
assertThat(tenants(scoped)).containsOnly("acme").hasSize(3);
|
||||
assertThat(tenants(injected)).contains("acme", "globex");
|
||||
assertThat(asValue).isEmpty();
|
||||
assertThat(asValueDouble).isEmpty();
|
||||
}
|
||||
|
||||
private static List<String> tenants(List<Document> docs) {
|
||||
return docs.stream().map(d -> String.valueOf(d.getMetadata().get("tenant_id"))).toList();
|
||||
}
|
||||
|
||||
private static void describe(Transcript t, String label, List<Document> docs) {
|
||||
t.line("%s: %d chunk(s)", label, docs.size());
|
||||
for (Document d : docs) {
|
||||
t.line(" tenant=%-6s page=%s \"%s...\"", d.getMetadata().get("tenant_id"), d.getMetadata().get("page_number"),
|
||||
d.getText().substring(0, 40).replace('\n', ' '));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package com.ankurm.rag.config;
|
||||
|
||||
import java.util.function.Function;
|
||||
|
||||
import com.ankurm.rag.config.RagProperties.Chunking;
|
||||
import com.ankurm.rag.config.RagProperties.Rerank;
|
||||
import com.ankurm.rag.config.RagProperties.Retrieval;
|
||||
import com.ankurm.rag.ingest.IngestionService;
|
||||
import com.ankurm.rag.ingest.IngestionTracker;
|
||||
import com.ankurm.rag.query.LlmReranker;
|
||||
import com.ankurm.rag.query.RagQueryService;
|
||||
import com.ankurm.rag.support.FakeChatModel;
|
||||
import com.ankurm.rag.support.HashingEmbeddingModel;
|
||||
import com.ankurm.rag.support.Stores;
|
||||
import io.micrometer.core.instrument.simple.SimpleMeterRegistry;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.document.DocumentTransformer;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.vectorstore.VectorStore;
|
||||
|
||||
/**
|
||||
* The whole pipeline, wired with the same {@link RagConfig} the application uses but with the fake
|
||||
* chat and embedding models, so a test can run it with no API key. It lives in the same package as
|
||||
* {@code RagConfig} because the {@code @Bean} methods are package-private.
|
||||
*/
|
||||
public final class TestPipeline {
|
||||
|
||||
public final FakeChatModel chat = new FakeChatModel();
|
||||
public final HashingEmbeddingModel embeddings = new HashingEmbeddingModel(Stores.DIMS);
|
||||
public final SimpleMeterRegistry metrics = new SimpleMeterRegistry();
|
||||
public final IngestionTracker tracker = new IngestionTracker();
|
||||
public final RagProperties properties;
|
||||
public final VectorStore store;
|
||||
public final RagConfig config = new RagConfig();
|
||||
public final DocumentTransformer chunker;
|
||||
public final LlmReranker reranker;
|
||||
public final IngestionService ingestion;
|
||||
public final RagQueryService queries;
|
||||
|
||||
private TestPipeline(RagProperties properties, Function<EmbeddingModel, VectorStore> storeFactory) {
|
||||
this.properties = properties;
|
||||
this.store = storeFactory.apply(embeddings);
|
||||
ChatClient.Builder builder = ChatClient.builder(chat);
|
||||
this.reranker = config.reranker(builder, properties, metrics);
|
||||
this.chunker = config.chunker(properties);
|
||||
this.ingestion = new IngestionService(store, chunker, tracker, metrics);
|
||||
var advisor = config.retrievalAdvisor(store, reranker, properties);
|
||||
this.queries = new RagQueryService(config.ragChatClient(ChatClient.builder(chat), advisor),
|
||||
config.factChecker(ChatClient.builder(chat)), metrics);
|
||||
}
|
||||
|
||||
public static RagProperties defaults() {
|
||||
return new RagProperties(new Retrieval(20, 0.0), new Rerank(5), new Chunking(512, 350));
|
||||
}
|
||||
|
||||
public static TestPipeline on(Function<EmbeddingModel, VectorStore> storeFactory) {
|
||||
return new TestPipeline(defaults(), storeFactory);
|
||||
}
|
||||
|
||||
public static TestPipeline on(RagProperties properties, Function<EmbeddingModel, VectorStore> storeFactory) {
|
||||
return new TestPipeline(properties, storeFactory);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
package com.ankurm.rag.support;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
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.Prompt;
|
||||
|
||||
/**
|
||||
* A scripted chat model. It plays three roles, told apart by what the prompt says, so one bean
|
||||
* serves the answering client, the reranker and the fact checker:
|
||||
*
|
||||
* <ul>
|
||||
* <li>a <b>rating</b> prompt gets the number of question words found in the passage, scaled to 0-10;</li>
|
||||
* <li>a <b>fact-check</b> prompt gets "yes" if every content word of the claim is in the document;</li>
|
||||
* <li>anything else is a <b>question</b>: the reply is the sentence of the supplied context that
|
||||
* shares the most words with the question, or the refusal if there is no context.</li>
|
||||
* </ul>
|
||||
*
|
||||
* <p>It records every prompt it receives, which is what the tests assert on: what Spring AI
|
||||
* <em>sent</em> is the part of a RAG pipeline that Spring AI is responsible for. What a real model
|
||||
* would <em>reply</em> is not tested anywhere in this repository.
|
||||
*/
|
||||
public class FakeChatModel implements ChatModel {
|
||||
|
||||
public static final String REFUSAL = "I don't have enough information in the provided documents.";
|
||||
|
||||
private static final Pattern QUESTION_LINE = Pattern.compile("QUESTION: (.*)");
|
||||
private static final Pattern PASSAGE_LINE = Pattern.compile("PASSAGE: (.*)", Pattern.DOTALL);
|
||||
private static final Pattern QUERY_LINE = Pattern.compile("Query: (.*)");
|
||||
private static final Pattern CONTEXT_BLOCK = Pattern.compile("Context information is below\\.\\s*-+\\s*(.*?)\\s*-+\\s*Given",
|
||||
Pattern.DOTALL);
|
||||
|
||||
private final List<Prompt> prompts = new CopyOnWriteArrayList<>();
|
||||
private final AtomicInteger calls = new AtomicInteger();
|
||||
|
||||
/** Milliseconds each call sleeps, to stand in for network latency. */
|
||||
private volatile long latencyMillis;
|
||||
/** When set, question replies are this text instead of a sentence taken from the context. */
|
||||
private volatile String forcedAnswer;
|
||||
/** When set, every fact-check reply is this text. */
|
||||
private volatile String forcedVerdict;
|
||||
/** When set, every rating reply is this text. */
|
||||
private volatile String forcedRating;
|
||||
/** When true, a fact-check call throws, as a network failure or a rate limit would. */
|
||||
private volatile boolean failFactChecks;
|
||||
|
||||
public List<Prompt> prompts() {
|
||||
return Collections.unmodifiableList(prompts);
|
||||
}
|
||||
|
||||
public int calls() {
|
||||
return calls.get();
|
||||
}
|
||||
|
||||
public void reset() {
|
||||
prompts.clear();
|
||||
calls.set(0);
|
||||
latencyMillis = 0;
|
||||
forcedAnswer = null;
|
||||
forcedVerdict = null;
|
||||
forcedRating = null;
|
||||
failFactChecks = false;
|
||||
}
|
||||
|
||||
public void latencyMillis(long millis) {
|
||||
this.latencyMillis = millis;
|
||||
}
|
||||
|
||||
public void forcedAnswer(String answer) {
|
||||
this.forcedAnswer = answer;
|
||||
}
|
||||
|
||||
public void forcedVerdict(String verdict) {
|
||||
this.forcedVerdict = verdict;
|
||||
}
|
||||
|
||||
public void failFactChecks(boolean fail) {
|
||||
this.failFactChecks = fail;
|
||||
}
|
||||
|
||||
public void forcedRating(String rating) {
|
||||
this.forcedRating = rating;
|
||||
}
|
||||
|
||||
@Override
|
||||
public ChatResponse call(Prompt prompt) {
|
||||
prompts.add(prompt);
|
||||
calls.incrementAndGet();
|
||||
if (latencyMillis > 0) {
|
||||
try {
|
||||
Thread.sleep(latencyMillis);
|
||||
} catch (InterruptedException e) {
|
||||
Thread.currentThread().interrupt();
|
||||
}
|
||||
}
|
||||
return new ChatResponse(List.of(new Generation(new AssistantMessage(reply(prompt.getContents())))));
|
||||
}
|
||||
|
||||
private String reply(String text) {
|
||||
if (text.contains("Rate how well the PASSAGE")) {
|
||||
return forcedRating != null ? forcedRating : rating(text);
|
||||
}
|
||||
if (text.contains("Evaluate whether or not the following claim is supported")) {
|
||||
if (failFactChecks) {
|
||||
throw new IllegalStateException("simulated: 429 Too Many Requests from the judge model");
|
||||
}
|
||||
return forcedVerdict != null ? forcedVerdict : verdict(text);
|
||||
}
|
||||
if (forcedAnswer != null) {
|
||||
return forcedAnswer;
|
||||
}
|
||||
return answer(text);
|
||||
}
|
||||
|
||||
private static String rating(String text) {
|
||||
Matcher q = QUESTION_LINE.matcher(text);
|
||||
Matcher p = PASSAGE_LINE.matcher(text);
|
||||
if (!q.find() || !p.find()) {
|
||||
return "0";
|
||||
}
|
||||
List<String> question = HashingEmbeddingModel.words(q.group(1));
|
||||
List<String> passage = HashingEmbeddingModel.words(p.group(1));
|
||||
long hits = question.stream().filter(passage::contains).count();
|
||||
return String.valueOf(question.isEmpty() ? 0 : Math.round(10.0 * hits / question.size()));
|
||||
}
|
||||
|
||||
private static String verdict(String text) {
|
||||
int doc = text.indexOf("Document:");
|
||||
int claim = text.indexOf("Claim:");
|
||||
List<String> document = HashingEmbeddingModel.words(text.substring(doc + 9, claim));
|
||||
List<String> claimWords = HashingEmbeddingModel.words(text.substring(claim + 6));
|
||||
return document.containsAll(claimWords) ? "yes" : "no";
|
||||
}
|
||||
|
||||
private static String answer(String text) {
|
||||
Matcher context = CONTEXT_BLOCK.matcher(text);
|
||||
if (!context.find() || context.group(1).isBlank()) {
|
||||
return REFUSAL;
|
||||
}
|
||||
Matcher query = QUERY_LINE.matcher(text);
|
||||
List<String> questionWords = HashingEmbeddingModel.words(query.find() ? query.group(1) : "");
|
||||
String best = REFUSAL;
|
||||
long bestHits = -1;
|
||||
for (String sentence : context.group(1).split("(?<=[.!?])\\s+")) {
|
||||
List<String> words = HashingEmbeddingModel.words(sentence);
|
||||
long hits = questionWords.stream().filter(words::contains).count();
|
||||
if (hits > bestHits) {
|
||||
bestHits = hits;
|
||||
best = sentence.strip();
|
||||
}
|
||||
}
|
||||
return best;
|
||||
}
|
||||
|
||||
/** Convenience for tests: the user text of every prompt whose text contains {@code marker}. */
|
||||
public List<String> promptsContaining(String marker) {
|
||||
List<String> out = new ArrayList<>();
|
||||
for (Prompt p : prompts) {
|
||||
if (p.getContents().contains(marker)) {
|
||||
out.add(p.getContents());
|
||||
}
|
||||
}
|
||||
return out;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package com.ankurm.rag.support;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Locale;
|
||||
import java.util.Set;
|
||||
|
||||
import org.springframework.ai.document.Document;
|
||||
import org.springframework.ai.embedding.Embedding;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.embedding.EmbeddingRequest;
|
||||
import org.springframework.ai.embedding.EmbeddingResponse;
|
||||
|
||||
/**
|
||||
* A deterministic stand-in for a real embedding model: each word is hashed into one of {@code dims}
|
||||
* slots and the vector is normalised. Two texts are close when they share words, which is enough to
|
||||
* make retrieval behave the way it should, and nothing like a real model's sense of meaning.
|
||||
*
|
||||
* <p>Nothing here calls a network, so a test that uses it can be run by anyone with no API key.
|
||||
* It also counts the texts it was asked to embed, which is how the ingestion tests show that
|
||||
* skipping an unchanged file spends nothing.
|
||||
*/
|
||||
public class HashingEmbeddingModel implements EmbeddingModel {
|
||||
|
||||
private static final Set<String> STOPWORDS = Set.of("a", "an", "the", "of", "to", "in", "is", "are", "and",
|
||||
"or", "for", "on", "at", "be", "by", "do", "does", "how", "what", "many", "can", "it", "its");
|
||||
|
||||
private final int dims;
|
||||
private int textsEmbedded;
|
||||
|
||||
public HashingEmbeddingModel(int dims) {
|
||||
this.dims = dims;
|
||||
}
|
||||
|
||||
public synchronized int textsEmbedded() {
|
||||
return textsEmbedded;
|
||||
}
|
||||
|
||||
@Override
|
||||
public EmbeddingResponse call(EmbeddingRequest request) {
|
||||
List<Embedding> out = new ArrayList<>();
|
||||
for (int i = 0; i < request.getInstructions().size(); i++) {
|
||||
out.add(new Embedding(vector(request.getInstructions().get(i)), i));
|
||||
}
|
||||
return new EmbeddingResponse(out);
|
||||
}
|
||||
|
||||
@Override
|
||||
public float[] embed(Document document) {
|
||||
return vector(document.getText());
|
||||
}
|
||||
|
||||
@Override
|
||||
public int dimensions() {
|
||||
return dims;
|
||||
}
|
||||
|
||||
/** The words the model keeps for a text, lower-cased with a plural "s" removed. */
|
||||
public static List<String> words(String text) {
|
||||
List<String> words = new ArrayList<>();
|
||||
for (String w : text.toLowerCase(Locale.ROOT).split("[^a-z0-9]+")) {
|
||||
if (w.isEmpty() || STOPWORDS.contains(w)) {
|
||||
continue;
|
||||
}
|
||||
words.add(w.length() > 3 && w.endsWith("s") ? w.substring(0, w.length() - 1) : w);
|
||||
}
|
||||
return words;
|
||||
}
|
||||
|
||||
private synchronized float[] vector(String text) {
|
||||
textsEmbedded++;
|
||||
float[] v = new float[dims];
|
||||
for (String w : words(text)) {
|
||||
v[Math.floorMod(w.hashCode() * 0x9E3779B1, dims)] += 1f;
|
||||
}
|
||||
double norm = 0;
|
||||
for (float x : v) {
|
||||
norm += x * x;
|
||||
}
|
||||
norm = Math.sqrt(norm);
|
||||
if (norm == 0) {
|
||||
v[0] = 1f;
|
||||
return v;
|
||||
}
|
||||
for (int i = 0; i < dims; i++) {
|
||||
v[i] /= (float) norm;
|
||||
}
|
||||
return v;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package com.ankurm.rag.support;
|
||||
|
||||
import java.io.ByteArrayOutputStream;
|
||||
import java.io.IOException;
|
||||
import java.util.List;
|
||||
|
||||
import org.apache.pdfbox.pdmodel.PDDocument;
|
||||
import org.apache.pdfbox.pdmodel.PDPage;
|
||||
import org.apache.pdfbox.pdmodel.PDPageContentStream;
|
||||
import org.apache.pdfbox.pdmodel.font.PDType1Font;
|
||||
import org.apache.pdfbox.pdmodel.font.Standard14Fonts;
|
||||
import org.springframework.core.io.ByteArrayResource;
|
||||
import org.springframework.core.io.Resource;
|
||||
|
||||
/**
|
||||
* Builds small handbook PDFs on the fly, one page per entry, so the tests read a real PDF through
|
||||
* the real {@code PagePdfDocumentReader} without a binary file in the repository.
|
||||
*/
|
||||
public final class SamplePdf {
|
||||
|
||||
/** Two companies with different rules, so a tenant filter has something to keep apart. */
|
||||
public static final List<List<String>> ACME = List.of(
|
||||
List.of("Acme Employee Handbook 2026", "",
|
||||
"3.5 Probationary Period. During the three month probationary period annual leave",
|
||||
"accrues at half the normal rate. Probation can be extended once by four weeks."),
|
||||
List.of("4.1 Annual Leave Entitlement. Full-time employees are entitled to 20 working days",
|
||||
"of annual leave per calendar year. Part-time employees receive leave pro rata.",
|
||||
"",
|
||||
"4.2 Leave Carryover. Unused annual leave may be carried over for a maximum of 5 days",
|
||||
"into the next calendar year and must be used by 31 March."),
|
||||
List.of("4.3 Requesting Leave. All leave requests must be submitted through the HR portal",
|
||||
"at least two weeks in advance for absences longer than three days.",
|
||||
"",
|
||||
"4.7 Sick Leave. Sick leave is separate from annual leave and is not deducted from it.",
|
||||
"A medical certificate is required after three consecutive days."),
|
||||
List.of("6.1 Expenses. Receipts are required for every expense above 50 euros.",
|
||||
"Claims must be filed within 30 days of the expense.",
|
||||
"",
|
||||
"7.2 Remote Work. Employees may work remotely up to two days per week",
|
||||
"with the agreement of their manager."));
|
||||
|
||||
public static final List<List<String>> GLOBEX = List.of(
|
||||
List.of("Globex Staff Manual 2026", "",
|
||||
"4.1 Annual Leave Entitlement. Full-time staff are entitled to 25 working days",
|
||||
"of annual leave per calendar year. Unused annual leave cannot be carried over.",
|
||||
"",
|
||||
"5.4 Overtime. Overtime must be approved in advance and is paid at 1.5 times the base rate."));
|
||||
|
||||
private SamplePdf() {
|
||||
}
|
||||
|
||||
public static Resource of(List<List<String>> pages) {
|
||||
try (PDDocument document = new PDDocument(); ByteArrayOutputStream bytes = new ByteArrayOutputStream()) {
|
||||
PDType1Font font = new PDType1Font(Standard14Fonts.FontName.HELVETICA);
|
||||
for (List<String> lines : pages) {
|
||||
PDPage page = new PDPage();
|
||||
document.addPage(page);
|
||||
try (PDPageContentStream stream = new PDPageContentStream(document, page)) {
|
||||
stream.setFont(font, 11);
|
||||
stream.beginText();
|
||||
stream.setLeading(16);
|
||||
stream.newLineAtOffset(50, 730);
|
||||
for (String line : lines) {
|
||||
stream.showText(line);
|
||||
stream.newLine();
|
||||
}
|
||||
stream.endText();
|
||||
}
|
||||
}
|
||||
document.save(bytes);
|
||||
return new ByteArrayResource(bytes.toByteArray());
|
||||
} catch (IOException e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package com.ankurm.rag.support;
|
||||
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.ai.vectorstore.SimpleVectorStore;
|
||||
import org.springframework.ai.vectorstore.pgvector.PgVectorStore;
|
||||
import org.springframework.jdbc.core.JdbcTemplate;
|
||||
import org.springframework.jdbc.datasource.DriverManagerDataSource;
|
||||
|
||||
/**
|
||||
* The two vector stores the tests use. {@link #simple} keeps vectors in a map in memory;
|
||||
* {@link #pg} talks to a real PostgreSQL with the pgvector extension.
|
||||
*
|
||||
* <p>The PostgreSQL location comes from {@code RAG_PG_URL} (default the docker-compose one,
|
||||
* {@code localhost:5432}); see {@code scripts/pg-up.sh} for starting one without Docker.
|
||||
*/
|
||||
public final class Stores {
|
||||
|
||||
public static final String PG_URL = env("RAG_PG_URL", "jdbc:postgresql://localhost:5432/ragdb");
|
||||
public static final String PG_USER = env("RAG_PG_USER", "raguser");
|
||||
public static final String PG_PASSWORD = env("RAG_PG_PASSWORD", "ragpass");
|
||||
/** Matches text-embedding-3-small, the model the shipped application.yml names. */
|
||||
public static final int DIMS = 1536;
|
||||
|
||||
private Stores() {
|
||||
}
|
||||
|
||||
public static SimpleVectorStore simple(EmbeddingModel embeddings) {
|
||||
return SimpleVectorStore.builder(embeddings).build();
|
||||
}
|
||||
|
||||
public static JdbcTemplate jdbc() {
|
||||
return new JdbcTemplate(new DriverManagerDataSource(PG_URL, PG_USER, PG_PASSWORD));
|
||||
}
|
||||
|
||||
/** A pgvector store with a freshly created, empty table and an HNSW cosine index. */
|
||||
public static PgVectorStore pg(JdbcTemplate jdbc, EmbeddingModel embeddings) {
|
||||
PgVectorStore store = PgVectorStore.builder(jdbc, embeddings)
|
||||
.dimensions(DIMS)
|
||||
.distanceType(PgVectorStore.PgDistanceType.COSINE_DISTANCE)
|
||||
.indexType(PgVectorStore.PgIndexType.HNSW)
|
||||
.removeExistingVectorStoreTable(true)
|
||||
.initializeSchema(true)
|
||||
.build();
|
||||
store.afterPropertiesSet();
|
||||
return store;
|
||||
}
|
||||
|
||||
private static String env(String name, String fallback) {
|
||||
String value = System.getenv(name);
|
||||
return value == null || value.isBlank() ? fallback : value;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package com.ankurm.rag.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 docs/output/} 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("docs", "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;
|
||||
}
|
||||
|
||||
public Transcript section(String heading) {
|
||||
out.println();
|
||||
out.println("--- " + heading + " ---");
|
||||
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