Skip to content
Merged
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
95 changes: 84 additions & 11 deletions core/src/main/java/com/google/adk/models/BedrockBaseLM.java
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
import com.google.genai.types.FunctionDeclaration;
import com.google.genai.types.GenerateContentConfig;
import com.google.genai.types.GenerateContentResponseUsageMetadata;
import com.google.genai.types.MediaModality;
import com.google.genai.types.ModalityTokenCount;
import com.google.genai.types.Part;
import com.google.genai.types.Schema;
import io.reactivex.rxjava3.core.Flowable;
Expand Down Expand Up @@ -626,6 +628,8 @@ private Flowable<LlmResponse> createRobustStreamingResponse(
final AtomicInteger inputTokens = new AtomicInteger(0);
final AtomicInteger outputTokens = new AtomicInteger(0);
final AtomicInteger totalTokens = new AtomicInteger(0);
final AtomicInteger promptAudioTokens = new AtomicInteger(0);
final AtomicInteger completionAudioTokens = new AtomicInteger(0);

return Flowable.generate(
() -> callLLMChatStream(modelId, messages, functions),
Expand All @@ -642,7 +646,12 @@ private Flowable<LlmResponse> createRobustStreamingResponse(
if (accumulatedText.length() > 0) {
// Create usage metadata from accumulated token counts
GenerateContentResponseUsageMetadata usageMetadata =
getUsageMetadata(inputTokens.get(), outputTokens.get(), totalTokens.get());
getUsageMetadata(
inputTokens.get(),
outputTokens.get(),
totalTokens.get(),
promptAudioTokens.get(),
completionAudioTokens.get());

LlmResponse.Builder finalResponseBuilder =
LlmResponse.builder()
Expand Down Expand Up @@ -693,6 +702,18 @@ private Flowable<LlmResponse> createRobustStreamingResponse(
int total = usage.getInt("totalTokens");
totalTokens.set(total);
}
if (usage.has("prompt_tokens_details")) {
JSONObject pDetails = usage.optJSONObject("prompt_tokens_details");
if (pDetails != null && pDetails.has("audio_tokens")) {
promptAudioTokens.set(pDetails.getInt("audio_tokens"));
}
}
if (usage.has("completion_tokens_details")) {
JSONObject cDetails = usage.optJSONObject("completion_tokens_details");
if (cDetails != null && cDetails.has("audio_tokens")) {
completionAudioTokens.set(cDetails.getInt("audio_tokens"));
}
}
}

JSONObject message = null;
Expand Down Expand Up @@ -792,7 +813,12 @@ private Flowable<LlmResponse> createRobustStreamingResponse(

// Create usage metadata from accumulated token counts
GenerateContentResponseUsageMetadata usageMetadata =
getUsageMetadata(inputTokens.get(), outputTokens.get(), totalTokens.get());
getUsageMetadata(
inputTokens.get(),
outputTokens.get(),
totalTokens.get(),
promptAudioTokens.get(),
completionAudioTokens.get());

// Handle function call completion
if (inFunctionCall.get() && functionCallName.length() > 0) {
Expand Down Expand Up @@ -1284,18 +1310,41 @@ public Flowable<JSONObject> generateContent(

// Add overloaded method for streaming token usage
private GenerateContentResponseUsageMetadata getUsageMetadata(
int promptTokens, int completionTokens, int totalTokens) {
int promptTokens,
int completionTokens,
int totalTokens,
int promptAudioTokens,
int completionAudioTokens) {
if (totalTokens > 0 || promptTokens > 0 || completionTokens > 0) {
logger.info(
"Streaming token counts: prompt={}, completion={}, total={}",
promptTokens,
completionTokens,
totalTokens);
return GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens)
.build();
GenerateContentResponseUsageMetadata.Builder builder =
GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens);

if (promptAudioTokens > 0) {
builder.promptTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(promptAudioTokens)
.build()));
}
if (completionAudioTokens > 0) {
builder.candidatesTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(completionAudioTokens)
.build()));
}

return builder.build();
}
return null;
}
Expand All @@ -1322,12 +1371,36 @@ private GenerateContentResponseUsageMetadata getUsageMetadata(JSONObject agentRe
promptTokens,
completionTokens,
totalTokens);
return Optional.of(
GenerateContentResponseUsageMetadata.Builder builder =
GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens)
.build());
.totalTokenCount(totalTokens);

if (usage.has("prompt_tokens_details")) {
JSONObject pDetails = usage.optJSONObject("prompt_tokens_details");
if (pDetails != null && pDetails.has("audio_tokens")) {
builder.promptTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(pDetails.getInt("audio_tokens"))
.build()));
}
}
if (usage.has("completion_tokens_details")) {
JSONObject cDetails = usage.optJSONObject("completion_tokens_details");
if (cDetails != null && cDetails.has("audio_tokens")) {
builder.candidatesTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(cDetails.getInt("audio_tokens"))
.build()));
}
}

return Optional.of(builder.build());
}
}
}
Expand Down
109 changes: 97 additions & 12 deletions core/src/main/java/com/google/adk/models/OllamaBaseLM.java
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@
import com.google.genai.types.FunctionDeclaration;
import com.google.genai.types.GenerateContentConfig;
import com.google.genai.types.GenerateContentResponseUsageMetadata;
import com.google.genai.types.MediaModality;
import com.google.genai.types.ModalityTokenCount;
import com.google.genai.types.Part;
import com.google.genai.types.Schema;
import io.reactivex.rxjava3.core.Flowable;
Expand Down Expand Up @@ -461,6 +463,8 @@ private Flowable<LlmResponse> createRobustStreamingResponse(
final AtomicInteger outputTokens = new AtomicInteger(0);
final AtomicLong promptEvalDuration = new AtomicLong(0);
final AtomicLong evalDuration = new AtomicLong(0);
final AtomicInteger promptAudioTokens = new AtomicInteger(0);
final AtomicInteger completionAudioTokens = new AtomicInteger(0);

return Flowable.generate(
() -> callLLMChatStream(modelId, messages, functions),
Expand Down Expand Up @@ -530,7 +534,33 @@ private Flowable<LlmResponse> createRobustStreamingResponse(
if (responseJson.optBoolean("done", false)) {
streamCompleted.set(true);

GenerateContentResponseUsageMetadata usageMetadata = getUsageMetadata(responseJson);
if (responseJson.has("prompt_eval_count")) {
inputTokens.set(responseJson.getInt("prompt_eval_count"));
}
if (responseJson.has("eval_count")) {
outputTokens.set(responseJson.getInt("eval_count"));
}
// Check for audio tokens if Ollama adds them in the future
if (responseJson.has("prompt_tokens_details")) {
JSONObject pDetails = responseJson.optJSONObject("prompt_tokens_details");
if (pDetails != null && pDetails.has("audio_tokens")) {
promptAudioTokens.set(pDetails.getInt("audio_tokens"));
}
}
if (responseJson.has("completion_tokens_details")) {
JSONObject cDetails = responseJson.optJSONObject("completion_tokens_details");
if (cDetails != null && cDetails.has("audio_tokens")) {
completionAudioTokens.set(cDetails.getInt("audio_tokens"));
}
}

GenerateContentResponseUsageMetadata usageMetadata =
getUsageMetadata(
inputTokens.get(),
outputTokens.get(),
inputTokens.get() + outputTokens.get(),
promptAudioTokens.get(),
completionAudioTokens.get());

if (accumulatedText.length() > 0 && !inFunctionCall.get()) {
LlmResponse.Builder aggregatedResponseBuilder =
Expand Down Expand Up @@ -612,13 +642,35 @@ private LlmResponse createTextResponse(String text, boolean partial) {
}

private GenerateContentResponseUsageMetadata getUsageMetadata(
int promptTokens, int completionTokens, int totalTokens) {
int promptTokens,
int completionTokens,
int totalTokens,
int promptAudioTokens,
int completionAudioTokens) {
if (totalTokens > 0 || promptTokens > 0 || completionTokens > 0) {
return GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens)
.build();
GenerateContentResponseUsageMetadata.Builder builder =
GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens);

if (promptAudioTokens > 0) {
builder.promptTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(promptAudioTokens)
.build()));
}
if (completionAudioTokens > 0) {
builder.candidatesTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(completionAudioTokens)
.build()));
}
return builder.build();
}
return null;
}
Expand All @@ -632,6 +684,8 @@ private GenerateContentResponseUsageMetadata getUsageMetadata(JSONObject agentRe
int promptTokens = 0;
int completionTokens = 0;
int totalTokens = 0;
int promptAudioTokens = 0;
int completionAudioTokens = 0;

if (agentResponse.has("prompt_eval_count")) {
promptTokens = agentResponse.getInt("prompt_eval_count");
Expand All @@ -642,17 +696,48 @@ private GenerateContentResponseUsageMetadata getUsageMetadata(JSONObject agentRe
}
totalTokens = promptTokens + completionTokens;

if (agentResponse.has("prompt_tokens_details")) {
JSONObject pDetails = agentResponse.optJSONObject("prompt_tokens_details");
if (pDetails != null && pDetails.has("audio_tokens")) {
promptAudioTokens = pDetails.getInt("audio_tokens");
}
}
if (agentResponse.has("completion_tokens_details")) {
JSONObject cDetails = agentResponse.optJSONObject("completion_tokens_details");
if (cDetails != null && cDetails.has("audio_tokens")) {
completionAudioTokens = cDetails.getInt("audio_tokens");
}
}

if (totalTokens > 0 || promptTokens > 0 || completionTokens > 0) {
logger.info(
"Ollama token counts: prompt={}, completion={}, total={}",
promptTokens,
completionTokens,
totalTokens);
return GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens)
.build();
GenerateContentResponseUsageMetadata.Builder builder =
GenerateContentResponseUsageMetadata.builder()
.promptTokenCount(promptTokens)
.candidatesTokenCount(completionTokens)
.totalTokenCount(totalTokens > 0 ? totalTokens : promptTokens + completionTokens);

if (promptAudioTokens > 0) {
builder.promptTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(promptAudioTokens)
.build()));
}
if (completionAudioTokens > 0) {
builder.candidatesTokensDetails(
ImmutableList.of(
ModalityTokenCount.builder()
.modality(MediaModality.Known.AUDIO)
.tokenCount(completionAudioTokens)
.build()));
}
return builder.build();
}
} catch (Exception e) {
logger.warn("Failed to parse token usage from Ollama response", e);
Expand Down
Loading
Loading