From 0abf1528dd6d6bb3a6747c4bb435b30bb4c3dd46 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Wed, 31 Jan 2024 11:17:08 +0100 Subject: [PATCH] Fix for Redis NPE when using FilterExpression #265 - Part of #265 --- .../RedisFilterExpressionConverter.java | 4 +- .../ai/vectorstore/RedisVectorStore.java | 82 ++++++++++--------- 2 files changed, 45 insertions(+), 41 deletions(-) diff --git a/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisFilterExpressionConverter.java b/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisFilterExpressionConverter.java index a2fdca587..40e3f5129 100644 --- a/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisFilterExpressionConverter.java +++ b/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisFilterExpressionConverter.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2023 the original author or authors. + * Copyright 2023-2024 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -104,7 +104,7 @@ public class RedisFilterExpressionConverter extends AbstractFilterExpressionConv private void doField(Expression expression, StringBuilder context) { Key key = (Key) expression.left(); doKey(key, context); - MetadataField field = metadataFields.getOrDefault(key.key(), MetadataField.tag(key.key())); + MetadataField field = this.metadataFields.getOrDefault(key.key(), MetadataField.tag(key.key())); Value value = (Value) expression.right(); switch (field.fieldType()) { case NUMERIC: diff --git a/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java b/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java index 5dff41909..331c6eca2 100644 --- a/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java +++ b/vector-stores/spring-ai-redis/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java @@ -1,5 +1,5 @@ /* - * Copyright 2023-2023 the original author or authors. + * Copyright 2023-2024 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -57,7 +57,7 @@ import redis.clients.jedis.search.schemafields.VectorField.VectorAlgorithm; * offers functionalities like adding, deleting, and performing similarity searches on * documents. * - * The store utilizes RedisJSON and RediSearch to handle JSON documents and to index and + * The store utilizes RedisJSON and RedisSearch to handle JSON documents and to index and * search vector data. It supports various vector algorithms (e.g., FLAT, HSNW) for * efficient similarity searches. Additionally, it allows for custom metadata fields in * the documents to be stored alongside the vector and content data. @@ -68,6 +68,7 @@ import redis.clients.jedis.search.schemafields.VectorField.VectorAlgorithm; * them. * * @author Julien Ruaux + * @author Christian Tzolov * @see VectorStore * @see RedisVectorStoreConfig * @see EmbeddingClient @@ -115,6 +116,20 @@ public class RedisVectorStore implements VectorStore, InitializingBean { private final List metadataFields; + private RedisVectorStoreConfig() { + this(builder()); + } + + private RedisVectorStoreConfig(Builder builder) { + this.uri = builder.uri; + this.indexName = builder.indexName; + this.prefix = builder.prefix; + this.contentFieldName = builder.contentFieldName; + this.embeddingFieldName = builder.embeddingFieldName; + this.vectorAlgorithm = builder.vectorAlgorithm; + this.metadataFields = builder.metadataFields; + } + /** * Start building a new configuration. * @return The entry point for creating a new configuration. @@ -132,16 +147,6 @@ public class RedisVectorStore implements VectorStore, InitializingBean { return builder().build(); } - private RedisVectorStoreConfig(Builder builder) { - this.uri = builder.uri; - this.indexName = builder.indexName; - this.prefix = builder.prefix; - this.contentFieldName = builder.contentFieldName; - this.embeddingFieldName = builder.embeddingFieldName; - this.vectorAlgorithm = builder.vectorAlgorithm; - this.metadataFields = builder.metadataFields; - } - public static class Builder { private String uri = DEFAULT_URI; @@ -290,22 +295,23 @@ public class RedisVectorStore implements VectorStore, InitializingBean { this.jedis = new JedisPooled(config.uri); this.embeddingClient = embeddingClient; this.config = config; + this.filterExpressionConverter = new RedisFilterExpressionConverter(this.config.metadataFields); } public JedisPooled getJedis() { - return jedis; + return this.jedis; } @Override public void add(List documents) { - Pipeline pipeline = jedis.pipelined(); + Pipeline pipeline = this.jedis.pipelined(); for (Document document : documents) { var embedding = this.embeddingClient.embed(document); document.setEmbedding(embedding); var fields = new HashMap(); - fields.put(config.embeddingFieldName, embedding); - fields.put(config.contentFieldName, document.getContent()); + fields.put(this.config.embeddingFieldName, embedding); + fields.put(this.config.contentFieldName, document.getContent()); fields.putAll(document.getMetadata()); pipeline.jsonSetWithEscape(key(document.getId()), JSON_SET_PATH, fields); } @@ -321,12 +327,12 @@ public class RedisVectorStore implements VectorStore, InitializingBean { } private String key(String id) { - return config.prefix + id; + return this.config.prefix + id; } @Override public Optional delete(List idList) { - Pipeline pipeline = jedis.pipelined(); + Pipeline pipeline = this.jedis.pipelined(); for (String id : idList) { pipeline.jsonDel(key(id)); } @@ -350,13 +356,13 @@ public class RedisVectorStore implements VectorStore, InitializingBean { String filter = nativeExpressionFilter(request); - String queryString = String.format(QUERY_FORMAT, filter, request.getTopK(), config.embeddingFieldName, + String queryString = String.format(QUERY_FORMAT, filter, request.getTopK(), this.config.embeddingFieldName, EMBEDDING_PARAM_NAME, DISTANCE_FIELD_NAME); List returnFields = new ArrayList<>(); - config.metadataFields.stream().map(MetadataField::name).forEach(returnFields::add); - returnFields.add(config.embeddingFieldName); - returnFields.add(config.contentFieldName); + this.config.metadataFields.stream().map(MetadataField::name).forEach(returnFields::add); + returnFields.add(this.config.embeddingFieldName); + returnFields.add(this.config.contentFieldName); returnFields.add(DISTANCE_FIELD_NAME); var embedding = toFloatArray(this.embeddingClient.embed(request.getQuery())); Query query = new Query(queryString).addParam(EMBEDDING_PARAM_NAME, RediSearchUtil.toByteArray(embedding)) @@ -364,7 +370,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean { .setSortBy(DISTANCE_FIELD_NAME, true) .dialect(2); - SearchResult result = jedis.ftSearch(config.indexName, query); + SearchResult result = this.jedis.ftSearch(this.config.indexName, query); return result.getDocuments() .stream() .filter(d -> similarityScore(d) >= request.getSimilarityThreshold()) @@ -373,9 +379,10 @@ public class RedisVectorStore implements VectorStore, InitializingBean { } private Document toDocument(redis.clients.jedis.search.Document doc) { - var id = doc.getId().substring(config.prefix.length()); - var content = doc.hasProperty(config.contentFieldName) ? doc.getString(config.contentFieldName) : null; - Map metadata = config.metadataFields.stream() + var id = doc.getId().substring(this.config.prefix.length()); + var content = doc.hasProperty(this.config.contentFieldName) ? doc.getString(this.config.contentFieldName) + : null; + Map metadata = this.config.metadataFields.stream() .map(MetadataField::name) .filter(doc::hasProperty) .collect(Collectors.toMap(Function.identity(), doc::getString)); @@ -391,44 +398,41 @@ public class RedisVectorStore implements VectorStore, InitializingBean { if (request.getFilterExpression() == null) { return "*"; } - return "(" + filterExpressionConverter.convertExpression(request.getFilterExpression()) + ")"; + return "(" + this.filterExpressionConverter.convertExpression(request.getFilterExpression()) + ")"; } @Override public void afterPropertiesSet() { // If index already exists don't do anything - if (jedis.ftList().contains(config.indexName)) { + if (this.jedis.ftList().contains(this.config.indexName)) { return; } - String response = jedis.ftCreate(config.indexName, - FTCreateParams.createParams().on(IndexDataType.JSON).addPrefix(config.prefix), schemaFields()); + String response = this.jedis.ftCreate(this.config.indexName, + FTCreateParams.createParams().on(IndexDataType.JSON).addPrefix(this.config.prefix), schemaFields()); if (!RESPONSE_OK.test(response)) { String message = MessageFormat.format("Could not create index: {0}", response); throw new RuntimeException(message); } - - filterExpressionConverter = new RedisFilterExpressionConverter(config.metadataFields); - } private Iterable schemaFields() { Map vectorAttrs = new HashMap<>(); - vectorAttrs.put("DIM", embeddingClient.dimensions()); + vectorAttrs.put("DIM", this.embeddingClient.dimensions()); vectorAttrs.put("DISTANCE_METRIC", DEFAULT_DISTANCE_METRIC); vectorAttrs.put("TYPE", VECTOR_TYPE_FLOAT32); List fields = new ArrayList<>(); - fields.add(TextField.of(jsonPath(config.contentFieldName)).as(config.contentFieldName).weight(1.0)); + fields.add(TextField.of(jsonPath(this.config.contentFieldName)).as(this.config.contentFieldName).weight(1.0)); fields.add(VectorField.builder() - .fieldName(jsonPath(config.embeddingFieldName)) + .fieldName(jsonPath(this.config.embeddingFieldName)) .algorithm(vectorAlgorithm()) .attributes(vectorAttrs) - .as(config.embeddingFieldName) + .as(this.config.embeddingFieldName) .build()); - if (!CollectionUtils.isEmpty(config.metadataFields)) { - for (MetadataField field : config.metadataFields) { + if (!CollectionUtils.isEmpty(this.config.metadataFields)) { + for (MetadataField field : this.config.metadataFields) { fields.add(schemaField(field)); } }