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