diff --git a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java index d6d87c2d8..b7711cf23 100644 --- a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java +++ b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java @@ -122,6 +122,8 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { private static final String COLUMN_CONTENT = "content"; + private static final String COLUMN_DISTANCE = "distance"; + private ObjectMapper objectMapper; public DocumentRowMapper(ObjectMapper objectMapper) { @@ -132,10 +134,14 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { public Document mapRow(ResultSet rs, int rowNum) throws SQLException { String id = rs.getString(COLUMN_ID); String content = rs.getString(COLUMN_CONTENT); - PGobject metadata = rs.getObject(COLUMN_METADATA, PGobject.class); + PGobject pgMetadata = rs.getObject(COLUMN_METADATA, PGobject.class); PGobject embedding = rs.getObject(COLUMN_EMBEDDING, PGobject.class); + Float distance = rs.getFloat(COLUMN_DISTANCE); - Document document = new Document(id, content, toMap(metadata)); + Map metadata = toMap(pgMetadata); + metadata.put(COLUMN_DISTANCE, distance); + + Document document = new Document(id, content, metadata); document.setEmbedding(toDoubleList(embedding)); return document; @@ -231,8 +237,8 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { @Override public List similaritySearch(String query, int k) { PGvector queryEmbedding = getQueryEmbedding(query); - return this.jdbcTemplate.query("SELECT * FROM " + VECTOR_TABLE_NAME + " ORDER BY embedding " - + this.getDistanceType().operator + " ? LIMIT ?", new DocumentRowMapper(this.objectMapper), + return this.jdbcTemplate.query("SELECT *, embedding " + this.comparisonOperator() + " ? AS distance FROM " + + VECTOR_TABLE_NAME + " ORDER BY distance LIMIT ?", new DocumentRowMapper(this.objectMapper), queryEmbedding, k); } @@ -240,14 +246,14 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { public List similaritySearch(String query, int k, double threshold) { PGvector queryEmbedding = getQueryEmbedding(query); return this.jdbcTemplate.query( - "SELECT * FROM " + VECTOR_TABLE_NAME + " WHERE embedding " + this.getDistanceType().operator - + " ? < ? ORDER BY embedding " + this.getDistanceType().operator + " ? LIMIT ? ", - new DocumentRowMapper(this.objectMapper), queryEmbedding, threshold, queryEmbedding, k); + "SELECT *, embedding " + this.comparisonOperator() + " ? AS distance FROM " + VECTOR_TABLE_NAME + + " WHERE embedding " + this.comparisonOperator() + " ? < ? ORDER BY distance LIMIT ? ", + new DocumentRowMapper(this.objectMapper), queryEmbedding, queryEmbedding, threshold, k); } public List embeddingDistance(String query) { return this.jdbcTemplate.query( - "SELECT embedding " + this.getDistanceType().operator + " ? AS distance FROM " + VECTOR_TABLE_NAME, + "SELECT embedding " + this.comparisonOperator() + " ? AS distance FROM " + VECTOR_TABLE_NAME, new RowMapper() { @Override @Nullable @@ -263,6 +269,10 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { return new PGvector(toFloatArray(embedding)); } + private String comparisonOperator() { + return this.getDistanceType().operator; + } + // --------------------------------------------------------------------------------- // SmartLifecycle // --------------------------------------------------------------------------------- diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java index edce3bc3f..e35845e5a 100644 --- a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java @@ -17,6 +17,7 @@ package org.springframework.ai.vectorstore; import java.util.Collections; +import java.util.Iterator; import java.util.List; import java.util.UUID; import java.util.stream.Collectors; @@ -43,6 +44,7 @@ import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Primary; import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.util.CollectionUtils; import static org.assertj.core.api.Assertions.assertThat; @@ -69,7 +71,7 @@ public class PgVectorStoreIT { private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() .withUserConfiguration(TestApplication.class) .withPropertyValues("spring.datasource.type=com.zaxxer.hikari.HikariDataSource", - "spring.ai.openai.apiKey=" + System.getenv("SPRING_AI_OPENAI_API_KEY"), + "spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"), // JdbcTemplate configuration String.format("app.datasource.url=jdbc:postgresql://localhost:%d/%s", @@ -92,7 +94,7 @@ public class PgVectorStoreIT { assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId()); assertThat(resultDoc.getText()).isEqualTo( "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); + assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance"); // Remove all documents from the store vectorStore.delete(documents.stream().map(doc -> doc.getId()).collect(Collectors.toList())); @@ -121,7 +123,7 @@ public class PgVectorStoreIT { Document resultDoc = results.get(0); assertThat(resultDoc.getId()).isEqualTo(document.getId()); assertThat(resultDoc.getText()).isEqualTo("Spring AI rocks!!"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta1", "meta1")); + assertThat(resultDoc.getMetadata()).containsKeys("meta1", "distance"); Document sameIdDocument = new Document(document.getId(), "The World is Big and Salvation Lurks Around the Corner", @@ -135,8 +137,7 @@ public class PgVectorStoreIT { resultDoc = results.get(0); assertThat(resultDoc.getId()).isEqualTo(document.getId()); assertThat(resultDoc.getText()).isEqualTo("The World is Big and Salvation Lurks Around the Corner"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); - + assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance"); }); } @@ -149,7 +150,13 @@ public class PgVectorStoreIT { vectorStore.add(documents); - assertThat(vectorStore.similaritySearch("Great", 5, 1.0)).hasSize(3); + List fullResult = vectorStore.similaritySearch("Great", 5, 1.0); + + assertThat(fullResult).hasSize(3); + + assertThat(isSortedByDistance(fullResult)).isTrue(); + + fullResult.stream().forEach(doc -> System.out.println(doc.getMetadata().get("distance"))); List embeddingDistance = ((PgVectorStore) vectorStore).embeddingDistance("Great"); @@ -160,11 +167,33 @@ public class PgVectorStoreIT { assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId()); assertThat(resultDoc.getText()).isEqualTo( "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); + assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance"); }); } + private static boolean isSortedByDistance(List docs) { + + List distances = docs.stream() + .map(doc -> (Float) doc.getMetadata().get("distance")) + .collect(Collectors.toList()); + + if (CollectionUtils.isEmpty(distances) || distances.size() == 1) { + return true; + } + + Iterator iter = distances.iterator(); + Float current, previous = iter.next(); + while (iter.hasNext()) { + current = iter.next(); + if (previous > current) { + return false; + } + previous = current; + } + return true; + } + @SpringBootConfiguration @EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class }) public static class TestApplication {