Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 41 additions & 1 deletion core/src/main/java/com/google/adk/models/Claude.java
Original file line number Diff line number Diff line change
Expand Up @@ -16,14 +16,18 @@

package com.google.adk.models;

import static java.nio.charset.StandardCharsets.UTF_8;

import com.anthropic.client.AnthropicClient;
import com.anthropic.models.messages.ContentBlock;
import com.anthropic.models.messages.ContentBlockParam;
import com.anthropic.models.messages.Message;
import com.anthropic.models.messages.MessageCreateParams;
import com.anthropic.models.messages.MessageParam;
import com.anthropic.models.messages.MessageParam.Role;
import com.anthropic.models.messages.RedactedThinkingBlockParam;
import com.anthropic.models.messages.TextBlockParam;
import com.anthropic.models.messages.ThinkingBlockParam;
import com.anthropic.models.messages.Tool;
import com.anthropic.models.messages.ToolChoice;
import com.anthropic.models.messages.ToolChoiceAuto;
Expand All @@ -32,6 +36,7 @@
import com.anthropic.models.messages.ToolUseBlockParam;
import com.fasterxml.jackson.core.type.TypeReference;
import com.google.adk.JsonBaseModel;
import com.google.common.base.Utf8;
import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
import com.google.genai.types.*;
Expand All @@ -44,6 +49,7 @@
import java.util.Objects;
import java.util.Optional;
import java.util.stream.Collectors;
import org.jspecify.annotations.Nullable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

Expand Down Expand Up @@ -173,7 +179,29 @@ private MessageParam contentToAnthropicMessageParam(Content content) {
.build();
}

private ContentBlockParam partToAnthropicMessageBlock(Part part) {
private @Nullable ContentBlockParam partToAnthropicMessageBlock(Part part) {
if (part.thought().orElse(false)) {
// Signatures this class stores are UTF-8; a binary one (e.g. Gemini's) is not Claude's.
Optional<String> signature =
part.thoughtSignature()
.filter(bytes -> bytes.length > 0 && Utf8.isWellFormed(bytes))
.map(bytes -> new String(bytes, UTF_8));
if (signature.isPresent()) {
// Unlike ADK Python, empty text (display "omitted") stays thinking, not redacted_thinking.
return part.text().isPresent()
? ContentBlockParam.ofThinking(
ThinkingBlockParam.builder()
.thinking(part.text().get())
.signature(signature.get())
.build())
: ContentBlockParam.ofRedactedThinking(
RedactedThinkingBlockParam.builder().data(signature.get()).build());
}
// Other thoughts go out as text (ADK Python sends thinking), or not at all if empty.
if (part.text().map(String::isEmpty).orElse(false)) {
return null;
}
}
if (part.text().isPresent()) {
return ContentBlockParam.ofText(TextBlockParam.builder().text(part.text().get()).build());
} else if (part.functionCall().isPresent()) {
Expand Down Expand Up @@ -387,6 +415,18 @@ private Part anthropicContentBlockToPart(ContentBlock block) {
.convert(new TypeReference<Map<String, Object>>() {}))
.build())
.build();
} else if (block.isThinking()) {
return Part.builder()
.text(block.asThinking().thinking())
.thought(true)
.thoughtSignature(block.asThinking().signature().getBytes(UTF_8))
.build();
} else if (block.isRedactedThinking()) {
// Keep the encrypted data in thoughtSignature so the block can be sent back unchanged.
return Part.builder()
.thought(true)
.thoughtSignature(block.asRedactedThinking().data().getBytes(UTF_8))
.build();
}
throw new UnsupportedOperationException("Not supported yet.");
}
Expand Down
182 changes: 182 additions & 0 deletions core/src/test/java/com/google/adk/models/ClaudeTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -17,44 +17,62 @@
package com.google.adk.models;

import static com.google.common.truth.Truth.assertThat;
import static java.nio.charset.StandardCharsets.UTF_8;
import static org.junit.Assert.assertEquals;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;

import com.anthropic.client.AnthropicClient;
import com.anthropic.core.JsonValue;
import com.anthropic.models.messages.ContentBlock;
import com.anthropic.models.messages.ContentBlockParam;
import com.anthropic.models.messages.DirectCaller;
import com.anthropic.models.messages.Message;
import com.anthropic.models.messages.MessageCreateParams;
import com.anthropic.models.messages.RedactedThinkingBlock;
import com.anthropic.models.messages.TextBlock;
import com.anthropic.models.messages.ThinkingBlock;
import com.anthropic.models.messages.Tool;
import com.anthropic.models.messages.ToolResultBlockParam;
import com.anthropic.models.messages.ToolUseBlock;
import com.anthropic.models.messages.Usage;
import com.anthropic.services.blocking.MessageService;
import com.fasterxml.jackson.core.type.TypeReference;
import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
import com.google.genai.types.Content;
import com.google.genai.types.FunctionDeclaration;
import com.google.genai.types.FunctionResponse;
import com.google.genai.types.Part;
import com.google.genai.types.Schema;
import java.lang.reflect.Method;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import org.junit.Before;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.junit.runners.JUnit4;
import org.mockito.ArgumentCaptor;
import org.mockito.Mockito;

@RunWith(JUnit4.class)
public final class ClaudeTest {

private Claude claude;
private MessageService messageService;
private Method partToAnthropicMessageBlockMethod;
private Method functionDeclarationToAnthropicToolMethod;

@Before
public void setUp() throws Exception {
AnthropicClient mockClient = Mockito.mock(AnthropicClient.class);
messageService = Mockito.mock(MessageService.class);
when(mockClient.messages()).thenReturn(messageService);
claude = new Claude("claude-3-opus", mockClient);

// Access private method for testing the extraction logic
Expand All @@ -74,6 +92,38 @@ private static Map<String, Object> inputSchemaProperties(Tool tool) {
return properties.convert(new TypeReference<Map<String, Object>>() {});
}

private static Message message(ContentBlock... blocks) {
Message message = mock(Message.class);
when(message.content()).thenReturn(ImmutableList.copyOf(blocks));
return message;
}

private static ContentBlock thinkingBlock(String thinking, String signature) {
return ContentBlock.ofThinking(
ThinkingBlock.builder().thinking(thinking).signature(signature).build());
}

private static ContentBlock redactedThinkingBlock(String data) {
return ContentBlock.ofRedactedThinking(RedactedThinkingBlock.builder().data(data).build());
}

private static ContentBlock textBlock(String text) {
return ContentBlock.ofText(
TextBlock.builder().text(text).citations(ImmutableList.of()).build());
}

private static Content userText(String text) {
return Content.builder().role("user").parts(Part.fromText(text)).build();
}

private static LlmRequest request(Content... contents) {
return LlmRequest.builder().contents(ImmutableList.copyOf(contents)).build();
}

private static String signatureOf(Part part) {
return new String(part.thoughtSignature().get(), UTF_8);
}

@Test
public void testPartToAnthropicMessageBlock_mcpTool_legacyTextOutputKey() throws Exception {
Map<String, Object> responseData =
Expand Down Expand Up @@ -260,4 +310,136 @@ public void functionDeclarationToAnthropicTool_preservesRefsAndDefs() throws Exc
Map<String, Object> defs = defsValue.convert(new TypeReference<Map<String, Object>>() {});
assertThat(defs).containsKey("Pet");
}

@Test
public void generateContent_thinkingBlocks_becomeThoughtParts() {
// With display "omitted" (the Claude 5 default) a thinking block has empty text.
Message reply =
message(
thinkingBlock("", "signature"), redactedThinkingBlock("redacted-data"), textBlock("4"));
when(messageService.create(any(MessageCreateParams.class))).thenReturn(reply);

LlmResponse response = claude.generateContent(request(userText("2+2?")), false).blockingFirst();

List<Part> parts = response.content().get().parts().get();
assertThat(parts).hasSize(3);
assertThat(parts.get(0).thought()).hasValue(true);
assertThat(parts.get(0).text()).hasValue("");
assertThat(signatureOf(parts.get(0))).isEqualTo("signature");
assertThat(parts.get(1).thought()).hasValue(true);
assertThat(parts.get(1).text()).isEmpty();
assertThat(signatureOf(parts.get(1))).isEqualTo("redacted-data");
assertThat(parts.get(2).thought()).isEmpty();
assertThat(parts.get(2).text()).hasValue("4");
}

@Test
public void generateContent_toolTurn_sendsThinkingBlocksBackUnchanged() {
Message toolUseReply =
message(
thinkingBlock("", "signature-1"),
redactedThinkingBlock("redacted-data"),
textBlock("Let me check."),
thinkingBlock("Checking the weather.", "signature-2"),
ContentBlock.ofToolUse(
ToolUseBlock.builder()
.id("toolu_1")
.name("getWeather")
.input(JsonValue.from(ImmutableMap.of("city", "Seoul")))
.caller(DirectCaller.builder().build())
.build()));
Message finalReply = message(textBlock("It is sunny."));
when(messageService.create(any(MessageCreateParams.class)))
.thenReturn(toolUseReply, finalReply);
Content modelTurn =
claude
.generateContent(request(userText("Weather in Seoul?")), false)
.blockingFirst()
.content()
.get();
Content toolResult =
Content.builder()
.role("user")
.parts(
Part.builder()
.functionResponse(
FunctionResponse.builder()
.id("toolu_1")
.name("getWeather")
.response(ImmutableMap.of("result", "sunny"))
.build())
.build())
.build();

claude
.generateContent(request(userText("Weather in Seoul?"), modelTurn, toolResult), false)
.blockingFirst();

ArgumentCaptor<MessageCreateParams> params = ArgumentCaptor.forClass(MessageCreateParams.class);
verify(messageService, times(2)).create(params.capture());
List<ContentBlockParam> blocks =
params.getAllValues().get(1).messages().get(1).content().asBlockParams();
assertThat(blocks).hasSize(5);
assertThat(blocks.get(0).asThinking().thinking()).isEmpty();
assertThat(blocks.get(0).asThinking().signature()).isEqualTo("signature-1");
assertThat(blocks.get(1).asRedactedThinking().data()).isEqualTo("redacted-data");
assertThat(blocks.get(2).asText().text()).isEqualTo("Let me check.");
assertThat(blocks.get(3).asThinking().thinking()).isEqualTo("Checking the weather.");
assertThat(blocks.get(3).asThinking().signature()).isEqualTo("signature-2");
assertThat(blocks.get(4).asToolUse().id()).isEqualTo("toolu_1");
}

@Test
public void generateContent_emptyThoughtWithoutSignature_isNotSent() {
Message reply = message(textBlock("ok"));
when(messageService.create(any(MessageCreateParams.class))).thenReturn(reply);
Content modelTurn =
Content.builder()
.role("model")
.parts(
Part.builder().text("").thought(true).build(),
Part.builder().text("").thought(true).thoughtSignature(new byte[0]).build(),
Part.fromText("Hello."))
.build();

claude
.generateContent(request(userText("Hi"), modelTurn, userText("Bye")), false)
.blockingFirst();

ArgumentCaptor<MessageCreateParams> params = ArgumentCaptor.forClass(MessageCreateParams.class);
verify(messageService).create(params.capture());
// Neither thought has a signature to send back, and Anthropic rejects empty text blocks.
List<ContentBlockParam> blocks = params.getValue().messages().get(1).content().asBlockParams();
assertThat(blocks).hasSize(1);
assertThat(blocks.get(0).asText().text()).isEqualTo("Hello.");
}

@Test
public void generateContent_thoughtWithBinarySignature_isSentAsText() {
Message reply = message(textBlock("ok"));
when(messageService.create(any(MessageCreateParams.class))).thenReturn(reply);
// Another model (e.g. Gemini) stores a binary signature, which Anthropic could not decode.
Content modelTurn =
Content.builder()
.role("model")
.parts(
Part.builder()
.text("Looking up the weather.")
.thought(true)
.thoughtSignature(new byte[] {(byte) 0xC3, 0x28})
.build(),
Part.fromText("It is sunny."))
.build();

claude
.generateContent(request(userText("Hi"), modelTurn, userText("Thanks")), false)
.blockingFirst();

ArgumentCaptor<MessageCreateParams> params = ArgumentCaptor.forClass(MessageCreateParams.class);
verify(messageService).create(params.capture());
List<ContentBlockParam> blocks = params.getValue().messages().get(1).content().asBlockParams();
assertThat(blocks).hasSize(2);
assertThat(blocks.get(0).asText().text()).isEqualTo("Looking up the weather.");
assertThat(blocks.get(1).asText().text()).isEqualTo("It is sunny.");
}
}
Loading