diff --git a/core/src/main/java/com/google/adk/models/GeminiUtil.java b/core/src/main/java/com/google/adk/models/GeminiUtil.java index 5986c3b5e..d0c8f0546 100644 --- a/core/src/main/java/com/google/adk/models/GeminiUtil.java +++ b/core/src/main/java/com/google/adk/models/GeminiUtil.java @@ -35,9 +35,20 @@ /** Request / Response utilities for {@link Gemini}. */ public final class GeminiUtil { + /** + * Text of the user turn appended when a request has no contents, so that the model acts on the + * system instruction. Same wording as ADK Python and ADK TypeScript. + */ + public static final String HANDLE_SYSTEM_INSTRUCTION_MESSAGE = + "Handle the requests as specified in the System Instruction."; + + /** + * Text of the user turn appended when the last content is not from the user, so that the model + * keeps producing output. Same wording as ADK Python, ADK TypeScript and ADK Go. + */ public static final String CONTINUE_OUTPUT_MESSAGE = - "Continue output. DO NOT look at this line. ONLY look at the content before this line and" - + " system instruction."; + "Continue processing previous requests as instructed. Exit or provide a summary if no more" + + " outputs are needed."; private GeminiUtil() {} @@ -193,9 +204,10 @@ private static Part removeClientFunctionCallIdFromPart(Part part) { * Ensures that the content is conducive to prompting a model response by ensuring the last * content part is from the user. * - *

If the list is empty or the last message is not from the user, a new "user" content part - * with a {@link #CONTINUE_OUTPUT_MESSAGE} is appended to the list. This is necessary to prompt - * the model to generate a response. + *

If the list is empty, a new "user" content part with {@link + * #HANDLE_SYSTEM_INSTRUCTION_MESSAGE} is appended. If the last message is not from the user, a + * new "user" content part with {@link #CONTINUE_OUTPUT_MESSAGE} is appended. This is necessary to + * prompt the model to generate a response. * * @param contents The original list of {@link Content}. * @return A list of {@link Content} where the last element is guaranteed to be from the "user". @@ -204,11 +216,10 @@ static List ensureModelResponse(List contents) { // Last content must be from the user, otherwise the model won't respond. if (contents.isEmpty() || !Ascii.equalsIgnoreCase(Iterables.getLast(contents).role().orElse(""), Role.USER)) { + String text = + contents.isEmpty() ? HANDLE_SYSTEM_INSTRUCTION_MESSAGE : CONTINUE_OUTPUT_MESSAGE; Content userContent = - Content.builder() - .parts(ImmutableList.of(Part.fromText(CONTINUE_OUTPUT_MESSAGE))) - .role(Role.USER) - .build(); + Content.builder().parts(ImmutableList.of(Part.fromText(text))).role(Role.USER).build(); return Stream.concat(contents.stream(), Stream.of(userContent)).collect(toImmutableList()); } return contents; diff --git a/core/src/test/java/com/google/adk/models/GeminiUtilTest.java b/core/src/test/java/com/google/adk/models/GeminiUtilTest.java index b0943aa50..dca88890e 100644 --- a/core/src/test/java/com/google/adk/models/GeminiUtilTest.java +++ b/core/src/test/java/com/google/adk/models/GeminiUtilTest.java @@ -37,8 +37,16 @@ @RunWith(JUnit4.class) public final class GeminiUtilTest { + // Same wording as the user turns that ADK Python appends in + // BaseLlm._maybe_append_user_content. + private static final Content SYSTEM_INSTRUCTION_CONTENT = + Content.fromParts( + Part.fromText("Handle the requests as specified in the System Instruction.")); private static final Content CONTINUE_CONTENT = - Content.fromParts(Part.fromText(GeminiUtil.CONTINUE_OUTPUT_MESSAGE)); + Content.fromParts( + Part.fromText( + "Continue processing previous requests as instructed. Exit or provide a summary if" + + " no more outputs are needed.")); @Test public void getPart0FromLlmResponse_noContent_returnsEmpty() { @@ -334,12 +342,12 @@ public void sanitizeRequestForGeminiApi_multipleContents_sanitizesAll() { } @Test - public void ensureModelResponse_emptyList_appendsContinueMessage() { + public void ensureModelResponse_emptyList_appendsSystemInstructionMessage() { ImmutableList contents = ImmutableList.of(); List result = GeminiUtil.ensureModelResponse(contents); - assertThat(result).containsExactly(CONTINUE_CONTENT); + assertThat(result).containsExactly(SYSTEM_INSTRUCTION_CONTENT); } @Test @@ -405,12 +413,13 @@ public void ensureModelResponse_lastContentIsNotUser_appendsContinueMessage() { } @Test - public void prepareGenenerateContentRequest_emptyRequest_returnsRequestWithContinueContent() { + public void + prepareGenenerateContentRequest_emptyRequest_returnsRequestWithSystemInstructionContent() { LlmRequest request = LlmRequest.builder().build(); LlmRequest result = GeminiUtil.prepareGenenerateContentRequest(request, true); - assertThat(result.contents()).containsExactly(CONTINUE_CONTENT); + assertThat(result.contents()).containsExactly(SYSTEM_INSTRUCTION_CONTENT); assertThat(result.config()).isEmpty(); }