From 98ca3e2a8fe4228b5c556e14dbdeb0b5c5ece0ae Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Wed, 8 Nov 2023 18:16:14 +0100 Subject: [PATCH] Add Chroma VectorStore support - Implement ChromaApi client, based on Chroma REST API. - Implement ChromaVectorStore, including support for filter expression conversion. - Common VectorStoreUtil class to share to/from Float/Double list/array convertion as well as Json/Map convertions. - Add ITs including for Basic Auth and Token autheticatios. - Add ChromaApi security support for BasicAuth and Token. - Fix an issue with Text filter expression parser, related to double-quoted identifiers. - Add Chroma README.md. - Add Chroma boot autoconfiguration Resolves# #86 --- pom.xml | 1 + .../filter/FilterExpressionTextParser.java | 7 +- .../AbstractFilterExpressionConverter.java | 27 +- .../ChromaFilterExpressionConverter.java | 29 + .../converter/FilterExpressionConverter.java | 32 + .../MilvusFilterExpressionConverter.java | 9 +- .../PgVectorFilterExpressionConverter.java | 4 +- .../PineconeFilterExpressionConverter.java | 11 +- .../PrintFilterExpressionConverter.java | 4 +- .../filter/FilterExpressionBuilderTests.java | 2 - .../FilterExpressionTextParserTests.java | 12 + .../MilvusFilterExpressionConverterTests.java | 33 +- ...gVectorFilterExpressionConverterTests.java | 30 +- ...ineconeFilterExpressionConverterTests.java | 33 +- spring-ai-spring-boot-autoconfigure/pom.xml | 8 + .../chroma/ChromaApiProperties.java | 79 +++ .../ChromaVectorStoreAutoConfiguration.java | 72 ++ .../chroma/ChromaVectorStoreProperties.java | 40 ++ ...ot.autoconfigure.AutoConfiguration.imports | 1 + .../ChromaVectorStoreAutoConfigurationIT.java | 100 +++ .../PgVectorStoreAutoConfigurationIT.java | 3 +- vector-stores/spring-ai-chroma/README.md | 154 +++++ vector-stores/spring-ai-chroma/pom.xml | 85 +++ .../experimental/ai/chroma/ChromaApi.java | 380 +++++++++++ .../ai/vectorsore/ChromaVectorStore.java | 161 +++++ .../experimental/ai/vectorsore/JsonUtils.java | 84 +++ .../experimental/ai/chroma/ChromaApiIT.java | 186 ++++++ .../vectorstore/BasicAuthChromaWhereIT.java | 141 ++++ .../ai/vectorstore/ChromaVectorStoreIT.java | 240 +++++++ .../TokenSecuredChromaWhereIT.java | 174 +++++ .../src/test/resources/api.yaml | 630 ++++++++++++++++++ .../src/test/resources/server.htpasswd | 2 + .../ai/vectorstore/MilvusVectorStore.java | 5 +- .../ai/vectorstore/PgVectorStore.java | 5 +- .../ai/vectorstore/PgVectorStoreIT.java | 12 +- .../ai/vectorstore/PineconeVectorStore.java | 6 +- 36 files changed, 2735 insertions(+), 67 deletions(-) create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/ChromaFilterExpressionConverter.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/FilterExpressionConverter.java create mode 100644 spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaApiProperties.java create mode 100644 spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java create mode 100644 spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreProperties.java create mode 100644 spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java create mode 100644 vector-stores/spring-ai-chroma/README.md create mode 100644 vector-stores/spring-ai-chroma/pom.xml create mode 100644 vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/chroma/ChromaApi.java create mode 100644 vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/ChromaVectorStore.java create mode 100644 vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/JsonUtils.java create mode 100644 vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/chroma/ChromaApiIT.java create mode 100644 vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/BasicAuthChromaWhereIT.java create mode 100644 vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/ChromaVectorStoreIT.java create mode 100644 vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/TokenSecuredChromaWhereIT.java create mode 100644 vector-stores/spring-ai-chroma/src/test/resources/api.yaml create mode 100644 vector-stores/spring-ai-chroma/src/test/resources/server.htpasswd diff --git a/pom.xml b/pom.xml index d327df2f7..9a6cc37df 100644 --- a/pom.xml +++ b/pom.xml @@ -32,6 +32,7 @@ document-readers/tika-reader embedding-clients/transformers-embedding vector-stores/spring-ai-pinecone + vector-stores/spring-ai-chroma diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParser.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParser.java index 5494ccd9a..dddbf50ba 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParser.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParser.java @@ -191,11 +191,14 @@ public class FilterExpressionTextParser { @Override public Filter.Operand visitTextConstant(FiltersParser.TextConstantContext ctx) { - var twiceQuotedText = ctx.getText(); - String onceQuotedText = twiceQuotedText.substring(1, twiceQuotedText.length() - 1); + String onceQuotedText = removeOuterQuotes(ctx.getText()); return new Filter.Value(onceQuotedText); } + private String removeOuterQuotes(String in) { + return in.substring(1, in.length() - 1); + } + @Override public Filter.Operand visitIntegerConstant(FiltersParser.IntegerConstantContext ctx) { return new Filter.Value(Integer.valueOf(ctx.getText())); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/AbstractFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/AbstractFilterExpressionConverter.java index fbf5246d3..160d02068 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/AbstractFilterExpressionConverter.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/AbstractFilterExpressionConverter.java @@ -19,6 +19,7 @@ package org.springframework.ai.vectorstore.filter.converter; import java.util.List; import org.springframework.ai.vectorstore.filter.Filter; +import org.springframework.ai.vectorstore.filter.Filter.Expression; import org.springframework.ai.vectorstore.filter.Filter.ExpressionType; import org.springframework.ai.vectorstore.filter.Filter.Group; import org.springframework.ai.vectorstore.filter.Filter.Operand; @@ -27,16 +28,20 @@ import org.springframework.util.Assert; /** * @author Christian Tzolov */ -public abstract class AbstractFilterExpressionConverter { +public abstract class AbstractFilterExpressionConverter implements FilterExpressionConverter { - public String convert(Operand operand) { - Assert.notNull(operand, "Operand can't be null"); + @Override + public String convertExpression(Expression expression) { + return this.convertOperand(expression); + } + + protected String convertOperand(Operand operand) { var context = new StringBuilder(); - this.convert(operand, context); + this.convertOperand(operand, context); return context.toString(); } - protected void convert(Operand operand, StringBuilder context) { + protected void convertOperand(Operand operand, StringBuilder context) { if (operand instanceof Filter.Group group) { this.doGroup(group, context); @@ -88,7 +93,7 @@ public abstract class AbstractFilterExpressionConverter { protected void doGroup(Group group, StringBuilder context) { this.doStartGroup(group, context); - this.convert(group.content(), context); + this.convertOperand(group.content(), context); this.doEndGroup(group, context); } @@ -110,4 +115,14 @@ public abstract class AbstractFilterExpressionConverter { context.append(","); } + // Utilities + protected boolean hasOuterQuotes(String str) { + str = str.trim(); + return (str.startsWith("\"") && str.endsWith("\"")) || (str.startsWith("'") && str.endsWith("'")); + } + + protected String removeOuterQuotes(String in) { + return in.substring(1, in.length() - 1); + } + } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/ChromaFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/ChromaFilterExpressionConverter.java new file mode 100644 index 000000000..30cc5006d --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/ChromaFilterExpressionConverter.java @@ -0,0 +1,29 @@ +/* + * Copyright 2023-2023 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.filter.converter; + +import org.springframework.ai.vectorstore.filter.Filter; + +/** + * Converts {@link Filter.Expression} into Chroma metadata filter expression format. + * (https://docs.trychroma.com/usage-guide#using-where-filters) + * + * @author Christian Tzolov + */ +public class ChromaFilterExpressionConverter extends PineconeFilterExpressionConverter { + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/FilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/FilterExpressionConverter.java new file mode 100644 index 000000000..ece17dc0c --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/FilterExpressionConverter.java @@ -0,0 +1,32 @@ +/* + * Copyright 2023-2023 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.filter.converter; + +import org.springframework.ai.vectorstore.filter.Filter; + +/** + * Converters a generic, portable {@link Filter.Expression} into a + * {@link org.springframework.ai.vectorstore.VectorStore} specific expression language + * format. + * + * @author Christian Tzolov + */ +public interface FilterExpressionConverter { + + public String convertExpression(Filter.Expression expression); + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverter.java index cbf012548..3a32552a3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverter.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverter.java @@ -31,9 +31,9 @@ public class MilvusFilterExpressionConverter extends AbstractFilterExpressionCon @Override protected void doExpression(Expression exp, StringBuilder context) { - this.convert(exp.left(), context); + this.convertOperand(exp.left(), context); context.append(getOperationSymbol(exp)); - this.convert(exp.right(), context); + this.convertOperand(exp.right(), context); } private String getOperationSymbol(Expression exp) { @@ -65,12 +65,13 @@ public class MilvusFilterExpressionConverter extends AbstractFilterExpressionCon @Override protected void doGroup(Group group, StringBuilder context) { - this.convert(new Expression(ExpressionType.AND, group.content(), group.content()), context); // trick + this.convertOperand(new Expression(ExpressionType.AND, group.content(), group.content()), context); // trick } @Override protected void doKey(Key key, StringBuilder context) { - context.append("metadata[\"" + key.key() + "\"]"); + var identifier = (hasOuterQuotes(key.key())) ? removeOuterQuotes(key.key()) : key.key(); + context.append("metadata[\"" + identifier + "\"]"); } } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverter.java index 230a71df0..5eeabe96a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverter.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverter.java @@ -30,9 +30,9 @@ public class PgVectorFilterExpressionConverter extends AbstractFilterExpressionC @Override protected void doExpression(Expression expression, StringBuilder context) { - this.convert(expression.left(), context); + this.convertOperand(expression.left(), context); context.append(getOperationSymbol(expression)); - this.convert(expression.right(), context); + this.convertOperand(expression.right(), context); } private String getOperationSymbol(Expression exp) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverter.java index b49f8b062..eb4614c61 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverter.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverter.java @@ -35,16 +35,16 @@ public class PineconeFilterExpressionConverter extends AbstractFilterExpressionC if (exp.type() == ExpressionType.AND || exp.type() == ExpressionType.OR) { context.append(getOperationSymbol(exp)); context.append("["); - this.convert(exp.left(), context); + this.convertOperand(exp.left(), context); context.append(","); - this.convert(exp.right(), context); + this.convertOperand(exp.right(), context); context.append("]"); } else { - this.convert(exp.left(), context); + this.convertOperand(exp.left(), context); context.append("{"); context.append(getOperationSymbol(exp)); - this.convert(exp.right(), context); + this.convertOperand(exp.right(), context); context.append("}"); } context.append("}"); @@ -57,7 +57,8 @@ public class PineconeFilterExpressionConverter extends AbstractFilterExpressionC @Override protected void doKey(Key key, StringBuilder context) { - context.append("\"" + key.key() + "\": "); + var identifier = (hasOuterQuotes(key.key())) ? removeOuterQuotes(key.key()) : key.key(); + context.append("\"" + identifier + "\": "); } } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PrintFilterExpressionConverter.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PrintFilterExpressionConverter.java index 00b4e07c2..6860cc7ff 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PrintFilterExpressionConverter.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/filter/converter/PrintFilterExpressionConverter.java @@ -28,9 +28,9 @@ import org.springframework.ai.vectorstore.filter.Filter.Key; public class PrintFilterExpressionConverter extends AbstractFilterExpressionConverter { public void doExpression(Expression expression, StringBuilder context) { - this.convert(expression.left(), context); + this.convertOperand(expression.left(), context); context.append(" " + expression.type() + " "); - this.convert(expression.right(), context); + this.convertOperand(expression.right(), context); } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionBuilderTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionBuilderTests.java index 5a2bbf4b4..c05cba98b 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionBuilderTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionBuilderTests.java @@ -18,7 +18,6 @@ package org.springframework.ai.vectorstore.filter; import java.util.List; -import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.springframework.ai.vectorstore.filter.Filter.Expression; @@ -44,7 +43,6 @@ public class FilterExpressionBuilderTests { @Test public void testEQ() { - Expression expression = b.eq("country", "BG").build(); // country == "BG" assertThat(b.eq("country", "BG").build()).isEqualTo(new Expression(EQ, new Key("country"), new Value("BG"))); } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParserTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParserTests.java index 8e82e11ba..057ae86cc 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParserTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/FilterExpressionTextParserTests.java @@ -123,4 +123,16 @@ public class FilterExpressionTextParserTests { assertThat(parser.getCache().get("WHERE " + expText)).isEqualTo(exp); } + @Test + public void testIdentifiers() { + Expression exp = parser.parse("'country.1' == 'BG'"); + assertThat(exp).isEqualTo(new Expression(EQ, new Key("'country.1'"), new Value("BG"))); + + exp = parser.parse("'country_1_2_3' == 'BG'"); + assertThat(exp).isEqualTo(new Expression(EQ, new Key("'country_1_2_3'"), new Value("BG"))); + + exp = parser.parse("\"country 1 2 3\" == 'BG'"); + assertThat(exp).isEqualTo(new Expression(EQ, new Key("\"country 1 2 3\""), new Value("BG"))); + } + } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverterTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverterTests.java index c439316da..ab8d38cfb 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverterTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/MilvusFilterExpressionConverterTests.java @@ -40,12 +40,12 @@ import static org.springframework.ai.vectorstore.filter.Filter.ExpressionType.OR */ public class MilvusFilterExpressionConverterTests { - MilvusFilterExpressionConverter converter = new MilvusFilterExpressionConverter(); + FilterExpressionConverter converter = new MilvusFilterExpressionConverter(); @Test public void testEQ() { // country == "BG" - String vectorExpr = converter.convert(new Expression(EQ, new Key("country"), new Value("BG"))); + String vectorExpr = converter.convertExpression(new Expression(EQ, new Key("country"), new Value("BG"))); assertThat(vectorExpr).isEqualTo("metadata[\"country\"] == \"BG\""); } @@ -53,7 +53,7 @@ public class MilvusFilterExpressionConverterTests { public void tesEqAndGte() { // genre == "drama" AND year >= 2020 String vectorExpr = converter - .convert(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), + .convertExpression(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), new Expression(GTE, new Key("year"), new Value(2020)))); assertThat(vectorExpr).isEqualTo("metadata[\"genre\"] == \"drama\" && metadata[\"year\"] >= 2020"); } @@ -61,17 +61,18 @@ public class MilvusFilterExpressionConverterTests { @Test public void tesIn() { // genre in ["comedy", "documentary", "drama"] - String vectorExpr = converter - .convert(new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); + String vectorExpr = converter.convertExpression( + new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); assertThat(vectorExpr).isEqualTo("metadata[\"genre\"] in [\"comedy\",\"documentary\",\"drama\"]"); } @Test public void testNe() { // year >= 2020 OR country == "BG" AND city != "Sofia" - String vectorExpr = converter.convert(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), - new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), - new Expression(NE, new Key("city"), new Value("Sofia"))))); + String vectorExpr = converter + .convertExpression(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), + new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), + new Expression(NE, new Key("city"), new Value("Sofia"))))); assertThat(vectorExpr).isEqualTo( "metadata[\"year\"] >= 2020 || metadata[\"country\"] == \"BG\" && metadata[\"city\"] != \"Sofia\""); } @@ -79,7 +80,7 @@ public class MilvusFilterExpressionConverterTests { @Test public void testGroup() { // (year >= 2020 OR country == "BG") AND city NIN ["Sofia", "Plovdiv"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Group(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), new Expression(EQ, new Key("country"), new Value("BG")))), new Expression(NIN, new Key("city"), new Value(List.of("Sofia", "Plovdiv"))))); @@ -90,7 +91,7 @@ public class MilvusFilterExpressionConverterTests { @Test public void tesBoolean() { // isOpen == true AND year >= 2020 AND country IN ["BG", "NL", "US"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Expression(AND, new Expression(EQ, new Key("isOpen"), new Value(true)), new Expression(GTE, new Key("year"), new Value(2020))), new Expression(IN, new Key("country"), new Value(List.of("BG", "NL", "US"))))); @@ -103,10 +104,20 @@ public class MilvusFilterExpressionConverterTests { public void testDecimal() { // temperature >= -15.6 && temperature <= +20.13 String vectorExpr = converter - .convert(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), + .convertExpression(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), new Expression(LTE, new Key("temperature"), new Value(20.13)))); assertThat(vectorExpr).isEqualTo("metadata[\"temperature\"] >= -15.6 && metadata[\"temperature\"] <= 20.13"); } + @Test + public void testComplexIdentifiers() { + String vectorExpr = converter + .convertExpression(new Expression(EQ, new Key("\"country 1 2 3\""), new Value("BG"))); + assertThat(vectorExpr).isEqualTo("metadata[\"country 1 2 3\"] == \"BG\""); + + vectorExpr = converter.convertExpression(new Expression(EQ, new Key("'country 1 2 3'"), new Value("BG"))); + assertThat(vectorExpr).isEqualTo("metadata[\"country 1 2 3\"] == \"BG\""); + } + } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverterTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverterTests.java index 4ad8bb9a1..cac2dd2c5 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverterTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PgVectorFilterExpressionConverterTests.java @@ -40,12 +40,12 @@ import static org.springframework.ai.vectorstore.filter.Filter.ExpressionType.OR */ public class PgVectorFilterExpressionConverterTests { - PgVectorFilterExpressionConverter converter = new PgVectorFilterExpressionConverter(); + FilterExpressionConverter converter = new PgVectorFilterExpressionConverter(); @Test public void testEQ() { // country == "BG" - String vectorExpr = converter.convert(new Expression(EQ, new Key("country"), new Value("BG"))); + String vectorExpr = converter.convertExpression(new Expression(EQ, new Key("country"), new Value("BG"))); assertThat(vectorExpr).isEqualTo("$.country == \"BG\""); } @@ -53,7 +53,7 @@ public class PgVectorFilterExpressionConverterTests { public void tesEqAndGte() { // genre == "drama" AND year >= 2020 String vectorExpr = converter - .convert(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), + .convertExpression(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), new Expression(GTE, new Key("year"), new Value(2020)))); assertThat(vectorExpr).isEqualTo("$.genre == \"drama\" && $.year >= 2020"); } @@ -61,24 +61,25 @@ public class PgVectorFilterExpressionConverterTests { @Test public void tesIn() { // genre in ["comedy", "documentary", "drama"] - String vectorExpr = converter - .convert(new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); + String vectorExpr = converter.convertExpression( + new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); assertThat(vectorExpr).isEqualTo("$.genre in [\"comedy\",\"documentary\",\"drama\"]"); } @Test public void testNe() { // year >= 2020 OR country == "BG" AND city != "Sofia" - String vectorExpr = converter.convert(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), - new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), - new Expression(NE, new Key("city"), new Value("Sofia"))))); + String vectorExpr = converter + .convertExpression(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), + new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), + new Expression(NE, new Key("city"), new Value("Sofia"))))); assertThat(vectorExpr).isEqualTo("$.year >= 2020 || $.country == \"BG\" && $.city != \"Sofia\""); } @Test public void testGroup() { // (year >= 2020 OR country == "BG") AND city NIN ["Sofia", "Plovdiv"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Group(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), new Expression(EQ, new Key("country"), new Value("BG")))), new Expression(NIN, new Key("city"), new Value(List.of("Sofia", "Plovdiv"))))); @@ -89,7 +90,7 @@ public class PgVectorFilterExpressionConverterTests { @Test public void tesBoolean() { // isOpen == true AND year >= 2020 AND country IN ["BG", "NL", "US"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Expression(AND, new Expression(EQ, new Key("isOpen"), new Value(true)), new Expression(GTE, new Key("year"), new Value(2020))), new Expression(IN, new Key("country"), new Value(List.of("BG", "NL", "US"))))); @@ -101,10 +102,17 @@ public class PgVectorFilterExpressionConverterTests { public void testDecimal() { // temperature >= -15.6 && temperature <= +20.13 String vectorExpr = converter - .convert(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), + .convertExpression(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), new Expression(LTE, new Key("temperature"), new Value(20.13)))); assertThat(vectorExpr).isEqualTo("$.temperature >= -15.6 && $.temperature <= 20.13"); } + @Test + public void testComplexIdentifiers() { + String vectorExpr = converter + .convertExpression(new Expression(EQ, new Key("\"country 1 2 3\""), new Value("BG"))); + assertThat(vectorExpr).isEqualTo("$.\"country 1 2 3\" == \"BG\""); + } + } diff --git a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverterTests.java b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverterTests.java index 3b37fc50e..45a9277b3 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverterTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/vectorstore/filter/converter/PineconeFilterExpressionConverterTests.java @@ -40,12 +40,12 @@ import static org.springframework.ai.vectorstore.filter.Filter.ExpressionType.OR */ public class PineconeFilterExpressionConverterTests { - PineconeFilterExpressionConverter converter = new PineconeFilterExpressionConverter(); + FilterExpressionConverter converter = new PineconeFilterExpressionConverter(); @Test public void testEQ() { // country == "BG" - String vectorExpr = converter.convert(new Expression(EQ, new Key("country"), new Value("BG"))); + String vectorExpr = converter.convertExpression(new Expression(EQ, new Key("country"), new Value("BG"))); assertThat(vectorExpr).isEqualTo("{\"country\": {\"$eq\": \"BG\"}}"); } @@ -53,7 +53,7 @@ public class PineconeFilterExpressionConverterTests { public void tesEqAndGte() { // genre == "drama" AND year >= 2020 String vectorExpr = converter - .convert(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), + .convertExpression(new Expression(AND, new Expression(EQ, new Key("genre"), new Value("drama")), new Expression(GTE, new Key("year"), new Value(2020)))); assertThat(vectorExpr) .isEqualTo("{\"$and\": [{\"genre\": {\"$eq\": \"drama\"}},{\"year\": {\"$gte\": 2020}}]}"); @@ -62,17 +62,18 @@ public class PineconeFilterExpressionConverterTests { @Test public void tesIn() { // genre in ["comedy", "documentary", "drama"] - String vectorExpr = converter - .convert(new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); + String vectorExpr = converter.convertExpression( + new Expression(IN, new Key("genre"), new Value(List.of("comedy", "documentary", "drama")))); assertThat(vectorExpr).isEqualTo("{\"genre\": {\"$in\": [\"comedy\",\"documentary\",\"drama\"]}}"); } @Test public void testNe() { // year >= 2020 OR country == "BG" AND city != "Sofia" - String vectorExpr = converter.convert(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), - new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), - new Expression(NE, new Key("city"), new Value("Sofia"))))); + String vectorExpr = converter + .convertExpression(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), + new Expression(AND, new Expression(EQ, new Key("country"), new Value("BG")), + new Expression(NE, new Key("city"), new Value("Sofia"))))); assertThat(vectorExpr).isEqualTo( "{\"$or\": [{\"year\": {\"$gte\": 2020}},{\"$and\": [{\"country\": {\"$eq\": \"BG\"}},{\"city\": {\"$ne\": \"Sofia\"}}]}]}"); } @@ -80,7 +81,7 @@ public class PineconeFilterExpressionConverterTests { @Test public void testGroup() { // (year >= 2020 OR country == "BG") AND city NIN ["Sofia", "Plovdiv"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Group(new Expression(OR, new Expression(GTE, new Key("year"), new Value(2020)), new Expression(EQ, new Key("country"), new Value("BG")))), new Expression(NIN, new Key("city"), new Value(List.of("Sofia", "Plovdiv"))))); @@ -91,7 +92,7 @@ public class PineconeFilterExpressionConverterTests { @Test public void tesBoolean() { // isOpen == true AND year >= 2020 AND country IN ["BG", "NL", "US"] - String vectorExpr = converter.convert(new Expression(AND, + String vectorExpr = converter.convertExpression(new Expression(AND, new Expression(AND, new Expression(EQ, new Key("isOpen"), new Value(true)), new Expression(GTE, new Key("year"), new Value(2020))), new Expression(IN, new Key("country"), new Value(List.of("BG", "NL", "US"))))); @@ -104,11 +105,21 @@ public class PineconeFilterExpressionConverterTests { public void testDecimal() { // temperature >= -15.6 && temperature <= +20.13 String vectorExpr = converter - .convert(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), + .convertExpression(new Expression(AND, new Expression(GTE, new Key("temperature"), new Value(-15.6)), new Expression(LTE, new Key("temperature"), new Value(20.13)))); assertThat(vectorExpr) .isEqualTo("{\"$and\": [{\"temperature\": {\"$gte\": -15.6}},{\"temperature\": {\"$lte\": 20.13}}]}"); } + @Test + public void testComplexIdentifiers() { + String vectorExpr = converter + .convertExpression(new Expression(EQ, new Key("\"country 1 2 3\""), new Value("BG"))); + assertThat(vectorExpr).isEqualTo("{\"country 1 2 3\": {\"$eq\": \"BG\"}}"); + + vectorExpr = converter.convertExpression(new Expression(EQ, new Key("'country 1 2 3'"), new Value("BG"))); + assertThat(vectorExpr).isEqualTo("{\"country 1 2 3\": {\"$eq\": \"BG\"}}"); + } + } diff --git a/spring-ai-spring-boot-autoconfigure/pom.xml b/spring-ai-spring-boot-autoconfigure/pom.xml index 52d128847..98345a519 100644 --- a/spring-ai-spring-boot-autoconfigure/pom.xml +++ b/spring-ai-spring-boot-autoconfigure/pom.xml @@ -94,6 +94,14 @@ true + + + org.springframework.experimental.ai + spring-ai-chroma-store + ${project.parent.version} + true + + org.springframework.boot diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaApiProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaApiProperties.java new file mode 100644 index 000000000..7fd599995 --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaApiProperties.java @@ -0,0 +1,79 @@ +/* + * Copyright 2023-2023 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.autoconfigure.vectorstore.chroma; + +import org.springframework.boot.context.properties.ConfigurationProperties; + +/** + * @author Christian Tzolov + */ +@ConfigurationProperties(ChromaApiProperties.CONFIG_PREFIX) +public class ChromaApiProperties { + + public static final String CONFIG_PREFIX = "spring.ai.vectorstore.chroma.client"; + + private String host = "http://localhost"; + + private int port = 8000; + + private String keyToken; + + private String username; + + private String password; + + public String getHost() { + return host; + } + + public void setHost(String baseUrl) { + this.host = baseUrl; + } + + public int getPort() { + return port; + } + + public void setPort(int port) { + this.port = port; + } + + public String getKeyToken() { + return keyToken; + } + + public void setKeyToken(String keyToken) { + this.keyToken = keyToken; + } + + public String getUsername() { + return username; + } + + public void setUsername(String username) { + this.username = username; + } + + public String getPassword() { + return password; + } + + public void setPassword(String password) { + this.password = password; + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java new file mode 100644 index 000000000..096ccc1f3 --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java @@ -0,0 +1,72 @@ +/* + * Copyright 2023-2023 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.autoconfigure.vectorstore.chroma; + +import com.fasterxml.jackson.databind.ObjectMapper; + +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.boot.autoconfigure.AutoConfiguration; +import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; +import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; +import org.springframework.boot.context.properties.EnableConfigurationProperties; +import org.springframework.context.annotation.Bean; +import org.springframework.experimental.ai.chroma.ChromaApi; +import org.springframework.experimental.ai.vectorsore.ChromaVectorStore; +import org.springframework.util.StringUtils; +import org.springframework.web.client.RestTemplate; + +/** + * @author Christian Tzolov + */ +@AutoConfiguration +@ConditionalOnClass({ EmbeddingClient.class, RestTemplate.class, ChromaVectorStore.class, ObjectMapper.class }) +@EnableConfigurationProperties({ ChromaApiProperties.class, ChromaVectorStoreProperties.class }) +public class ChromaVectorStoreAutoConfiguration { + + @Bean + @ConditionalOnMissingBean + public RestTemplate restTemplate() { + return new RestTemplate(); + } + + @Bean + @ConditionalOnMissingBean + public ChromaApi chromaApi(ChromaApiProperties apiProperties, RestTemplate restTemplate) { + + String chromaUrl = String.format("%s:%s", apiProperties.getHost(), apiProperties.getPort()); + + var chromaApi = new ChromaApi(chromaUrl, restTemplate, new ObjectMapper()); + + if (StringUtils.hasText(apiProperties.getKeyToken())) { + chromaApi.withKeyToken(apiProperties.getKeyToken()); + } + else if (StringUtils.hasText(apiProperties.getUsername()) && StringUtils.hasText(apiProperties.getPassword())) { + chromaApi.withBasicAuthCredentials(apiProperties.getUsername(), apiProperties.getPassword()); + } + + return chromaApi; + } + + @Bean + @ConditionalOnMissingBean + public VectorStore vectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi, + ChromaVectorStoreProperties storeProperties) { + return new ChromaVectorStore(embeddingClient, chromaApi, storeProperties.getCollectionName()); + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreProperties.java new file mode 100644 index 000000000..1d6b5525f --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreProperties.java @@ -0,0 +1,40 @@ +/* + * Copyright 2023-2023 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.autoconfigure.vectorstore.chroma; + +import org.springframework.boot.context.properties.ConfigurationProperties; +import org.springframework.experimental.ai.vectorsore.ChromaVectorStore; + +/** + * @author Christian Tzolov + */ +@ConfigurationProperties(ChromaVectorStoreProperties.CONFIG_PREFIX) +public class ChromaVectorStoreProperties { + + public static final String CONFIG_PREFIX = "spring.ai.vectorstore.chroma.store"; + + private String collectionName = ChromaVectorStore.DEFAULT_COLLECTION_NAME; + + public String getCollectionName() { + return collectionName; + } + + public void setCollectionName(String collectionName) { + this.collectionName = collectionName; + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports b/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports index 2c6f7d20a..5d3c956f0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports +++ b/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports @@ -5,3 +5,4 @@ org.springframework.ai.autoconfigure.vectorstore.pinecone.PineconeVectorStoreAut org.springframework.ai.autoconfigure.vectorstore.milvus.MilvusVectorStoreAutoConfiguration org.springframework.ai.autoconfigure.embedding.transformer.TransformersEmbeddingClientAutoConfiguration org.springframework.ai.autoconfigure.huggingface.HuggingfaceAutoConfiguration +org.springframework.ai.autoconfigure.vectorstore.chroma.ChromaVectorStoreAutoConfiguration diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java new file mode 100644 index 000000000..30bba69cf --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java @@ -0,0 +1,100 @@ +/* + * Copyright 2023-2023 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.autoconfigure.vectorstore.chroma; + +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.Test; +import org.testcontainers.containers.GenericContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; + +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.TransformersEmbeddingClient; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * @author Christian Tzolov + */ +@Testcontainers +public class ChromaVectorStoreAutoConfigurationIT { + + @Container + static GenericContainer chromaContainer = new GenericContainer<>("ghcr.io/chroma-core/chroma:0.4.15") + .withExposedPorts(8000); + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withConfiguration(AutoConfigurations.of(ChromaVectorStoreAutoConfiguration.class)) + .withUserConfiguration(Config.class) + .withPropertyValues("spring.ai.vectorstore.chroma.client.host=http://localhost", + "spring.ai.vectorstore.chroma.client.port=" + chromaContainer.getMappedPort(8000), + "spring.ai.vectorstore.chroma.store.collectionName=TestCollection"); + + @Test + public void addAndSearchWithFilters() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + var bgDocument = new Document("The World is Big and Salvation Lurks Around the Corner", + Map.of("country", "Bulgaria")); + var nlDocument = new Document("The World is Big and Salvation Lurks Around the Corner", + Map.of("country", "Netherland")); + + vectorStore.add(List.of(bgDocument, nlDocument)); + + var request = SearchRequest.query("The World").withTopK(5); + + List results = vectorStore.similaritySearch(request); + assertThat(results).hasSize(2); + + results = vectorStore + .similaritySearch(request.withSimilarityThresholdAll().withFilterExpression("country == 'Bulgaria'")); + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo(bgDocument.getId()); + + results = vectorStore + .similaritySearch(request.withSimilarityThresholdAll().withFilterExpression("country == 'Netherland'")); + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo(nlDocument.getId()); + + // Remove all documents from the store + vectorStore.delete(List.of(bgDocument, nlDocument).stream().map(doc -> doc.getId()).toList()); + }); + } + + @Configuration(proxyBeanMethods = false) + static class Config { + + @Bean + public EmbeddingClient embeddingClient() { + return new TransformersEmbeddingClient(); + } + + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java index 7ef80d715..33bdca50a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java @@ -22,12 +22,10 @@ import java.util.List; import java.util.Map; import org.junit.jupiter.api.Test; -import org.springframework.core.io.DefaultResourceLoader; import org.testcontainers.containers.GenericContainer; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; -import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; import org.springframework.ai.embedding.TransformersEmbeddingClient; @@ -39,6 +37,7 @@ import org.springframework.boot.autoconfigure.jdbc.JdbcTemplateAutoConfiguration import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; +import org.springframework.core.io.DefaultResourceLoader; import static org.assertj.core.api.Assertions.assertThat; diff --git a/vector-stores/spring-ai-chroma/README.md b/vector-stores/spring-ai-chroma/README.md new file mode 100644 index 000000000..f40c33ff6 --- /dev/null +++ b/vector-stores/spring-ai-chroma/README.md @@ -0,0 +1,154 @@ +# Chroma VectorStore + +This readme will walk you through setting up the Chroma VectorStore to store document embeddings and perform similarity searches. + + + +## What is Chroma? + +[Chroma](https://docs.trychroma.com/) is the open-source embedding database. It gives you the tools to store document embeddings, content and metadata and to search through those embeddings including metadata filtering. + +## Prerequisites + +1. OpenAI Account: Create an account at [OpenAI Signup](https://platform.openai.com/signup) and generate the token at [API Keys](https://platform.openai.com/account/api-keys). + +2. Access to ChromeDB. The [setup local ChromaDB](#appendix_a) appendix show how to setup a DB locally with a Docker container. + +On startup the `ChromaVectorStore` creates the required collection if one is not provisioned already. + +## Configuration + +To set up ChromaVectorStore, you'll need to provide your OpenAI API Key. Set it as an environment variable like so: + +```bash +export SPRING_AI_OPENAI_API_KEY='Your_OpenAI_API_Key' +``` + +## Dependencies + +Add these dependencies to your project: + +1. OpenAI: Required for calculating embeddings. + + ```xml + + org.springframework.experimental.ai + spring-ai-openai-spring-boot-starter + 0.7.0-SNAPSHOT + + ``` + +2. Chroma VectorStore. + + ```xml + + org.springframework.experimental.ai + spring-ai-chroma-store + 0.7.0-SNAPSHOT + + ``` + +## Sample Code + +Create an `RestTemplate` instance with proper ChromaDB authorization configurations and Use it to create `ChromaApi` instance: + +```java +@Bean +public RestTemplate restTemplate() { + return new RestTemplate(); +} + +@Bean +public ChromaApi chromaApi(RestTemplate restTemplate) { + String chromaUrl = "http://localhost:8000"; + ChromaApi chromaApi = ChromaApi(chromaUrl, restTemplate); + return chromaApi; +} +``` +> [!NOTE] +> For ChromaDB secured with [Static API Token Authentication](https://docs.trychroma.com/usage-guide#static-api-token-authentication) use the `ChromaApi#withKeyToken()` method to set your credentials. Check the `ChromaWhereIT` for an example. + +> [!NOTE] +> For ChromaDB secured with [Basic Authentication](https://docs.trychroma.com/usage-guide#basic-authentication) use the `ChromaApi#withBasicAuth(, )` method to set your credentials. Check the `BasicAuthChromaWhereIT` for an example. + + +Integrate with OpenAI's embeddings by adding the Spring Boot OpenAI starter to your project. +This provides you with an implementation of the Embeddings client: + +```java +@Bean +public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); +} +``` + +In your main code, create some documents + +```java +List documents = List.of( + new Document("Spring AI rocks!! Spring AI rocks!! Spring AI rocks!! Spring AI rocks!! Spring AI rocks!!", Map.of("meta1", "meta1")), + new Document("The World is Big and Salvation Lurks Around the Corner"), + new Document("You walk forward facing the past and you turn back toward the future.", Map.of("meta2", "meta2"))); +``` + +Add the documents to your vector store: + +```java +vectorStore.add(List.of(document)); +``` + +And finally, retrieve documents similar to a query: + +```java +List results = vectorStore.similaritySearch("Spring", 5); +``` + +If all goes well, you should retrieve the document containing the text "Spring AI rocks!!". + +### Metadata filtering + +You can leverage the generic, portable [metadata filters](https://docs.spring.io/spring-ai/reference/api/vectordbs.html#_metadata_filters) with ChromaVector store as well. + +For example you can use either the text expression language: + +```java +vectorStore.similaritySearch("The World", TOP_K, SIMILARITY_THRESHOLD, + "author in ['john', 'jill'] && article_type == 'blog'"); +``` + +or programmatically using the `Filter.Expression` DSL: + +```java +FilterExpressionBuilder b = new FilterExpressionBuilder(); + +vectorStore.similaritySearch("The World", TOP_K, SIMILARITY_THRESHOLD, + b.and( + b.in(List.of("john", "jill")), + b.eq("article_type", "blog")).build()); +``` + +NOTE: Those (portable) filter expressions get automatically converted into the proprietary Chroma `where` [filter expressions](https://docs.trychroma.com/usage-guide#using-where-filters). + +For example this portable filter expression: + +```sql +author in ['john', 'jill'] && article_type == 'blog' +``` + +is converted into the proprietary Chroma format: + +```json +{"$and":[ + {"author": {"$in": ["john", "jill"]}}, + {"article_type":{"$eq":"blog"}}] +}" +``` + + +## Appendix A: Run Chroma Locally + +``` +docker run -it --rm --name chroma -p 8000:8000 ghcr.io/chroma-core/chroma:0.4.15 +``` + +starts a chroma store at diff --git a/vector-stores/spring-ai-chroma/pom.xml b/vector-stores/spring-ai-chroma/pom.xml new file mode 100644 index 000000000..b6874ab3c --- /dev/null +++ b/vector-stores/spring-ai-chroma/pom.xml @@ -0,0 +1,85 @@ + + + 4.0.0 + + org.springframework.experimental.ai + spring-ai + 0.7.1-SNAPSHOT + ../../pom.xml + + spring-ai-chroma-store + jar + Spring AI Chroma Vector Store + Spring AI Chroma Vector Store + https://github.com/spring-projects-experimental/spring-ai + + + https://github.com/spring-projects-experimental/spring-ai + git://github.com/spring-projects-experimental/spring-ai.git + git@github.com:spring-projects-experimental/spring-ai.git + + + + + org.springframework.experimental.ai + spring-ai-core + ${parent.version} + + + + + org.springframework.experimental.ai + spring-ai-openai + ${parent.version} + test + + + + org.springframework.boot + spring-boot-starter-test + test + + + + org.testcontainers + testcontainers + ${testcontainers.version} + test + + + + org.testcontainers + junit-jupiter + ${testcontainers.version} + test + + + + + + + + diff --git a/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/chroma/ChromaApi.java b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/chroma/ChromaApi.java new file mode 100644 index 000000000..d0a2aa31c --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/chroma/ChromaApi.java @@ -0,0 +1,380 @@ +/* + * Copyright 2023-2023 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.experimental.ai.chroma; + +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +import com.fasterxml.jackson.annotation.JsonProperty; +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; + +import org.springframework.experimental.ai.chroma.ChromaApi.QueryRequest.Include; +import org.springframework.http.HttpEntity; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.MediaType; +import org.springframework.http.client.support.BasicAuthenticationInterceptor; +import org.springframework.util.CollectionUtils; +import org.springframework.util.StringUtils; +import org.springframework.web.client.HttpServerErrorException; +import org.springframework.web.client.RestTemplate; + +/** + * Single-class Chroma API implementation based on the (unofficial) Chroma REST API. + * + * @author Christian Tzolov + */ +public class ChromaApi { + + // Regular expression pattern that looks for a message inside the ValueError(...). + private static Pattern VALUE_ERROR_PATTERN = Pattern.compile("ValueError\\('([^']*)'\\)"); + + private final String baseUrl; + + private final RestTemplate restTemplate; + + private final ObjectMapper objectMapper; + + private String keyToken; + + public ChromaApi(String baseUrl, RestTemplate restTemplate) { + this(baseUrl, restTemplate, new ObjectMapper()); + } + + public ChromaApi(String baseUrl, RestTemplate restTemplate, ObjectMapper objectMapper) { + this.baseUrl = baseUrl; + this.restTemplate = restTemplate; + this.objectMapper = objectMapper; + } + + /** + * Configure access to ChromaDB secured with static API Token Authentication: + * https://docs.trychroma.com/usage-guide#static-api-token-authentication + * @param keyToken Chroma static API Token Authentication. (Optional) + */ + public ChromaApi withKeyToken(String keyToken) { + this.keyToken = keyToken; + return this; + } + + /** + * Configure access to ChromaDB secured with Basic Authentication: + * https://docs.trychroma.com/usage-guide#basic-authentication + * @param username Credentials username. + * @param password Credentials password. + */ + public ChromaApi withBasicAuthCredentials(String username, String password) { + this.restTemplate.getInterceptors().add(new BasicAuthenticationInterceptor(username, username)); + return this; + } + + /** + * Chroma embedding collection. + * + * @param id Collection Id. + * @param name The name of the collection. + * @param metadata Metadata associated with the collection. + */ + public record Collection(String id, String name, Map metadata) { + } + + /** + * Request to create a new collection with the given name and metadata. + * + * @param name The name of the collection to create. + * @param metadata Optional metadata to associate with the collection. + */ + public record CreateCollectionRequest(String name, Map metadata) { + public CreateCollectionRequest(String name) { + this(name, new HashMap<>(Map.of("hnsw:space", "cosine"))); + } + } + + /** + * Add embeddings to the chroma data store. + * + * @param ids The ids of the embeddings to add. + * @param embeddings The embeddings to add. + * @param metadata The metadata to associate with the embeddings. When querying, you + * can filter on this metadata. + * @param documents The documents contents to associate with the embeddings. + */ + public record AddEmbeddingsRequest(List ids, List embeddings, + @JsonProperty("metadatas") List> metadata, List documents) { + + // Convenance for adding a single embedding. + public AddEmbeddingsRequest(String id, float[] embedding, Map metadata, String document) { + this(List.of(id), List.of(embedding), List.of(metadata), List.of(document)); + } + } + + /** + * Request to delete embedding from a collection. + * + * @param ids The ids of the embeddings to delete. (Optional) + * @param where Condition to filter items to delete based on metadata values. + * (Optional) + */ + public record DeleteEmbeddingsRequest(List ids, Map where) { + public DeleteEmbeddingsRequest(List ids) { + this(ids, Map.of()); + } + } + + /** + * Get embeddings from a collection. + * + * @param ids IDs of the embeddings to get. + * @param where Condition to filter results based on metadata values. + * @param limit Limit on the number of collection embeddings to get. + * @param offset Offset on the embeddings to get. + * @param include A list of what to include in the results. Can contain "embeddings", + * "metadatas", "documents", "distances". Ids are always included. Defaults to + * [metadatas, documents, distances]. + */ + public record GetEmbeddingsRequest(List ids, Map where, int limit, int offset, + List include) { + + public GetEmbeddingsRequest(List ids) { + this(ids, Map.of(), 10, 0, Include.all); + } + + public GetEmbeddingsRequest(List ids, Map where) { + this(ids, where, 10, 0, Include.all); + } + + public GetEmbeddingsRequest(List ids, Map where, int limit, int offset) { + this(ids, where, limit, offset, Include.all); + } + } + + /** + * Object containing the get embedding results. + * + * @param ids List of document ids. One for each returned document. + * @param embeddings List of document embeddings. One for each returned document. + * @param documents List of document contents. One for each returned document. + * @param metadata List of document metadata. One for each returned document. + */ + public record GetEmbeddingResponse(List ids, List> embeddings, List documents, + @JsonProperty("metadatas") List> metadata) { + } + + /** + * Request to get the nResults nearest neighbor embeddings for provided + * queryEmbeddings. + * + * @param queryEmbeddings The embeddings to get the closes neighbors of. + * @param nResults The number of neighbors to return for each query_embedding or + * query_texts. + * @param where Condition to filter results based on metadata values. + * @param include A list of what to include in the results. Can contain "embeddings", + * "metadatas", "documents", "distances". Ids are always included. Defaults to + * [metadatas, documents, distances]. + */ + public record QueryRequest(@JsonProperty("query_embeddings") List> queryEmbeddings, + @JsonProperty("n_results") int nResults, Map where, List include) { + + public enum Include { + + metadatas, documents, distances, embeddings; + + public static final List all = List.of(metadatas, documents, distances, embeddings); + + } + + /** + * Convenience to query for a single embedding instead of a batch of embeddings. + */ + public QueryRequest(List queryEmbedding, int nResults) { + this(List.of(queryEmbedding), nResults, Map.of(), Include.all); + } + + public QueryRequest(List queryEmbedding, int nResults, Map where) { + this(List.of(queryEmbedding), nResults, where, Include.all); + } + } + + /** + * A QueryResponse object containing the query results. + * + * @param ids List of list of document ids. One for each returned document. + * @param embeddings List of list of document embeddings. One for each returned + * document. + * @param documents List of list of document contents. One for each returned document. + * @param metadata List of list of document metadata. One for each returned document. + * @param distances List of list of search distances. One for each returned document. + */ + public record QueryResponse(List> ids, List>> embeddings, + List> documents, @JsonProperty("metadatas") List>> metadata, + List> distances) { + } + + /** + * Single query embedding response. + */ + public record Embedding(String id, List embedding, String document, Map metadata, + Double distances) { + } + + public List toEmbeddingResponseList(QueryResponse queryResponse) { + List result = new ArrayList<>(); + + if (queryResponse != null && !CollectionUtils.isEmpty(queryResponse.ids())) { + for (int i = 0; i < queryResponse.ids().get(0).size(); i++) { + result.add(new Embedding(queryResponse.ids().get(0).get(i), queryResponse.embeddings().get(0).get(i), + queryResponse.documents().get(0).get(i), queryResponse.metadata().get(0).get(i), + queryResponse.distances().get(0).get(i))); + } + } + + return result; + } + + // + // Chroma Client API (https://docs.trychroma.com/js_reference/Client) + // + + public Collection createCollection(CreateCollectionRequest createCollectionRequest) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections", HttpMethod.POST, + this.getHttpEntityFor(createCollectionRequest), Collection.class) + .getBody(); + } + + /** + * Delete a collection with the given name. + * @param collectionName the name of the collection to delete. + * + */ + public void deleteCollection(String collectionName) { + + this.restTemplate.exchange(this.baseUrl + "/api/v1/collections/{collection_name}", HttpMethod.DELETE, + new HttpEntity<>(httpHeaders()), Void.class, collectionName); + } + + public Collection getCollection(String collectionName) { + + try { + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_name}", HttpMethod.GET, + new HttpEntity<>(httpHeaders()), Collection.class, collectionName) + .getBody(); + } + catch (HttpServerErrorException e) { + String msg = this.getValueErrorMessage(e.getMessage()); + if (String.format("Collection %s does not exist.", collectionName).equals(msg)) { + return null; + } + throw new RuntimeException(msg, e); + } + } + + private static class CollectionList extends ArrayList { + + } + + public List listCollections() { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/", HttpMethod.GET, new HttpEntity<>(httpHeaders()), + CollectionList.class) + .getBody(); + } + + // + // Chroma Collection API (https://docs.trychroma.com/js_reference/Collection) + // + + public Boolean upsertEmbeddings(String collectionId, AddEmbeddingsRequest embedding) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_id}/upsert", HttpMethod.POST, + this.getHttpEntityFor(embedding), Boolean.class, collectionId) + .getBody(); + } + + public List deleteEmbeddings(String collectionId, DeleteEmbeddingsRequest deleteRequest) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_id}/delete", HttpMethod.POST, + this.getHttpEntityFor(deleteRequest), List.class, collectionId) + .getBody(); + } + + public Long countEmbeddings(String collectionId) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_id}/count", HttpMethod.GET, + new HttpEntity<>(httpHeaders()), Long.class, collectionId) + .getBody(); + } + + public QueryResponse queryCollection(String collectionId, QueryRequest queryRequest) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_id}/query", HttpMethod.POST, + this.getHttpEntityFor(queryRequest), QueryResponse.class, collectionId) + .getBody(); + } + + public GetEmbeddingResponse getEmbeddings(String collectionId, GetEmbeddingsRequest getEmbeddingsRequest) { + + return this.restTemplate + .exchange(this.baseUrl + "/api/v1/collections/{collection_id}/get", HttpMethod.POST, + this.getHttpEntityFor(getEmbeddingsRequest), GetEmbeddingResponse.class, collectionId) + .getBody(); + } + + // Utils + public Map where(String text) { + try { + return this.objectMapper.readValue(text, Map.class); + } + catch (JsonProcessingException e) { + throw new RuntimeException(e); + } + } + + private HttpEntity getHttpEntityFor(T body) { + return new HttpEntity<>(body, httpHeaders()); + } + + private HttpHeaders httpHeaders() { + HttpHeaders headers = new HttpHeaders(); + headers.setContentType(MediaType.APPLICATION_JSON); + if (StringUtils.hasText(this.keyToken)) { + headers.setBearerAuth(this.keyToken); + } + return headers; + } + + private String getValueErrorMessage(String logString) { + if (!StringUtils.hasText(logString)) { + return ""; + } + Matcher m = VALUE_ERROR_PATTERN.matcher(logString); + return (m.find()) ? m.group(1) : ""; + } + +} diff --git a/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/ChromaVectorStore.java b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/ChromaVectorStore.java new file mode 100644 index 000000000..872819c0b --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/ChromaVectorStore.java @@ -0,0 +1,161 @@ +/* + * Copyright 2023-2023 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.experimental.ai.vectorsore; + +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.ai.vectorstore.filter.converter.ChromaFilterExpressionConverter; +import org.springframework.ai.vectorstore.filter.converter.FilterExpressionConverter; +import org.springframework.beans.factory.InitializingBean; +import org.springframework.experimental.ai.chroma.ChromaApi; +import org.springframework.experimental.ai.chroma.ChromaApi.AddEmbeddingsRequest; +import org.springframework.experimental.ai.chroma.ChromaApi.DeleteEmbeddingsRequest; +import org.springframework.experimental.ai.chroma.ChromaApi.Embedding; +import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; +import org.springframework.util.StringUtils; + +/** + * @author Christian Tzolov + */ +public class ChromaVectorStore implements VectorStore, InitializingBean { + + public static final String DISTANCE_FIELD_NAME = "distance"; + + public static final String DEFAULT_COLLECTION_NAME = "SpringAiCollection"; + + public static final double SIMILARITY_THRESHOLD_ALL = 0.0; + + public static final int DEFAULT_TOP_K = 4; + + private final EmbeddingClient embeddingClient; + + private final ChromaApi chromaApi; + + private final String collectionName; + + private FilterExpressionConverter filterExpressionConverter; + + private String collectionId; + + public ChromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { + this(embeddingClient, chromaApi, DEFAULT_COLLECTION_NAME); + } + + public ChromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi, String collectionName) { + this.embeddingClient = embeddingClient; + this.chromaApi = chromaApi; + this.collectionName = collectionName; + this.filterExpressionConverter = new ChromaFilterExpressionConverter(); + } + + public void setFilterExpressionConverter(FilterExpressionConverter filterExpressionConverter) { + Assert.notNull(filterExpressionConverter, "FilterExpressionConverter should not be null."); + this.filterExpressionConverter = filterExpressionConverter; + } + + @Override + public void add(List documents) { + Assert.notNull(documents, "Documents must not be null"); + if (CollectionUtils.isEmpty(documents)) { + return; + } + + List ids = new ArrayList<>(); + List> metadatas = new ArrayList<>(); + List contents = new ArrayList<>(); + List embeddings = new ArrayList<>(); + + for (Document document : documents) { + ids.add(document.getId()); + metadatas.add(document.getMetadata()); + contents.add(document.getContent()); + document.setEmbedding(this.embeddingClient.embed(document)); + embeddings.add(JsonUtils.toFloatArray(document.getEmbedding())); + } + + var success = this.chromaApi.upsertEmbeddings(this.collectionId, + new AddEmbeddingsRequest(ids, embeddings, metadatas, contents)); + + if (!success) { + throw new RuntimeException("Unsuccessful storing!"); + } + } + + @Override + public Optional delete(List idList) { + Assert.notNull(idList, "Document id list must not be null"); + List deletedIds = this.chromaApi.deleteEmbeddings(this.collectionId, + new DeleteEmbeddingsRequest(idList)); + return Optional.of(deletedIds.size() == idList.size()); + } + + @Override + public List similaritySearch(SearchRequest request) { + + String nativeFilterExpression = (request.getFilterExpression() != null) + ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; + + String query = request.getQuery(); + Assert.notNull(query, "Query string must not be null"); + + List embedding = this.embeddingClient.embed(query); + Map where = (StringUtils.hasText(nativeFilterExpression)) + ? JsonUtils.jsonToMap(nativeFilterExpression) : Map.of(); + var queryRequest = new ChromaApi.QueryRequest(JsonUtils.toFloatList(embedding), request.getTopK(), where); + var queryResponse = this.chromaApi.queryCollection(this.collectionId, queryRequest); + var embeddings = this.chromaApi.toEmbeddingResponseList(queryResponse); + + List responseDocuments = new ArrayList<>(); + + for (Embedding chromaEmbedding : embeddings) { + float distance = chromaEmbedding.distances().floatValue(); + if ((1 - distance) >= request.getSimilarityThreshold()) { + String id = chromaEmbedding.id(); + String content = chromaEmbedding.document(); + Map metadata = chromaEmbedding.metadata(); + if (metadata == null) { + metadata = new HashMap<>(); + } + metadata.put(DISTANCE_FIELD_NAME, distance); + Document document = new Document(id, content, metadata); + document.setEmbedding(JsonUtils.toDouble(chromaEmbedding.embedding())); + responseDocuments.add(document); + } + } + + return responseDocuments; + } + + @Override + public void afterPropertiesSet() throws Exception { + var collection = this.chromaApi.getCollection(this.collectionName); + if (collection == null) { + collection = this.chromaApi.createCollection(new ChromaApi.CreateCollectionRequest(this.collectionName)); + } + this.collectionId = collection.id(); + } + +} diff --git a/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/JsonUtils.java b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/JsonUtils.java new file mode 100644 index 000000000..0e624b1bc --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/main/java/org/springframework/experimental/ai/vectorsore/JsonUtils.java @@ -0,0 +1,84 @@ +/* + * Copyright 2023-2023 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.experimental.ai.vectorsore; + +import java.util.List; +import java.util.Map; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; + +/** + * @author Christian Tzolov + */ + +/** + * Utility class for JSON processing. Provides methods for converting JSON strings to maps + * and lists, and for converting between lists of different numeric types. + * + * @author Christian Tzolov + */ +public class JsonUtils { + + /** + * Converts a JSON string to a map. + * @param jsonText the JSON string to convert + * @return the map representation of the JSON string + * @throws RuntimeException if an error occurs during the conversion + */ + public static Map jsonToMap(String jsonText) { + try { + return (Map) new ObjectMapper().readValue(jsonText, Map.class); + } + catch (JsonProcessingException e) { + throw new RuntimeException(e); + } + } + + /** + * Converts a list of doubles to a list of floats. + * @param embeddingDouble the list of doubles to convert + * @return the list of floats + */ + public static List toFloatList(List embeddingDouble) { + return embeddingDouble.stream().map(Number::floatValue).toList(); + } + + /** + * Converts a list of doubles to a float array. + * @param embeddingDouble the list of doubles to convert + * @return the float array + */ + public static float[] toFloatArray(List embeddingDouble) { + float[] embeddingFloat = new float[embeddingDouble.size()]; + int i = 0; + for (Double d : embeddingDouble) { + embeddingFloat[i++] = d.floatValue(); + } + return embeddingFloat; + } + + /** + * Converts a list of floats to a list of doubles. + * @param floats the list of floats to convert + * @return the list of doubles + */ + public static List toDouble(List floats) { + return floats.stream().map(f -> f.doubleValue()).toList(); + } + +} diff --git a/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/chroma/ChromaApiIT.java b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/chroma/ChromaApiIT.java new file mode 100644 index 000000000..1b95393e0 --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/chroma/ChromaApiIT.java @@ -0,0 +1,186 @@ +/* + * Copyright 2023-2023 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.experimental.ai.chroma; + +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.testcontainers.containers.GenericContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; + +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.context.annotation.Bean; +import org.springframework.experimental.ai.chroma.ChromaApi.AddEmbeddingsRequest; +import org.springframework.experimental.ai.chroma.ChromaApi.Collection; +import org.springframework.experimental.ai.chroma.ChromaApi.GetEmbeddingsRequest; +import org.springframework.experimental.ai.chroma.ChromaApi.QueryRequest; +import org.springframework.web.client.RestTemplate; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * @author Christian Tzolov + */ +@SpringBootTest +@Testcontainers +public class ChromaApiIT { + + @Container + static GenericContainer chromaContainer = new GenericContainer<>("ghcr.io/chroma-core/chroma:0.4.15") + .withExposedPorts(8000); + + @Autowired + ChromaApi chroma; + + @BeforeEach + public void beforeEach() { + chroma.listCollections().stream().forEach(c -> chroma.deleteCollection(c.name())); + } + + @Test + public void testClient() { + var newCollection = chroma.createCollection(new ChromaApi.CreateCollectionRequest("TestCollection")); + assertThat(newCollection).isNotNull(); + assertThat(newCollection.name()).isEqualTo("TestCollection"); + + var getCollection = chroma.getCollection("TestCollection"); + assertThat(getCollection).isNotNull(); + assertThat(getCollection.name()).isEqualTo("TestCollection"); + assertThat(getCollection.id()).isEqualTo(newCollection.id()); + + List collections = chroma.listCollections(); + assertThat(collections).hasSize(1); + assertThat(collections.get(0).id()).isEqualTo(newCollection.id()); + + chroma.deleteCollection(newCollection.name()); + assertThat(chroma.listCollections()).hasSize(0); + } + + @Test + public void testCollection() { + var newCollection = chroma.createCollection(new ChromaApi.CreateCollectionRequest("TestCollection")); + assertThat(chroma.countEmbeddings(newCollection.id())).isEqualTo(0); + + var addEmbeddingRequest = new AddEmbeddingsRequest(List.of("id1", "id2"), + List.of(new float[] { 1f, 1f, 1f }, new float[] { 2f, 2f, 2f }), + List.of(Map.of(), Map.of("key1", "value1", "key2", true, "key3", 23.4)), + List.of("Hello World", "Big World")); + + var success = chroma.upsertEmbeddings(newCollection.id(), addEmbeddingRequest); + + assertThat(success).isTrue(); + + var addEmbeddingRequest2 = new AddEmbeddingsRequest("id3", new float[] { 3f, 3f, 3f }, + Map.of("key1", "value1", "key2", true, "key3", 23.4), "Big World"); + + chroma.upsertEmbeddings(newCollection.id(), addEmbeddingRequest2); + + assertThat(chroma.countEmbeddings(newCollection.id())).isEqualTo(3); + + // Update existing embedding. + chroma.upsertEmbeddings(newCollection.id(), new AddEmbeddingsRequest("id3", new float[] { 6f, 6f, 6f }, + Map.of("key1", "value2", "key2", false, "key4", 23.4), "Small World")); + + var result = chroma.getEmbeddings(newCollection.id(), new GetEmbeddingsRequest(List.of("id2"))); + assertThat(result.ids().get(0)).isEqualTo("id2"); + + result = chroma.getEmbeddings(newCollection.id(), new GetEmbeddingsRequest(List.of(), chroma.where(""" + { "key2" : { "$eq": true} } + """))); + + assertThat(result.ids()).containsExactlyInAnyOrder("id2"); + } + + @Test + public void testQueryWhere() { + + var collection = chroma.createCollection(new ChromaApi.CreateCollectionRequest("TestCollection")); + + var add1 = new AddEmbeddingsRequest("id1", new float[] { 1f, 1f, 1f }, + Map.of("country", "BG", "active", true, "price", 23.4, "year", 2020), + "The World is Big and Salvation Lurks Around the Corner"); + + var add2 = new AddEmbeddingsRequest("id2", new float[] { 1f, 1f, 1f }, Map.of("country", "NL"), + "The World is Big and Salvation Lurks Around the Corner"); + + var add3 = new AddEmbeddingsRequest("id3", new float[] { 1f, 1f, 1f }, + Map.of("country", "BG", "active", false, "price", 40.1, "year", 2023), + "The World is Big and Salvation Lurks Around the Corner"); + + chroma.upsertEmbeddings(collection.id(), add1); + chroma.upsertEmbeddings(collection.id(), add2); + chroma.upsertEmbeddings(collection.id(), add3); + + assertThat(chroma.countEmbeddings(collection.id())).isEqualTo(3); + + var queryResult = chroma.queryCollection(collection.id(), new QueryRequest(List.of(1f, 1f, 1f), 3)); + + assertThat(queryResult.ids().get(0)).hasSize(3); + assertThat(queryResult.ids().get(0)).containsExactlyInAnyOrder("id1", "id2", "id3"); + + var chromaEmbeddings = chroma.toEmbeddingResponseList(queryResult); + + assertThat(chromaEmbeddings).hasSize(3); + assertThat(chromaEmbeddings).hasSize(3); + + queryResult = chroma.queryCollection(collection.id(), new QueryRequest(List.of(1f, 1f, 1f), 3, chroma.where(""" + { + "$and" : [ + {"country" : { "$eq": "BG"}}, + {"year" : { "$gte": 2020}} + ] + } + """))); + assertThat(queryResult.ids().get(0)).hasSize(2); + assertThat(queryResult.ids().get(0)).containsExactlyInAnyOrder("id1", "id3"); + + queryResult = chroma.queryCollection(collection.id(), new QueryRequest(List.of(1f, 1f, 1f), 3, chroma.where(""" + { + "$and" : [ + {"country" : { "$eq": "BG"}}, + {"year" : { "$gte": 2020}}, + {"active" : { "$eq": true}} + ] + } + """))); + assertThat(queryResult.ids().get(0)).hasSize(1); + assertThat(queryResult.ids().get(0)).containsExactlyInAnyOrder("id1"); + } + + @SpringBootConfiguration + public static class Config { + + @Bean + public RestTemplate restTemplate() { + return new RestTemplate(); + } + + @Bean + public ChromaApi chromaApi(RestTemplate restTemplate) { + + int port = chromaContainer.getMappedPort(8000); + return new ChromaApi("http://localhost:" + port, restTemplate); + } + + } + +} diff --git a/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/BasicAuthChromaWhereIT.java b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/BasicAuthChromaWhereIT.java new file mode 100644 index 000000000..9e789943b --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/BasicAuthChromaWhereIT.java @@ -0,0 +1,141 @@ +/* + * Copyright 2023-2023 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.experimental.ai.vectorstore; + +import java.time.Duration; +import java.util.List; +import java.util.Map; + +import com.theokanning.openai.client.OpenAiApi; +import com.theokanning.openai.service.OpenAiService; +import org.junit.jupiter.api.Test; +import org.testcontainers.containers.GenericContainer; +import org.testcontainers.images.builder.Transferable; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; +import retrofit2.Retrofit; +import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory; +import retrofit2.converter.jackson.JacksonConverterFactory; + +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.autoconfigure.AutoConfigurations; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.experimental.ai.chroma.ChromaApi; +import org.springframework.experimental.ai.vectorsore.ChromaVectorStore; +import org.springframework.web.client.RestTemplate; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * ChromaDB with Basic Authentication: + * https://docs.trychroma.com/usage-guide#basic-authentication + * + * The scr/test/resource/server.htpasswd file is generated with: + * htpasswd -Bbn admin admin > server.htpasswd + * + * @author Christian Tzolov + */ +@Testcontainers +public class BasicAuthChromaWhereIT { + + public static String CHROMA_SERVER_URL = "http://localhost:"; + + /** + * ChromaDB with Basic Authentication: + * https://docs.trychroma.com/usage-guide#basic-authentication + */ + @Container + static GenericContainer chromaContainer = new GenericContainer<>("ghcr.io/chroma-core/chroma:0.4.15") + .withEnv("CHROMA_SERVER_AUTH_CREDENTIALS_FILE", "server.htpasswd") + .withEnv("CHROMA_SERVER_AUTH_CREDENTIALS_PROVIDER", + "chromadb.auth.providers.HtpasswdFileServerAuthCredentialsProvider") + .withEnv("CHROMA_SERVER_AUTH_PROVIDER", "chromadb.auth.basic.BasicAuthServerProvider") + .withCopyToContainer(Transferable.of("src/test/resources/server.htpasswd"), "server.htpasswd") + .withExposedPorts(8000); + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withUserConfiguration(TestApplication.class) + .withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY")); + + @Test + public void withInFiltersExpressions1() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + vectorStore.add(List.of(new Document("1", "Article by john", Map.of("author", "john")), + new Document("2", "Article by Jack", Map.of("author", "jack")), + new Document("3", "Article by Jill", Map.of("author", "jill")))); + + String query = "Give me articles by john"; + + List results = vectorStore.similaritySearch(SearchRequest.query(query).withTopK(5)); + assertThat(results).hasSize(3); + + results = vectorStore.similaritySearch(SearchRequest.query(query) + .withTopK(5) + .withSimilarityThresholdAll() + .withFilterExpression("author in ['john', 'jill']")); + + assertThat(results).hasSize(2); + assertThat(results.stream().map(d -> d.getId()).toList()).containsExactlyInAnyOrder("1", "3"); + }); + } + + @SpringBootConfiguration + public static class TestApplication { + + @Bean + public RestTemplate restTemplate() { + return new RestTemplate(); + } + + @Bean + public ChromaApi chromaApi(RestTemplate restTemplate) { + int port = chromaContainer.getMappedPort(8000); + return new ChromaApi(CHROMA_SERVER_URL + port, restTemplate).withBasicAuthCredentials("admin", "admin"); + } + + @Bean + public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + } + + @Bean + public EmbeddingClient embeddingClient() { + + Retrofit retrofit = new Retrofit.Builder().baseUrl("https://api.openai.com") + .client(OpenAiService.defaultClient(System.getenv("OPENAI_API_KEY"), Duration.ofSeconds(60))) + .addConverterFactory(JacksonConverterFactory.create(OpenAiService.defaultObjectMapper())) + .addCallAdapterFactory(RxJava2CallAdapterFactory.create()) + .build(); + + OpenAiApi api = retrofit.create(OpenAiApi.class); + + return new OpenAiEmbeddingClient(new OpenAiService(api), "text-embedding-ada-002"); + } + + } + +} diff --git a/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/ChromaVectorStoreIT.java b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/ChromaVectorStoreIT.java new file mode 100644 index 000000000..7bddd4e98 --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/ChromaVectorStoreIT.java @@ -0,0 +1,240 @@ +/* + * Copyright 2023-2023 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.experimental.ai.vectorstore; + +import java.time.Duration; +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.UUID; + +import com.theokanning.openai.client.OpenAiApi; +import com.theokanning.openai.service.OpenAiService; +import org.junit.jupiter.api.Test; +import org.testcontainers.containers.GenericContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; +import retrofit2.Retrofit; +import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory; +import retrofit2.converter.jackson.JacksonConverterFactory; + +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.experimental.ai.chroma.ChromaApi; +import org.springframework.experimental.ai.vectorsore.ChromaVectorStore; +import org.springframework.web.client.RestTemplate; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * @author Christian Tzolov + */ +@Testcontainers +public class ChromaVectorStoreIT { + + @Container + static GenericContainer chromaContainer = new GenericContainer<>("ghcr.io/chroma-core/chroma:0.4.15") + .withExposedPorts(8000); + + List documents = List.of( + new Document("Spring AI rocks!! Spring AI rocks!! Spring AI rocks!! Spring AI rocks!! Spring AI rocks!!", + Collections.singletonMap("meta1", "meta1")), + new Document("Hello World Hello World Hello World Hello World Hello World Hello World Hello World"), + new Document( + "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression", + Collections.singletonMap("meta2", "meta2"))); + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withUserConfiguration(TestApplication.class) + .withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY")); + + @Test + public void addAndSearch() { + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + vectorStore.add(documents); + + List results = vectorStore.similaritySearch(SearchRequest.query("Great").withTopK(1)); + + assertThat(results).hasSize(1); + Document resultDoc = results.get(0); + assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId()); + assertThat(resultDoc.getContent()).isEqualTo( + "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); + assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance"); + + // Remove all documents from the store + vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList()); + + List results2 = vectorStore.similaritySearch(SearchRequest.query("Great").withTopK(1)); + assertThat(results2).hasSize(0); + }); + } + + @Test + public void addAndSearchWithFilters() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + var bgDocument = new Document("The World is Big and Salvation Lurks Around the Corner", + Map.of("country", "Bulgaria")); + var nlDocument = new Document("The World is Big and Salvation Lurks Around the Corner", + Map.of("country", "Netherland")); + + vectorStore.add(List.of(bgDocument, nlDocument)); + + var request = SearchRequest.query("The World").withTopK(5); + + List results = vectorStore.similaritySearch(request); + assertThat(results).hasSize(2); + + results = vectorStore + .similaritySearch(request.withSimilarityThresholdAll().withFilterExpression("country == 'Bulgaria'")); + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo(bgDocument.getId()); + + results = vectorStore + .similaritySearch(request.withSimilarityThresholdAll().withFilterExpression("country == 'Netherland'")); + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo(nlDocument.getId()); + + // Remove all documents from the store + vectorStore.delete(List.of(bgDocument, nlDocument).stream().map(doc -> doc.getId()).toList()); + }); + } + + @Test + public void documentUpdateTest() { + + // Note ,using OpenAI to calculate embeddings + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + Document document = new Document(UUID.randomUUID().toString(), "Spring AI rocks!!", + Collections.singletonMap("meta1", "meta1")); + + vectorStore.add(List.of(document)); + + List results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(5)); + + assertThat(results).hasSize(1); + Document resultDoc = results.get(0); + assertThat(resultDoc.getId()).isEqualTo(document.getId()); + assertThat(resultDoc.getContent()).isEqualTo("Spring AI rocks!!"); + assertThat(resultDoc.getMetadata()).containsKey("meta1"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); + + Document sameIdDocument = new Document(document.getId(), + "The World is Big and Salvation Lurks Around the Corner", + Collections.singletonMap("meta2", "meta2")); + + vectorStore.add(List.of(sameIdDocument)); + + results = vectorStore.similaritySearch(SearchRequest.query("FooBar").withTopK(5)); + + assertThat(results).hasSize(1); + resultDoc = results.get(0); + assertThat(resultDoc.getId()).isEqualTo(document.getId()); + assertThat(resultDoc.getContent()).isEqualTo("The World is Big and Salvation Lurks Around the Corner"); + assertThat(resultDoc.getMetadata()).containsKey("meta2"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); + + // Remove all documents from the store + vectorStore.delete(List.of(document.getId())); + }); + } + + @Test + public void searchThresholdTest() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + vectorStore.add(documents); + + var request = SearchRequest.query("Great").withTopK(5); + List fullResult = vectorStore.similaritySearch(request.withSimilarityThresholdAll()); + + List distances = fullResult.stream().map(doc -> (Float) doc.getMetadata().get("distance")).toList(); + + assertThat(distances).hasSize(3); + + float threshold = (distances.get(0) + distances.get(1)) / 2; + + List results = vectorStore.similaritySearch(request.withSimilarityThreshold(1 - threshold)); + + assertThat(results).hasSize(1); + Document resultDoc = results.get(0); + assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId()); + assertThat(resultDoc.getContent()).isEqualTo( + "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); + assertThat(resultDoc.getMetadata()).containsKey("meta2"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); + + // Remove all documents from the store + vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList()); + }); + } + + @SpringBootConfiguration + public static class TestApplication { + + @Bean + public RestTemplate restTemplate() { + return new RestTemplate(); + } + + @Bean + public ChromaApi chromaApi(RestTemplate restTemplate) { + int port = chromaContainer.getMappedPort(8000); + return new ChromaApi("http://localhost:" + port, restTemplate); + } + + @Bean + public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + } + + @Bean + public EmbeddingClient embeddingClient() { + + Retrofit retrofit = new Retrofit.Builder().baseUrl("https://api.openai.com") + .client(OpenAiService.defaultClient(System.getenv("OPENAI_API_KEY"), Duration.ofSeconds(60))) + .addConverterFactory(JacksonConverterFactory.create(OpenAiService.defaultObjectMapper())) + .addCallAdapterFactory(RxJava2CallAdapterFactory.create()) + .build(); + + OpenAiApi api = retrofit.create(OpenAiApi.class); + + return new OpenAiEmbeddingClient(new OpenAiService(api), "text-embedding-ada-002"); + } + + } + +} diff --git a/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/TokenSecuredChromaWhereIT.java b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/TokenSecuredChromaWhereIT.java new file mode 100644 index 000000000..b7d7922f4 --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/java/org/springframework/experimental/ai/vectorstore/TokenSecuredChromaWhereIT.java @@ -0,0 +1,174 @@ +/* + * Copyright 2023-2023 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.experimental.ai.vectorstore; + +import java.time.Duration; +import java.util.List; +import java.util.Map; + +import com.theokanning.openai.client.OpenAiApi; +import com.theokanning.openai.service.OpenAiService; +import org.junit.jupiter.api.Test; +import org.testcontainers.containers.GenericContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; +import retrofit2.Retrofit; +import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory; +import retrofit2.converter.jackson.JacksonConverterFactory; + +import org.springframework.ai.document.Document; +import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; +import org.springframework.ai.vectorstore.SearchRequest; +import org.springframework.ai.vectorstore.VectorStore; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import org.springframework.context.annotation.Bean; +import org.springframework.experimental.ai.chroma.ChromaApi; +import org.springframework.experimental.ai.vectorsore.ChromaVectorStore; +import org.springframework.web.client.RestTemplate; + +import static org.assertj.core.api.Assertions.assertThat; + +/** + * ChromaDB with static API Token Authentication: + * https://docs.trychroma.com/usage-guide#static-api-token-authentication + * + * Test cases are based on the Chroma: + * https://docs.trychroma.com/usage-guide#using-where-filters and the related + * https://github.com/chroma-core/chroma/blob/main/examples/basic_functionality/in_not_in_filtering.ipynb + * + * @author Christian Tzolov + */ +@Testcontainers +public class TokenSecuredChromaWhereIT { + + public static String CHROMA_SERVER_URL = "http://localhost:"; + + public static String CHROMA_SERVER_AUTH_CREDENTIALS = "test-token"; + + /** + * ChromaDB with static API Token Authentication: + * https://docs.trychroma.com/usage-guide#static-api-token-authentication + */ + @Container + static GenericContainer chromaContainer = new GenericContainer<>("ghcr.io/chroma-core/chroma:0.4.15") + .withEnv("CHROMA_SERVER_AUTH_CREDENTIALS", CHROMA_SERVER_AUTH_CREDENTIALS) + .withEnv("CHROMA_SERVER_AUTH_CREDENTIALS_PROVIDER", + "chromadb.auth.token.TokenConfigServerAuthCredentialsProvider") + .withEnv("CHROMA_SERVER_AUTH_PROVIDER", "chromadb.auth.token.TokenAuthServerProvider") + + .withExposedPorts(8000); + + private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() + .withUserConfiguration(TestApplication.class) + .withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY")); + + @Test + public void withInFiltersExpressions1() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + vectorStore.add(List.of(new Document("1", "Article by john", Map.of("author", "john")), + new Document("2", "Article by Jack", Map.of("author", "jack")), + new Document("3", "Article by Jill", Map.of("author", "jill")))); + + var request = SearchRequest.query("Give me articles by john").withTopK(5); + + List results = vectorStore.similaritySearch(request); + assertThat(results).hasSize(3); + + results = vectorStore.similaritySearch( + request.withSimilarityThresholdAll().withFilterExpression("author in ['john', 'jill']")); + + assertThat(results).hasSize(2); + assertThat(results.stream().map(d -> d.getId()).toList()).containsExactlyInAnyOrder("1", "3"); + }); + } + + @Test + public void withInFiltersExpressions() { + + contextRunner.run(context -> { + + VectorStore vectorStore = context.getBean(VectorStore.class); + + vectorStore + .add(List.of(new Document("1", "Article by john", Map.of("author", "john", "article_type", "blog")), + new Document("2", "Article by Jack", Map.of("author", "jack", "article_type", "social")), + new Document("3", "Article by Jill", Map.of("author", "jill", "article_type", "paper")))); + + var request = SearchRequest.query("Give me articles by john").withTopK(5); + + List results = vectorStore.similaritySearch(request); + assertThat(results).hasSize(3); + + results = vectorStore.similaritySearch(request.withSimilarityThresholdAll() + .withFilterExpression("author in ['john', 'jill'] && 'article_type' == 'blog'")); + + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo("1"); + + results = vectorStore.similaritySearch(request.withSimilarityThresholdAll() + .withFilterExpression("author in ['john'] || 'article_type' == 'paper'")); + + assertThat(results).hasSize(2); + + assertThat(results.stream().map(d -> d.getId()).toList()).containsExactlyInAnyOrder("1", "3"); + }); + } + + @SpringBootConfiguration + public static class TestApplication { + + @Bean + public RestTemplate restTemplate() { + return new RestTemplate(); + } + + @Bean + public ChromaApi chromaApi(RestTemplate restTemplate) { + int port = chromaContainer.getMappedPort(8000); + var chromaApi = new ChromaApi(CHROMA_SERVER_URL + port, restTemplate); + chromaApi.withKeyToken(CHROMA_SERVER_AUTH_CREDENTIALS); + return chromaApi; + } + + @Bean + public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + } + + @Bean + public EmbeddingClient embeddingClient() { + + Retrofit retrofit = new Retrofit.Builder().baseUrl("https://api.openai.com") + .client(OpenAiService.defaultClient(System.getenv("OPENAI_API_KEY"), Duration.ofSeconds(60))) + .addConverterFactory(JacksonConverterFactory.create(OpenAiService.defaultObjectMapper())) + .addCallAdapterFactory(RxJava2CallAdapterFactory.create()) + .build(); + + OpenAiApi api = retrofit.create(OpenAiApi.class); + + return new OpenAiEmbeddingClient(new OpenAiService(api), "text-embedding-ada-002"); + } + + } + +} diff --git a/vector-stores/spring-ai-chroma/src/test/resources/api.yaml b/vector-stores/spring-ai-chroma/src/test/resources/api.yaml new file mode 100644 index 000000000..924c8a046 --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/resources/api.yaml @@ -0,0 +1,630 @@ +openapi: "3.0.0" +info: + title: FastAPI + version: 0.1.0 +paths: + /api/v1: + get: + summary: Root + operationId: root + responses: + '200': + description: Successful Response + content: + application/json: + schema: + additionalProperties: + type: integer + type: object + title: Response Root Api V1 Get + /api/v1/reset: + post: + summary: Reset + operationId: reset + responses: + '200': + description: Successful Response + content: + application/json: + schema: + type: boolean + title: Response Reset Api V1 Reset Post + /api/v1/version: + get: + summary: Version + operationId: version + responses: + '200': + description: Successful Response + content: + application/json: + schema: + type: string + title: Response Version Api V1 Version Get + /api/v1/heartbeat: + get: + summary: Heartbeat + operationId: heartbeat + responses: + '200': + description: Successful Response + content: + application/json: + schema: + additionalProperties: + type: number + type: object + title: Response Heartbeat Api V1 Heartbeat Get + /api/v1/raw_sql: + post: + summary: Raw Sql + operationId: raw_sql + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/RawSql' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections: + get: + summary: List Collections + operationId: list_collections + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + post: + summary: Create Collection + operationId: create_collection + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/CreateCollection' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/add: + post: + summary: Add + operationId: add + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/AddEmbedding' + required: true + responses: + '201': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/update: + post: + summary: Update + operationId: update + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/UpdateEmbedding' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/upsert: + post: + summary: Upsert + operationId: upsert + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/AddEmbedding' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/get: + post: + summary: Get + operationId: get + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/GetEmbedding' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} #TODO add actual GetResult Body + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/delete: + post: + summary: Delete + operationId: delete + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/DeleteEmbedding' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/count: + get: + summary: Count + operationId: count + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + responses: + '200': + description: Successful Response + content: + application/json: + schema: + type: integer + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}/query: + post: + summary: Get Nearest Neighbors + operationId: get_nearest_neighbors + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/QueryEmbedding' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_name}/create_index: + post: + summary: Create Index + operationId: create_index + parameters: + - required: true + schema: + type: string + title: Collection Name + name: collection_name + in: path + responses: + '200': + description: Successful Response + content: + application/json: + schema: + type: boolean + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_name}: + get: + summary: Get Collection + operationId: get_collection + parameters: + - required: true + schema: + type: string + title: Collection Name + name: collection_name + in: path + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + delete: + summary: Delete Collection + operationId: delete_collection + parameters: + - required: true + schema: + type: string + title: Collection Name + name: collection_name + in: path + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' + /api/v1/collections/{collection_id}: + put: + summary: Update Collection + operationId: update_collection + parameters: + - required: true + schema: + type: string + title: Collection Id + name: collection_id + in: path + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/UpdateCollection' + required: true + responses: + '200': + description: Successful Response + content: + application/json: + schema: {} + '422': + description: Validation Error + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPValidationError' +components: + schemas: + AddEmbedding: + properties: + embeddings: + items: {} + type: array + title: Embeddings + metadatas: + items: + type: object + additionalProperties: true + type: array + title: Metadatas + documents: + items: + type: string + type: array + title: Documents + ids: + items: + type: string + type: array + title: Ids + increment_index: + type: boolean + title: Increment Index + default: true + type: object + required: + - ids + title: AddEmbedding + CreateCollection: + properties: + name: + type: string + title: Name + metadata: + type: object + title: Metadata + get_or_create: + type: boolean + title: Get Or Create + default: false + type: object + required: + - name + title: CreateCollection + DeleteEmbedding: + properties: + ids: + items: + type: string + type: array + title: Ids + where: + type: object + title: Where + where_document: + type: object + title: Where Document + type: object + title: DeleteEmbedding + GetEmbedding: + properties: + ids: + items: + type: string + type: array + title: Ids + where: + type: object + title: Where + where_document: + type: object + title: Where Document + sort: + type: string + title: Sort + limit: + type: integer + title: Limit + offset: + type: integer + title: Offset + include: + items: + anyOf: + - type: string + enum: + - documents + - type: string + enum: + - embeddings + - type: string + enum: + - metadatas + - type: string + enum: + - distances + type: array + title: Include + default: + - metadatas + - documents + type: object + title: GetEmbedding + HTTPValidationError: + properties: + detail: + items: + $ref: '#/components/schemas/ValidationError' + type: array + title: Detail + type: object + title: HTTPValidationError + QueryEmbedding: + properties: + where: + type: object + title: Where + additionalProperties: true + default: {} + where_document: + type: object + title: Where Document + additionalProperties: true + default: {} + query_embeddings: + items: {} + type: array + additionalProperties: true + title: Query Embeddings + n_results: + type: integer + title: N Results + default: 10 + include: + items: + type: string + enum: + - documents + - embeddings + - metadatas + - distances + type: array + title: Include + default: + - metadatas + - documents + - distances + type: object + required: + - query_embeddings + title: QueryEmbedding + RawSql: + properties: + raw_sql: + type: string + title: Raw Sql + type: object + required: + - raw_sql + title: RawSql + UpdateCollection: + properties: + new_name: + type: string + title: New Name + new_metadata: + type: object + title: New Metadata + type: object + title: UpdateCollection + UpdateEmbedding: + properties: + embeddings: + items: {} + type: array + title: Embeddings + metadatas: + items: + type: object + type: array + title: Metadatas + documents: + items: + type: string + type: array + title: Documents + ids: + items: + type: string + type: array + title: Ids + increment_index: + type: boolean + title: Increment Index + default: true + type: object + required: + - ids + title: UpdateEmbedding + ValidationError: + properties: + loc: + items: + anyOf: + - type: string + - type: integer + type: array + title: Location + msg: + type: string + title: Message + type: + type: string + title: Error Type + type: object + required: + - loc + - msg + - type + title: ValidationError \ No newline at end of file diff --git a/vector-stores/spring-ai-chroma/src/test/resources/server.htpasswd b/vector-stores/spring-ai-chroma/src/test/resources/server.htpasswd new file mode 100644 index 000000000..ea052f5fc --- /dev/null +++ b/vector-stores/spring-ai-chroma/src/test/resources/server.htpasswd @@ -0,0 +1,2 @@ +admin:$2y$05$fM.6b629s3L6L8RcA.kqg.BmxEzwB9t4MpGux62MEXNMJ9M7w8CY2 + diff --git a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java index 0f15f8641..ff81cc1e4 100644 --- a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java +++ b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java @@ -54,6 +54,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.filter.converter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.MilvusFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; import org.springframework.util.Assert; @@ -88,7 +89,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { public static final List SEARCH_OUTPUT_FIELDS = Arrays.asList(DOC_ID_FIELD_NAME, CONTENT_FIELD_NAME, METADATA_FIELD_NAME); - public final MilvusFilterExpressionConverter filterExpressionConverter = new MilvusFilterExpressionConverter(); + public final FilterExpressionConverter filterExpressionConverter = new MilvusFilterExpressionConverter(); private final MilvusServiceClient milvusClient; @@ -324,7 +325,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { public List similaritySearch(SearchRequest request) { String nativeFilterExpressions = (request.getFilterExpression() != null) - ? this.filterExpressionConverter.convert(request.getFilterExpression()) : ""; + ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; Assert.notNull(request.getQuery(), "Query string must not be null"); 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 9739a5649..d755c2ede 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 @@ -33,6 +33,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.filter.converter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.PgVectorFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; import org.springframework.jdbc.core.JdbcTemplate; @@ -56,7 +57,7 @@ public class PgVectorStore implements VectorStore, InitializingBean { public static final String VECTOR_TABLE_NAME = "vector_store"; - public final PgVectorFilterExpressionConverter filterExpressionConverter = new PgVectorFilterExpressionConverter(); + public final FilterExpressionConverter filterExpressionConverter = new PgVectorFilterExpressionConverter(); private final JdbcTemplate jdbcTemplate; @@ -273,7 +274,7 @@ public class PgVectorStore implements VectorStore, InitializingBean { public List similaritySearch(SearchRequest request) { String nativeFilterExpression = (request.getFilterExpression() != null) - ? this.filterExpressionConverter.convert(request.getFilterExpression()) : ""; + ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; String jsonPathFilter = ""; 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 6e433937f..60ec87a9e 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 @@ -34,7 +34,6 @@ import org.junit.Assert; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.ValueSource; -import org.springframework.core.io.DefaultResourceLoader; import org.testcontainers.containers.GenericContainer; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; @@ -42,7 +41,6 @@ import retrofit2.Retrofit; import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory; import retrofit2.converter.jackson.JacksonConverterFactory; -import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; @@ -58,6 +56,7 @@ import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.ApplicationContext; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Primary; +import org.springframework.core.io.DefaultResourceLoader; import org.springframework.jdbc.core.JdbcTemplate; import org.springframework.util.CollectionUtils; @@ -145,7 +144,7 @@ public class PgVectorStoreIT { VectorStore vectorStore = context.getBean(VectorStore.class); var bgDocument = new Document("The World is Big and Salvation Lurks Around the Corner", - Map.of("country", "BG", "year", 2020)); + Map.of("country", "BG", "year", 2020, "foo bar 1", "bar.foo")); var nlDocument = new Document("The World is Big and Salvation Lurks Around the Corner", Map.of("country", "NL")); var bgDocument2 = new Document("The World is Big and Salvation Lurks Around the Corner", @@ -183,6 +182,13 @@ public class PgVectorStoreIT { assertThat(results.get(0).getId()).isIn(bgDocument.getId(), nlDocument.getId()); assertThat(results.get(1).getId()).isIn(bgDocument.getId(), nlDocument.getId()); + results = vectorStore.similaritySearch(SearchRequest.query("The World") + .withTopK(5) + .withSimilarityThresholdAll() + .withFilterExpression("\"foo bar 1\" == 'bar.foo'")); + assertThat(results).hasSize(1); + assertThat(results.get(0).getId()).isEqualTo(bgDocument.getId()); + try { vectorStore.similaritySearch(searchRequest.withFilterExpression("country == NL")); Assert.fail("Invalid filter expression should have been cached!"); diff --git a/vector-stores/spring-ai-pinecone/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java b/vector-stores/spring-ai-pinecone/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java index 49b851da7..b320d8df3 100644 --- a/vector-stores/spring-ai-pinecone/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java +++ b/vector-stores/spring-ai-pinecone/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java @@ -36,6 +36,8 @@ import io.pinecone.proto.Vector; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.vectorstore.filter.Filter; +import org.springframework.ai.vectorstore.filter.converter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.PineconeFilterExpressionConverter; import org.springframework.util.Assert; import org.springframework.util.StringUtils; @@ -53,7 +55,7 @@ public class PineconeVectorStore implements VectorStore { private static final String DISTANCE_METADATA_FIELD_NAME = "distance"; - public final PineconeFilterExpressionConverter filterExpressionConverter = new PineconeFilterExpressionConverter(); + public final FilterExpressionConverter filterExpressionConverter = new PineconeFilterExpressionConverter(); private final EmbeddingClient embeddingClient; @@ -307,7 +309,7 @@ public class PineconeVectorStore implements VectorStore { // Filter.Expression filterExpression) { String nativeExpressionFilters = (request.getFilterExpression() != null) - ? this.filterExpressionConverter.convert(request.getFilterExpression()) : ""; + ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; List queryEmbedding = this.embeddingClient.embed(request.getQuery());