From 9d207e37bed5809f9f4ad1c31ff658f613222d6f Mon Sep 17 00:00:00 2001 From: Ilayaperumal Gopinathan Date: Mon, 25 Nov 2024 15:00:52 +0000 Subject: [PATCH] Refactor SimpleVectorStore - Remove SimpleVectorStore's dependency on deprecated embeddings from Document object - Create a custom Content object that represents the SimpleVectorStore's contents and embedding - Add tests --- .../ai/vectorstore/SimpleVectorStore.java | 41 ++- .../vectorstore/SimpleVectorStoreContent.java | 165 +++++++++++ .../SimpleVectorStoreSimilarityTests.java | 49 ++++ .../vectorstore/SimpleVectorStoreTests.java | 263 ++++++++++++++++++ 4 files changed, 504 insertions(+), 14 deletions(-) create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStoreContent.java create mode 100644 spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreSimilarityTests.java create mode 100644 spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreTests.java diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java index f2e558eb6..c94143920 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java @@ -67,6 +67,7 @@ import org.springframework.core.io.Resource; * @author Mark Pollack * @author Christian Tzolov * @author Sebastien Deleuze + * @author Ilayaperumal Gopinathan */ public class SimpleVectorStore extends AbstractObservationVectorStore { @@ -74,7 +75,7 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { private final ObjectMapper objectMapper; - protected Map store = new ConcurrentHashMap<>(); + protected Map store = new ConcurrentHashMap<>(); protected EmbeddingModel embeddingModel; @@ -94,11 +95,17 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { @Override public void doAdd(List documents) { + Objects.requireNonNull(documents, "Documents list cannot be null"); + if (documents.isEmpty()) { + throw new IllegalArgumentException("Documents list cannot be empty"); + } + for (Document document : documents) { logger.info("Calling EmbeddingModel for document id = {}", document.getId()); float[] embedding = this.embeddingModel.embed(document); - document.setEmbedding(embedding); - this.store.put(document.getId(), document); + SimpleVectorStoreContent storeContent = new SimpleVectorStoreContent(document.getId(), + document.getContent(), document.getMetadata(), embedding); + this.store.put(document.getId(), storeContent); } } @@ -120,12 +127,12 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { float[] userQueryEmbedding = getUserQueryEmbedding(request.getQuery()); return this.store.values() .stream() - .map(entry -> new Similarity(entry.getId(), + .map(entry -> new Similarity(entry, EmbeddingMath.cosineSimilarity(userQueryEmbedding, entry.getEmbedding()))) .filter(s -> s.score >= request.getSimilarityThreshold()) .sorted(Comparator.comparingDouble(s -> s.score).reversed()) .limit(request.getTopK()) - .map(s -> this.store.get(s.key)) + .map(s -> s.getDocument()) .toList(); } @@ -176,12 +183,11 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { * @param file the file to load the vector store content */ public void load(File file) { - TypeReference> typeRef = new TypeReference<>() { + TypeReference> typeRef = new TypeReference<>() { }; try { - Map deserializedMap = this.objectMapper.readValue(file, typeRef); - this.store = deserializedMap; + this.store = this.objectMapper.readValue(file, typeRef); } catch (IOException ex) { throw new RuntimeException(ex); @@ -193,12 +199,11 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { * @param resource the resource to load the vector store content */ public void load(Resource resource) { - TypeReference> typeRef = new TypeReference<>() { + TypeReference> typeRef = new TypeReference<>() { }; try { - Map deserializedMap = this.objectMapper.readValue(resource.getInputStream(), typeRef); - this.store = deserializedMap; + this.store = this.objectMapper.readValue(resource.getInputStream(), typeRef); } catch (IOException ex) { throw new RuntimeException(ex); @@ -232,15 +237,23 @@ public class SimpleVectorStore extends AbstractObservationVectorStore { public static class Similarity { - private String key; + private SimpleVectorStoreContent content; private double score; - public Similarity(String key, double score) { - this.key = key; + public Similarity(SimpleVectorStoreContent content, double score) { + this.content = content; this.score = score; } + Document getDocument() { + return Document.builder() + .withId(this.content.getId()) + .withContent(this.content.getContent()) + .withMetadata(this.content.getMetadata()) + .build(); + } + } public final class EmbeddingMath { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStoreContent.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStoreContent.java new file mode 100644 index 000000000..c54a96d00 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStoreContent.java @@ -0,0 +1,165 @@ +/* + * 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.vectorstore; + +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; + +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonProperty; + +import org.springframework.ai.document.id.IdGenerator; +import org.springframework.ai.document.id.RandomIdGenerator; +import org.springframework.ai.model.Content; +import org.springframework.util.Assert; + +/** + * An immutable {@link Content} implementation representing content, metadata, and its + * embeddings. This class is thread-safe and all its fields are final and deeply + * immutable. The embedding vector is required for all instances of this class. + */ +public final class SimpleVectorStoreContent implements Content { + + private final String id; + + private final String content; + + private final Map metadata; + + private final float[] embedding; + + /** + * Creates a new instance with the given content, empty metadata, and embedding + * vector. + * @param content the content text, must not be null + * @param embedding the embedding vector, must not be null + */ + @JsonCreator(mode = JsonCreator.Mode.PROPERTIES) + public SimpleVectorStoreContent(@JsonProperty("content") String content, + @JsonProperty("embedding") float[] embedding) { + this(content, new HashMap<>(), embedding); + } + + /** + * Creates a new instance with the given content, metadata, and embedding vector. + * @param content the content text, must not be null + * @param metadata the metadata map, must not be null + * @param embedding the embedding vector, must not be null + */ + public SimpleVectorStoreContent(String content, Map metadata, float[] embedding) { + this(content, metadata, new RandomIdGenerator(), embedding); + } + + /** + * Creates a new instance with the given content, metadata, custom ID generator, and + * embedding vector. + * @param content the content text, must not be null + * @param metadata the metadata map, must not be null + * @param idGenerator the ID generator to use, must not be null + * @param embedding the embedding vector, must not be null + */ + public SimpleVectorStoreContent(String content, Map metadata, IdGenerator idGenerator, + float[] embedding) { + this(idGenerator.generateId(content, metadata), content, metadata, embedding); + } + + /** + * Creates a new instance with all fields specified. + * @param id the unique identifier, must not be empty + * @param content the content text, must not be null + * @param metadata the metadata map, must not be null + * @param embedding the embedding vector, must not be null + * @throws IllegalArgumentException if any parameter is null or if id is empty + */ + public SimpleVectorStoreContent(String id, String content, Map metadata, float[] embedding) { + Assert.hasText(id, "id must not be null or empty"); + Assert.notNull(content, "content must not be null"); + Assert.notNull(metadata, "metadata must not be null"); + Assert.notNull(embedding, "embedding must not be null"); + Assert.isTrue(embedding.length > 0, "embedding vector must not be empty"); + + this.id = id; + this.content = content; + this.metadata = Collections.unmodifiableMap(new HashMap<>(metadata)); + this.embedding = Arrays.copyOf(embedding, embedding.length); + } + + /** + * Creates a new instance with an updated embedding vector. + * @param embedding the new embedding vector, must not be null + * @return a new instance with the updated embedding + * @throws IllegalArgumentException if embedding is null or empty + */ + public SimpleVectorStoreContent withEmbedding(float[] embedding) { + Assert.notNull(embedding, "embedding must not be null"); + Assert.isTrue(embedding.length > 0, "embedding vector must not be empty"); + return new SimpleVectorStoreContent(this.id, this.content, this.metadata, embedding); + } + + public String getId() { + return this.id; + } + + @Override + public String getContent() { + return this.content; + } + + @Override + public Map getMetadata() { + return this.metadata; + } + + /** + * Returns a defensive copy of the embedding vector. + * @return a new array containing the embedding vector + */ + public float[] getEmbedding() { + return Arrays.copyOf(this.embedding, this.embedding.length); + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (o == null || getClass() != o.getClass()) + return false; + + SimpleVectorStoreContent that = (SimpleVectorStoreContent) o; + return Objects.equals(this.id, that.id) && Objects.equals(this.content, that.content) + && Objects.equals(this.metadata, that.metadata) && Arrays.equals(this.embedding, that.embedding); + } + + @Override + public int hashCode() { + int result = Objects.hashCode(this.id); + result = 31 * result + Objects.hashCode(this.content); + result = 31 * result + Objects.hashCode(this.metadata); + result = 31 * result + Arrays.hashCode(this.embedding); + return result; + } + + @Override + public String toString() { + return "SimpleVectorStoreContent{" + "id='" + this.id + '\'' + ", content='" + this.content + '\'' + + ", metadata=" + this.metadata + ", embedding=" + Arrays.toString(embedding) + '}'; + } + +} diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreSimilarityTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreSimilarityTests.java new file mode 100644 index 000000000..3447e5f7c --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreSimilarityTests.java @@ -0,0 +1,49 @@ +/* + * 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.vectorstore; + +import java.util.HashMap; +import java.util.Map; + +import org.junit.Test; + +import org.springframework.ai.document.Document; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * @author Ilayaperumal Gopinathan + */ +public class SimpleVectorStoreSimilarityTests { + + @Test + public void testSimilarity() { + Map metadata = new HashMap<>(); + metadata.put("foo", "bar"); + float[] testEmbedding = new float[] { 1.0f, 2.0f, 3.0f }; + + SimpleVectorStoreContent storeContent = new SimpleVectorStoreContent("1", "hello, how are you?", metadata, + testEmbedding); + SimpleVectorStore.Similarity similarity = new SimpleVectorStore.Similarity(storeContent, 0.6d); + Document document = similarity.getDocument(); + assertThat(document).isNotNull(); + assertThat(document.getId()).isEqualTo("1"); + assertThat(document.getContent()).isEqualTo("hello, how are you?"); + assertThat(document.getMetadata().get("foo")).isEqualTo("bar"); + } + +} diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreTests.java new file mode 100644 index 000000000..9e14abbd8 --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/SimpleVectorStoreTests.java @@ -0,0 +1,263 @@ +/* + * 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. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.vectorstore; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.CleanupMode; +import org.junit.jupiter.api.io.TempDir; +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.core.io.Resource; + +import java.io.File; +import java.io.IOException; +import java.nio.file.Path; +import java.util.Arrays; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.stream.Collectors; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +class SimpleVectorStoreTests { + + @TempDir(cleanup = CleanupMode.ON_SUCCESS) + Path tempDir; + + private SimpleVectorStore vectorStore; + + private EmbeddingModel mockEmbeddingModel; + + @BeforeEach + void setUp() { + mockEmbeddingModel = mock(EmbeddingModel.class); + when(mockEmbeddingModel.dimensions()).thenReturn(3); + when(mockEmbeddingModel.embed(any(String.class))).thenReturn(new float[] { 0.1f, 0.2f, 0.3f }); + when(mockEmbeddingModel.embed(any(Document.class))).thenReturn(new float[] { 0.1f, 0.2f, 0.3f }); + + vectorStore = new SimpleVectorStore(mockEmbeddingModel); + } + + @Test + void shouldAddAndRetrieveDocument() { + Document doc = Document.builder() + .withId("1") + .withContent("test content") + .withMetadata(Map.of("key", "value")) + .build(); + + vectorStore.add(List.of(doc)); + + List results = vectorStore.similaritySearch("test content"); + assertThat(results).hasSize(1).first().satisfies(result -> { + assertThat(result.getId()).isEqualTo("1"); + assertThat(result.getContent()).isEqualTo("test content"); + assertThat(result.getMetadata()).containsEntry("key", "value"); + }); + } + + @Test + void shouldAddMultipleDocuments() { + List docs = Arrays.asList(Document.builder().withId("1").withContent("first").build(), + Document.builder().withId("2").withContent("second").build()); + + vectorStore.add(docs); + + List results = vectorStore.similaritySearch("first"); + assertThat(results).hasSize(2).extracting(Document::getId).containsExactlyInAnyOrder("1", "2"); + } + + @Test + void shouldHandleEmptyDocumentList() { + assertThatThrownBy(() -> vectorStore.add(Collections.emptyList())).isInstanceOf(IllegalArgumentException.class) + .hasMessage("Documents list cannot be empty"); + } + + @Test + void shouldHandleNullDocumentList() { + assertThatThrownBy(() -> vectorStore.add(null)).isInstanceOf(NullPointerException.class) + .hasMessage("Documents list cannot be null"); + } + + @Test + void shouldDeleteDocuments() { + Document doc = Document.builder().withId("1").withContent("test content").build(); + + vectorStore.add(List.of(doc)); + assertThat(vectorStore.similaritySearch("test")).hasSize(1); + + vectorStore.delete(List.of("1")); + assertThat(vectorStore.similaritySearch("test")).isEmpty(); + } + + @Test + void shouldHandleDeleteOfNonexistentDocument() { + vectorStore.delete(List.of("nonexistent-id")); + // Should not throw exception and return true + assertThat(vectorStore.delete(List.of("nonexistent-id")).get()).isTrue(); + } + + @Test + void shouldPerformSimilaritySearchWithThreshold() { + // Configure mock to return different embeddings for different queries + when(mockEmbeddingModel.embed("query")).thenReturn(new float[] { 0.9f, 0.9f, 0.9f }); + + Document doc = Document.builder().withId("1").withContent("test content").build(); + + vectorStore.add(List.of(doc)); + + SearchRequest request = SearchRequest.query("query").withSimilarityThreshold(0.99f).withTopK(5); + + List results = vectorStore.similaritySearch(request); + assertThat(results).isEmpty(); + } + + @Test + void shouldSaveAndLoadVectorStore() throws IOException { + Document doc = Document.builder() + .withId("1") + .withContent("test content") + .withMetadata(new HashMap<>(Map.of("key", "value"))) + .build(); + + vectorStore.add(List.of(doc)); + + File saveFile = tempDir.resolve("vector-store.json").toFile(); + vectorStore.save(saveFile); + + SimpleVectorStore loadedStore = new SimpleVectorStore(mockEmbeddingModel); + loadedStore.load(saveFile); + + List results = loadedStore.similaritySearch("test content"); + assertThat(results).hasSize(1).first().satisfies(result -> { + assertThat(result.getId()).isEqualTo("1"); + assertThat(result.getContent()).isEqualTo("test content"); + assertThat(result.getMetadata()).containsEntry("key", "value"); + }); + } + + @Test + void shouldHandleLoadFromInvalidResource() throws IOException { + Resource mockResource = mock(Resource.class); + when(mockResource.getInputStream()).thenThrow(new IOException("Resource not found")); + + assertThatThrownBy(() -> vectorStore.load(mockResource)).isInstanceOf(RuntimeException.class) + .hasCauseInstanceOf(IOException.class) + .hasMessageContaining("Resource not found"); + } + + @Test + void shouldHandleSaveToInvalidLocation() { + File invalidFile = new File("/invalid/path/file.json"); + + assertThatThrownBy(() -> vectorStore.save(invalidFile)).isInstanceOf(RuntimeException.class) + .hasCauseInstanceOf(IOException.class); + } + + @Test + void shouldHandleConcurrentOperations() throws InterruptedException { + int numThreads = 10; + Thread[] threads = new Thread[numThreads]; + + for (int i = 0; i < numThreads; i++) { + final String id = String.valueOf(i); + threads[i] = new Thread(() -> { + Document doc = Document.builder().withId(id).withContent("content " + id).build(); + vectorStore.add(List.of(doc)); + }); + threads[i].start(); + } + + for (Thread thread : threads) { + thread.join(); + } + + SearchRequest request = SearchRequest.query("test").withTopK(numThreads); + + List results = vectorStore.similaritySearch(request); + + assertThat(results).hasSize(numThreads); + + // Verify all documents were properly added + Set resultIds = results.stream().map(Document::getId).collect(Collectors.toSet()); + + Set expectedIds = new java.util.HashSet<>(); + for (int i = 0; i < numThreads; i++) { + expectedIds.add(String.valueOf(i)); + } + + assertThat(resultIds).containsExactlyInAnyOrderElementsOf(expectedIds); + + // Verify content integrity + results.forEach(doc -> assertThat(doc.getContent()).isEqualTo("content " + doc.getId())); + } + + @Test + void shouldRejectInvalidSimilarityThreshold() { + assertThatThrownBy(() -> SearchRequest.query("test").withSimilarityThreshold(2.0f)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Similarity threshold must be in [0,1] range."); + } + + @Test + void shouldRejectNegativeTopK() { + assertThatThrownBy(() -> SearchRequest.query("test").withTopK(-1)).isInstanceOf(IllegalArgumentException.class) + .hasMessage("TopK should be positive."); + } + + @Test + void shouldHandleCosineSimilarityEdgeCases() { + float[] zeroVector = new float[] { 0f, 0f, 0f }; + float[] normalVector = new float[] { 1f, 1f, 1f }; + + assertThatThrownBy(() -> SimpleVectorStore.EmbeddingMath.cosineSimilarity(zeroVector, normalVector)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Vectors cannot have zero norm"); + } + + @Test + void shouldHandleVectorLengthMismatch() { + float[] vector1 = new float[] { 1f, 2f }; + float[] vector2 = new float[] { 1f, 2f, 3f }; + + assertThatThrownBy(() -> SimpleVectorStore.EmbeddingMath.cosineSimilarity(vector1, vector2)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Vectors lengths must be equal"); + } + + @Test + void shouldHandleNullVectors() { + float[] vector = new float[] { 1f, 2f, 3f }; + + assertThatThrownBy(() -> SimpleVectorStore.EmbeddingMath.cosineSimilarity(null, vector)) + .isInstanceOf(RuntimeException.class) + .hasMessage("Vectors must not be null"); + + assertThatThrownBy(() -> SimpleVectorStore.EmbeddingMath.cosineSimilarity(vector, null)) + .isInstanceOf(RuntimeException.class) + .hasMessage("Vectors must not be null"); + } + +}