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
29 changes: 20 additions & 9 deletions core/src/main/java/com/google/adk/models/GeminiUtil.java
Original file line number Diff line number Diff line change
Expand Up @@ -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() {}

Expand Down Expand Up @@ -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.
*
* <p>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.
* <p>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".
Expand All @@ -204,11 +216,10 @@ static List<Content> ensureModelResponse(List<Content> 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;
Expand Down
19 changes: 14 additions & 5 deletions core/src/test/java/com/google/adk/models/GeminiUtilTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down Expand Up @@ -334,12 +342,12 @@ public void sanitizeRequestForGeminiApi_multipleContents_sanitizesAll() {
}

@Test
public void ensureModelResponse_emptyList_appendsContinueMessage() {
public void ensureModelResponse_emptyList_appendsSystemInstructionMessage() {
ImmutableList<Content> contents = ImmutableList.of();

List<Content> result = GeminiUtil.ensureModelResponse(contents);

assertThat(result).containsExactly(CONTINUE_CONTENT);
assertThat(result).containsExactly(SYSTEM_INSTRUCTION_CONTENT);
}

@Test
Expand Down Expand Up @@ -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();
}

Expand Down
Loading