Abstract API for AI model clients

* An abstract API for AI model clients
 * Providing portable client request options while still allowing vendor specific options when required.  Implemented only for StabilityAI/OpenAI ImageClient
 * Support for text->image for openai and stabilityai.

  Partial fix for #27 :  Text To Image and Fixes #266 and Fixes #261
This commit is contained in:
Mark Pollack
2024-01-15 00:06:24 -05:00
committed by Christian Tzolov
parent 08fa0e393c
commit 243cef976c
162 changed files with 3691 additions and 630 deletions

View File

@@ -32,18 +32,18 @@ import com.azure.ai.openai.models.ContentFilterResultsForPrompt;
import com.azure.core.util.IterableStream;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import reactor.core.publisher.Flux;
import org.springframework.ai.azure.openai.metadata.AzureOpenAiGenerationMetadata;
import org.springframework.ai.azure.openai.metadata.AzureOpenAiChatResponseMetadata;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.PromptMetadata;
import org.springframework.ai.metadata.PromptMetadata.PromptFilterMetadata;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.PromptMetadata.PromptFilterMetadata;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.messages.Message;
import org.springframework.util.Assert;
/**
@@ -134,7 +134,7 @@ public class AzureOpenAiChatClient implements ChatClient, StreamingChatClient {
}
@Override
public String generate(String text) {
public String call(String text) {
ChatRequestMessage azureChatMessage = new ChatRequestUserMessage(text);
@@ -160,7 +160,7 @@ public class AzureOpenAiChatClient implements ChatClient, StreamingChatClient {
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
ChatCompletionsOptions options = toAzureChatCompletionsOptions(prompt);
options.setStream(false);
@@ -174,11 +174,12 @@ public class AzureOpenAiChatClient implements ChatClient, StreamingChatClient {
List<Generation> generations = chatCompletions.getChoices()
.stream()
.map(choice -> new Generation(choice.getMessage().getContent())
.withChoiceMetadata(generateChoiceMetadata(choice)))
.withGenerationMetadata(generateChoiceMetadata(choice)))
.toList();
return new ChatResponse(generations, AzureOpenAiGenerationMetadata.from(chatCompletions))
.withPromptMetadata(generatePromptMetadata(chatCompletions));
PromptMetadata promptFilterMetadata = generatePromptMetadata(chatCompletions);
return new ChatResponse(generations,
AzureOpenAiChatResponseMetadata.from(chatCompletions, promptFilterMetadata));
}
@Override
@@ -199,14 +200,17 @@ public class AzureOpenAiChatClient implements ChatClient, StreamingChatClient {
.flatMap(List::stream)
.map(choice -> {
var content = (choice.getDelta() != null) ? choice.getDelta().getContent() : null;
var generation = new Generation(content).withChoiceMetadata(generateChoiceMetadata(choice));
var generation = new Generation(content).withGenerationMetadata(generateChoiceMetadata(choice));
return new ChatResponse(List.of(generation));
}));
}
private ChatCompletionsOptions toAzureChatCompletionsOptions(Prompt prompt) {
List<ChatRequestMessage> azureMessages = prompt.getMessages().stream().map(this::fromSpringAiMessage).toList();
List<ChatRequestMessage> azureMessages = prompt.getInstructions()
.stream()
.map(this::fromSpringAiMessage)
.toList();
ChatCompletionsOptions options = new ChatCompletionsOptions(azureMessages);
@@ -233,8 +237,8 @@ public class AzureOpenAiChatClient implements ChatClient, StreamingChatClient {
}
private ChoiceMetadata generateChoiceMetadata(ChatChoice choice) {
return ChoiceMetadata.from(String.valueOf(choice.getFinishReason()), choice.getContentFilterResults());
private ChatGenerationMetadata generateChoiceMetadata(ChatChoice choice) {
return ChatGenerationMetadata.from(String.valueOf(choice.getFinishReason()), choice.getContentFilterResults());
}
private PromptMetadata generatePromptMetadata(ChatCompletions chatCompletions) {

View File

@@ -18,38 +18,44 @@ package org.springframework.ai.azure.openai.metadata;
import com.azure.ai.openai.models.ChatCompletions;
import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**
* {@link GenerationMetadata} implementation for
* {@link ChatResponseMetadata} implementation for
* {@literal Microsoft Azure OpenAI Service}.
*
* @author John Blum
* @see org.springframework.ai.metadata.GenerationMetadata
* @see ChatResponseMetadata
* @since 0.7.1
*/
public class AzureOpenAiGenerationMetadata implements GenerationMetadata {
public class AzureOpenAiChatResponseMetadata implements ChatResponseMetadata {
protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }";
@SuppressWarnings("all")
public static AzureOpenAiGenerationMetadata from(ChatCompletions chatCompletions) {
public static AzureOpenAiChatResponseMetadata from(ChatCompletions chatCompletions,
PromptMetadata promptFilterMetadata) {
Assert.notNull(chatCompletions, "Azure OpenAI ChatCompletions must not be null");
String id = chatCompletions.getId();
AzureOpenAiUsage usage = AzureOpenAiUsage.from(chatCompletions);
AzureOpenAiGenerationMetadata generationMetadata = new AzureOpenAiGenerationMetadata(id, usage);
return generationMetadata;
AzureOpenAiChatResponseMetadata chatResponseMetadata = new AzureOpenAiChatResponseMetadata(id, usage,
promptFilterMetadata);
return chatResponseMetadata;
}
private final String id;
private final Usage usage;
protected AzureOpenAiGenerationMetadata(String id, AzureOpenAiUsage usage) {
private final PromptMetadata promptMetadata;
protected AzureOpenAiChatResponseMetadata(String id, AzureOpenAiUsage usage, PromptMetadata promptMetadata) {
this.id = id;
this.usage = usage;
this.promptMetadata = promptMetadata;
}
public String getId() {
@@ -61,6 +67,11 @@ public class AzureOpenAiGenerationMetadata implements GenerationMetadata {
return this.usage;
}
@Override
public PromptMetadata getPromptMetadata() {
return this.promptMetadata;
}
@Override
public String toString() {
return AI_METADATA_STRING.formatted(getClass().getTypeName(), getId(), getUsage(), getRateLimit());

View File

@@ -19,7 +19,7 @@ package org.springframework.ai.azure.openai.metadata;
import com.azure.ai.openai.models.ChatCompletions;
import com.azure.ai.openai.models.CompletionsUsage;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**

View File

@@ -14,14 +14,15 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.context.SpringBootTest;
@@ -53,8 +54,8 @@ class AzureOpenAiChatClientIT {
UserMessage userMessage = new UserMessage("Generate the names of 5 famous pirates.");
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = chatClient.generate(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
ChatResponse response = chatClient.call(prompt);
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Test
@@ -70,9 +71,9 @@ class AzureOpenAiChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = chatClient.generate(prompt).getGeneration();
Generation generation = chatClient.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -89,9 +90,9 @@ class AzureOpenAiChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = chatClient.generate(prompt).getGeneration();
Generation generation = chatClient.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -108,9 +109,9 @@ class AzureOpenAiChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = chatClient.generate(prompt).getGeneration();
Generation generation = chatClient.call(prompt).getResult();
ActorsFilms actorsFilms = outputParser.parse(generation.getContent());
ActorsFilms actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isNotNull();
}
@@ -129,9 +130,9 @@ class AzureOpenAiChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = chatClient.generate(prompt).getGeneration();
Generation generation = chatClient.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
System.out.println(actorsFilms);
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
@@ -154,9 +155,10 @@ class AzureOpenAiChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.filter(Objects::nonNull)
.collect(Collectors.joining());

View File

@@ -28,12 +28,13 @@ import org.springframework.ai.azure.openai.AzureOpenAiChatClient;
import org.springframework.ai.azure.openai.MockAzureOpenAiTestConfiguration;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.PromptMetadata;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.PromptMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.context.SpringBootTest;
@@ -75,14 +76,15 @@ class AzureOpenAiChatClientMetadataTests {
Prompt prompt = new Prompt("Can I fly like a bird?");
ChatResponse response = this.aiClient.generate(prompt);
ChatResponse response = this.aiClient.call(prompt);
assertThat(response).isNotNull();
Generation generation = response.getGeneration();
Generation generation = response.getResult();
assertThat(generation).isNotNull()
.extracting(Generation::getContent)
.extracting(Generation::getOutput)
.extracting(AssistantMessage::getContent)
.isEqualTo("No! You will actually land with a resounding thud. This is the way!");
assertPromptMetadata(response);
@@ -92,7 +94,7 @@ class AzureOpenAiChatClientMetadataTests {
private void assertPromptMetadata(ChatResponse response) {
PromptMetadata promptMetadata = response.getPromptMetadata();
PromptMetadata promptMetadata = response.getMetadata().getPromptMetadata();
assertThat(promptMetadata).isNotNull();
@@ -106,12 +108,12 @@ class AzureOpenAiChatClientMetadataTests {
private void assertGenerationMetadata(ChatResponse response) {
GenerationMetadata generationMetadata = response.getGenerationMetadata();
ChatResponseMetadata chatResponseMetadata = response.getMetadata();
assertThat(generationMetadata).isNotNull();
assertThat(generationMetadata.getRateLimit()).isEqualTo(RateLimit.NULL);
assertThat(chatResponseMetadata).isNotNull();
assertThat(chatResponseMetadata.getRateLimit()).isEqualTo(RateLimit.NULL);
Usage usage = generationMetadata.getUsage();
Usage usage = chatResponseMetadata.getUsage();
assertThat(usage).isNotNull();
assertThat(usage).isNotEqualTo(Usage.NULL);
@@ -122,11 +124,11 @@ class AzureOpenAiChatClientMetadataTests {
private void assertChoiceMetadata(Generation generation) {
ChoiceMetadata choiceMetadata = generation.getChoiceMetadata();
ChatGenerationMetadata chatGenerationMetadata = generation.getMetadata();
assertThat(choiceMetadata).isNotNull();
assertThat(choiceMetadata.getFinishReason()).isEqualTo("stop");
assertContentFilterResults(choiceMetadata.getContentFilterMetadata());
assertThat(chatGenerationMetadata).isNotNull();
assertThat(chatGenerationMetadata.getFinishReason()).isEqualTo("stop");
assertContentFilterResults(chatGenerationMetadata.getContentFilterMetadata());
}
private void assertContentFilterResultsForPrompt(ContentFilterResultDetailsForPrompt contentFilterResultForPrompt,

View File

@@ -17,7 +17,7 @@
package org.springframework.ai.bedrock;
import org.springframework.ai.bedrock.api.AbstractBedrockApi.AmazonBedrockInvocationMetrics;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.util.Assert;
/**

View File

@@ -19,8 +19,8 @@ package org.springframework.ai.bedrock;
import java.util.List;
import java.util.stream.Collectors;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.MessageType;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.MessageType;
/**
* Converts a list of messages to a prompt for bedrock models.

View File

@@ -20,6 +20,7 @@ import java.util.List;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import reactor.core.publisher.Flux;
import org.springframework.ai.bedrock.MessageToPromptConverter;
@@ -28,12 +29,11 @@ import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.Anth
import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatResponse;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.prompt.Prompt;
/**
* Java {@link ChatClient} and {@link StreamingChatClient} for the Bedrock Anthropic chat
* model.
* generative.
*
* @author Christian Tzolov
* @since 0.8.0
@@ -89,8 +89,8 @@ public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClie
}
@Override
public ChatResponse generate(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
public ChatResponse call(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
AnthropicChatRequest request = AnthropicChatRequest.builder(promptValue)
.withTemperature(this.temperature)
@@ -109,7 +109,7 @@ public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClie
@Override
public Flux<ChatResponse> generateStream(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
AnthropicChatRequest request = AnthropicChatRequest.builder(promptValue)
.withTemperature(this.temperature)
@@ -126,8 +126,8 @@ public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClie
String stopReason = response.stopReason() != null ? response.stopReason() : null;
var generation = new Generation(response.completion());
if (response.amazonBedrockInvocationMetrics() != null) {
generation = generation
.withChoiceMetadata(ChoiceMetadata.from(stopReason, response.amazonBedrockInvocationMetrics()));
generation = generation.withGenerationMetadata(
ChatGenerationMetadata.from(stopReason, response.amazonBedrockInvocationMetrics()));
}
return new ChatResponse(List.of(generation));
});

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.bedrock.cohere;
import java.util.List;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import reactor.core.publisher.Flux;
import org.springframework.ai.bedrock.BedrockUsage;
@@ -32,9 +33,8 @@ import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChat
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.prompt.Prompt;
/**
* @author Christian Tzolov
@@ -112,7 +112,7 @@ public class BedrockCohereChatClient implements ChatClient, StreamingChatClient
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
CohereChatResponse response = this.chatApi.chatCompletion(this.createRequest(prompt, false));
List<Generation> generations = response.generations().stream().map(g -> {
return new Generation(g.text());
@@ -127,15 +127,15 @@ public class BedrockCohereChatClient implements ChatClient, StreamingChatClient
if (g.isFinished()) {
String finishReason = g.finishReason().name();
Usage usage = BedrockUsage.from(g.amazonBedrockInvocationMetrics());
return new ChatResponse(
List.of(new Generation("").withChoiceMetadata(ChoiceMetadata.from(finishReason, usage))));
return new ChatResponse(List
.of(new Generation("").withGenerationMetadata(ChatGenerationMetadata.from(finishReason, usage))));
}
return new ChatResponse(List.of(new Generation(g.text())));
});
}
private CohereChatRequest createRequest(Prompt prompt, boolean stream) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
return CohereChatRequest.builder(promptValue)
.withTemperature(this.temperature)

View File

@@ -28,13 +28,13 @@ import org.springframework.ai.bedrock.llama2.api.Llama2ChatBedrockApi.Llama2Chat
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.prompt.Prompt;
/**
* Java {@link ChatClient} and {@link StreamingChatClient} for the Bedrock Llama2 chat
* model.
* generative.
*
* @author Christian Tzolov
* @since 0.8.0
@@ -69,8 +69,8 @@ public class BedrockLlama2ChatClient implements ChatClient, StreamingChatClient
}
@Override
public ChatResponse generate(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
public ChatResponse call(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
var request = Llama2ChatRequest.builder(promptValue)
.withTemperature(this.temperature)
@@ -80,14 +80,14 @@ public class BedrockLlama2ChatClient implements ChatClient, StreamingChatClient
Llama2ChatResponse response = this.chatApi.chatCompletion(request);
return new ChatResponse(List.of(new Generation(response.generation())
.withChoiceMetadata(ChoiceMetadata.from(response.stopReason().name(), extractUsage(response)))));
return new ChatResponse(List.of(new Generation(response.generation()).withGenerationMetadata(
ChatGenerationMetadata.from(response.stopReason().name(), extractUsage(response)))));
}
@Override
public Flux<ChatResponse> generateStream(Prompt prompt) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
var request = Llama2ChatRequest.builder(promptValue)
.withTemperature(this.temperature)
@@ -100,7 +100,7 @@ public class BedrockLlama2ChatClient implements ChatClient, StreamingChatClient
return fluxResponse.map(response -> {
String stopReason = response.stopReason() != null ? response.stopReason().name() : null;
return new ChatResponse(List.of(new Generation(response.generation())
.withChoiceMetadata(ChoiceMetadata.from(stopReason, extractUsage(response)))));
.withGenerationMetadata(ChatGenerationMetadata.from(stopReason, extractUsage(response)))));
});
}

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.bedrock.titan;
import java.util.List;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import reactor.core.publisher.Flux;
import org.springframework.ai.bedrock.MessageToPromptConverter;
@@ -29,9 +30,8 @@ import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatRes
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.chat.prompt.Prompt;
/**
* @author Christian Tzolov
@@ -74,7 +74,7 @@ public class BedrockTitanChatClient implements ChatClient, StreamingChatClient {
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
TitanChatResponse response = this.chatApi.chatCompletion(this.createRequest(prompt, false));
List<Generation> generations = response.results().stream().map(result -> {
return new Generation(result.outputText());
@@ -91,12 +91,13 @@ public class BedrockTitanChatClient implements ChatClient, StreamingChatClient {
if (chunk.amazonBedrockInvocationMetrics() != null) {
String completionReason = chunk.completionReason().name();
generation = generation
.withChoiceMetadata(ChoiceMetadata.from(completionReason, chunk.amazonBedrockInvocationMetrics()));
generation = generation.withGenerationMetadata(
ChatGenerationMetadata.from(completionReason, chunk.amazonBedrockInvocationMetrics()));
}
else if (chunk.inputTextTokenCount() != null && chunk.totalOutputTextTokenCount() != null) {
String completionReason = chunk.completionReason().name();
generation = generation.withChoiceMetadata(ChoiceMetadata.from(completionReason, extractUsage(chunk)));
generation = generation
.withGenerationMetadata(ChatGenerationMetadata.from(completionReason, extractUsage(chunk)));
}
return new ChatResponse(List.of(generation));
@@ -104,7 +105,7 @@ public class BedrockTitanChatClient implements ChatClient, StreamingChatClient {
}
private TitanChatRequest createRequest(Prompt prompt, boolean stream) {
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getMessages());
final String promptValue = MessageToPromptConverter.create().toPrompt(prompt.getInstructions());
return TitanChatRequest.builder(promptValue)
.withTemperature(this.temperature)

View File

@@ -9,6 +9,7 @@ import com.fasterxml.jackson.databind.ObjectMapper;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.messages.AssistantMessage;
import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider;
import software.amazon.awssdk.regions.Region;
@@ -17,11 +18,11 @@ import org.springframework.ai.chat.Generation;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringBootConfiguration;
@@ -52,9 +53,9 @@ class BedrockAnthropicChatClientIT {
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
ChatResponse response = client.call(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Test
@@ -70,9 +71,9 @@ class BedrockAnthropicChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -88,9 +89,9 @@ class BedrockAnthropicChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -112,9 +113,9 @@ class BedrockAnthropicChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}
@@ -137,9 +138,10 @@ class BedrockAnthropicChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -8,6 +8,7 @@ import java.util.stream.Collectors;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.messages.AssistantMessage;
import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider;
import software.amazon.awssdk.regions.Region;
@@ -18,11 +19,11 @@ import org.springframework.ai.chat.Generation;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringBootConfiguration;
@@ -53,8 +54,8 @@ class BedrockCohereChatClientIT {
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource);
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice));
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
ChatResponse response = client.call(prompt);
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Test
@@ -70,9 +71,9 @@ class BedrockCohereChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -89,9 +90,9 @@ class BedrockCohereChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -112,9 +113,9 @@ class BedrockCohereChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}
@@ -137,9 +138,10 @@ class BedrockCohereChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -10,6 +10,7 @@ import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.messages.AssistantMessage;
import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider;
import software.amazon.awssdk.regions.Region;
@@ -19,11 +20,11 @@ import org.springframework.ai.chat.Generation;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringBootConfiguration;
@@ -54,9 +55,9 @@ class BedrockLlama2ChatClientIT {
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
ChatResponse response = client.call(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Disabled("TODO: Fix the parser instructions to return the correct format")
@@ -73,9 +74,9 @@ class BedrockLlama2ChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -91,9 +92,9 @@ class BedrockLlama2ChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -116,9 +117,9 @@ class BedrockLlama2ChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}
@@ -142,9 +143,10 @@ class BedrockLlama2ChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -10,6 +10,7 @@ import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.messages.AssistantMessage;
import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider;
import software.amazon.awssdk.regions.Region;
@@ -19,11 +20,11 @@ import org.springframework.ai.chat.Generation;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringBootConfiguration;
@@ -54,8 +55,8 @@ class BedrockTitanChatClientIT {
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource);
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice));
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
ChatResponse response = client.call(prompt);
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Disabled("TODO: Fix the parser instructions to return the correct format")
@@ -72,9 +73,9 @@ class BedrockTitanChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -93,9 +94,9 @@ class BedrockTitanChatClientIT {
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -117,9 +118,9 @@ class BedrockTitanChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}
@@ -143,9 +144,10 @@ class BedrockTitanChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -32,7 +32,7 @@ import org.springframework.ai.huggingface.model.AllOfGenerateResponseDetails;
import org.springframework.ai.huggingface.model.GenerateParameters;
import org.springframework.ai.huggingface.model.GenerateRequest;
import org.springframework.ai.huggingface.model.GenerateResponse;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.prompt.Prompt;
/**
* An implementation of {@link ChatClient} that interfaces with HuggingFace Inference
@@ -86,7 +86,7 @@ public class HuggingfaceChatClient implements ChatClient {
* @return ChatResponse containing the generated text and other related details.
*/
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
GenerateRequest generateRequest = new GenerateRequest();
generateRequest.setInputs(prompt.getContents());
GenerateParameters generateParameters = new GenerateParameters();

View File

@@ -22,7 +22,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.huggingface.HuggingfaceChatClient;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
@@ -47,8 +47,8 @@ public class ClientIT {
[/INST]
""";
Prompt prompt = new Prompt(mistral7bInstruct);
ChatResponse chatResponse = huggingfaceChatClient.generate(prompt);
assertThat(chatResponse.getGeneration().getContent()).isNotEmpty();
ChatResponse chatResponse = huggingfaceChatClient.call(prompt);
assertThat(chatResponse.getResult().getOutput().getContent()).isNotEmpty();
String expectedResponse = """
```json
{
@@ -57,9 +57,9 @@ public class ClientIT {
"address": "#1 Samuel St."
}
```""";
assertThat(chatResponse.getGeneration().getContent()).isEqualTo(expectedResponse);
assertThat(chatResponse.getGeneration().getProperties()).containsKey("generated_tokens");
assertThat(chatResponse.getGeneration().getProperties()).containsEntry("generated_tokens", 39);
assertThat(chatResponse.getResult().getOutput().getContent()).isEqualTo(expectedResponse);
assertThat(chatResponse.getResult().getOutput().getProperties()).containsKey("generated_tokens");
assertThat(chatResponse.getResult().getOutput().getProperties()).containsEntry("generated_tokens", 39);
}

View File

@@ -21,6 +21,12 @@
</properties>
<dependencies>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-core</artifactId>

View File

@@ -16,25 +16,30 @@
package org.springframework.ai.ollama;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.stream.Collectors;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import reactor.core.publisher.Flux;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.ollama.api.OllamaApi;
import org.springframework.ai.ollama.api.OllamaOptions;
import org.springframework.ai.ollama.api.OllamaApi.ChatRequest;
import org.springframework.ai.ollama.api.OllamaApi.Message.Role;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.MessageType;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.MessageType;
/**
* {@link ChatClient} implementation for {@literal Ollma}.
@@ -58,6 +63,8 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient {
private Map<String, Object> clientOptions;
private final static ObjectMapper OBJECT_MAPPER = new ObjectMapper();
public OllamaChatClient(OllamaApi chatApi) {
this.chatApi = chatApi;
}
@@ -78,12 +85,13 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient {
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
OllamaApi.ChatResponse response = this.chatApi.chat(request(prompt, this.model, false));
var generator = new Generation(response.message().content());
if (response.promptEvalCount() != null && response.evalCount() != null) {
generator = generator.withChoiceMetadata(ChoiceMetadata.from("unknown", extractUsage(response)));
generator = generator
.withGenerationMetadata(ChatGenerationMetadata.from("unknown", extractUsage(response)));
}
return new ChatResponse(List.of(generator));
}
@@ -97,7 +105,8 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient {
Generation generation = (chunk.message() != null) ? new Generation(chunk.message().content())
: new Generation("");
if (Boolean.TRUE.equals(chunk.done())) {
generation = generation.withChoiceMetadata(ChoiceMetadata.from("unknown", extractUsage(chunk)));
generation = generation
.withGenerationMetadata(ChatGenerationMetadata.from("unknown", extractUsage(chunk)));
}
return new ChatResponse(List.of(generation));
});
@@ -120,20 +129,57 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient {
private OllamaApi.ChatRequest request(Prompt prompt, String model, boolean stream) {
List<OllamaApi.Message> ollamaMessages = prompt.getMessages()
List<OllamaApi.Message> ollamaMessages = prompt.getInstructions()
.stream()
.filter(message -> message.getMessageType() == MessageType.USER
|| message.getMessageType() == MessageType.ASSISTANT)
.map(m -> OllamaApi.Message.builder(toRole(m)).withContent(m.getContent()).build())
.toList();
// runtime options
Map<String, Object> promptOptions = objectToMap(prompt.getOptions());
Map<String, Object> clientOptionsToUse = merge(promptOptions, this.clientOptions, HashMap.class);
return ChatRequest.builder(model)
.withStream(stream)
.withMessages(ollamaMessages)
.withOptions(this.clientOptions)
.withOptions(clientOptionsToUse)
.build();
}
public static Map<String, Object> objectToMap(Object source) {
try {
String json = OBJECT_MAPPER.writeValueAsString(source);
return OBJECT_MAPPER.readValue(json, new TypeReference<Map<String, Object>>() {
});
}
catch (JsonProcessingException e) {
throw new RuntimeException(e);
}
}
public static <T> T mapToClass(Map<String, Object> source, Class<T> clazz) {
try {
String json = OBJECT_MAPPER.writeValueAsString(source);
return OBJECT_MAPPER.readValue(json, clazz);
}
catch (JsonProcessingException e) {
throw new RuntimeException(e);
}
}
public static <T> T merge(Object source, Object target, Class<T> clazz) {
Map<String, Object> sourceMap = objectToMap(source);
Map<String, Object> targetMap = objectToMap(target);
targetMap.putAll(sourceMap.entrySet()
.stream()
.filter(e -> e.getValue() != null)
.collect(Collectors.toMap(e -> e.getKey(), e -> e.getValue())));
return mapToClass(targetMap, clazz);
}
private OllamaApi.Message.Role toRole(Message message) {
switch (message.getMessageType()) {

View File

@@ -25,6 +25,7 @@ import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.ai.chat.ChatOptions;
/**
* Helper class for creating strongly-typed Ollama options.
@@ -38,7 +39,7 @@ import com.fasterxml.jackson.databind.ObjectMapper;
* Types</a>
*/
@JsonInclude(Include.NON_NULL)
public class OllamaOptions {
public class OllamaOptions implements ChatOptions {
// @formatter:off
/**

View File

@@ -11,6 +11,8 @@ import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.ChatOptionsBuilder;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
@@ -22,11 +24,11 @@ import org.springframework.ai.ollama.api.OllamaOptions;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.context.SpringBootTest;
@@ -51,7 +53,7 @@ class OllamaChatClientIT {
@BeforeAll
public static void beforeAll() throws IOException, InterruptedException {
logger.info("Start pulling the '" + MODEL + " ' model ... would take several minutes ...");
logger.info("Start pulling the '" + MODEL + " ' generative ... would take several minutes ...");
ollamaContainer.execInContainer("ollama", "pull", MODEL);
logger.info(MODEL + " pulling competed!");
@@ -72,9 +74,16 @@ class OllamaChatClientIT {
UserMessage userMessage = new UserMessage("Tell me about 5 famous pirates from the Golden Age of Piracy.");
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
assertThat(response.getGeneration().getContent()).contains("Blackbeard");
// portable/generic options
var chatOptionsBuilder = ChatOptionsBuilder.builder();
// ollama specific options
var ollamaOptions = new OllamaOptions().withLowVRAM(true);
Prompt prompt = new Prompt(List.of(userMessage, systemMessage),
chatOptionsBuilder.withTemperature(0.7f).build());
ChatResponse response = client.call(prompt);
assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard");
}
@Disabled("TODO: Fix the parser instructions to return the correct format")
@@ -91,9 +100,9 @@ class OllamaChatClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -112,9 +121,9 @@ class OllamaChatClientIT {
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -136,9 +145,9 @@ class OllamaChatClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}
@@ -162,9 +171,10 @@ class OllamaChatClientIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -36,7 +36,7 @@ class OllamaEmbeddingClientIT {
@BeforeAll
public static void beforeAll() throws IOException, InterruptedException {
logger.info("Start pulling the 'orca-mini' model (3GB) ... would take several minutes ...");
logger.info("Start pulling the 'orca-mini' generative (3GB) ... would take several minutes ...");
ollamaContainer.execInContainer("ollama", "pull", "orca-mini");
logger.info("orca-mini pulling competed!");

View File

@@ -57,7 +57,7 @@ public class OllamaApiIT {
@BeforeAll
public static void beforeAll() throws IOException, InterruptedException {
logger.info("Start pulling the 'orca-mini' model (3GB) ... would take several minutes ...");
logger.info("Start pulling the 'orca-mini' generative (3GB) ... would take several minutes ...");
ollamaContainer.execInContainer("ollama", "pull", "orca-mini");
logger.info("orca-mini pulling competed!");

View File

@@ -25,7 +25,7 @@ import static org.assertj.core.api.Assertions.assertThat;
/**
* @author Christian Tzolov
*/
public class OllamaOptionsTests {
public class OllamaModelOptionsTests {
@Test
public void testOptions() {

View File

@@ -28,17 +28,17 @@ import reactor.core.publisher.Flux;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion;
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage;
import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiClientErrorException;
import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException;
import org.springframework.ai.openai.metadata.OpenAiGenerationMetadata;
import org.springframework.ai.openai.metadata.OpenAiChatResponseMetadata;
import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.messages.Message;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
@@ -94,10 +94,10 @@ public class OpenAiChatClient implements ChatClient, StreamingChatClient {
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
return this.retryTemplate.execute(ctx -> {
List<Message> messages = prompt.getMessages();
List<Message> messages = prompt.getInstructions();
List<ChatCompletionMessage> chatCompletionMessages = messages.stream()
.map(m -> new ChatCompletionMessage(m.getContent(),
@@ -118,18 +118,18 @@ public class OpenAiChatClient implements ChatClient, StreamingChatClient {
List<Generation> generations = chatCompletion.choices().stream().map(choice -> {
return new Generation(choice.message().content(), Map.of("role", choice.message().role().name()))
.withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null));
.withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null));
}).toList();
return new ChatResponse(generations,
OpenAiGenerationMetadata.from(completionEntity.getBody()).withRateLimit(rateLimits));
OpenAiChatResponseMetadata.from(completionEntity.getBody()).withRateLimit(rateLimits));
});
}
@Override
public Flux<ChatResponse> generateStream(Prompt prompt) {
return this.retryTemplate.execute(ctx -> {
List<Message> messages = prompt.getMessages();
List<Message> messages = prompt.getInstructions();
List<ChatCompletionMessage> chatCompletionMessages = messages.stream()
.map(m -> new ChatCompletionMessage(m.getContent(),
@@ -153,7 +153,7 @@ public class OpenAiChatClient implements ChatClient, StreamingChatClient {
var generation = new Generation(choice.delta().content(), Map.of("role", roleMap.get(chunkId)));
if (choice.finishReason() != null) {
generation = generation
.withChoiceMetadata(ChoiceMetadata.from(choice.finishReason().name(), null));
.withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null));
}
return generation;
}).toList();

View File

@@ -13,6 +13,7 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai;
import java.time.Duration;

View File

@@ -0,0 +1,148 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.image.*;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.openai.api.*;
import org.springframework.ai.openai.metadata.OpenAiImageGenerationMetadata;
import org.springframework.ai.openai.metadata.OpenAiImageResponseMetadata;
import org.springframework.http.ResponseEntity;
import org.springframework.retry.support.RetryTemplate;
import org.springframework.util.Assert;
import java.time.Duration;
import java.util.List;
public class OpenAiImageClient implements ImageClient {
private final Logger logger = LoggerFactory.getLogger(getClass());
private OpenAiImageOptions options;
private final OpenAiImageApi openAiImageApi;
public final RetryTemplate retryTemplate = RetryTemplate.builder()
.maxAttempts(10)
.retryOn(OpenAiApi.OpenAiApiException.class)
.exponentialBackoff(Duration.ofMillis(2000), 5, Duration.ofMillis(3 * 60000))
.build();
public OpenAiImageClient(OpenAiImageApi openAiImageApi) {
Assert.notNull(openAiImageApi, "OpenAiImageApi must not be null");
this.openAiImageApi = openAiImageApi;
}
public OpenAiImageOptions getOptions() {
return options;
}
@Override
public ImageResponse call(ImagePrompt imagePrompt) {
return this.retryTemplate.execute(ctx -> {
ImageOptions runtimeOptions = imagePrompt.getOptions();
OpenAiImageOptions imageOptionsToUse = updateImageOptions(imagePrompt.getOptions());
// Merge the runtime options passed via the prompt with the
// StabilityAiImageClient
// options configured via Autoconfiguration.
// Runtime options overwrite StabilityAiImageClient options
OpenAiImageOptions optionsToUse = ModelOptionsUtils.merge(runtimeOptions, this.options,
OpenAiImageOptionsImpl.class);
// Copy the org.springframework.ai.model derived ImagePrompt and ImageOptions
// data
// types to the data types used in OpenAiImageApi
String instructions = imagePrompt.getInstructions().get(0).getText();
String size;
if (imageOptionsToUse.getWidth() != null && imageOptionsToUse.getHeight() != null) {
size = imageOptionsToUse.getWidth() + "x" + imageOptionsToUse.getHeight();
}
else {
size = null;
}
OpenAiImageApi.OpenAiImageRequest openAiImageRequest = new OpenAiImageApi.OpenAiImageRequest(instructions,
imageOptionsToUse.getModel(), imageOptionsToUse.getN(), imageOptionsToUse.getQuality(), size,
imageOptionsToUse.getResponseFormat(), imageOptionsToUse.getStyle(), imageOptionsToUse.getUser());
// Make the request
ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity = this.openAiImageApi
.createImage(openAiImageRequest);
// Convert to org.springframework.ai.model derived ImageResponse data type
return convertResponse(imageResponseEntity, openAiImageRequest);
});
}
private ImageResponse convertResponse(ResponseEntity<OpenAiImageApi.OpenAiImageResponse> imageResponseEntity,
OpenAiImageApi.OpenAiImageRequest openAiImageRequest) {
OpenAiImageApi.OpenAiImageResponse imageApiResponse = imageResponseEntity.getBody();
if (imageApiResponse == null) {
logger.warn("No image response returned for request: {}", openAiImageRequest);
return new ImageResponse(List.of());
}
List<ImageGeneration> imageGenerationList = imageApiResponse.data().stream().map(entry -> {
return new ImageGeneration(new Image(entry.url(), entry.b64Json()),
new OpenAiImageGenerationMetadata(entry.revisedPrompt()));
}).toList();
ImageResponseMetadata openAiImageResponseMetadata = OpenAiImageResponseMetadata.from(imageApiResponse);
return new ImageResponse(imageGenerationList, openAiImageResponseMetadata);
}
private OpenAiImageOptions updateImageOptions(ImageOptions runtimeImageOptions) {
OpenAiImageOptionsBuilder openAiImageOptionsBuilder = OpenAiImageOptionsBuilder.builder();
if (runtimeImageOptions != null) {
// Handle portable image options
if (runtimeImageOptions.getN() != null) {
openAiImageOptionsBuilder.withN(runtimeImageOptions.getN());
}
if (runtimeImageOptions.getModel() != null) {
openAiImageOptionsBuilder.withModel(runtimeImageOptions.getModel());
}
if (runtimeImageOptions.getResponseFormat() != null) {
openAiImageOptionsBuilder.withResponseFormat(runtimeImageOptions.getResponseFormat());
}
if (runtimeImageOptions.getWidth() != null) {
openAiImageOptionsBuilder.withWidth(runtimeImageOptions.getWidth());
}
if (runtimeImageOptions.getHeight() != null) {
openAiImageOptionsBuilder.withHeight(runtimeImageOptions.getHeight());
}
// Handle OpenAI specific image options
if (runtimeImageOptions instanceof OpenAiImageOptions) {
OpenAiImageOptions runtimeOpenAiImageOptions = (OpenAiImageOptions) runtimeImageOptions;
if (runtimeOpenAiImageOptions.getQuality() != null) {
openAiImageOptionsBuilder.withQuality(runtimeOpenAiImageOptions.getQuality());
}
if (runtimeOpenAiImageOptions.getStyle() != null) {
openAiImageOptionsBuilder.withStyle(runtimeOpenAiImageOptions.getStyle());
}
if (runtimeOpenAiImageOptions.getUser() != null) {
openAiImageOptionsBuilder.withUser(runtimeOpenAiImageOptions.getUser());
}
}
}
OpenAiImageOptions updatedOpenAiImageOptions = openAiImageOptionsBuilder.build();
return updatedOpenAiImageOptions;
}
}

View File

@@ -69,7 +69,7 @@ public class OpenAiApi {
}
/**
* Create an new chat completion api.
* Create a new chat completion api.
*
* @param baseUrl api base URL.
* @param openAiToken OpenAI apiKey.

View File

@@ -0,0 +1,226 @@
package org.springframework.ai.openai.api;
import java.io.IOException;
import java.util.List;
import java.util.Objects;
import java.util.function.Consumer;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiClientErrorException;
import org.springframework.ai.openai.api.OpenAiApi.OpenAiApiException;
import org.springframework.ai.openai.api.OpenAiApi.ResponseError;
import org.springframework.http.HttpHeaders;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.util.Assert;
import org.springframework.web.client.ResponseErrorHandler;
import org.springframework.web.client.RestClient;
public class OpenAiImageApi {
private static final String DEFAULT_BASE_URL = "https://api.openai.com";
public static final String DEFAULT_IMAGE_MODEL = "dall-e-2";
// Assuming RestClient and WebClient are properly defined somewhere
private final RestClient restClient;
private final ObjectMapper objectMapper;
/**
* Create a new OpenAI Image api with base URL set to https://api.openai.com
* @param openAiToken OpenAI apiKey.
*/
public OpenAiImageApi(String openAiToken) {
this(DEFAULT_BASE_URL, openAiToken, RestClient.builder());
}
public OpenAiImageApi(String baseUrl, String openAiToken, RestClient.Builder restClientBuilder) {
this.objectMapper = new ObjectMapper().configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false);
Consumer<HttpHeaders> jsonContentHeaders = headers -> {
headers.setBearerAuth(openAiToken);
headers.setContentType(MediaType.APPLICATION_JSON);
};
var responseErrorHandler = new ResponseErrorHandler() {
@Override
public boolean hasError(ClientHttpResponse response) throws IOException {
return response.getStatusCode().isError();
}
@Override
public void handleError(ClientHttpResponse response) throws IOException {
if (response.getStatusCode().isError()) {
if (response.getStatusCode().is4xxClientError()) {
throw new OpenAiApiClientErrorException(String.format("%s - %s",
response.getStatusCode().value(),
OpenAiImageApi.this.objectMapper.readValue(response.getBody(), ResponseError.class)));
}
throw new OpenAiApiException(String.format("%s - %s", response.getStatusCode().value(),
OpenAiImageApi.this.objectMapper.readValue(response.getBody(), ResponseError.class)));
}
}
};
this.restClient = restClientBuilder.baseUrl(baseUrl)
.defaultHeaders(jsonContentHeaders)
.defaultStatusHandler(responseErrorHandler)
.build();
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public static class OpenAiImageRequest {
@JsonProperty("prompt")
private String prompt;
@JsonProperty("model")
private String model = DEFAULT_IMAGE_MODEL;
@JsonProperty("n")
private Integer n;
@JsonProperty("quality")
private String quality;
@JsonProperty("response_format")
private String responseFormat;
@JsonProperty("size")
private String size;
@JsonProperty("style")
private String style;
@JsonProperty("user")
private String user;
public OpenAiImageRequest() {
}
public OpenAiImageRequest(String prompt, String model, Integer n, String quality, String size,
String responseFormat, String style, String user) {
this.prompt = prompt;
this.model = model;
this.n = n;
this.quality = quality;
this.size = size;
this.responseFormat = responseFormat;
this.style = style;
this.user = user;
}
public String getPrompt() {
return prompt;
}
public void setPrompt(String prompt) {
this.prompt = prompt;
}
public String getModel() {
return model;
}
public void setModel(String model) {
this.model = model;
}
public Integer getN() {
return n;
}
public void setN(Integer n) {
this.n = n;
}
public String getQuality() {
return quality;
}
public void setQuality(String quality) {
this.quality = quality;
}
public String getSize() {
return size;
}
public void setSize(String size) {
this.size = size;
}
public String getResponseFormat() {
return responseFormat;
}
public void setResponseFormat(String responseFormat) {
this.responseFormat = responseFormat;
}
public String getStyle() {
return style;
}
public void setStyle(String style) {
this.style = style;
}
public String getUser() {
return user;
}
public void setUser(String user) {
this.user = user;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof OpenAiImageRequest that))
return false;
return Objects.equals(prompt, that.prompt) && Objects.equals(model, that.model) && Objects.equals(n, that.n)
&& Objects.equals(quality, that.quality) && Objects.equals(size, that.size)
&& Objects.equals(responseFormat, that.responseFormat) && Objects.equals(style, that.style)
&& Objects.equals(user, that.user);
}
@Override
public int hashCode() {
return Objects.hash(prompt, model, n, quality, size, responseFormat, style, user);
}
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public record OpenAiImageResponse(@JsonProperty("created") Long created, @JsonProperty("data") List<Data> data) {
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public record Data(@JsonProperty("url") String url, @JsonProperty("b64_json") String b64Json,
@JsonProperty("revised_prompt") String revisedPrompt) {
}
public ResponseEntity<OpenAiImageResponse> createImage(OpenAiImageRequest openAiImageRequest) {
Assert.notNull(openAiImageRequest, "Image request cannot be null.");
Assert.hasLength(openAiImageRequest.getPrompt(), "Prompt cannot be empty.");
return this.restClient.post()
.uri("v1/images/generations")
.body(openAiImageRequest)
.retrieve()
.toEntity(OpenAiImageResponse.class);
}
}

View File

@@ -0,0 +1,13 @@
package org.springframework.ai.openai.api;
import org.springframework.ai.image.ImageOptions;
public interface OpenAiImageOptions extends ImageOptions {
String getQuality();
String getStyle();
String getUser();
}

View File

@@ -0,0 +1,59 @@
package org.springframework.ai.openai.api;
public class OpenAiImageOptionsBuilder {
private final OpenAiImageOptionsImpl options;
private OpenAiImageOptionsBuilder() {
this.options = new OpenAiImageOptionsImpl();
}
public static OpenAiImageOptionsBuilder builder() {
return new OpenAiImageOptionsBuilder();
}
public OpenAiImageOptionsBuilder withN(Integer n) {
options.setN(n);
return this;
}
public OpenAiImageOptionsBuilder withModel(String model) {
options.setModel(model);
return this;
}
public OpenAiImageOptionsBuilder withQuality(String quality) {
options.setQuality(quality);
return this;
}
public OpenAiImageOptionsBuilder withResponseFormat(String responseFormat) {
options.setResponseFormat(responseFormat);
return this;
}
public OpenAiImageOptionsBuilder withWidth(Integer width) {
options.setWidth(width);
return this;
}
public OpenAiImageOptionsBuilder withHeight(Integer height) {
options.setHeight(height);
return this;
}
public OpenAiImageOptionsBuilder withStyle(String style) {
options.setStyle(style);
return this;
}
public OpenAiImageOptionsBuilder withUser(String user) {
options.setUser(user);
return this;
}
public OpenAiImageOptions build() {
return options;
}
}

View File

@@ -0,0 +1,93 @@
package org.springframework.ai.openai.api;
public class OpenAiImageOptionsImpl implements OpenAiImageOptions {
private Integer n;
private String model;
private String quality;
private String responseFormat;
private Integer width;
private Integer height;
private String style;
private String user;
@Override
public Integer getN() {
return n;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public String getModel() {
return model;
}
public void setModel(String model) {
this.model = model;
}
@Override
public String getQuality() {
return quality;
}
public void setQuality(String quality) {
this.quality = quality;
}
@Override
public String getResponseFormat() {
return responseFormat;
}
public void setResponseFormat(String responseFormat) {
this.responseFormat = responseFormat;
}
@Override
public Integer getWidth() {
return width;
}
public void setWidth(Integer width) {
this.width = width;
}
@Override
public Integer getHeight() {
return height;
}
public void setHeight(Integer height) {
this.height = height;
}
@Override
public String getStyle() {
return style;
}
public void setStyle(String style) {
this.style = style;
}
@Override
public String getUser() {
return user;
}
public void setUser(String user) {
this.user = user;
}
}

View File

@@ -16,31 +16,31 @@
package org.springframework.ai.openai.metadata;
import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
/**
* {@link GenerationMetadata} implementation for {@literal OpenAI}.
* {@link ChatResponseMetadata} implementation for {@literal OpenAI}.
*
* @author John Blum
* @see org.springframework.ai.metadata.GenerationMetadata
* @see org.springframework.ai.metadata.RateLimit
* @see org.springframework.ai.metadata.Usage
* @see ChatResponseMetadata
* @see RateLimit
* @see Usage
* @since 0.7.0
*/
public class OpenAiGenerationMetadata implements GenerationMetadata {
public class OpenAiChatResponseMetadata implements ChatResponseMetadata {
protected static final String AI_METADATA_STRING = "{ @type: %1$s, id: %2$s, usage: %3$s, rateLimit: %4$s }";
public static OpenAiGenerationMetadata from(OpenAiApi.ChatCompletion result) {
public static OpenAiChatResponseMetadata from(OpenAiApi.ChatCompletion result) {
Assert.notNull(result, "OpenAI ChatCompletionResult must not be null");
OpenAiUsage usage = OpenAiUsage.from(result.usage());
OpenAiGenerationMetadata generationMetadata = new OpenAiGenerationMetadata(result.id(), usage);
return generationMetadata;
OpenAiChatResponseMetadata chatResponseMetadata = new OpenAiChatResponseMetadata(result.id(), usage);
return chatResponseMetadata;
}
private final String id;
@@ -50,11 +50,11 @@ public class OpenAiGenerationMetadata implements GenerationMetadata {
private final Usage usage;
protected OpenAiGenerationMetadata(String id, OpenAiUsage usage) {
protected OpenAiChatResponseMetadata(String id, OpenAiUsage usage) {
this(id, usage, null);
}
protected OpenAiGenerationMetadata(String id, OpenAiUsage usage, @Nullable OpenAiRateLimit rateLimit) {
protected OpenAiChatResponseMetadata(String id, OpenAiUsage usage, @Nullable OpenAiRateLimit rateLimit) {
this.id = id;
this.usage = usage;
this.rateLimit = rateLimit;
@@ -77,7 +77,7 @@ public class OpenAiGenerationMetadata implements GenerationMetadata {
return usage != null ? usage : Usage.NULL;
}
public OpenAiGenerationMetadata withRateLimit(RateLimit rateLimit) {
public OpenAiChatResponseMetadata withRateLimit(RateLimit rateLimit) {
this.rateLimit = rateLimit;
return this;
}

View File

@@ -0,0 +1,54 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai.metadata;
import org.springframework.ai.image.ImageGenerationMetadata;
import java.util.Objects;
public class OpenAiImageGenerationMetadata implements ImageGenerationMetadata {
private String revisedPrompt;
public OpenAiImageGenerationMetadata(String revisedPrompt) {
this.revisedPrompt = revisedPrompt;
}
public String getRevisedPrompt() {
return revisedPrompt;
}
@Override
public String toString() {
return "OpenAiImageGenerationMetadata{" + "revisedPrompt='" + revisedPrompt + '\'' + '}';
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof OpenAiImageGenerationMetadata that))
return false;
return Objects.equals(revisedPrompt, that.revisedPrompt);
}
@Override
public int hashCode() {
return Objects.hash(revisedPrompt);
}
}

View File

@@ -0,0 +1,62 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai.metadata;
import org.springframework.ai.image.ImageResponseMetadata;
import org.springframework.ai.openai.api.OpenAiImageApi;
import org.springframework.util.Assert;
import java.util.Objects;
public class OpenAiImageResponseMetadata implements ImageResponseMetadata {
private final Long created;
public static OpenAiImageResponseMetadata from(OpenAiImageApi.OpenAiImageResponse openAiImageResponse) {
Assert.notNull(openAiImageResponse, "OpenAiImageResponse must not be null");
return new OpenAiImageResponseMetadata(openAiImageResponse.created());
}
protected OpenAiImageResponseMetadata(Long created) {
this.created = created;
}
@Override
public Long created() {
return this.created;
}
@Override
public String toString() {
return "OpenAiImageResponseMetadata{" + "created=" + created + '}';
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof OpenAiImageResponseMetadata that))
return false;
return Objects.equals(created, that.created);
}
@Override
public int hashCode() {
return Objects.hash(created);
}
}

View File

@@ -18,7 +18,7 @@ package org.springframework.ai.openai.metadata;
import java.time.Duration;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.chat.metadata.RateLimit;
/**
* {@link RateLimit} implementation for {@literal OpenAI}.

View File

@@ -16,7 +16,7 @@
package org.springframework.ai.openai.metadata;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.Usage;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.util.Assert;

View File

@@ -26,7 +26,7 @@ import java.util.regex.Pattern;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.chat.metadata.RateLimit;
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion;
import org.springframework.ai.openai.metadata.OpenAiRateLimit;
import org.springframework.http.ResponseEntity;

View File

@@ -2,6 +2,7 @@ package org.springframework.ai.openai;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.api.OpenAiImageApi;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.context.annotation.Bean;
import org.springframework.util.StringUtils;
@@ -11,9 +12,12 @@ public class OpenAiTestConfiguration {
@Bean
public OpenAiApi openAiApi() {
String apiKey = getApiKey();
OpenAiApi openAiService = new OpenAiApi(apiKey);
return openAiService;
return new OpenAiApi(getApiKey());
}
@Bean
public OpenAiImageApi openAiImageApi() {
return new OpenAiImageApi(getApiKey());
}
private String getApiKey() {
@@ -32,6 +36,13 @@ public class OpenAiTestConfiguration {
return openAiChatClient;
}
@Bean
public OpenAiImageClient openAiImageClient(OpenAiImageApi imageApi) {
OpenAiImageClient openAiImageClient = new OpenAiImageClient(imageApi);
// openAiImageClient.setModel("foobar");
return openAiImageClient;
}
@Bean
public EmbeddingClient openAiEmbeddingClient(OpenAiApi api) {
return new OpenAiEmbeddingClient(api);

View File

@@ -15,10 +15,10 @@ import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.OpenAiChatClient;
import org.springframework.ai.openai.OpenAiEmbeddingClient;
import org.springframework.ai.openai.testutils.AbstractIT;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.reader.JsonReader;
import org.springframework.ai.transformer.splitter.TokenTextSplitter;
import org.springframework.ai.vectorstore.SimpleVectorStore;
@@ -90,10 +90,10 @@ public class AcmeIT extends AbstractIT {
// Create the prompt ad-hoc for now, need to put in system message and user
// message via ChatPromptTemplate or some other message building mechanic;
logger.info("Asking AI model to reply to question.");
logger.info("Asking AI generative to reply to question.");
Prompt prompt = new Prompt(List.of(systemMessage, userMessage));
logger.info("AI responded.");
ChatResponse response = chatClient.generate(prompt);
ChatResponse response = chatClient.call(prompt);
evaluateQuestionAndAnswer(userQuery, response, true);
}

View File

@@ -10,16 +10,17 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.testutils.AbstractIT;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.core.convert.support.DefaultConversionService;
@@ -41,9 +42,9 @@ class OpenAiChatClientIT extends AbstractIT {
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource);
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate"));
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = openAiChatClient.generate(prompt);
assertThat(response.getGenerations()).hasSize(1);
assertThat(response.getGenerations().get(0).getContent()).contains("Blackbeard");
ChatResponse response = openAiChatClient.call(prompt);
assertThat(response.getResults()).hasSize(1);
assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard");
// needs fine tuning... evaluateQuestionAndAnswer(request, response, false);
}
@@ -60,9 +61,9 @@ class OpenAiChatClientIT extends AbstractIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.openAiChatClient.generate(prompt).getGeneration();
Generation generation = this.openAiChatClient.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -79,9 +80,9 @@ class OpenAiChatClientIT extends AbstractIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = openAiChatClient.generate(prompt).getGeneration();
Generation generation = openAiChatClient.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -98,9 +99,9 @@ class OpenAiChatClientIT extends AbstractIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = openAiChatClient.generate(prompt).getGeneration();
Generation generation = openAiChatClient.call(prompt).getResult();
ActorsFilms actorsFilms = outputParser.parse(generation.getContent());
ActorsFilms actorsFilms = outputParser.parse(generation.getOutput().getContent());
}
record ActorsFilmsRecord(String actor, List<String> movies) {
@@ -118,9 +119,9 @@ class OpenAiChatClientIT extends AbstractIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = openAiChatClient.generate(prompt).getGeneration();
Generation generation = openAiChatClient.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
System.out.println(actorsFilms);
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
@@ -143,9 +144,10 @@ class OpenAiChatClientIT extends AbstractIT {
.collectList()
.block()
.stream()
.map(ChatResponse::getGenerations)
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getContent)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);

View File

@@ -22,15 +22,11 @@ import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.PromptMetadata;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.metadata.Usage;
import org.springframework.ai.chat.metadata.*;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.OpenAiChatClient;
import org.springframework.ai.openai.metadata.support.OpenAiApiResponseHeaders;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.autoconfigure.web.client.RestClientTest;
@@ -52,8 +48,8 @@ import static org.springframework.test.web.client.response.MockRestResponseCreat
* @author Christian Tzolov
* @since 0.7.0
*/
@RestClientTest(OpenAiChatClientWithGenerationMetadataTests.Config.class)
public class OpenAiChatClientWithGenerationMetadataTests {
@RestClientTest(OpenAiChatClientWithChatResponseMetadataTests.Config.class)
public class OpenAiChatClientWithChatResponseMetadataTests {
private static String TEST_API_KEY = "sk-1234567890";
@@ -75,22 +71,22 @@ public class OpenAiChatClientWithGenerationMetadataTests {
Prompt prompt = new Prompt("Reach for the sky.");
ChatResponse response = this.openAiChatClient.generate(prompt);
ChatResponse response = this.openAiChatClient.call(prompt);
assertThat(response).isNotNull();
GenerationMetadata generationMetadata = response.getGenerationMetadata();
ChatResponseMetadata chatResponseMetadata = response.getMetadata();
assertThat(generationMetadata).isNotNull();
assertThat(chatResponseMetadata).isNotNull();
Usage usage = generationMetadata.getUsage();
Usage usage = chatResponseMetadata.getUsage();
assertThat(usage).isNotNull();
assertThat(usage.getPromptTokens()).isEqualTo(9L);
assertThat(usage.getGenerationTokens()).isEqualTo(12L);
assertThat(usage.getTotalTokens()).isEqualTo(21L);
RateLimit rateLimit = generationMetadata.getRateLimit();
RateLimit rateLimit = chatResponseMetadata.getRateLimit();
Duration expectedRequestsReset = Duration.ofDays(2L)
.plus(Duration.ofHours(16L))
@@ -109,16 +105,16 @@ public class OpenAiChatClientWithGenerationMetadataTests {
assertThat(rateLimit.getTokensRemaining()).isEqualTo(112_358L);
assertThat(rateLimit.getTokensReset()).isEqualTo(expectedTokensReset);
PromptMetadata promptMetadata = response.getPromptMetadata();
PromptMetadata promptMetadata = response.getMetadata().getPromptMetadata();
assertThat(promptMetadata).isNotNull();
assertThat(promptMetadata).isEmpty();
response.getGenerations().forEach(generation -> {
ChoiceMetadata choiceMetadata = generation.getChoiceMetadata();
assertThat(choiceMetadata).isNotNull();
assertThat(choiceMetadata.getFinishReason()).isEqualTo("STOP");
assertThat(choiceMetadata.<Object>getContentFilterMetadata()).isNull();
response.getResults().forEach(generation -> {
ChatGenerationMetadata chatGenerationMetadata = generation.getMetadata();
assertThat(chatGenerationMetadata).isNotNull();
assertThat(chatGenerationMetadata.getFinishReason()).isEqualTo("STOP");
assertThat(chatGenerationMetadata.<Object>getContentFilterMetadata()).isNull();
});
}

View File

@@ -38,7 +38,7 @@ class EmbeddingIT {
EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World"));
assertThat(embeddingResponse.getData()).hasSize(1);
assertThat(embeddingResponse.getData().get(0).getEmbedding()).isNotEmpty();
assertThat(embeddingResponse.getMetadata()).containsEntry("model", "text-embedding-ada-002-v2");
assertThat(embeddingResponse.getMetadata()).containsEntry("generative", "text-embedding-ada-002-v2");
assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2);
assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2);

View File

@@ -0,0 +1,60 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai.image;
import org.assertj.core.api.Assertions;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.image.*;
import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.metadata.OpenAiImageGenerationMetadata;
import org.springframework.ai.openai.testutils.AbstractIT;
import org.springframework.boot.test.context.SpringBootTest;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest(classes = OpenAiTestConfiguration.class)
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class OpenAiImageClientIT extends AbstractIT {
@Test
void imageAsUrlTest() {
var options = ImageOptionsBuilder.builder().withHeight(256).withWidth(256).build();
ImagePrompt imagePrompt = new ImagePrompt("Create an image of a mini golden doodle dog.", options);
ImageResponse imageResponse = openaiImageClient.call(imagePrompt);
assertThat(imageResponse.getResults()).hasSize(1);
ImageResponseMetadata imageResponseMetadata = imageResponse.getMetadata();
assertThat(imageResponseMetadata.created()).isPositive();
var generation = imageResponse.getResult();
Image image = generation.getOutput();
assertThat(image.getUrl()).isNotEmpty();
assertThat(image.getB64Json()).isNull();
var imageGenerationMetadata = generation.getMetadata();
Assertions.assertThat(imageGenerationMetadata).isInstanceOf(OpenAiImageGenerationMetadata.class);
OpenAiImageGenerationMetadata openAiImageGenerationMetadata = (OpenAiImageGenerationMetadata) imageGenerationMetadata;
assertThat(openAiImageGenerationMetadata).isNotNull();
assertThat(openAiImageGenerationMetadata.getRevisedPrompt()).isNull();
}
}

View File

@@ -0,0 +1,142 @@
/*
* Copyright 2023-2023 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.openai.image;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
import org.springframework.ai.image.ImageGeneration;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.image.ImageResponse;
import org.springframework.ai.image.ImageResponseMetadata;
import org.springframework.ai.openai.OpenAiImageClient;
import org.springframework.ai.openai.api.OpenAiImageApi;
import org.springframework.ai.openai.metadata.support.OpenAiApiResponseHeaders;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.test.autoconfigure.web.client.RestClientTest;
import org.springframework.context.annotation.Bean;
import org.springframework.http.HttpHeaders;
import org.springframework.http.HttpMethod;
import org.springframework.http.MediaType;
import org.springframework.test.web.client.MockRestServiceServer;
import org.springframework.web.client.RestClient;
import java.util.List;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.test.web.client.match.MockRestRequestMatchers.*;
import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess;
/**
* @author John Blum
* @author Christian Tzolov
* @since 0.7.0
*/
@RestClientTest(OpenAiImageClientWithImageResponseMetadataTests.Config.class)
public class OpenAiImageClientWithImageResponseMetadataTests {
private static String TEST_API_KEY = "sk-1234567890";
@Autowired
private OpenAiImageClient openAiImageClient;
@Autowired
private MockRestServiceServer server;
@AfterEach
void resetMockServer() {
server.reset();
}
@Test
void aiResponseContainsImageResponseMetadata() {
prepareMock();
ImagePrompt prompt = new ImagePrompt("Create an image of a mini golden doodle dog.");
ImageResponse response = this.openAiImageClient.call(prompt);
assertThat(response).isNotNull();
List<ImageGeneration> imageGenerations = response.getResults();
assertThat(imageGenerations).isNotNull();
assertThat(imageGenerations).hasSize(2);
ImageResponseMetadata imageResponseMetadata = response.getMetadata();
assertThat(imageResponseMetadata).isNotNull();
Long created = imageResponseMetadata.created();
assertThat(created).isNotNull();
assertThat(created).isEqualTo(1589478378);
ImageResponseMetadata responseMetadata = response.getMetadata();
assertThat(responseMetadata).isNotNull();
}
private void prepareMock() {
HttpHeaders httpHeaders = new HttpHeaders();
httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_LIMIT_HEADER.getName(), "4000");
httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_REMAINING_HEADER.getName(), "999");
httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_RESET_HEADER.getName(), "2d16h15m29s");
httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_LIMIT_HEADER.getName(), "725000");
httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_REMAINING_HEADER.getName(), "112358");
httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_RESET_HEADER.getName(), "27h55s451ms");
server.expect(requestTo("v1/images/generations"))
.andExpect(method(HttpMethod.POST))
.andExpect(header(HttpHeaders.AUTHORIZATION, "Bearer " + TEST_API_KEY))
.andRespond(withSuccess(getJson(), MediaType.APPLICATION_JSON).headers(httpHeaders));
}
private String getJson() {
return """
{
"created": 1589478378,
"data": [
{
"url": "https://upload.wikimedia.org/wikipedia/commons/4/4e/Mini_Golden_Doodle.jpg"
},
{
"url": "https://upload.wikimedia.org/wikipedia/commons/8/85/Goldendoodle_puppy_Marty.jpg"
}
]
}
""";
}
@SpringBootConfiguration
static class Config {
@Bean
public OpenAiImageApi imageGenerationApi(RestClient.Builder builder) {
return new OpenAiImageApi("", TEST_API_KEY, builder);
}
@Bean
public OpenAiImageClient openAiImageClient(OpenAiImageApi openAiImageApi) {
return new OpenAiImageClient(openAiImageApi);
}
}
}

View File

@@ -9,10 +9,11 @@ import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.StreamingChatClient;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.SystemMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.image.ImageClient;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.core.io.Resource;
@@ -27,6 +28,9 @@ public abstract class AbstractIT {
@Autowired
protected ChatClient openAiChatClient;
@Autowired
protected ImageClient openaiImageClient;
@Autowired
protected StreamingChatClient openStreamingChatClient;
@@ -44,7 +48,7 @@ public abstract class AbstractIT {
protected void evaluateQuestionAndAnswer(String question, ChatResponse response, boolean factBased) {
assertThat(response).isNotNull();
String answer = response.getGeneration().getContent();
String answer = response.getResult().getOutput().getContent();
logger.info("Question: " + question);
logger.info("Answer:" + answer);
PromptTemplate userPromptTemplate = new PromptTemplate(userEvaluatorResource,
@@ -58,12 +62,12 @@ public abstract class AbstractIT {
}
Message userMessage = userPromptTemplate.createMessage();
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
String yesOrNo = openAiChatClient.generate(prompt).getGeneration().getContent();
String yesOrNo = openAiChatClient.call(prompt).getResult().getOutput().getContent();
logger.info("Is Answer related to question: " + yesOrNo);
if (yesOrNo.equalsIgnoreCase("no")) {
SystemMessage notRelatedSystemMessage = new SystemMessage(qaEvaluatorNotRelatedResource);
prompt = new Prompt(List.of(userMessage, notRelatedSystemMessage));
String reasonForFailure = openAiChatClient.generate(prompt).getGeneration().getContent();
String reasonForFailure = openAiChatClient.call(prompt).getResult().getOutput().getContent();
fail(reasonForFailure);
}
else {

View File

@@ -65,7 +65,7 @@ public class MetadataTransformerIT {
Document document2 = new Document(
"The Spring Framework is divided into modules. Applications can choose which modules"
+ " they need. At the heart are the modules of the core container, including a configuration model and a "
+ " they need. At the heart are the modules of the core container, including a configuration generative and a "
+ "dependency injection mechanism. Beyond that, the Spring Framework provides foundational support "
+ " for different application architectures, including messaging, transactional data and persistence, "
+ "and web. It also includes the Servlet-based Spring MVC web framework and, in parallel, the Spring "

View File

@@ -0,0 +1,60 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance" xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/maven-v4_0_0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai</artifactId>
<version>0.8.0-SNAPSHOT</version>
<relativePath>../../pom.xml</relativePath>
</parent>
<artifactId>spring-ai-stability-ai</artifactId>
<packaging>jar</packaging>
<name>Spring AI Stability AI</name>
<description>Stability AI support</description>
<url>https://github.com/spring-projects/spring-ai</url>
<scm>
<url>https://github.com/spring-projects/spring-ai</url>
<connection>git://github.com/spring-projects/spring-ai.git</connection>
<developerConnection>git@github.com:spring-projects/spring-ai.git</developerConnection>
</scm>
<dependencies>
<!-- production dependencies -->
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-core</artifactId>
<version>${project.parent.version}</version>
</dependency>
<dependency>
<groupId>org.springframework</groupId>
<artifactId>spring-web</artifactId>
<version>${spring-framework.version}</version>
</dependency>
<!-- Spring Framework -->
<dependency>
<groupId>org.springframework</groupId>
<artifactId>spring-context-support</artifactId>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-logging</artifactId>
</dependency>
<!-- test dependencies -->
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-test</artifactId>
<version>${project.version}</version>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -0,0 +1,158 @@
package org.springframework.ai.stabilityai;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.image.*;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptions;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptionsBuilder;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptionsImpl;
import org.springframework.util.Assert;
import java.util.List;
import java.util.stream.Collectors;
public class StabilityAiImageClient implements ImageClient {
private final Logger logger = LoggerFactory.getLogger(getClass());
private StabilityAiImageOptions options;
private final StabilityAiApi stabilityAiApi;
public StabilityAiImageClient(StabilityAiApi stabilityAiApi) {
this(stabilityAiApi, StabilityAiImageOptionsBuilder.builder().build());
}
public StabilityAiImageClient(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options) {
Assert.notNull(stabilityAiApi, "StabilityAiApi must not be null");
Assert.notNull(options, "StabilityAiImageOptions must not be null");
this.stabilityAiApi = stabilityAiApi;
this.options = options;
}
public StabilityAiImageOptions getOptions() {
return options;
}
/**
* Calls the StabilityAiImageClient with the given StabilityAiImagePrompt and returns
* the ImageResponse. This overloaded call method lets you pass the full set of Prompt
* instructions that StabilityAI supports.
* @param imagePrompt the StabilityAiImagePrompt containing the prompt and image model
* options
* @return the ImageResponse generated by the StabilityAiImageClient
*/
public ImageResponse call(ImagePrompt imagePrompt) {
ImageOptions runtimeOptions = imagePrompt.getOptions();
// Merge the runtime options passed via the prompt with the StabilityAiImageClient
// options configured via Autoconfiguration.
// Runtime options overwrite StabilityAiImageClient options
StabilityAiImageOptions optionsToUse = ModelOptionsUtils.merge(runtimeOptions, this.options,
StabilityAiImageOptionsImpl.class);
// Copy the org.springframework.ai.model derived ImagePrompt and ImageOptions data
// types to the data types used in StabilityAiApi
StabilityAiApi.GenerateImageRequest generateImageRequest = getGenerateImageRequest(imagePrompt, optionsToUse);
// Make the request
StabilityAiApi.GenerateImageResponse generateImageResponse = this.stabilityAiApi
.generateImage(generateImageRequest);
// Convert to org.springframework.ai.model derived ImageResponse data type
return convertResponse(generateImageResponse);
}
private static StabilityAiApi.GenerateImageRequest getGenerateImageRequest(ImagePrompt stabilityAiImagePrompt,
StabilityAiImageOptions optionsToUse) {
StabilityAiApi.GenerateImageRequest.Builder builder = new StabilityAiApi.GenerateImageRequest.Builder();
StabilityAiApi.GenerateImageRequest generateImageRequest = builder
.withTextPrompts(stabilityAiImagePrompt.getInstructions()
.stream()
.map(message -> new StabilityAiApi.GenerateImageRequest.TextPrompts(message.getText(),
message.getWeight()))
.collect(Collectors.toList()))
.withHeight(optionsToUse.getHeight())
.withWidth(optionsToUse.getWidth())
.withCfgScale(optionsToUse.getCfgScale())
.withClipGuidancePreset(optionsToUse.getClipGuidancePreset())
.withSampler(optionsToUse.getSampler())
.withSamples(optionsToUse.getSamples())
.withSeed(optionsToUse.getSeed())
.withSteps(optionsToUse.getSteps())
.withStylePreset(optionsToUse.getStylePreset())
.build();
return generateImageRequest;
}
private ImageResponse convertResponse(StabilityAiApi.GenerateImageResponse generateImageResponse) {
List<ImageGeneration> imageGenerationList = generateImageResponse.artifacts().stream().map(entry -> {
return new ImageGeneration(new Image(null, entry.base64()),
new StabilityAiImageGenerationMetadata(entry.finishReason(), entry.seed()));
}).toList();
return new ImageResponse(imageGenerationList, ImageResponseMetadata.NULL);
}
private StabilityAiImageOptions convertOptions(ImageOptions runtimeOptions) {
StabilityAiImageOptionsBuilder builder = StabilityAiImageOptionsBuilder.builder();
if (runtimeOptions == null) {
return builder.build();
}
if (runtimeOptions.getN() != null) {
builder.withN(runtimeOptions.getN());
}
if (runtimeOptions.getModel() != null) {
builder.withModel(runtimeOptions.getModel());
}
if (runtimeOptions.getResponseFormat() != null) {
builder.withResponseFormat(runtimeOptions.getResponseFormat());
}
if (runtimeOptions.getWidth() != null) {
builder.withWidth(runtimeOptions.getWidth());
}
if (runtimeOptions.getHeight() != null) {
builder.withHeight(runtimeOptions.getHeight());
}
if (runtimeOptions instanceof StabilityAiImageOptions) {
StabilityAiImageOptions stabilityAiImageOptions = (StabilityAiImageOptions) runtimeOptions;
if (stabilityAiImageOptions.getCfgScale() != null) {
builder.withCfgScale(stabilityAiImageOptions.getCfgScale());
}
if (stabilityAiImageOptions.getClipGuidancePreset() != null) {
builder.withClipGuidancePreset(stabilityAiImageOptions.getClipGuidancePreset());
}
if (stabilityAiImageOptions.getSampler() != null) {
builder.withSampler(stabilityAiImageOptions.getSampler());
}
if (stabilityAiImageOptions.getSeed() != null) {
builder.withSeed(stabilityAiImageOptions.getSeed());
}
if (stabilityAiImageOptions.getSteps() != null) {
builder.withSteps(stabilityAiImageOptions.getSteps());
}
if (stabilityAiImageOptions.getStylePreset() != null) {
builder.withStylePreset(stabilityAiImageOptions.getStylePreset());
}
}
return builder.build();
}
private ImagePrompt createUpdatedPrompt(ImagePrompt prompt) {
ImageOptions runtimeImageModelOptions = prompt.getOptions();
ImageOptionsBuilder imageOptionsBuilder = ImageOptionsBuilder.builder();
if (runtimeImageModelOptions != null) {
if (runtimeImageModelOptions.getModel() != null) {
imageOptionsBuilder.withModel(runtimeImageModelOptions.getModel());
}
}
ImageOptions updatedImageModelOptions = imageOptionsBuilder.build();
return new ImagePrompt(prompt.getInstructions(), updatedImageModelOptions);
}
}

View File

@@ -0,0 +1,45 @@
package org.springframework.ai.stabilityai;
import org.springframework.ai.image.ImageGenerationMetadata;
import java.util.Objects;
public class StabilityAiImageGenerationMetadata implements ImageGenerationMetadata {
private String finishReason;
private Long seed;
public StabilityAiImageGenerationMetadata(String finishReason, Long seed) {
this.finishReason = finishReason;
this.seed = seed;
}
public String getFinishReason() {
return finishReason;
}
public Long getSeed() {
return seed;
}
@Override
public String toString() {
return "StabilityAiImageGenerationMetadata{" + "finishReason='" + finishReason + '\'' + ", seed=" + seed + '}';
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof StabilityAiImageGenerationMetadata that))
return false;
return Objects.equals(finishReason, that.finishReason) && Objects.equals(seed, that.seed);
}
@Override
public int hashCode() {
return Objects.hash(finishReason, seed);
}
}

View File

@@ -0,0 +1,212 @@
package org.springframework.ai.stabilityai.api;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.http.HttpHeaders;
import org.springframework.http.MediaType;
import org.springframework.http.client.ClientHttpResponse;
import org.springframework.util.Assert;
import org.springframework.web.client.ResponseErrorHandler;
import org.springframework.web.client.RestClient;
import java.io.IOException;
import java.util.List;
import java.util.function.Consumer;
public class StabilityAiApi {
public static final String DEFAULT_IMAGE_MODEL = "stable-diffusion-v1-6";
public static final String DEFAULT_BASE_URL = "https://api.stability.ai/v1";
private final RestClient restClient;
private final String apiKey;
private final String model;
/**
* Create a new StabilityAI API.
* @param apiKey StabilityAI apiKey.
*/
public StabilityAiApi(String apiKey) {
this(apiKey, DEFAULT_IMAGE_MODEL, DEFAULT_BASE_URL, RestClient.builder());
}
public StabilityAiApi(String apiKey, String model) {
this(apiKey, model, DEFAULT_BASE_URL, RestClient.builder());
}
public StabilityAiApi(String apiKey, String model, String baseUrl) {
this(apiKey, model, baseUrl, RestClient.builder());
}
/**
* Create a new StabilityAI API.
* @param apiKey StabilityAI apiKey.
* @param model StabilityAI model.
* @param baseUrl api base URL.
* @param restClientBuilder RestClient builder.
*/
public StabilityAiApi(String apiKey, String model, String baseUrl, RestClient.Builder restClientBuilder) {
this.model = model;
this.apiKey = apiKey;
Consumer<HttpHeaders> jsonContentHeaders = headers -> {
headers.setBearerAuth(apiKey);
headers.setAccept(List.of(MediaType.APPLICATION_JSON)); // base64 in JSON +
// metadata or return
// image in bytes.
headers.setContentType(MediaType.APPLICATION_JSON);
};
ResponseErrorHandler responseErrorHandler = new ResponseErrorHandler() {
@Override
public boolean hasError(ClientHttpResponse response) throws IOException {
return response.getStatusCode().isError();
}
@Override
public void handleError(ClientHttpResponse response) throws IOException {
if (response.getStatusCode().isError()) {
throw new RuntimeException(String.format("%s - %s", response.getStatusCode().value(),
new ObjectMapper().readValue(response.getBody(), ResponseError.class)));
}
}
};
this.restClient = restClientBuilder.baseUrl(baseUrl)
.defaultHeaders(jsonContentHeaders)
.defaultStatusHandler(responseErrorHandler)
.build();
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public record ResponseError(@JsonProperty("id") String id, @JsonProperty("name") String name,
@JsonProperty("message") String message
) {
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public record GenerateImageRequest(@JsonProperty("text_prompts") List<TextPrompts> textPrompts,
@JsonProperty("height") Integer height, @JsonProperty("width") Integer width,
@JsonProperty("cfg_scale") Float cfgScale, @JsonProperty("clip_guidance_preset") String clipGuidancePreset,
@JsonProperty("sampler") String sampler, @JsonProperty("samples") Integer samples,
@JsonProperty("seed") Long seed, @JsonProperty("steps") Integer steps,
@JsonProperty("style_present") String stylePreset) {
@JsonInclude(JsonInclude.Include.NON_NULL)
public record TextPrompts(@JsonProperty("text") String text, @JsonProperty("weight") Float weight) {
}
public static Builder builder() {
return new Builder();
}
public static class Builder {
List<TextPrompts> textPrompts;
Integer height;
Integer width;
Float cfgScale;
String clipGuidancePreset;
String sampler;
Integer samples;
Long seed;
Integer steps;
String stylePreset;
public Builder() {
}
public Builder withTextPrompts(List<TextPrompts> textPrompts) {
this.textPrompts = textPrompts;
return this;
}
public Builder withHeight(Integer height) {
this.height = height;
return this;
}
public Builder withWidth(Integer width) {
this.width = width;
return this;
}
public Builder withCfgScale(Float cfgScale) {
this.cfgScale = cfgScale;
return this;
}
public Builder withClipGuidancePreset(String clipGuidancePreset) {
this.clipGuidancePreset = clipGuidancePreset;
return this;
}
public Builder withSampler(String sampler) {
this.sampler = sampler;
return this;
}
public Builder withSamples(Integer samples) {
this.samples = samples;
return this;
}
public Builder withSeed(Long seed) {
this.seed = seed;
return this;
}
public Builder withSteps(Integer steps) {
this.steps = steps;
return this;
}
public Builder withStylePreset(String stylePreset) {
this.stylePreset = stylePreset;
return this;
}
public GenerateImageRequest build() {
return new GenerateImageRequest(textPrompts, height, width, cfgScale, clipGuidancePreset, sampler,
samples, seed, steps, stylePreset);
}
}
}
@JsonInclude(JsonInclude.Include.NON_NULL)
public record GenerateImageResponse(@JsonProperty("result") String result,
@JsonProperty("artifacts") List<Artifacts> artifacts) {
public record Artifacts(@JsonProperty("seed") long seed, @JsonProperty("base64") String base64,
@JsonProperty("finishReason") String finishReason) {
}
}
public GenerateImageResponse generateImage(GenerateImageRequest request) {
Assert.notNull(request, "The request body can not be null.");
return this.restClient.post()
.uri("/generation/{model}/text-to-image", this.model)
.body(request)
.retrieve()
.body(GenerateImageResponse.class);
}
}

View File

@@ -0,0 +1,23 @@
package org.springframework.ai.stabilityai.api;
import org.springframework.ai.image.ImageOptions;
public interface StabilityAiImageOptions extends ImageOptions {
Float getCfgScale();
String getClipGuidancePreset();
String getSampler();
Integer getSamples();
Long getSeed();
Integer getSteps();
String getStylePreset();
// extras json object...
}

View File

@@ -0,0 +1,79 @@
package org.springframework.ai.stabilityai.api;
public class StabilityAiImageOptionsBuilder {
private StabilityAiImageOptionsImpl options;
private StabilityAiImageOptionsBuilder() {
options = new StabilityAiImageOptionsImpl();
}
public static StabilityAiImageOptionsBuilder builder() {
return new StabilityAiImageOptionsBuilder();
}
public StabilityAiImageOptionsBuilder withN(Integer n) {
options.setN(n);
return this;
}
public StabilityAiImageOptionsBuilder withModel(String model) {
options.setModel(model);
return this;
}
public StabilityAiImageOptionsBuilder withWidth(Integer width) {
options.setWidth(width);
return this;
}
public StabilityAiImageOptionsBuilder withHeight(Integer height) {
options.setHeight(height);
return this;
}
public StabilityAiImageOptionsBuilder withResponseFormat(String responseFormat) {
options.setResponseFormat(responseFormat);
return this;
}
public StabilityAiImageOptionsBuilder withCfgScale(Float cfgScale) {
options.setCfgScale(cfgScale);
return this;
}
public StabilityAiImageOptionsBuilder withClipGuidancePreset(String clipGuidancePreset) {
options.setClipGuidancePreset(clipGuidancePreset);
return this;
}
public StabilityAiImageOptionsBuilder withSampler(String sampler) {
options.setSampler(sampler);
return this;
}
public StabilityAiImageOptionsBuilder withSeed(Long seed) {
options.setSeed(seed);
return this;
}
public StabilityAiImageOptionsBuilder withSteps(Integer steps) {
options.setSteps(steps);
return this;
}
public StabilityAiImageOptionsBuilder withSamples(Integer samples) {
options.setSamples(samples);
return this;
}
public StabilityAiImageOptionsBuilder withStylePreset(String stylePreset) {
options.setStylePreset(stylePreset);
return this;
}
public StabilityAiImageOptions build() {
return options;
}
}

View File

@@ -0,0 +1,140 @@
package org.springframework.ai.stabilityai.api;
public class StabilityAiImageOptionsImpl implements StabilityAiImageOptions {
private Integer n;
private String model;
private Integer width;
private Integer height;
private String responseFormat;
private Float cfgScale;
private String clipGuidancePreset;
private String sampler;
private Integer samples;
private Long seed;
private Integer steps;
private String stylePreset;
public StabilityAiImageOptionsImpl() {
}
@Override
public Integer getN() {
return this.n;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public String getModel() {
return this.model;
}
public void setModel(String model) {
this.model = model;
}
@Override
public Integer getWidth() {
return this.width;
}
public void setWidth(Integer width) {
this.width = width;
}
@Override
public Integer getHeight() {
return this.height;
}
public void setHeight(Integer height) {
this.height = height;
}
@Override
public String getResponseFormat() {
return this.responseFormat;
}
public void setResponseFormat(String responseFormat) {
this.responseFormat = responseFormat;
}
@Override
public Float getCfgScale() {
return this.cfgScale;
}
public void setCfgScale(Float cfgScale) {
this.cfgScale = cfgScale;
}
@Override
public String getClipGuidancePreset() {
return this.clipGuidancePreset;
}
public void setClipGuidancePreset(String clipGuidancePreset) {
this.clipGuidancePreset = clipGuidancePreset;
}
@Override
public String getSampler() {
return this.sampler;
}
public void setSampler(String sampler) {
this.sampler = sampler;
}
@Override
public Integer getSamples() {
return this.samples;
}
public void setSamples(Integer samples) {
this.samples = samples;
}
@Override
public Long getSeed() {
return this.seed;
}
public void setSeed(Long seed) {
this.seed = seed;
}
@Override
public Integer getSteps() {
return this.steps;
}
public void setSteps(Integer steps) {
this.steps = steps;
}
@Override
public String getStylePreset() {
return this.stylePreset;
}
public void setStylePreset(String stylePreset) {
this.stylePreset = stylePreset;
}
}

View File

@@ -0,0 +1,64 @@
package org.springframework.ai.stabilityai;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import java.io.File;
import java.io.FileOutputStream;
import java.io.IOException;
import java.util.Base64;
import java.util.List;
import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "STABILITYAI_API_KEY", matches = ".*")
public class StabilityAiApiIT {
StabilityAiApi stabilityAiApi = new StabilityAiApi(System.getenv("STABILITYAI_API_KEY"));
@Test
void generateImage() throws IOException {
List<StabilityAiApi.GenerateImageRequest.TextPrompts> textPrompts = List
.of(new StabilityAiApi.GenerateImageRequest.TextPrompts(
"A light cream colored mini golden doodle holding a sign that says 'Heading to BARCADE !'", 0.5f));
var builder = StabilityAiApi.GenerateImageRequest.builder()
.withTextPrompts(textPrompts)
.withHeight(1024)
.withWidth(1024)
.withCfgScale(7f)
.withSamples(1)
.withSeed(123L)
.withSteps(30)
.withStylePreset("photographic");
StabilityAiApi.GenerateImageRequest request = builder.build();
StabilityAiApi.GenerateImageResponse response = stabilityAiApi.generateImage(request);
assertThat(response).isNotNull();
List<StabilityAiApi.GenerateImageResponse.Artifacts> artifacts = response.artifacts();
writeToFile(artifacts);
assertThat(artifacts).hasSize(1);
var firstArtifact = artifacts.get(0);
assertThat(firstArtifact.base64()).isNotEmpty();
assertThat(firstArtifact.seed()).isPositive();
assertThat(firstArtifact.finishReason()).isEqualTo("SUCCESS");
}
private static void writeToFile(List<StabilityAiApi.GenerateImageResponse.Artifacts> artifacts) throws IOException {
int counter = 0;
String systemTempDir = System.getProperty("java.io.tmpdir");
for (StabilityAiApi.GenerateImageResponse.Artifacts artifact : artifacts) {
counter++;
byte[] imageBytes = Base64.getDecoder().decode(artifact.base64());
String fileName = String.format("dog%d.png", counter);
String filePath = systemTempDir + File.separator + fileName;
File file = new File(filePath);
try (FileOutputStream fos = new FileOutputStream(file)) {
fos.write(imageBytes);
}
}
}
}

View File

@@ -0,0 +1,50 @@
package org.springframework.ai.stabilityai;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.image.*;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import java.io.File;
import java.io.FileNotFoundException;
import java.io.FileOutputStream;
import java.io.IOException;
import java.util.Base64;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest(classes = StabilityAiImageTestConfiguration.class)
@EnabledIfEnvironmentVariable(named = "STABILITYAI_API_KEY", matches = ".*")
public class StabilityAiImageClientIT {
@Autowired
protected ImageClient stabilityAiImageClient;
@Test
void imageAsBase64Test() throws IOException {
ImagePrompt imagePrompt = new ImagePrompt(
"A light cream colored mini golden doodle holding a sign that says 'I want to go with you on vacation!'");
ImageResponse imageResponse = this.stabilityAiImageClient.call(imagePrompt);
ImageGeneration imageGeneration = imageResponse.getResult();
Image image = imageGeneration.getOutput();
assertThat(image.getB64Json()).isNotEmpty();
writeFile(image);
}
private static void writeFile(Image image) throws IOException {
byte[] imageBytes = Base64.getDecoder().decode(image.getB64Json());
String systemTempDir = System.getProperty("java.io.tmpdir");
String filePath = systemTempDir + File.separator + "dog.png";
File file = new File(filePath);
try (FileOutputStream fos = new FileOutputStream(file)) {
fos.write(imageBytes);
}
}
}

View File

@@ -0,0 +1,30 @@
package org.springframework.ai.stabilityai;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.context.annotation.Bean;
import org.springframework.util.StringUtils;
@SpringBootConfiguration
public class StabilityAiImageTestConfiguration {
@Bean
public StabilityAiApi stabilityAiApi() {
return new StabilityAiApi(getApiKey());
}
@Bean
StabilityAiImageClient stabilityAiImageClient(StabilityAiApi stabilityAiApi) {
return new StabilityAiImageClient(stabilityAiApi);
}
private String getApiKey() {
String apiKey = System.getenv("STABILITYAI_API_KEY");
if (!StringUtils.hasText(apiKey)) {
throw new IllegalArgumentException(
"You must provide an API key. Put it in an environment variable under the name STABILITYAI_API_KEY");
}
return apiKey;
}
}

View File

@@ -56,7 +56,7 @@ public class ResourceCacheService {
private List<String> excludedUriSchemas = new ArrayList<>(List.of("file", "classpath"));
public ResourceCacheService() {
this(new File(System.getProperty("java.io.tmpdir"), "spring-ai-onnx-model").getAbsolutePath());
this(new File(System.getProperty("java.io.tmpdir"), "spring-ai-onnx-generative").getAbsolutePath());
}
public ResourceCacheService(String rootCacheDirectory) {

View File

@@ -41,10 +41,10 @@ public class TransformersEmbeddingClient extends AbstractEmbeddingClient impleme
private static final Log logger = LogFactory.getLog(TransformersEmbeddingClient.class);
// ONNX tokenizer for the all-MiniLM-L6-v2 model
// ONNX tokenizer for the all-MiniLM-L6-v2 generative
public final static String DEFAULT_ONNX_TOKENIZER_URI = "https://raw.githubusercontent.com/spring-projects/spring-ai/main/models/spring-ai-transformers/src/main/resources/onnx/all-MiniLM-L6-v2/tokenizer.json";
// ONNX model for all-MiniLM-L6-v2 pre-trained transformer:
// ONNX generative for all-MiniLM-L6-v2 pre-trained transformer:
// https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2
public final static String DEFAULT_ONNX_MODEL_URI = "https://github.com/spring-projects/spring-ai/raw/main/models/spring-ai-transformers/src/main/resources/onnx/all-MiniLM-L6-v2/model.onnx";
@@ -70,7 +70,7 @@ public class TransformersEmbeddingClient extends AbstractEmbeddingClient impleme
private OrtEnvironment environment;
/**
* Runtime session that wraps the ONNX model and enables inference calls.
* Runtime session that wraps the ONNX generative and enables inference calls.
*/
private OrtSession session;
@@ -181,7 +181,7 @@ public class TransformersEmbeddingClient extends AbstractEmbeddingClient impleme
logger.info("Model output names: " + onnxModelOutputs.stream().collect(Collectors.joining(", ")));
Assert.isTrue(onnxModelOutputs.contains(this.modelOutputName),
"The model output names doesn't contain expected: " + this.modelOutputName);
"The generative output names doesn't contain expected: " + this.modelOutputName);
}
private Resource getCachedResource(Resource resource) {

View File

@@ -58,7 +58,7 @@ public class ONNXSample {
public static void main(String[] args) throws Exception {
String TOKENIZER_URI = "classpath:/onnx/tokenizer.json";
String MODEL_URI = "classpath:/onnx/model.onnx";
String MODEL_URI = "classpath:/onnx/generative.onnx";
var tokenizerResource = new DefaultResourceLoader().getResource(TOKENIZER_URI);
var modelResource = new DefaultResourceLoader().getResource(MODEL_URI);

View File

@@ -22,8 +22,8 @@ import java.util.stream.Collectors;
import org.springframework.ai.chat.ChatClient;
import org.springframework.ai.chat.ChatResponse;
import org.springframework.ai.chat.Generation;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.MessageType;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.messages.MessageType;
import org.springframework.ai.vertex.api.VertexAiApi;
import org.springframework.ai.vertex.api.VertexAiApi.GenerateMessageRequest;
import org.springframework.ai.vertex.api.VertexAiApi.GenerateMessageResponse;
@@ -71,15 +71,15 @@ public class VertexAiChatClient implements ChatClient {
}
@Override
public ChatResponse generate(Prompt prompt) {
public ChatResponse call(Prompt prompt) {
String vertexContext = prompt.getMessages()
String vertexContext = prompt.getInstructions()
.stream()
.filter(m -> m.getMessageType() == MessageType.SYSTEM)
.map(m -> m.getContent())
.collect(Collectors.joining("\n"));
List<VertexAiApi.Message> vertexMessages = prompt.getMessages()
List<VertexAiApi.Message> vertexMessages = prompt.getInstructions()
.stream()
.filter(m -> m.getMessageType() == MessageType.USER || m.getMessageType() == MessageType.ASSISTANT)
.map(m -> new VertexAiApi.Message(m.getMessageType().getValue(), m.getContent()))

View File

@@ -112,7 +112,7 @@ public class VertexAiApi {
private final String embeddingModel;
/**
* Create an new chat completion api.
* Create a new chat completion api.
* @param apiKey vertex apiKey.
*/
public VertexAiApi(String apiKey) {
@@ -120,7 +120,7 @@ public class VertexAiApi {
}
/**
* Create an new chat completion api.
* Create a new chat completion api.
* @param baseUrl api base URL.
* @param apiKey vertex apiKey.
* @param model vertex model.

View File

@@ -78,7 +78,7 @@ public class VertexAiApiTests {
List.of(new VertexAiApi.GenerateMessageResponse.ContentFilter(BlockedReason.SAFETY, "reason")));
server
.expect(requestToUriTemplate("/models/{model}:generateMessage?key={apiKey}",
.expect(requestToUriTemplate("/models/{generative}:generateMessage?key={apiKey}",
VertexAiApi.DEFAULT_GENERATE_MODEL, TEST_API_KEY))
.andExpect(method(HttpMethod.POST))
.andExpect(content().json(objectMapper.writeValueAsString(request)))
@@ -99,8 +99,8 @@ public class VertexAiApiTests {
Embedding expectedEmbedding = new Embedding(List.of(0.1, 0.2, 0.3));
server
.expect(requestToUriTemplate("/models/{model}:embedText?key={apiKey}", VertexAiApi.DEFAULT_EMBEDDING_MODEL,
TEST_API_KEY))
.expect(requestToUriTemplate("/models/{generative}:embedText?key={apiKey}",
VertexAiApi.DEFAULT_EMBEDDING_MODEL, TEST_API_KEY))
.andExpect(method(HttpMethod.POST))
.andExpect(content().json(objectMapper.writeValueAsString(Map.of("text", text))))
.andRespond(withSuccess(objectMapper.writeValueAsString(Map.of("embedding", expectedEmbedding)),
@@ -122,7 +122,7 @@ public class VertexAiApiTests {
new Embedding(List.of(0.4, 0.5, 0.6)));
server
.expect(requestToUriTemplate("/models/{model}:batchEmbedText?key={apiKey}",
.expect(requestToUriTemplate("/models/{generative}:batchEmbedText?key={apiKey}",
VertexAiApi.DEFAULT_EMBEDDING_MODEL, TEST_API_KEY))
.andExpect(method(HttpMethod.POST))
.andExpect(content().json(objectMapper.writeValueAsString(Map.of("texts", texts))))

View File

@@ -12,11 +12,11 @@ import org.springframework.ai.chat.Generation;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
import org.springframework.ai.parser.MapOutputParser;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.SystemPromptTemplate;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.vertex.VertexAiChatClient;
import org.springframework.ai.vertex.api.VertexAiApi;
import org.springframework.beans.factory.annotation.Autowired;
@@ -48,8 +48,8 @@ class VertexAiChatGenerationClientIT {
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource);
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice));
Prompt prompt = new Prompt(List.of(userMessage, systemMessage));
ChatResponse response = client.generate(prompt);
assertThat(response.getGeneration().getContent()).contains("Bartholomew");
ChatResponse response = client.call(prompt);
assertThat(response.getResult().getOutput().getContent()).contains("Bartholomew");
}
// @Test
@@ -65,9 +65,9 @@ class VertexAiChatGenerationClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "ice cream flavors.", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = this.client.generate(prompt).getGeneration();
Generation generation = this.client.call(prompt).getResult();
List<String> list = outputParser.parse(generation.getContent());
List<String> list = outputParser.parse(generation.getOutput().getContent());
assertThat(list).hasSize(5);
}
@@ -84,9 +84,9 @@ class VertexAiChatGenerationClientIT {
PromptTemplate promptTemplate = new PromptTemplate(template,
Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
Map<String, Object> result = outputParser.parse(generation.getContent());
Map<String, Object> result = outputParser.parse(generation.getOutput().getContent());
assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9));
}
@@ -106,9 +106,9 @@ class VertexAiChatGenerationClientIT {
""";
PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format));
Prompt prompt = new Prompt(promptTemplate.createMessage());
Generation generation = client.generate(prompt).getGeneration();
Generation generation = client.call(prompt).getResult();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getContent());
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getOutput().getContent());
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}

View File

@@ -14,14 +14,15 @@
<modules>
<module>spring-ai-core</module>
<module>models/spring-ai-transformers</module>
<module>models/spring-ai-postgresml</module>
<module>models/spring-ai-bedrock</module>
<module>models/spring-ai-azure-openai</module>
<module>models/spring-ai-huggingface</module>
<module>models/spring-ai-ollama</module>
<module>models/spring-ai-openai</module>
<module>models/spring-ai-vertex-ai</module>
<module>models/spring-ai-transformers</module>
<module>models/spring-ai-postgresml</module>
<module>models/spring-ai-stabilityai</module>
<module>spring-ai-test</module>
<module>spring-ai-spring-boot-autoconfigure</module>
<module>spring-ai-spring-boot-starters/spring-ai-starter-openai</module>

View File

@@ -16,17 +16,18 @@
package org.springframework.ai.chat;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.model.ModelClient;
@FunctionalInterface
public interface ChatClient {
public interface ChatClient extends ModelClient<Prompt, ChatResponse> {
default String generate(String message) {
default String call(String message) {
Prompt prompt = new Prompt(new UserMessage(message));
return generate(prompt).getGeneration().getContent();
return call(prompt).getResult().getOutput().getContent();
}
ChatResponse generate(Prompt prompt);
ChatResponse call(Prompt prompt);
}

View File

@@ -0,0 +1,40 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.chat;
import org.springframework.ai.model.ModelOptions;
/**
* portable options
*/
public interface ChatOptions extends ModelOptions {
// determine portable optionsb
Float getTemperature();
void setTemperature(Float temperature);
Float getTopP();
void setTopP(Float topP);
Integer getTopK();
void setTopK(Integer topK);
}

View File

@@ -0,0 +1,89 @@
/*
* Copyright 2024-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.chat;
public class ChatOptionsBuilder {
private class ChatOptionsImpl implements ChatOptions {
private Float temperature;
private Float topP;
private Integer topK;
@Override
public Float getTemperature() {
return temperature;
}
@Override
public void setTemperature(Float temperature) {
this.temperature = temperature;
}
@Override
public Float getTopP() {
return topP;
}
@Override
public void setTopP(Float topP) {
this.topP = topP;
}
@Override
public Integer getTopK() {
return topK;
}
@Override
public void setTopK(Integer topK) {
this.topK = topK;
}
}
private final ChatOptionsImpl options = new ChatOptionsImpl();
private ChatOptionsBuilder() {
}
public static ChatOptionsBuilder builder() {
return new ChatOptionsBuilder();
}
public ChatOptionsBuilder withTemperature(Float temperature) {
options.setTemperature(temperature);
return this;
}
public ChatOptionsBuilder withTopP(Float topP) {
options.setTopP(topP);
return this;
}
public ChatOptionsBuilder withTopK(Integer topK) {
options.setTopK(topK);
return this;
}
public ChatOptions build() {
return options;
}
}

View File

@@ -17,42 +17,39 @@ package org.springframework.ai.chat;
import java.util.List;
import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.PromptMetadata;
import org.springframework.lang.Nullable;
import org.springframework.ai.model.ModelResponse;
import org.springframework.ai.chat.metadata.ChatResponseMetadata;
/**
* The chat completion (e.g. generation) response returned by an AI provider.
*/
public class ChatResponse {
public class ChatResponse implements ModelResponse<Generation> {
private final GenerationMetadata metadata;
private final ChatResponseMetadata chatResponseMetadata;
/**
* List of generated messages returned by the AI provider.
*/
private final List<Generation> generations;
private PromptMetadata promptMetadata;
/**
* Construct a new {@link ChatResponse} instance without metadata.
* @param generations the {@link List} of {@link Generation} returned by the AI
* provider.
*/
public ChatResponse(List<Generation> generations) {
this(generations, GenerationMetadata.NULL);
this(generations, ChatResponseMetadata.NULL);
}
/**
* Construct a new {@link ChatResponse} instance.
* @param generations the {@link List} of {@link Generation} returned by the AI
* provider.
* @param metadata {@link GenerationMetadata} containing information about the use of
* the AI provider's API.
* @param chatResponseMetadata {@link ChatResponseMetadata} containing information
* about the use of the AI provider's API.
*/
public ChatResponse(List<Generation> generations, GenerationMetadata metadata) {
this.metadata = metadata;
public ChatResponse(List<Generation> generations, ChatResponseMetadata chatResponseMetadata) {
this.chatResponseMetadata = chatResponseMetadata;
this.generations = List.copyOf(generations);
}
@@ -63,45 +60,25 @@ public class ChatResponse {
* multiple output {@link Generation generations}.
* @return the {@link List} of {@link Generation generated outputs}.
*/
public List<Generation> getGenerations() {
@Override
public List<Generation> getResults() {
return this.generations;
}
/**
* @return Returns the first {@link Generation} in the generations list.
*/
public Generation getGeneration() {
public Generation getResult() {
return this.generations.get(0);
}
/**
* @return Returns {@link GenerationMetadata} containing information about the use of
* the AI provider's API.
* @return Returns {@link ChatResponseMetadata} containing information about the use
* of the AI provider's API.
*/
public GenerationMetadata getGenerationMetadata() {
return this.metadata;
}
/**
* @return {@link PromptMetadata} containing information on prompt processing by the
* AI.
*/
public PromptMetadata getPromptMetadata() {
PromptMetadata promptMetadata = this.promptMetadata;
return promptMetadata != null ? promptMetadata : PromptMetadata.empty();
}
/**
* Builder method used to include {@link PromptMetadata} returned in the AI response
* when processing the prompt.
* @param promptMetadata {@link PromptMetadata} returned by the AI in the response
* when processing the prompt.
* @return this {@link ChatResponse}.
* @see #getPromptMetadata()
*/
public ChatResponse withPromptMetadata(@Nullable PromptMetadata promptMetadata) {
this.promptMetadata = promptMetadata;
return this;
public ChatResponseMetadata getMetadata() {
return this.chatResponseMetadata;
}
}

View File

@@ -16,46 +16,65 @@
package org.springframework.ai.chat;
import java.util.Collections;
import java.util.Map;
import java.util.Objects;
import org.springframework.ai.metadata.ChoiceMetadata;
import org.springframework.ai.prompt.messages.AbstractMessage;
import org.springframework.ai.prompt.messages.MessageType;
import org.springframework.ai.model.ModelResult;
import org.springframework.ai.chat.metadata.ChatGenerationMetadata;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.lang.Nullable;
/**
* Represents a response returned by the AI.
*/
public class Generation extends AbstractMessage {
public class Generation implements ModelResult<AssistantMessage> {
private ChoiceMetadata choiceMetadata;
private AssistantMessage assistantMessage;
private ChatGenerationMetadata chatGenerationMetadata;
public Generation(String text) {
this(text, Collections.emptyMap());
this.assistantMessage = new AssistantMessage(text);
}
public Generation(String content, Map<String, Object> properties) {
super(MessageType.ASSISTANT, content, properties);
public Generation(String text, Map<String, Object> properties) {
this.assistantMessage = new AssistantMessage(text, properties);
}
public Generation(String content, Map<String, Object> properties, MessageType type) {
super(type, content, properties);
@Override
public AssistantMessage getOutput() {
return this.assistantMessage;
}
public ChoiceMetadata getChoiceMetadata() {
ChoiceMetadata choiceMetadata = this.choiceMetadata;
return choiceMetadata != null ? choiceMetadata : ChoiceMetadata.NULL;
public ChatGenerationMetadata getMetadata() {
ChatGenerationMetadata chatGenerationMetadata = this.chatGenerationMetadata;
return chatGenerationMetadata != null ? chatGenerationMetadata : ChatGenerationMetadata.NULL;
}
public Generation withChoiceMetadata(@Nullable ChoiceMetadata choiceMetadata) {
this.choiceMetadata = choiceMetadata;
public Generation withGenerationMetadata(@Nullable ChatGenerationMetadata chatGenerationMetadata) {
this.chatGenerationMetadata = chatGenerationMetadata;
return this;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof Generation that))
return false;
return Objects.equals(assistantMessage, that.assistantMessage)
&& Objects.equals(chatGenerationMetadata, that.chatGenerationMetadata);
}
@Override
public int hashCode() {
return Objects.hash(assistantMessage, chatGenerationMetadata);
}
@Override
public String toString() {
return "Generation{" + "text='" + content + '\'' + ", info=" + properties + '}';
return "Generation{" + "assistantMessage=" + assistantMessage + ", chatGenerationMetadata="
+ chatGenerationMetadata + '}';
}
}

View File

@@ -18,7 +18,7 @@ package org.springframework.ai.chat;
import reactor.core.publisher.Flux;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.chat.prompt.Prompt;
@FunctionalInterface
public interface StreamingChatClient {

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import org.springframework.core.io.Resource;
import org.springframework.util.StreamUtils;
@@ -31,7 +31,7 @@ public abstract class AbstractMessage implements Message {
protected String content;
/**
* Additional options for the message to influence the response, not a model map.
* Additional options for the message to influence the response, not a generative map.
*/
protected Map<String, Object> properties = new HashMap<>();

View File

@@ -14,14 +14,14 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import java.util.Map;
/**
* Lets the model know the content was generated as a response to the user. This role
* indicates messages that the model has previously generated in the conversation. By
* including assistant messages in the series, you provide context to the model about
* Lets the generative know the content was generated as a response to the user. This role
* indicates messages that the generative has previously generated in the conversation. By
* including assistant messages in the series, you provide context to the generative about
* prior exchanges in the conversation.
*/
public class AssistantMessage extends AbstractMessage {

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import java.util.Map;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import java.util.Map;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import java.util.Map;

View File

@@ -13,7 +13,7 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
public enum MessageType {

View File

@@ -14,15 +14,16 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import org.springframework.core.io.Resource;
/**
* A message of the type 'system' passed as input. The system message gives high level
* instructions for the conversation. This role typically provides high-level instructions
* for the conversation. For example, you might use a system message to instruct the model
* to behave like a certain character or to provide answers in a specific format.
* for the conversation. For example, you might use a system message to instruct the
* generative to behave like a certain character or to provide answers in a specific
* format.
*/
public class SystemMessage extends AbstractMessage {

View File

@@ -14,14 +14,14 @@
* limitations under the License.
*/
package org.springframework.ai.prompt.messages;
package org.springframework.ai.chat.messages;
import org.springframework.core.io.Resource;
/**
* A message of the type 'user' passed as input Messages with the user role are from the
* end-user or developer. They represent questions, prompts, or any input that you want
* the model to respond to.
* the generative to respond to.
*/
public class UserMessage extends AbstractMessage {

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
import java.time.Duration;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
/**
* Abstract base class used as a foundation for implementing {@link Usage}.

View File

@@ -14,8 +14,9 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
import org.springframework.ai.model.ResultMetadata;
import org.springframework.lang.Nullable;
/**
@@ -25,21 +26,21 @@ import org.springframework.lang.Nullable;
* @author John Blum
* @since 0.7.0
*/
public interface ChoiceMetadata {
public interface ChatGenerationMetadata extends ResultMetadata {
ChoiceMetadata NULL = ChoiceMetadata.from(null, null);
ChatGenerationMetadata NULL = ChatGenerationMetadata.from(null, null);
/**
* Factory method used to construct a new {@link ChoiceMetadata} from the given
* {@link String finish reason} and content filter metadata.
* Factory method used to construct a new {@link ChatGenerationMetadata} from the
* given {@link String finish reason} and content filter metadata.
* @param finishReason {@link String} contain the reason for the choice completion.
* @param contentFilterMetadata underlying AI provider metadata for filtering applied
* to generation content.
* @return a new {@link ChoiceMetadata} from the given {@link String finish reason}
* and content filter metadata.
* @return a new {@link ChatGenerationMetadata} from the given {@link String finish
* reason} and content filter metadata.
*/
static ChoiceMetadata from(String finishReason, Object contentFilterMetadata) {
return new ChoiceMetadata() {
static ChatGenerationMetadata from(String finishReason, Object contentFilterMetadata) {
return new ChatGenerationMetadata() {
@Override
@SuppressWarnings("unchecked")

View File

@@ -14,7 +14,9 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
import org.springframework.ai.model.ResponseMetadata;
/**
* Abstract Data Type (ADT) modeling common AI provider metadata returned in an AI
@@ -23,9 +25,9 @@ package org.springframework.ai.metadata;
* @author John Blum
* @since 0.7.0
*/
public interface GenerationMetadata {
public interface ChatResponseMetadata extends ResponseMetadata {
GenerationMetadata NULL = new GenerationMetadata() {
ChatResponseMetadata NULL = new ChatResponseMetadata() {
};
/**
@@ -46,4 +48,8 @@ public interface GenerationMetadata {
return Usage.NULL;
}
default PromptMetadata getPromptMetadata() {
return PromptMetadata.empty();
}
}

View File

@@ -13,7 +13,7 @@
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
import java.util.Arrays;
import java.util.Optional;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
import java.time.Duration;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.metadata;
package org.springframework.ai.chat.metadata;
/**
* Abstract Data Type (ADT) encapsulating metadata on the usage of an AI provider's API

View File

@@ -0,0 +1,14 @@
/**
* The org.sf.ai.chat package represents the bounded context for the Chat Model within the
* AI generative model domain. This package extends the core domain defined in
* org.sf.ai.generative, providing implementations specific to chat-based generative AI
* interactions.
*
* In line with Domain-Driven Design principles, this package includes implementations of
* entities and value objects specific to the chat context, such as ChatPrompt and
* ChatResponse, adhering to the ubiquitous language of chat interactions in AI models.
*
* This bounded context is designed to encapsulate all aspects of chat-based AI
* functionalities, maintaining a clear boundary from other contexts within the AI domain.
*/
package org.springframework.ai.chat;

View File

@@ -14,9 +14,9 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.AssistantMessage;
import org.springframework.ai.chat.messages.AssistantMessage;
import java.util.Map;

View File

@@ -14,9 +14,9 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.chat.messages.Message;
import java.util.ArrayList;
import java.util.List;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
public class FunctionPromptTemplate extends PromptTemplate {

View File

@@ -14,19 +14,23 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.model.ModelOptions;
import org.springframework.ai.model.ModelRequest;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import java.util.Collections;
import java.util.List;
import java.util.Objects;
public class Prompt {
public class Prompt implements ModelRequest<List<Message>> {
private final List<Message> messages;
private ModelOptions modelOptions;
public Prompt(String contents) {
this(new UserMessage(contents));
}
@@ -39,36 +43,53 @@ public class Prompt {
this.messages = messages;
}
public Prompt(String contents, ModelOptions modelOptions) {
this(new UserMessage(contents), modelOptions);
}
public Prompt(Message message, ModelOptions modelOptions) {
this(Collections.singletonList(message), modelOptions);
}
public Prompt(List<Message> messages, ModelOptions modelOptions) {
this.messages = messages;
this.modelOptions = modelOptions;
}
public String getContents() {
StringBuilder sb = new StringBuilder();
for (Message message : getMessages()) {
for (Message message : getInstructions()) {
sb.append(message.getContent());
}
return sb.toString();
}
public List<Message> getMessages() {
public ModelOptions getOptions() {
return modelOptions;
}
@Override
public List<Message> getInstructions() {
return this.messages;
}
@Override
public String toString() {
return "Prompt{" + "messages=" + messages + '}';
return "Prompt{" + "messages=" + messages + ", modelOptions=" + modelOptions + '}';
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (o == null || getClass() != o.getClass())
if (!(o instanceof Prompt prompt))
return false;
Prompt prompt = (Prompt) o;
return Objects.equals(messages, prompt.messages);
return Objects.equals(messages, prompt.messages) && Objects.equals(modelOptions, prompt.modelOptions);
}
@Override
public int hashCode() {
return Objects.hash(messages);
return Objects.hash(messages, modelOptions);
}
}

View File

@@ -14,13 +14,13 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.antlr.runtime.Token;
import org.antlr.runtime.TokenStream;
import org.springframework.ai.parser.OutputParser;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.UserMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.core.io.Resource;
import org.springframework.util.StreamUtils;
import org.stringtemplate.v4.ST;
@@ -189,7 +189,7 @@ public class PromptTemplate implements PromptTemplateActions, PromptTemplateMess
return new Prompt(render(model));
}
protected Set<String> getInputVariables() {
public Set<String> getInputVariables() {
TokenStream tokens = this.st.impl.tokens;
return IntStream.range(0, tokens.range())
.mapToObj(tokens::get)

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import java.util.Map;

View File

@@ -1,6 +1,6 @@
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.chat.messages.Message;
import java.util.List;
import java.util.Map;

View File

@@ -1,6 +1,6 @@
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.chat.messages.Message;
import java.util.Map;

View File

@@ -1,4 +1,4 @@
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import java.util.Map;

View File

@@ -14,10 +14,10 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.SystemMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.core.io.Resource;
import java.util.Map;

View File

@@ -14,7 +14,7 @@
* limitations under the License.
*/
package org.springframework.ai.prompt;
package org.springframework.ai.chat.prompt;
public enum TemplateFormat {

View File

@@ -71,7 +71,7 @@ public class DefaultContentFormatter implements ContentFormatter {
private final List<String> excludedInferenceMetadataKeys;
/**
* Metadata keys that are excluded from text for the embed model.
* Metadata keys that are excluded from text for the embed generative.
*/
private final List<String> excludedEmbedMetadataKeys;
@@ -157,7 +157,8 @@ public class DefaultContentFormatter implements ContentFormatter {
}
/**
* Configures the excluded Inference metadata keys to filter out from the model.
* Configures the excluded Inference metadata keys to filter out from the
* generative.
* @param excludedInferenceMetadataKeys Excluded inference metadata keys to use.
* @return this builder
*/
@@ -174,7 +175,7 @@ public class DefaultContentFormatter implements ContentFormatter {
}
/**
* Configures the excluded Embed metadata keys to filter out from the model.
* Configures the excluded Embed metadata keys to filter out from the generative.
* @param excludedEmbedMetadataKeys Excluded Embed metadata keys to use.
* @return this builder
*/

View File

@@ -38,7 +38,8 @@ public interface EmbeddingClient {
EmbeddingResponse embedForResponse(List<String> texts);
/**
* @return the number of dimensions of the embedded vectors. It is model specific.
* @return the number of dimensions of the embedded vectors. It is generative
* specific.
*/
default int dimensions() {
return embed("Test String").size();

Some files were not shown because too many files have changed in this diff Show More