Use jdbcTemplate.batchUpdate for PgVectorStore add() implementation

- remove unused variable
This commit is contained in:
Giani Segatto
2023-12-28 22:57:50 +00:00
committed by Mark Pollack
parent a892ea961f
commit 2a604262d8

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.vectorstore;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.SQLException;
import java.util.List;
@@ -36,8 +37,11 @@ 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.BatchPreparedStatementSetter;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.core.RowMapper;
import org.springframework.jdbc.core.SqlTypeValue;
import org.springframework.jdbc.core.StatementCreatorUtils;
import org.springframework.lang.Nullable;
import org.springframework.util.StringUtils;
@@ -176,12 +180,7 @@ public class PgVectorStore implements VectorStore, InitializingBean {
private List<Double> toDoubleList(PGobject embedding) throws SQLException {
float[] floatArray = new PGvector(embedding.getValue()).toArray();
List<Double> doubleEmbedding = IntStream.range(0, floatArray.length)
.mapToDouble(i -> floatArray[i])
.boxed()
.toList();
return doubleEmbedding;
return IntStream.range(0, floatArray.length).mapToDouble(i -> floatArray[i]).boxed().toList();
}
private Map<String, Object> toMap(PGobject pgObject) {
@@ -224,21 +223,36 @@ public class PgVectorStore implements VectorStore, InitializingBean {
@Override
public void add(List<Document> documents) {
for (Document document : documents) {
List<Double> embedding = this.embeddingClient.embed(document);
document.setEmbedding(embedding);
UUID id = UUID.fromString(document.getId());
String content = document.getContent();
Map<String, Object> metadata = document.getMetadata();
PGvector pgEmbedding = new PGvector(toFloatArray(embedding));
int size = documents.size();
this.jdbcTemplate.update(
"INSERT INTO " + VECTOR_TABLE_NAME
+ " (id, content, metadata, embedding) VALUES (?, ?, ?::jsonb, ?) " + "ON CONFLICT (id) DO "
+ "UPDATE SET content = ? , metadata = ?::jsonb , embedding = ? ",
id, content, toJson(metadata), pgEmbedding, content, toJson(metadata), pgEmbedding);
}
this.jdbcTemplate.batchUpdate(
"INSERT INTO " + VECTOR_TABLE_NAME + " (id, content, metadata, embedding) VALUES (?, ?, ?::jsonb, ?) "
+ "ON CONFLICT (id) DO " + "UPDATE SET content = ? , metadata = ?::jsonb , embedding = ? ",
new BatchPreparedStatementSetter() {
@Override
public void setValues(PreparedStatement ps, int i) throws SQLException {
var document = documents.get(i);
var content = document.getContent();
var json = toJson(document.getMetadata());
var pGvector = new PGvector(toFloatArray(embeddingClient.embed(document)));
StatementCreatorUtils.setParameterValue(ps, 1, SqlTypeValue.TYPE_UNKNOWN,
UUID.fromString(document.getId()));
StatementCreatorUtils.setParameterValue(ps, 2, SqlTypeValue.TYPE_UNKNOWN, content);
StatementCreatorUtils.setParameterValue(ps, 3, SqlTypeValue.TYPE_UNKNOWN, json);
StatementCreatorUtils.setParameterValue(ps, 4, SqlTypeValue.TYPE_UNKNOWN, pGvector);
StatementCreatorUtils.setParameterValue(ps, 5, SqlTypeValue.TYPE_UNKNOWN, content);
StatementCreatorUtils.setParameterValue(ps, 6, SqlTypeValue.TYPE_UNKNOWN, json);
StatementCreatorUtils.setParameterValue(ps, 7, SqlTypeValue.TYPE_UNKNOWN, pGvector);
}
@Override
public int getBatchSize() {
return size;
}
});
}
private String toJson(Map<String, Object> map) {