diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-29 06:51:03 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-29 06:51:03 +0800 |
| commit | 865be2bf75bb62b6aeee059f684400b4b9938a15 (patch) | |
| tree | b6296ccdde2af49251b02087b9f51c5239cc4d2b /src/client/bedrock.rs | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client/bedrock.rs')
| -rw-r--r-- | src/client/bedrock.rs | 51 |
1 files changed, 37 insertions, 14 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 5f0a385..dd94a41 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,7 +1,8 @@ -use super::claude::claude_build_body; +use super::claude::{claude_build_body, claude_extract_completion}; use super::{ - catch_error, generate_prompt, BedrockClient, Client, ExtraConfig, Model, ModelConfig, - PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, + catch_error, generate_prompt, BedrockClient, Client, CompletionStats, ExtraConfig, Model, + ModelConfig, PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, + LLAMA3_PROMPT_FORMAT, }; use crate::utils::PromptKind; @@ -40,7 +41,11 @@ pub struct BedrockConfig { impl Client for BedrockClient { client_common_fns!(); - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)> { let model_category = ModelCategory::from_str(&self.model.name)?; let builder = self.request_builder(client, data, &model_category)?; send_message(builder, &model_category).await @@ -124,7 +129,10 @@ impl BedrockClient { } } -async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) -> Result<String> { +async fn send_message( + builder: RequestBuilder, + model_category: &ModelCategory, +) -> Result<(String, CompletionStats)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -133,15 +141,11 @@ async fn send_message(builder: RequestBuilder, model_category: &ModelCategory) - catch_error(&data, status.as_u16())?; } - let output = match model_category { - ModelCategory::Anthropic => data["content"][0]["text"].as_str(), - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => data["generation"].as_str(), - ModelCategory::Mistral => data["outputs"][0]["text"].as_str(), - }; - - let output = output.ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(output.to_string()) + match model_category { + ModelCategory::Anthropic => claude_extract_completion(&data), + ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => llama_extract_completion(&data), + ModelCategory::Mistral => mistral_extrat_completion(&data), + } } async fn send_message_streaming( @@ -271,6 +275,25 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> { Ok(body) } +fn llama_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["generation"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + let stats = CompletionStats { + id: None, + input_tokens: data["prompt_token_count"].as_u64(), + output_tokens: data["generation_token_count"].as_u64(), + }; + Ok((text.to_string(), stats)) +} + +fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionStats)> { + let text = data["outputs"][0]["text"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + Ok((text.to_string(), CompletionStats::default())) +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ModelCategory { Anthropic, |
