diff --git a/vector-stores/spring-ai-mariadb-store/src/main/java/org/springframework/ai/vectorstore/mariadb/MariaDBVectorStore.java b/vector-stores/spring-ai-mariadb-store/src/main/java/org/springframework/ai/vectorstore/mariadb/MariaDBVectorStore.java index 387e53bd5..90ab43ec8 100644 --- a/vector-stores/spring-ai-mariadb-store/src/main/java/org/springframework/ai/vectorstore/mariadb/MariaDBVectorStore.java +++ b/vector-stores/spring-ai-mariadb-store/src/main/java/org/springframework/ai/vectorstore/mariadb/MariaDBVectorStore.java @@ -51,6 +51,7 @@ import org.springframework.util.StringUtils; * vector index will be auto-created if not available. * * @author Diego Dupin + * @author Ilayaperumal Gopinathan * @since 1.0.0 */ public class MariaDBVectorStore extends AbstractObservationVectorStore implements InitializingBean { @@ -192,21 +193,36 @@ public class MariaDBVectorStore extends AbstractObservationVectorStore implement @Override public void doAdd(List documents) { // Batch the documents based on the batching strategy - this.embeddingModel.embed(documents, EmbeddingOptionsBuilder.builder().build(), this.batchingStrategy); + List embeddings = this.embeddingModel.embed(documents, EmbeddingOptionsBuilder.builder().build(), + this.batchingStrategy); - List> batchedDocuments = batchDocuments(documents); + List> batchedDocuments = batchDocuments(documents, embeddings); batchedDocuments.forEach(this::insertOrUpdateBatch); } - private List> batchDocuments(List documents) { - List> batches = new ArrayList<>(); - for (int i = 0; i < documents.size(); i += this.maxDocumentBatchSize) { - batches.add(documents.subList(i, Math.min(i + this.maxDocumentBatchSize, documents.size()))); + private List> batchDocuments(List documents, List embeddings) { + List> batches = new ArrayList<>(); + List mariaDBDocuments = new ArrayList<>(documents.size()); + if (embeddings.size() == documents.size()) { + for (Document document : documents) { + mariaDBDocuments.add(new MariaDBDocument(document.getId(), document.getContent(), + document.getMetadata(), embeddings.get(documents.indexOf(document)))); + } + } + else { + for (Document document : documents) { + mariaDBDocuments + .add(new MariaDBDocument(document.getId(), document.getContent(), document.getMetadata(), null)); + } + } + + for (int i = 0; i < mariaDBDocuments.size(); i += this.maxDocumentBatchSize) { + batches.add(mariaDBDocuments.subList(i, Math.min(i + this.maxDocumentBatchSize, mariaDBDocuments.size()))); } return batches; } - private void insertOrUpdateBatch(List batch) { + private void insertOrUpdateBatch(List batch) { String sql = String.format( "INSERT INTO %s (%s, %s, %s, %s) VALUES (?, ?, ?, ?) " + "ON DUPLICATE KEY UPDATE %s = VALUES(%s) , %s = VALUES(%s) , %s = VALUES(%s)", @@ -219,10 +235,10 @@ public class MariaDBVectorStore extends AbstractObservationVectorStore implement @Override public void setValues(PreparedStatement ps, int i) throws SQLException { var document = batch.get(i); - ps.setObject(1, document.getId()); - ps.setString(2, document.getContent()); - ps.setString(3, toJson(document.getMetadata())); - ps.setObject(4, document.getEmbedding()); + ps.setObject(1, document.id()); + ps.setString(2, document.content()); + ps.setString(3, toJson(document.metadata())); + ps.setObject(4, document.embedding()); } @Override @@ -556,4 +572,15 @@ public class MariaDBVectorStore extends AbstractObservationVectorStore implement } + /** + * The representation of {@link Document} along with its embedding. + * + * @param id The id of the document + * @param content The content of the document + * @param metadata The metadata of the document + * @param embedding The vectors representing the content of the document + */ + public record MariaDBDocument(String id, String content, Map metadata, float[] embedding) { + } + } diff --git a/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreIT.java b/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreIT.java index 47bcdff99..2a49b665c 100644 --- a/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreIT.java +++ b/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreIT.java @@ -65,7 +65,6 @@ import org.testcontainers.junit.jupiter.Testcontainers; */ @Testcontainers @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -@Disabled("Failing after commit ebd29e0") public class MariaDBStoreIT { private static String schemaName = "testdb"; diff --git a/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreObservationIT.java b/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreObservationIT.java index 149f58f41..435a92c6e 100644 --- a/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreObservationIT.java +++ b/vector-stores/spring-ai-mariadb-store/src/test/java/org/springframework/ai/vectorstore/mariadb/MariaDBStoreObservationIT.java @@ -64,7 +64,6 @@ import org.testcontainers.junit.jupiter.Testcontainers; */ @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") @Testcontainers -@Disabled("Failing after commit ebd29e0") public class MariaDBStoreObservationIT { private static String schemaName = "testdb";