From 865be2bf75bb62b6aeee059f684400b4b9938a15 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 29 Apr 2024 06:51:03 +0800 Subject: feat: non-streaming returns completion stats (#456) --- src/client/vertexai.rs | 25 +++++++++++++++++++------ 1 file changed, 19 insertions(+), 6 deletions(-) (limited to 'src/client/vertexai.rs') diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 317ed2a..f4079f0 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,7 +1,7 @@ use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming}; use super::{ - catch_error, json_stream, message::*, patch_system_message, Client, ExtraConfig, Model, - ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient, + catch_error, json_stream, message::*, patch_system_message, Client, CompletionStats, + ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient, }; use crate::utils::PromptKind; @@ -81,7 +81,11 @@ impl VertexAIClient { impl Client for VertexAIClient { client_common_fns!(); - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result { + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)> { let model_category = ModelCategory::from_str(&self.model.name)?; self.prepare_access_token().await?; let builder = self.request_builder(client, data, &model_category)?; @@ -107,15 +111,14 @@ impl Client for VertexAIClient { } } -pub async fn gemini_send_message(builder: RequestBuilder) -> Result { +pub async fn gemini_send_message(builder: RequestBuilder) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; if status != 200 { catch_error(&data, status.as_u16())?; } - let output = gemini_extract_text(&data)?; - Ok(output.to_string()) + gemini_extract_completion_text(&data) } pub async fn gemini_send_message_streaming( @@ -138,6 +141,16 @@ pub async fn gemini_send_message_streaming( Ok(()) } +fn gemini_extract_completion_text(data: &Value) -> Result<(String, CompletionStats)> { + let text = gemini_extract_text(data)?; + let stats = CompletionStats { + id: None, + input_tokens: data["usageMetadata"]["promptTokenCount"].as_u64(), + output_tokens: data["usageMetadata"]["candidatesTokenCount"].as_u64(), + }; + Ok((text.to_string(), stats)) +} + fn gemini_extract_text(data: &Value) -> Result<&str> { match data["candidates"][0]["content"]["parts"][0]["text"].as_str() { Some(text) => Ok(text), -- cgit v1.2.3