feat: Enhance Anthropic integration with Thinking

- The `thinking` option is added to `AnthropicChatOptions` and `ChatCompletionRequest`.
- The `AnthropicApi` and `AnthropicChatModel` now handle `THINKING` and `REDACTED_THINKING` content blocks in responses.  New tests verify parsing of these blocks.
- Updated method signatures on ChatCompletionRequestBuilder, deprecating old builders with `with*` prefix in favor of those without.

Signed-off-by: Alexandros Pappas <apappascs@gmail.com>
This commit is contained in:
Alexandros Pappas
2025-02-25 18:10:23 +01:00
committed by Christian Tzolov
parent fbec267eca
commit 4f67959645
7 changed files with 414 additions and 67 deletions

View File

@@ -295,46 +295,49 @@ public class AnthropicChatModel implements ChatModel {
return new ChatResponse(List.of());
}
List<Generation> generations = chatCompletion.content()
.stream()
.filter(content -> content.type() != ContentBlock.Type.TOOL_USE)
.map(content -> new Generation(new AssistantMessage(content.text(), Map.of()),
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build()))
.toList();
List<Generation> allGenerations = new ArrayList<>(generations);
List<Generation> generations = new ArrayList<>();
List<AssistantMessage.ToolCall> toolCalls = new ArrayList<>();
for (ContentBlock content : chatCompletion.content()) {
switch (content.type()) {
case TEXT, TEXT_DELTA:
generations.add(new Generation(new AssistantMessage(content.text(), Map.of()),
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build()));
break;
case THINKING, THINKING_DELTA:
Map<String, Object> thinkingProperties = new HashMap<>();
thinkingProperties.put("signature", content.signature());
generations.add(new Generation(new AssistantMessage(content.thinking(), thinkingProperties),
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build()));
break;
case REDACTED_THINKING:
Map<String, Object> redactedProperties = new HashMap<>();
redactedProperties.put("data", content.data());
generations.add(new Generation(new AssistantMessage(null, redactedProperties),
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build()));
break;
case TOOL_USE:
var functionCallId = content.id();
var functionName = content.name();
var functionArguments = JsonParser.toJson(content.input());
toolCalls.add(
new AssistantMessage.ToolCall(functionCallId, "function", functionName, functionArguments));
break;
}
}
if (chatCompletion.stopReason() != null && generations.isEmpty()) {
Generation generation = new Generation(new AssistantMessage(null, Map.of()),
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build());
allGenerations.add(generation);
generations.add(generation);
}
List<ContentBlock> toolToUseList = chatCompletion.content()
.stream()
.filter(c -> c.type() == ContentBlock.Type.TOOL_USE)
.toList();
if (!CollectionUtils.isEmpty(toolToUseList)) {
List<AssistantMessage.ToolCall> toolCalls = new ArrayList<>();
for (ContentBlock toolToUse : toolToUseList) {
var functionCallId = toolToUse.id();
var functionName = toolToUse.name();
var functionArguments = JsonParser.toJson(toolToUse.input());
toolCalls
.add(new AssistantMessage.ToolCall(functionCallId, "function", functionName, functionArguments));
}
if (!CollectionUtils.isEmpty(toolCalls)) {
AssistantMessage assistantMessage = new AssistantMessage("", Map.of(), toolCalls);
Generation toolCallGeneration = new Generation(assistantMessage,
ChatGenerationMetadata.builder().finishReason(chatCompletion.stopReason()).build());
allGenerations.add(toolCallGeneration);
generations.add(toolCallGeneration);
}
return new ChatResponse(allGenerations, this.from(chatCompletion, usage));
return new ChatResponse(generations, this.from(chatCompletion, usage));
}
private ChatResponseMetadata from(AnthropicApi.ChatCompletionResponse result) {
@@ -506,7 +509,7 @@ public class AnthropicChatModel implements ChatModel {
List<ToolDefinition> toolDefinitions = this.toolCallingManager.resolveToolDefinitions(requestOptions);
if (!CollectionUtils.isEmpty(toolDefinitions)) {
request = ModelOptionsUtils.merge(request, this.defaultOptions, ChatCompletionRequest.class);
request = ChatCompletionRequest.from(request).withTools(getFunctionTools(toolDefinitions)).build();
request = ChatCompletionRequest.from(request).tools(getFunctionTools(toolDefinitions)).build();
}
return request;

View File

@@ -57,6 +57,7 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
private @JsonProperty("temperature") Double temperature;
private @JsonProperty("top_p") Double topP;
private @JsonProperty("top_k") Integer topK;
private @JsonProperty("thinking") ChatCompletionRequest.ThinkingConfig thinking;
/**
* Collection of {@link ToolCallback}s to be used for tool calling in the chat
@@ -103,6 +104,7 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
.temperature(fromOptions.getTemperature())
.topP(fromOptions.getTopP())
.topK(fromOptions.getTopK())
.thinking(fromOptions.getThinking())
.toolCallbacks(
fromOptions.getToolCallbacks() != null ? new ArrayList<>(fromOptions.getToolCallbacks()) : null)
.toolNames(fromOptions.getToolNames() != null ? new HashSet<>(fromOptions.getToolNames()) : null)
@@ -174,6 +176,14 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
this.topK = topK;
}
public ChatCompletionRequest.ThinkingConfig getThinking() {
return this.thinking;
}
public void setThinking(ChatCompletionRequest.ThinkingConfig thinking) {
this.thinking = thinking;
}
@Override
@JsonIgnore
public List<FunctionCallback> getToolCallbacks() {
@@ -308,7 +318,8 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
&& Objects.equals(this.metadata, that.metadata)
&& Objects.equals(this.stopSequences, that.stopSequences)
&& Objects.equals(this.temperature, that.temperature) && Objects.equals(this.topP, that.topP)
&& Objects.equals(this.topK, that.topK) && Objects.equals(this.toolCallbacks, that.toolCallbacks)
&& Objects.equals(this.topK, that.topK) && Objects.equals(this.thinking, that.thinking)
&& Objects.equals(this.toolCallbacks, that.toolCallbacks)
&& Objects.equals(this.toolNames, that.toolNames)
&& Objects.equals(this.internalToolExecutionEnabled, that.internalToolExecutionEnabled)
&& Objects.equals(this.toolContext, that.toolContext)
@@ -317,7 +328,7 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
@Override
public int hashCode() {
return Objects.hash(model, maxTokens, metadata, stopSequences, temperature, topP, topK, toolCallbacks,
return Objects.hash(model, maxTokens, metadata, stopSequences, temperature, topP, topK, thinking, toolCallbacks,
toolNames, internalToolExecutionEnabled, toolContext, httpHeaders);
}
@@ -365,6 +376,16 @@ public class AnthropicChatOptions implements ToolCallingChatOptions {
return this;
}
public Builder thinking(ChatCompletionRequest.ThinkingConfig thinking) {
this.options.thinking = thinking;
return this;
}
public Builder thinking(AnthropicApi.ThinkingType type, Integer budgetTokens) {
this.options.thinking = new ChatCompletionRequest.ThinkingConfig(type, budgetTokens);
return this;
}
public Builder toolCallbacks(List<FunctionCallback> toolCallbacks) {
this.options.setToolCallbacks(toolCallbacks);
return this;

View File

@@ -57,6 +57,7 @@ import org.springframework.web.reactive.function.client.WebClient;
* @author Mariusz Bernacki
* @author Thomas Vitale
* @author Jihoon Kim
* @author Alexandros Pappas
* @since 1.0.0
*/
public class AnthropicApi {
@@ -262,9 +263,9 @@ public class AnthropicApi {
* The claude-3-7-sonnet-latest model.
*/
CLAUDE_3_7_SONNET("claude-3-7-sonnet-latest"),
/**
* The claude-3-5-sonnet-20241022 model.
* The claude-3-5-sonnet-latest model.
*/
CLAUDE_3_5_SONNET("claude-3-5-sonnet-latest"),
@@ -347,6 +348,25 @@ public class AnthropicApi {
}
/**
* The thinking type.
*/
public enum ThinkingType {
/**
* Enabled thinking type.
*/
@JsonProperty("enabled")
ENABLED,
/**
* Disabled thinking type.
*/
@JsonProperty("disabled")
DISABLED
}
/**
* The event type of the streamed chunk.
*/
@@ -412,11 +432,8 @@ public class AnthropicApi {
@JsonSubTypes({ @JsonSubTypes.Type(value = ContentBlockStartEvent.class, name = "content_block_start"),
@JsonSubTypes.Type(value = ContentBlockDeltaEvent.class, name = "content_block_delta"),
@JsonSubTypes.Type(value = ContentBlockStopEvent.class, name = "content_block_stop"),
@JsonSubTypes.Type(value = PingEvent.class, name = "ping"),
@JsonSubTypes.Type(value = ErrorEvent.class, name = "error"),
@JsonSubTypes.Type(value = MessageStartEvent.class, name = "message_start"),
@JsonSubTypes.Type(value = MessageDeltaEvent.class, name = "message_delta"),
@JsonSubTypes.Type(value = MessageStopEvent.class, name = "message_stop") })
@@ -468,6 +485,8 @@ public class AnthropicApi {
* return tool_use content blocks that represent the model's use of those tools. You
* can then run those tools using the tool input generated by the model and then
* optionally return results back to the model using tool_result content blocks.
* @param thinking Configuration for the model's thinking mode. When enabled, the
* model can perform more in-depth reasoning before responding to a query.
*/
@JsonInclude(Include.NON_NULL)
public record ChatCompletionRequest(
@@ -482,17 +501,18 @@ public class AnthropicApi {
@JsonProperty("temperature") Double temperature,
@JsonProperty("top_p") Double topP,
@JsonProperty("top_k") Integer topK,
@JsonProperty("tools") List<Tool> tools) {
@JsonProperty("tools") List<Tool> tools,
@JsonProperty("thinking") ThinkingConfig thinking) {
// @formatter:on
public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens,
Double temperature, Boolean stream) {
this(model, messages, system, maxTokens, null, null, stream, temperature, null, null, null);
this(model, messages, system, maxTokens, null, null, stream, temperature, null, null, null, null);
}
public ChatCompletionRequest(String model, List<AnthropicMessage> messages, String system, Integer maxTokens,
List<String> stopSequences, Double temperature, Boolean stream) {
this(model, messages, system, maxTokens, null, stopSequences, stream, temperature, null, null, null);
this(model, messages, system, maxTokens, null, stopSequences, stream, temperature, null, null, null, null);
}
public static ChatCompletionRequestBuilder builder() {
@@ -516,6 +536,18 @@ public class AnthropicApi {
}
/**
* Configuration for the model's thinking mode.
*
* @param type The type of thinking mode. Currently, "enabled" is supported.
* @param budgetTokens The token budget available for the thinking process. Must
* be ≥1024 and less than max_tokens.
*/
@JsonInclude(Include.NON_NULL)
public record ThinkingConfig(@JsonProperty("type") ThinkingType type,
@JsonProperty("budget_tokens") Integer budgetTokens) {
}
}
public static final class ChatCompletionRequestBuilder {
@@ -542,6 +574,8 @@ public class AnthropicApi {
private List<Tool> tools;
private ChatCompletionRequest.ThinkingConfig thinking;
private ChatCompletionRequestBuilder() {
}
@@ -557,71 +591,209 @@ public class AnthropicApi {
this.topP = request.topP;
this.topK = request.topK;
this.tools = request.tools;
this.thinking = request.thinking;
}
/**
* @deprecated use {@link #model(ChatModel)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withModel(ChatModel model) {
this.model = model.getValue();
return this;
}
public ChatCompletionRequestBuilder model(ChatModel model) {
this.model = model.getValue();
return this;
}
/**
* @deprecated use {@link #model(String)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withModel(String model) {
this.model = model;
return this;
}
public ChatCompletionRequestBuilder model(String model) {
this.model = model;
return this;
}
/**
* @deprecated use {@link #messages(List)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withMessages(List<AnthropicMessage> messages) {
this.messages = messages;
return this;
}
public ChatCompletionRequestBuilder messages(List<AnthropicMessage> messages) {
this.messages = messages;
return this;
}
/**
* @deprecated use {@link #system(String)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withSystem(String system) {
this.system = system;
return this;
}
public ChatCompletionRequestBuilder system(String system) {
this.system = system;
return this;
}
/**
* @deprecated use {@link #maxTokens(Integer)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withMaxTokens(Integer maxTokens) {
this.maxTokens = maxTokens;
return this;
}
public ChatCompletionRequestBuilder maxTokens(Integer maxTokens) {
this.maxTokens = maxTokens;
return this;
}
/**
* @deprecated use {@link #metadata(ChatCompletionRequest.Metadata)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withMetadata(ChatCompletionRequest.Metadata metadata) {
this.metadata = metadata;
return this;
}
public ChatCompletionRequestBuilder metadata(ChatCompletionRequest.Metadata metadata) {
this.metadata = metadata;
return this;
}
/**
* @deprecated use {@link #stopSequences(List)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withStopSequences(List<String> stopSequences) {
this.stopSequences = stopSequences;
return this;
}
public ChatCompletionRequestBuilder stopSequences(List<String> stopSequences) {
this.stopSequences = stopSequences;
return this;
}
/**
* @deprecated use {@link #stream(Boolean)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withStream(Boolean stream) {
this.stream = stream;
return this;
}
public ChatCompletionRequestBuilder stream(Boolean stream) {
this.stream = stream;
return this;
}
/**
* @deprecated use {@link #temperature(Double)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withTemperature(Double temperature) {
this.temperature = temperature;
return this;
}
public ChatCompletionRequestBuilder temperature(Double temperature) {
this.temperature = temperature;
return this;
}
/**
* @deprecated use {@link #topP(Double)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withTopP(Double topP) {
this.topP = topP;
return this;
}
public ChatCompletionRequestBuilder topP(Double topP) {
this.topP = topP;
return this;
}
/**
* @deprecated use {@link #topK(Integer)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withTopK(Integer topK) {
this.topK = topK;
return this;
}
public ChatCompletionRequestBuilder topK(Integer topK) {
this.topK = topK;
return this;
}
/**
* @deprecated use {@link #tools(List)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withTools(List<Tool> tools) {
this.tools = tools;
return this;
}
public ChatCompletionRequestBuilder tools(List<Tool> tools) {
this.tools = tools;
return this;
}
/**
* @deprecated use {@link #thinking(ChatCompletionRequest.ThinkingConfig)}
* instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withThinking(ChatCompletionRequest.ThinkingConfig thinking) {
this.thinking = thinking;
return this;
}
public ChatCompletionRequestBuilder thinking(ChatCompletionRequest.ThinkingConfig thinking) {
this.thinking = thinking;
return this;
}
/**
* @deprecated use {@link #thinking(ThinkingType, Integer)} instead.
*/
@Deprecated(forRemoval = true, since = "1.0.0-M6")
public ChatCompletionRequestBuilder withThinking(ThinkingType type, Integer budgetTokens) {
this.thinking = new ChatCompletionRequest.ThinkingConfig(type, budgetTokens);
return this;
}
public ChatCompletionRequestBuilder thinking(ThinkingType type, Integer budgetTokens) {
this.thinking = new ChatCompletionRequest.ThinkingConfig(type, budgetTokens);
return this;
}
public ChatCompletionRequest build() {
return new ChatCompletionRequest(this.model, this.messages, this.system, this.maxTokens, this.metadata,
this.stopSequences, this.stream, this.temperature, this.topP, this.topK, this.tools);
this.stopSequences, this.stream, this.temperature, this.topP, this.topK, this.tools, this.thinking);
}
}
@@ -689,7 +861,14 @@ public class AnthropicApi {
// tool_result response only
@JsonProperty("tool_use_id") String toolUseId,
@JsonProperty("content") String content
@JsonProperty("content") String content,
// Thinking only
@JsonProperty("signature") String signature,
@JsonProperty("thinking") String thinking,
// Redacted Thinking only
@JsonProperty("data") String data
) {
// @formatter:on
@@ -708,7 +887,7 @@ public class AnthropicApi {
* @param source The source of the content.
*/
public ContentBlock(Type type, Source source) {
this(type, source, null, null, null, null, null, null, null);
this(type, source, null, null, null, null, null, null, null, null, null, null);
}
/**
@@ -716,7 +895,7 @@ public class AnthropicApi {
* @param source The source of the content.
*/
public ContentBlock(Source source) {
this(Type.IMAGE, source, null, null, null, null, null, null, null);
this(Type.IMAGE, source, null, null, null, null, null, null, null, null, null, null);
}
/**
@@ -724,7 +903,7 @@ public class AnthropicApi {
* @param text The text of the content.
*/
public ContentBlock(String text) {
this(Type.TEXT, null, text, null, null, null, null, null, null);
this(Type.TEXT, null, text, null, null, null, null, null, null, null, null, null);
}
// Tool result
@@ -735,7 +914,7 @@ public class AnthropicApi {
* @param content The content of the tool result.
*/
public ContentBlock(Type type, String toolUseId, String content) {
this(type, null, null, null, null, null, null, toolUseId, content);
this(type, null, null, null, null, null, null, toolUseId, content, null, null, null);
}
/**
@@ -746,7 +925,7 @@ public class AnthropicApi {
* @param index The index of the content block.
*/
public ContentBlock(Type type, Source source, String text, Integer index) {
this(type, source, text, index, null, null, null, null, null);
this(type, source, text, index, null, null, null, null, null, null, null, null);
}
// Tool use input JSON delta streaming
@@ -758,7 +937,7 @@ public class AnthropicApi {
* @param input The input of the tool use.
*/
public ContentBlock(Type type, String id, String name, Map<String, Object> input) {
this(type, null, null, null, id, name, input, null, null);
this(type, null, null, null, id, name, input, null, null, null, null, null);
}
/**
@@ -790,6 +969,22 @@ public class AnthropicApi {
@JsonProperty("text_delta")
TEXT_DELTA("text_delta"),
/**
* When using extended thinking with streaming enabled, youll receive
* thinking content via thinking_delta events. These deltas correspond to the
* thinking field of the thinking content blocks.
*/
@JsonProperty("thinking_delta")
THINKING_DELTA("thinking_delta"),
/**
* For thinking content, a special signature_delta event is sent just before
* the content_block_stop event. This signature is used to verify the
* integrity of the thinking block.
*/
@JsonProperty("signature_delta")
SIGNATURE_DELTA("signature_delta"),
/**
* Tool use input partial JSON delta streaming.
*/
@@ -806,7 +1001,19 @@ public class AnthropicApi {
* Document message.
*/
@JsonProperty("document")
DOCUMENT("document");
DOCUMENT("document"),
/**
* Thinking message.
*/
@JsonProperty("thinking")
THINKING("thinking"),
/**
* Redacted Thinking message.
*/
@JsonProperty("redacted_thinking")
REDACTED_THINKING("redacted_thinking");
public final String value;
@@ -1073,7 +1280,10 @@ public class AnthropicApi {
@JsonTypeInfo(use = JsonTypeInfo.Id.NAME, include = JsonTypeInfo.As.EXISTING_PROPERTY, property = "type",
visible = true)
@JsonSubTypes({ @JsonSubTypes.Type(value = ContentBlockDeltaText.class, name = "text_delta"),
@JsonSubTypes.Type(value = ContentBlockDeltaJson.class, name = "input_json_delta") })
@JsonSubTypes.Type(value = ContentBlockDeltaJson.class, name = "input_json_delta"),
@JsonSubTypes.Type(value = ContentBlockDeltaThinking.class, name = "thinking_delta"),
@JsonSubTypes.Type(value = ContentBlockDeltaSignature.class, name = "signature_delta")
})
public interface ContentBlockDeltaBody {
String type();
}
@@ -1099,6 +1309,26 @@ public class AnthropicApi {
@JsonProperty("type") String type,
@JsonProperty("partial_json") String partialJson) implements ContentBlockDeltaBody {
}
/**
* Thinking content block delta.
* @param type The content block type.
* @param thinking The thinking content.
*/
public record ContentBlockDeltaThinking(
@JsonProperty("type") String type,
@JsonProperty("thinking") String thinking) implements ContentBlockDeltaBody {
}
/**
* Signature content block delta.
* @param type The content block type.
* @param signature The signature content.
*/
public record ContentBlockDeltaSignature(
@JsonProperty("type") String type,
@JsonProperty("signature") String signature) implements ContentBlockDeltaBody {
}
}
// @formatter:on

View File

@@ -23,6 +23,7 @@ import java.util.List;
import java.util.Map;
import java.util.stream.Collectors;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.junit.jupiter.params.ParameterizedTest;
@@ -82,7 +83,6 @@ class AnthropicChatModelIT {
private static void validateChatResponseMetadata(ChatResponse response, String model) {
assertThat(response.getMetadata().getId()).isNotEmpty();
assertThat(response.getMetadata().getModel()).containsIgnoringCase(model);
assertThat(response.getMetadata().getUsage().getPromptTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getCompletionTokens()).isPositive();
assertThat(response.getMetadata().getUsage().getTotalTokens()).isPositive();
@@ -118,7 +118,7 @@ class AnthropicChatModelIT {
SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(this.systemResource);
Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate"));
Prompt prompt = new Prompt(List.of(userMessage, systemMessage),
AnthropicChatOptions.builder().model("claude-3-sonnet-20240229").build());
AnthropicChatOptions.builder().model("claude-3-5-sonnet-latest").build());
ChatResponse response = this.chatModel.call(prompt);
assertThat(response.getResult().getOutput().getText()).containsAnyOf("Blackbeard", "Bartholomew");
@@ -143,9 +143,6 @@ class AnthropicChatModelIT {
assertThat(streamingTokenUsage.getTotalTokens()).isGreaterThan(0);
assertThat(streamingTokenUsage.getPromptTokens()).isEqualTo(referenceTokenUsage.getPromptTokens());
// assertThat(streamingTokenUsage.getCompletionTokens()).isEqualTo(referenceTokenUsage.getCompletionTokens());
// assertThat(streamingTokenUsage.getTotalTokens()).isEqualTo(referenceTokenUsage.getTotalTokens());
}
@Test
@@ -357,7 +354,7 @@ class AnthropicChatModelIT {
@Test
void validateCallResponseMetadata() {
String model = AnthropicApi.ChatModel.CLAUDE_2_1.getName();
String model = AnthropicApi.ChatModel.CLAUDE_3_7_SONNET.getName();
// @formatter:off
ChatResponse response = ChatClient.create(this.chatModel).prompt()
.options(AnthropicChatOptions.builder().model(model).build())
@@ -384,13 +381,76 @@ class AnthropicChatModelIT {
logger.info(response.toString());
// Note, brittle test.
validateChatResponseMetadata(response, "claude-3-5-sonnet-20241022");
validateChatResponseMetadata(response, "claude-3-5-sonnet-latest");
}
record ActorsFilmsRecord(String actor, List<String> movies) {
}
@Test
void thinkingTest() {
UserMessage userMessage = new UserMessage(
"Are there an infinite number of prime numbers such that n mod 4 == 3?");
var promptOptions = AnthropicChatOptions.builder()
.model(AnthropicApi.ChatModel.CLAUDE_3_7_SONNET.getName())
.temperature(1.0) // temperature should be set to 1 when thinking is enabled
.maxTokens(8192)
.thinking(AnthropicApi.ThinkingType.ENABLED, 2048) // Must be ≥1024 && <
// max_tokens
.build();
ChatResponse response = this.chatModel.call(new Prompt(List.of(userMessage), promptOptions));
logger.info("Response: {}", response);
for (Generation generation : response.getResults()) {
AssistantMessage message = generation.getOutput();
if (message.getText() != null) { // text
assertThat(message.getText()).isNotBlank();
}
else if (message.getMetadata().containsKey("signature")) { // thinking
assertThat(message.getMetadata().get("signature")).isNotNull();
assertThat(message.getMetadata().get("thinking")).isNotNull();
}
else if (message.getMetadata().containsKey("data")) { // redacted thinking
assertThat(message.getMetadata().get("data")).isNotNull();
}
}
}
@Test
void testToolUseContentBlock() {
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo and Paris? Return the result in Celsius.");
List<Message> messages = new ArrayList<>(List.of(userMessage));
var promptOptions = AnthropicChatOptions.builder()
.model(AnthropicApi.ChatModel.CLAUDE_3_OPUS.getName())
.functionCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService())
.description(
"Get the weather in location. Return temperature in 36°F or 36°C format. Use multi-turn if needed.")
.inputType(MockWeatherService.Request.class)
.build()))
.build();
ChatResponse response = this.chatModel.call(new Prompt(messages, promptOptions));
logger.info("Response: {}", response);
for (Generation generation : response.getResults()) {
AssistantMessage message = generation.getOutput();
if (!message.getToolCalls().isEmpty()) {
assertThat(message.getToolCalls()).isNotEmpty();
AssistantMessage.ToolCall toolCall = message.getToolCalls().get(0);
assertThat(toolCall.id()).isNotBlank();
assertThat(toolCall.name()).isNotBlank();
assertThat(toolCall.arguments()).isNotBlank();
}
}
}
@SpringBootConfiguration
public static class Config {

View File

@@ -35,6 +35,7 @@ import static org.assertj.core.api.Assertions.assertThatThrownBy;
/**
* @author Christian Tzolov
* @author Jihoon Kim
* @author Alexandros Pappas
*/
@EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".+")
public class AnthropicApiIT {
@@ -55,6 +56,39 @@ public class AnthropicApiIT {
assertThat(response.getBody()).isNotNull();
}
@Test
void chatCompletionWithThinking() {
AnthropicMessage chatCompletionMessage = new AnthropicMessage(List.of(new ContentBlock("Tell me a Joke?")),
Role.USER);
ChatCompletionRequest request = ChatCompletionRequest.builder()
.model(AnthropicApi.ChatModel.CLAUDE_3_7_SONNET.getValue())
.messages(List.of(chatCompletionMessage))
.maxTokens(8192)
.temperature(1.0) // temperature should be set to 1 when thinking is enabled
.thinking(new ChatCompletionRequest.ThinkingConfig(AnthropicApi.ThinkingType.ENABLED, 2048))
.build();
ResponseEntity<ChatCompletionResponse> response = this.anthropicApi.chatCompletionEntity(request);
assertThat(response).isNotNull();
assertThat(response.getBody()).isNotNull();
List<ContentBlock> content = response.getBody().content();
for (ContentBlock block : content) {
if (block.type() == ContentBlock.Type.THINKING) {
assertThat(block.thinking()).isNotBlank();
assertThat(block.signature()).isNotBlank();
}
if (block.type() == ContentBlock.Type.REDACTED_THINKING) {
assertThat(block.data()).isNotBlank();
}
if (block.type() == ContentBlock.Type.TEXT) {
assertThat(block.text()).isNotBlank();
}
}
}
@Test
void chatCompletionStream() {

View File

@@ -102,11 +102,11 @@ public class AnthropicApiToolIT {
private ResponseEntity<ChatCompletionResponse> doCall(List<AnthropicMessage> messageConversation) {
ChatCompletionRequest chatCompletionRequest = ChatCompletionRequest.builder()
.withModel(AnthropicApi.ChatModel.CLAUDE_3_OPUS)
.withMessages(messageConversation)
.withMaxTokens(1500)
.withTemperature(0.8)
.withTools(this.tools)
.model(AnthropicApi.ChatModel.CLAUDE_3_OPUS)
.messages(messageConversation)
.maxTokens(1500)
.temperature(0.8)
.tools(this.tools)
.build();
ResponseEntity<ChatCompletionResponse> response = this.anthropicApi.chatCompletionEntity(chatCompletionRequest);

View File

@@ -284,8 +284,7 @@ class AnthropicChatClientIT {
}
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { "claude-3-opus-20240229", "claude-3-sonnet-20240229", "claude-3-haiku-20240307",
"claude-3-5-sonnet-20241022" })
@ValueSource(strings = { "claude-3-opus-latest", "claude-3-5-sonnet-latest", "claude-3-7-sonnet-latest" })
void multiModalityEmbeddedImage(String modelName) throws IOException {
// @formatter:off
@@ -303,8 +302,8 @@ class AnthropicChatClientIT {
@Disabled("Currently Anthropic API does not support external image URLs")
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { "claude-3-opus-20240229", "claude-3-sonnet-20240229", "claude-3-haiku-20240307",
"claude-3-5-sonnet-20241022" })
@ValueSource(strings = { "claude-3-opus-latest", "claude-3-5-sonnet-latest", "claude-3-haiku-latest",
"claude-3-7-sonnet-latest" })
void multiModalityImageUrl(String modelName) throws IOException {
// TODO: add url method that wrapps the checked exception.