Request-time filter expressions for RAG

When using the RetrievalAugmentationAdvisor with the VectorStoreDocumentRetriever, it’s now possible to provide a filter expression at request-time as an advisor context variable with key VectorStoreDocumentRetriever.FILTER_EXPRESSION.

Fixes gh-1776

Signed-off-by: Thomas Vitale <ThomasVitale@users.noreply.github.com>
This commit is contained in:
Thomas Vitale
2025-03-09 11:29:34 +01:00
committed by Ilayaperumal Gopinathan
parent 82b46d2182
commit 5a4e9f5108
6 changed files with 132 additions and 16 deletions

View File

@@ -24,8 +24,10 @@ import org.springframework.ai.rag.Query;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionTextParser;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
/**
* Retrieves documents from a vector store that are semantically similar to the input
@@ -48,6 +50,8 @@ import org.springframework.util.Assert;
*/
public final class VectorStoreDocumentRetriever implements DocumentRetriever {
public static final String FILTER_EXPRESSION = "vector_store_filter_expression";
private final VectorStore vectorStore;
private final Double similarityThreshold;
@@ -75,15 +79,24 @@ public final class VectorStoreDocumentRetriever implements DocumentRetriever {
@Override
public List<Document> retrieve(Query query) {
Assert.notNull(query, "query cannot be null");
var requestFilterExpression = computeRequestFilterExpression(query);
var searchRequest = SearchRequest.builder()
.query(query.text())
.filterExpression(this.filterExpression.get())
.filterExpression(requestFilterExpression)
.similarityThreshold(this.similarityThreshold)
.topK(this.topK)
.build();
return this.vectorStore.similaritySearch(searchRequest);
}
private Filter.Expression computeRequestFilterExpression(Query query) {
var contextFilterExpression = query.context().get(FILTER_EXPRESSION);
if (contextFilterExpression != null && StringUtils.hasText(contextFilterExpression.toString())) {
return new FilterExpressionTextParser().parse(contextFilterExpression.toString());
}
return this.filterExpression.get();
}
public static Builder builder() {
return new Builder();
}

View File

@@ -210,6 +210,30 @@ class VectorStoreDocumentRetrieverTests {
assertThat(result).hasSize(2).containsExactlyElementsOf(mockDocuments);
}
@Test
void retrieveWithQueryObjectAndRequestFilterExpression() {
var mockVectorStore = mock(VectorStore.class);
var documentRetriever = VectorStoreDocumentRetriever.builder().vectorStore(mockVectorStore).build();
var query = Query.builder()
.text("test query")
.context(Map.of(VectorStoreDocumentRetriever.FILTER_EXPRESSION, "location == 'Rivendell'"))
.build();
documentRetriever.retrieve(query);
// Verify the mock interaction
var searchRequestCaptor = ArgumentCaptor.forClass(SearchRequest.class);
verify(mockVectorStore).similaritySearch(searchRequestCaptor.capture());
// Verify the search request
var searchRequest = searchRequestCaptor.getValue();
assertThat(searchRequest.getQuery()).isEqualTo("test query");
assertThat(searchRequest.getSimilarityThreshold()).isEqualTo(SearchRequest.SIMILARITY_THRESHOLD_ACCEPT_ALL);
assertThat(searchRequest.getTopK()).isEqualTo(SearchRequest.DEFAULT_TOP_K);
assertThat(searchRequest.getFilterExpression())
.isEqualTo(new FilterExpressionBuilder().eq("location", "Rivendell").build());
}
static final class TenantContextHolder {
private static final ThreadLocal<String> tenantIdentifier = new ThreadLocal<>();