Neo4j module: Determine default embedding dimension from model.

In cases where no custom size is set, derive the size by the given
embedding model.

Had to migrate the embeddingDimension Spring Boot property from int to Integer
to introduce the null check for the fluent config. Everything else would
have been just noisy.

Auto-cherry-pick to 1.0.x

Fixes #977

Signed-off-by: Gerrit Meier <meistermeier@gmail.com>
This commit is contained in:
Gerrit Meier
2025-01-15 16:09:37 +01:00
committed by Ilayaperumal Gopinathan
parent 5f23dcade3
commit cb97d9c4e2
4 changed files with 36 additions and 6 deletions

View File

@@ -68,7 +68,8 @@ public class Neo4jVectorStoreAutoConfiguration {
.customObservationConvention(customObservationConvention.getIfAvailable(() -> null))
.batchingStrategy(batchingStrategy)
.databaseName(properties.getDatabaseName())
.embeddingDimension(properties.getEmbeddingDimension())
.embeddingDimension(properties.getEmbeddingDimension() != null ? properties.getEmbeddingDimension()
: embeddingModel.dimensions())
.distanceType(properties.getDistanceType())
.label(properties.getLabel())
.embeddingProperty(properties.getEmbeddingProperty())

View File

@@ -33,7 +33,7 @@ public class Neo4jVectorStoreProperties extends CommonVectorStoreProperties {
private String databaseName;
private int embeddingDimension = Neo4jVectorStore.DEFAULT_EMBEDDING_DIMENSION;
private Integer embeddingDimension;
private Neo4jVectorStore.Neo4jDistanceType distanceType = Neo4jVectorStore.Neo4jDistanceType.COSINE;
@@ -57,7 +57,7 @@ public class Neo4jVectorStoreProperties extends CommonVectorStoreProperties {
this.databaseName = databaseName;
}
public int getEmbeddingDimension() {
public Integer getEmbeddingDimension() {
return this.embeddingDimension;
}

View File

@@ -136,6 +136,7 @@ public class Neo4jVectorStore extends AbstractObservationVectorStore implements
private static final Logger logger = LoggerFactory.getLogger(Neo4jVectorStore.class);
@Deprecated(forRemoval = true)
public static final int DEFAULT_EMBEDDING_DIMENSION = 1536;
public static final int DEFAULT_TRANSACTION_SIZE = 10_000;
@@ -189,7 +190,7 @@ public class Neo4jVectorStore extends AbstractObservationVectorStore implements
this.driver = builder.driver;
this.sessionConfig = builder.sessionConfig;
this.embeddingDimension = builder.embeddingDimension;
this.embeddingDimension = builder.embeddingDimension.orElseGet(() -> builder.getEmbeddingModel().dimensions());
this.distanceType = builder.distanceType;
this.embeddingProperty = SchemaNames.sanitize(builder.embeddingProperty).orElseThrow();
this.label = SchemaNames.sanitize(builder.label).orElseThrow();
@@ -404,7 +405,7 @@ public class Neo4jVectorStore extends AbstractObservationVectorStore implements
private SessionConfig sessionConfig = SessionConfig.defaultConfig();
private int embeddingDimension = DEFAULT_EMBEDDING_DIMENSION;
private Optional<Integer> embeddingDimension = Optional.empty();
private Neo4jDistanceType distanceType = Neo4jDistanceType.COSINE;
@@ -459,7 +460,7 @@ public class Neo4jVectorStore extends AbstractObservationVectorStore implements
*/
public Builder embeddingDimension(int dimension) {
Assert.isTrue(dimension >= 1, "Dimension has to be positive");
this.embeddingDimension = dimension;
this.embeddingDimension = Optional.of(dimension);
return this;
}

View File

@@ -31,6 +31,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.neo4j.driver.AuthTokens;
import org.neo4j.driver.Driver;
import org.neo4j.driver.GraphDatabase;
import org.springframework.context.annotation.Primary;
import org.testcontainers.containers.Neo4jContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
@@ -356,16 +357,43 @@ class Neo4jVectorStoreIT extends BaseVectorStoreTests {
});
}
@Test
void vectorIndexDimensionsDefaultAndOverwriteWorks() {
this.contextRunner.run(context -> {
var result = context.getBean(Driver.class)
.executableQuery(
"SHOW VECTOR INDEXES yield name, options return name, options['indexConfig']['vector.dimensions'] as dimensions")
.execute()
.records()
.stream()
.map(r -> r.get("name").asString() + r.get("dimensions").asInt())
.toList();
assertThat(result).containsExactlyInAnyOrder("secondIndex123", "spring-ai-document-index1536");
});
}
@SpringBootConfiguration
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class TestApplication {
@Bean
@Primary
public VectorStore vectorStore(Driver driver, EmbeddingModel embeddingModel) {
return Neo4jVectorStore.builder(driver, embeddingModel).initializeSchema(true).build();
}
@Bean
public VectorStore vectorStoreWithCustomDimension(Driver driver, EmbeddingModel embeddingModel) {
return Neo4jVectorStore.builder(driver, embeddingModel)
.initializeSchema(true)
.indexName("secondIndex")
.embeddingProperty("somethingElse")
.embeddingDimension(123)
.build();
}
@Bean
public Driver driver() {
return GraphDatabase.driver(neo4jContainer.getBoltUrl(),