Add PgIdType based schema generation for PgVectorStore

Signed-off-by: CChuYong <yeongmin1061@gmail.com>
This commit is contained in:
CChuYong
2025-03-14 00:01:29 +09:00
committed by Ilayaperumal Gopinathan
parent 49df62533d
commit c02430ea6b
2 changed files with 45 additions and 29 deletions

View File

@@ -420,7 +420,10 @@ public class PgVectorStore extends AbstractObservationVectorStore implements Ini
// Enable the PGVector, JSONB and UUID support.
this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS vector");
this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS hstore");
this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\"");
if (this.idType == PgIdType.UUID) {
this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\"");
}
this.jdbcTemplate.execute(String.format("CREATE SCHEMA IF NOT EXISTS %s", this.getSchemaName()));
@@ -431,12 +434,12 @@ public class PgVectorStore extends AbstractObservationVectorStore implements Ini
this.jdbcTemplate.execute(String.format("""
CREATE TABLE IF NOT EXISTS %s (
id uuid DEFAULT uuid_generate_v4() PRIMARY KEY,
id %s PRIMARY KEY,
content text,
metadata json,
embedding vector(%d)
)
""", this.getFullyQualifiedTableName(), this.embeddingDimensions()));
""", this.getFullyQualifiedTableName(), this.getColumnTypeName(), this.embeddingDimensions()));
if (this.createIndexMethod != PgIndexType.NONE) {
this.jdbcTemplate.execute(String.format("""
@@ -466,6 +469,16 @@ public class PgVectorStore extends AbstractObservationVectorStore implements Ini
return this.vectorIndexName;
}
private String getColumnTypeName() {
return switch (getIdType()) {
case UUID -> "uuid DEFAULT uuid_generate_v4()";
case TEXT -> "text";
case INTEGER -> "integer";
case SERIAL -> "serial";
case BIGSERIAL -> "bigserial";
};
}
int embeddingDimensions() {
// The manually set dimensions have precedence over the computed one.
if (this.dimensions > 0) {

View File

@@ -112,27 +112,6 @@ public class PgVectorStoreIT extends BaseVectorStoreTests {
}
}
private static void initSchema(ApplicationContext context) {
PgVectorStore vectorStore = context.getBean(PgVectorStore.class);
JdbcTemplate jdbcTemplate = context.getBean(JdbcTemplate.class);
// Enable the PGVector, JSONB and UUID support.
jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS vector");
jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS hstore");
jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS \"uuid-ossp\"");
jdbcTemplate.execute(String.format("CREATE SCHEMA IF NOT EXISTS %s", PgVectorStore.DEFAULT_SCHEMA_NAME));
jdbcTemplate.execute(String.format("""
CREATE TABLE IF NOT EXISTS %s.%s (
id text PRIMARY KEY,
content text,
metadata json,
embedding vector(%d)
)
""", PgVectorStore.DEFAULT_SCHEMA_NAME, PgVectorStore.DEFAULT_TABLE_NAME,
vectorStore.embeddingDimensions()));
}
private static void dropTable(ApplicationContext context) {
JdbcTemplate jdbcTemplate = context.getBean(JdbcTemplate.class);
jdbcTemplate.execute("DROP TABLE IF EXISTS vector_store");
@@ -218,14 +197,12 @@ public class PgVectorStoreIT extends BaseVectorStoreTests {
}
@Test
public void testToPgTypeWithNonUuidIdType() {
public void testToPgTypeWithTextIdType() {
this.contextRunner.withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + "COSINE_DISTANCE")
.withPropertyValues("test.spring.ai.vectorstore.pgvector.initializeSchema=" + false)
.withPropertyValues("test.spring.ai.vectorstore.pgvector.idType=" + "TEXT")
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
initSchema(context);
vectorStore.add(List.of(new Document("NOT_UUID", "TEXT", new HashMap<>())));
@@ -233,6 +210,34 @@ public class PgVectorStoreIT extends BaseVectorStoreTests {
});
}
@Test
public void testToPgTypeWithSerialIdType() {
this.contextRunner.withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + "COSINE_DISTANCE")
.withPropertyValues("test.spring.ai.vectorstore.pgvector.idType=" + "SERIAL")
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.add(List.of(new Document("1", "TEXT", new HashMap<>())));
dropTable(context);
});
}
@Test
public void testToPgTypeWithBigSerialIdType() {
this.contextRunner.withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + "COSINE_DISTANCE")
.withPropertyValues("test.spring.ai.vectorstore.pgvector.idType=" + "BIGSERIAL")
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.add(List.of(new Document("1", "TEXT", new HashMap<>())));
dropTable(context);
});
}
@Test
public void testBulkOperationWithUuidIdType() {
this.contextRunner.withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + "COSINE_DISTANCE")
@@ -256,11 +261,9 @@ public class PgVectorStoreIT extends BaseVectorStoreTests {
@Test
public void testBulkOperationWithNonUuidIdType() {
this.contextRunner.withPropertyValues("test.spring.ai.vectorstore.pgvector.distanceType=" + "COSINE_DISTANCE")
.withPropertyValues("test.spring.ai.vectorstore.pgvector.initializeSchema=" + false)
.withPropertyValues("test.spring.ai.vectorstore.pgvector.idType=" + "TEXT")
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
initSchema(context);
List<Document> documents = List.of(new Document("NON_UUID_1", "TEXT", new HashMap<>()),
new Document("NON_UUID_2", "TEXT", new HashMap<>()),