From caa8cb2021c0bfca267e1d9d056de49fb0cfd703 Mon Sep 17 00:00:00 2001 From: Toshiaki Maki Date: Thu, 28 Sep 2023 15:13:15 +0900 Subject: [PATCH] Add PostgresML support for EmbeddingClient - Implement dimensions method - Add MetadataMode support. Defaults to EMBED - Drop the pgml extension between tests. - Disable the PostgresMlEmbeddingClientIT by default. Resolves #33 --- .../pom.xml | 74 +++++++ .../embedding/PostgresMlEmbeddingClient.java | 181 ++++++++++++++++++ .../PostgresMlEmbeddingClientIT.java | 137 +++++++++++++ pom.xml | 1 + 4 files changed, 393 insertions(+) create mode 100644 embedding-clients/spring-ai-postgresml-embedding-client/pom.xml create mode 100644 embedding-clients/spring-ai-postgresml-embedding-client/src/main/java/org/springframework/ai/embedding/PostgresMlEmbeddingClient.java create mode 100644 embedding-clients/spring-ai-postgresml-embedding-client/src/test/java/org/springframework/ai/embedding/PostgresMlEmbeddingClientIT.java diff --git a/embedding-clients/spring-ai-postgresml-embedding-client/pom.xml b/embedding-clients/spring-ai-postgresml-embedding-client/pom.xml new file mode 100644 index 000000000..e6ef6f0ff --- /dev/null +++ b/embedding-clients/spring-ai-postgresml-embedding-client/pom.xml @@ -0,0 +1,74 @@ + + + 4.0.0 + + org.springframework.experimental.ai + spring-ai + 0.7.0-SNAPSHOT + ../../pom.xml + + spring-ai-postgresml-embedding-client + jar + Spring AI Embedding Client - PostgresML + Spring AI PostgresML Embedding Client + 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.postgresql + postgresql + runtime + + + + org.springframework + spring-jdbc + + + + + org.springframework.boot + spring-boot-starter-test + test + + + + org.springframework.boot + spring-boot-testcontainers + test + + + + org.testcontainers + junit-jupiter + test + + + + org.testcontainers + postgresql + test + + + + com.zaxxer + HikariCP + test + + + + + diff --git a/embedding-clients/spring-ai-postgresml-embedding-client/src/main/java/org/springframework/ai/embedding/PostgresMlEmbeddingClient.java b/embedding-clients/spring-ai-postgresml-embedding-client/src/main/java/org/springframework/ai/embedding/PostgresMlEmbeddingClient.java new file mode 100644 index 000000000..1f4907b73 --- /dev/null +++ b/embedding-clients/spring-ai-postgresml-embedding-client/src/main/java/org/springframework/ai/embedding/PostgresMlEmbeddingClient.java @@ -0,0 +1,181 @@ +package org.springframework.ai.embedding; + +import java.sql.Array; +import java.sql.PreparedStatement; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.ObjectMapper; + +import org.springframework.ai.document.Document; +import org.springframework.ai.document.MetadataMode; +import org.springframework.beans.factory.InitializingBean; +import org.springframework.jdbc.core.JdbcTemplate; +import org.springframework.jdbc.core.RowMapper; +import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; +import org.springframework.util.StringUtils; + +/** + * PostgresML EmbeddingClient + * + * @author Toshiaki Maki + */ +public class PostgresMlEmbeddingClient implements EmbeddingClient, InitializingBean { + + private final JdbcTemplate jdbcTemplate; + + private final String transformer; + + private final VectorType vectorType; + + private final String kwargs; + + private final AtomicInteger embeddingDimensions = new AtomicInteger(-1); + + private final MetadataMode metadataMode; + + public enum VectorType { + + PG_ARRAY("", null, (rs, i) -> { + Array embedding = rs.getArray("embedding"); + return Arrays.stream((Float[]) embedding.getArray()).map(Float::doubleValue).toList(); + }), PG_VECTOR("::vector", "vector", (rs, i) -> { + String embedding = rs.getString("embedding"); + return Arrays.stream((embedding.substring(1, embedding.length() - 1) + /* remove leading '[' and trailing ']' */.split(","))).map(Double::parseDouble).toList(); + }); + + private final String cast; + + private final String extensionName; + + private final RowMapper> rowMapper; + + VectorType(String cast, String extensionName, RowMapper> rowMapper) { + this.cast = cast; + this.extensionName = extensionName; + this.rowMapper = rowMapper; + } + + } + + /** + * a constructor + * @param jdbcTemplate JdbcTemplate + */ + public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate) { + this(jdbcTemplate, "distilbert-base-uncased"); + } + + /** + * a constructor + * @param jdbcTemplate JdbcTemplate + * @param transformer huggingface sentence-transformer name + */ + public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer) { + this(jdbcTemplate, transformer, VectorType.PG_ARRAY); + } + + /** + * a constructor + * @param jdbcTemplate JdbcTemplate + * @param transformer huggingface sentence-transformer name + * @param vectorType vector type in PostgreSQL + */ + public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType) { + this(jdbcTemplate, transformer, vectorType, Map.of(), MetadataMode.EMBED); + } + + /** + * a constructor + * @param jdbcTemplate JdbcTemplate + * @param transformer huggingface sentence-transformer name + * @param vectorType vector type in PostgreSQL + * @param kwargs optional arguments + */ + public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType, + Map kwargs, MetadataMode metadataMode) { + Assert.notNull(jdbcTemplate, "jdbc template must not be null."); + Assert.notNull(transformer, "transformer must not be null."); + Assert.notNull(vectorType, "vectorType must not be null."); + Assert.notNull(kwargs, "kwargs must not be null."); + Assert.notNull(metadataMode, "metadataMode must not be null."); + + this.jdbcTemplate = jdbcTemplate; + this.transformer = transformer; + this.vectorType = vectorType; + this.metadataMode = metadataMode; + try { + this.kwargs = new ObjectMapper().writeValueAsString(kwargs); + } + catch (JsonProcessingException e) { + throw new IllegalArgumentException(e); + } + } + + @Override + public List embed(String text) { + return this.jdbcTemplate.queryForObject( + "SELECT pgml.embed(?, ?, ?::JSONB)" + this.vectorType.cast + " AS embedding", this.vectorType.rowMapper, + this.transformer, text, this.kwargs); + } + + @Override + public List embed(Document document) { + return this.embed(document.getFormattedContent(this.metadataMode)); + } + + @Override + public List> embed(List texts) { + if (CollectionUtils.isEmpty(texts)) { + return List.of(); + } + return this.jdbcTemplate.query(connection -> { + PreparedStatement preparedStatement = connection.prepareStatement("SELECT pgml.embed(?, text, ?::JSONB)" + + vectorType.cast + " AS embedding FROM (SELECT unnest(?) AS text) AS texts"); + preparedStatement.setString(1, transformer); + preparedStatement.setString(2, kwargs); + preparedStatement.setArray(3, connection.createArrayOf("TEXT", texts.toArray(Object[]::new))); + return preparedStatement; + }, rs -> { + List> result = new ArrayList<>(); + while (rs.next()) { + result.add(vectorType.rowMapper.mapRow(rs, -1)); + } + return result; + }); + } + + @Override + public EmbeddingResponse embedForResponse(List texts) { + List data = new ArrayList<>(); + List> embed = this.embed(texts); + for (int i = 0; i < embed.size(); i++) { + data.add(new Embedding(embed.get(i), i)); + } + return new EmbeddingResponse(data, + Map.of("transformer", this.transformer, "vector-type", this.vectorType.name(), "kwargs", this.kwargs)); + } + + @Override + public int dimensions() { + if (this.embeddingDimensions.get() < 0) { + this.embeddingDimensions.set(EmbeddingUtil.dimensions(this, this.transformer)); + } + return this.embeddingDimensions.get(); + } + + @Override + public void afterPropertiesSet() { + this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS pgml"); + if (StringUtils.hasText(this.vectorType.extensionName)) { + this.jdbcTemplate.execute("CREATE EXTENSION IF NOT EXISTS " + this.vectorType.extensionName); + } + } + +} diff --git a/embedding-clients/spring-ai-postgresml-embedding-client/src/test/java/org/springframework/ai/embedding/PostgresMlEmbeddingClientIT.java b/embedding-clients/spring-ai-postgresml-embedding-client/src/test/java/org/springframework/ai/embedding/PostgresMlEmbeddingClientIT.java new file mode 100644 index 000000000..e4324a3ed --- /dev/null +++ b/embedding-clients/spring-ai-postgresml-embedding-client/src/test/java/org/springframework/ai/embedding/PostgresMlEmbeddingClientIT.java @@ -0,0 +1,137 @@ +package org.springframework.ai.embedding; + +import java.time.Duration; +import java.time.temporal.ChronoUnit; +import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import org.testcontainers.containers.PostgreSQLContainer; +import org.testcontainers.containers.wait.strategy.LogMessageWaitStrategy; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; +import org.testcontainers.utility.DockerImageName; + +import org.springframework.ai.document.Document; +import org.springframework.ai.document.MetadataMode; +import org.springframework.ai.embedding.PostgresMlEmbeddingClient.VectorType; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.autoconfigure.SpringBootApplication; +import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase; +import org.springframework.boot.test.autoconfigure.jdbc.JdbcTest; +import org.springframework.boot.testcontainers.service.connection.ServiceConnection; +import org.springframework.jdbc.core.JdbcTemplate; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.ai.embedding.PostgresMlEmbeddingClient.VectorType.PG_ARRAY; +import static org.springframework.ai.embedding.PostgresMlEmbeddingClient.VectorType.PG_VECTOR; + +/** + * @author Toshiaki Maki + */ + +@JdbcTest(properties = "logging.level.sql=TRACE") +@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE) +@Testcontainers +@Disabled("Disabled from automatic execution, as it requires an excessive amount of memory (over 9GB)!") +class PostgresMlEmbeddingClientIT { + + @Container + @ServiceConnection + static PostgreSQLContainer postgres = new PostgreSQLContainer<>( + DockerImageName.parse("ghcr.io/postgresml/postgresml:2.7.3").asCompatibleSubstituteFor("postgres")) + .withCommand("sleep", "infinity") + .withLabel("org.springframework.boot.service-connection", "postgres") + .withUsername("postgresml") + .withPassword("postgresml") + .withDatabaseName("postgresml") + .waitingFor(new LogMessageWaitStrategy().withRegEx(".*Starting dashboard.*\\s") + .withStartupTimeout(Duration.of(60, ChronoUnit.SECONDS))); + + @Autowired + JdbcTemplate jdbcTemplate; + + @AfterEach + void dropPgmlExtension() { + this.jdbcTemplate.execute("DROP EXTENSION IF EXISTS pgml"); + } + + @Test + void embed() { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate); + embeddingClient.afterPropertiesSet(); + List embed = embeddingClient.embed("Hello World!"); + assertThat(embed).hasSize(768); + // embeddingClient.dropPgmlExtension(); + } + + @Test + void embedWithPgVector() { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + "distilbert-base-uncased", PG_VECTOR); + embeddingClient.afterPropertiesSet(); + List embed = embeddingClient.embed(new Document("Hello World!")); + assertThat(embed).hasSize(768); + // embeddingClient.dropPgmlExtension(); + } + + @Test + void embedWithDifferentModel() { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + "intfloat/e5-small"); + embeddingClient.afterPropertiesSet(); + List embed = embeddingClient.embed(new Document("Hello World!")); + assertThat(embed).hasSize(384); + // embeddingClient.dropPgmlExtension(); + } + + @Test + void embedWithKwargs() { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + "distilbert-base-uncased", PG_ARRAY, Map.of("device", "cpu"), MetadataMode.EMBED); + embeddingClient.afterPropertiesSet(); + List embed = embeddingClient.embed(new Document("Hello World!")); + assertThat(embed).hasSize(768); + // embeddingClient.dropPgmlExtension(); + } + + @ParameterizedTest + @ValueSource(strings = { "PG_ARRAY", "PG_VECTOR" }) + void embedForResponse(String vectorType) { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + "distilbert-base-uncased", VectorType.valueOf(vectorType)); + embeddingClient.afterPropertiesSet(); + EmbeddingResponse embeddingResponse = embeddingClient + .embedForResponse(List.of("Hello World!", "Spring AI!", "LLM!")); + assertThat(embeddingResponse).isNotNull(); + assertThat(embeddingResponse.getData()).hasSize(3); + assertThat(embeddingResponse.getMetadata()).containsExactlyEntriesOf( + Map.of("transformer", "distilbert-base-uncased", "vector-type", vectorType, "kwargs", "{}")); + assertThat(embeddingResponse.getData().get(0).getIndex()).isEqualTo(0); + assertThat(embeddingResponse.getData().get(0).getEmbedding()).hasSize(768); + assertThat(embeddingResponse.getData().get(1).getIndex()).isEqualTo(1); + assertThat(embeddingResponse.getData().get(1).getEmbedding()).hasSize(768); + assertThat(embeddingResponse.getData().get(2).getIndex()).isEqualTo(2); + assertThat(embeddingResponse.getData().get(2).getEmbedding()).hasSize(768); + // embeddingClient.dropPgmlExtension(); + } + + @Test + void dimensions() { + PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate); + embeddingClient.afterPropertiesSet(); + assertThat(embeddingClient.dimensions()).isEqualTo(768); + // cached + assertThat(embeddingClient.dimensions()).isEqualTo(768); + } + + @SpringBootApplication + public static class TestApplication { + + } + +} \ No newline at end of file diff --git a/pom.xml b/pom.xml index 28548238c..628d31dcc 100644 --- a/pom.xml +++ b/pom.xml @@ -22,6 +22,7 @@ vector-stores/spring-ai-pgvector-store vector-stores/spring-ai-milvus-store vector-stores/spring-ai-neo4j-store + embedding-clients/spring-ai-postgresml-embedding-client