embeddings = vertexAiApi.batchEmbedText(List.of("Hello, how are you?", "I am fine, thank you!"));
+```
diff --git a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/api/VertexAiApi.java b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/api/VertexAiApi.java
index c60353589..1d30a8842 100644
--- a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/api/VertexAiApi.java
+++ b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/api/VertexAiApi.java
@@ -49,7 +49,7 @@ import org.springframework.web.client.RestClient;
* Supported models:
*
*
- * name=models/chat-bison-001,
+ * name=models/chat-bison-001,
* version=001,
* displayName=Chat Bison,
* description=Chat-optimized generative language model.,
@@ -60,7 +60,7 @@ import org.springframework.web.client.RestClient;
* topP=0.95,
* topK=40
*
- * name=models/text-bison-001,
+ * name=models/text-bison-001,
* version=001,
* displayName=Text Bison,
* description=Model targeted for text generation.,
@@ -71,7 +71,7 @@ import org.springframework.web.client.RestClient;
* topP=0.95,
* topK=40
*
- * name=models/embedding-gecko-001,
+ * name=models/embedding-gecko-001,
* version=001,
* displayName=Embedding Gecko, description=Obtain a distributed representation of a text.,
* inputTokenLimit=1024,
@@ -98,7 +98,10 @@ public class VertexAiApi {
*/
public static final String DEFAULT_EMBEDDING_MODEL = "embedding-gecko-001";
- private static final String DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com/v1beta3";
+ /**
+ * The default base URL for accessing the Vertex AI API.
+ */
+ public static final String DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com/v1beta3";
private final RestClient restClient;
@@ -582,33 +585,5 @@ public class VertexAiApi {
}
}
}
-
- /**
- * Main method to test the VertexAiApi.
- * @param args blank.
- */
- public static void main(String[] args) {
- VertexAiApi vertexAiApi = new VertexAiApi(System.getenv("PALM_API_KEY"));
-
- var prompt = new MessagePrompt(List.of(new Message("0", "Hello, how are you?")));
-
- GenerateMessageRequest request = new GenerateMessageRequest(prompt);
-
- GenerateMessageResponse response = vertexAiApi.generateMessage(request);
-
- System.out.println(response);
-
- System.out.println(vertexAiApi.embedText("Hello, how are you?"));
-
- System.out.println(vertexAiApi.batchEmbedText(List.of("Hello, how are you?", "I am fine, thank you!")));
-
- System.out.println(vertexAiApi.countMessageTokens(prompt));
-
- System.out.println(vertexAiApi.listModels());
-
- System.out.println(vertexAiApi.listModels().stream().map(vertexAiApi::getModel).toList());
-
- }
-
}
// @formatter:on
\ No newline at end of file
diff --git a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClient.java b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClient.java
index bb18bdf7e..bcb505684 100644
--- a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClient.java
+++ b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClient.java
@@ -16,20 +16,20 @@
package org.springframework.ai.vertex.embedding;
-import java.util.ArrayList;
import java.util.List;
import java.util.Map;
+import java.util.concurrent.atomic.AtomicInteger;
import org.springframework.ai.document.Document;
+import org.springframework.ai.embedding.AbstractEmbeddingClient;
import org.springframework.ai.embedding.Embedding;
-import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.vertex.api.VertexAiApi;
/**
* @author Christian Tzolov
*/
-public class VertexAiEmbeddingClient implements EmbeddingClient {
+public class VertexAiEmbeddingClient extends AbstractEmbeddingClient {
private final VertexAiApi vertexAiApi;
@@ -56,11 +56,10 @@ public class VertexAiEmbeddingClient implements EmbeddingClient {
@Override
public EmbeddingResponse embedForResponse(List texts) {
List vertexEmbeddings = this.vertexAiApi.batchEmbedText(texts);
- int index = 0;
- List embeddings = new ArrayList<>();
- for (VertexAiApi.Embedding vertexEmbedding : vertexEmbeddings) {
- embeddings.add(new Embedding(vertexEmbedding.value(), index++));
- }
+ AtomicInteger indexCounter = new AtomicInteger(0);
+ List embeddings = vertexEmbeddings.stream()
+ .map(vm -> new Embedding(vm.value(), indexCounter.getAndIncrement()))
+ .toList();
return new EmbeddingResponse(embeddings, Map.of());
}
diff --git a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClient.java b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatClient.java
similarity index 84%
rename from spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClient.java
rename to spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatClient.java
index 19c6ff5c9..a231e9ec4 100644
--- a/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClient.java
+++ b/spring-ai-vertex-ai/src/main/java/org/springframework/ai/vertex/generation/VertexAiChatClient.java
@@ -34,7 +34,7 @@ import org.springframework.util.CollectionUtils;
/**
* @author Christian Tzolov
*/
-public class VertexAiChatGenerationClient implements AiClient {
+public class VertexAiChatClient implements AiClient {
private final VertexAiApi vertexAiApi;
@@ -42,14 +42,30 @@ public class VertexAiChatGenerationClient implements AiClient {
private Float topP;
+ private Integer topK;
+
private Integer candidateCount;
- private Integer maxTokens;
-
- public VertexAiChatGenerationClient(VertexAiApi vertexAiApi) {
+ public VertexAiChatClient(VertexAiApi vertexAiApi) {
this.vertexAiApi = vertexAiApi;
}
+ public void setTemperature(Float temperature) {
+ this.temperature = temperature;
+ }
+
+ public void setTopK(Integer candidateCount) {
+ this.topK = candidateCount;
+ }
+
+ public void setTopP(Float topP) {
+ this.topP = topP;
+ }
+
+ public void setCandidateCount(Integer maxTokens) {
+ this.candidateCount = maxTokens;
+ }
+
@Override
public AiResponse generate(Prompt prompt) {
@@ -70,7 +86,7 @@ public class VertexAiChatGenerationClient implements AiClient {
var vertexPrompt = new MessagePrompt(vertexContext, vertexMessages);
GenerateMessageRequest request = new GenerateMessageRequest(vertexPrompt, this.temperature, this.candidateCount,
- this.topP, this.maxTokens);
+ this.topP, this.topK);
GenerateMessageResponse response = this.vertexAiApi.generateMessage(request);
diff --git a/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClientIT.java b/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClientIT.java
index 78ebbfc30..99a0f10ed 100644
--- a/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClientIT.java
+++ b/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/embedding/VertexAiEmbeddingClientIT.java
@@ -30,6 +30,20 @@ class VertexAiEmbeddingClientIT {
assertThat(embeddingClient.dimensions()).isEqualTo(768);
}
+ @Test
+ void batchEmbedding() {
+ assertThat(embeddingClient).isNotNull();
+ EmbeddingResponse embeddingResponse = embeddingClient
+ .embedForResponse(List.of("Hello World", "World is big and salvation is near"));
+ assertThat(embeddingResponse.getData()).hasSize(2);
+ assertThat(embeddingResponse.getData().get(0).getEmbedding()).isNotEmpty();
+ assertThat(embeddingResponse.getData().get(0).getIndex()).isEqualTo(0);
+ assertThat(embeddingResponse.getData().get(1).getEmbedding()).isNotEmpty();
+ assertThat(embeddingResponse.getData().get(1).getIndex()).isEqualTo(1);
+
+ assertThat(embeddingClient.dimensions()).isEqualTo(768);
+ }
+
@SpringBootConfiguration
public static class TestConfiguration {
diff --git a/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClientIT.java b/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClientIT.java
index 9c5a61006..fc3bb2cfc 100644
--- a/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClientIT.java
+++ b/spring-ai-vertex-ai/src/test/java/org/springframework/ai/vertex/generation/VertexAiChatGenerationClientIT.java
@@ -33,7 +33,7 @@ import static org.assertj.core.api.Assertions.assertThat;
class VertexAiChatGenerationClientIT {
@Autowired
- private VertexAiChatGenerationClient client;
+ private VertexAiChatClient client;
@Value("classpath:/prompts/system-message.st")
private Resource systemResource;
@@ -121,8 +121,8 @@ class VertexAiChatGenerationClientIT {
}
@Bean
- public VertexAiChatGenerationClient vertexAiEmbedding(VertexAiApi vertexAiApi) {
- return new VertexAiChatGenerationClient(vertexAiApi);
+ public VertexAiChatClient vertexAiEmbedding(VertexAiApi vertexAiApi) {
+ return new VertexAiChatClient(vertexAiApi);
}
}
diff --git a/spring-ai-vertex-ai/src/test/resources/Google Generative AI - PaLM2 REST API.jpg b/spring-ai-vertex-ai/src/test/resources/Google Generative AI - PaLM2 REST API.jpg
new file mode 100644
index 000000000..dba271cf4
Binary files /dev/null and b/spring-ai-vertex-ai/src/test/resources/Google Generative AI - PaLM2 REST API.jpg differ