From 37a0cd08a92f07ef24e39bdab7c8bced5b59c146 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 29 Apr 2024 08:33:17 +0800 Subject: refactor: rename some structs (#457) --- src/client/bedrock.rs | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) (limited to 'src/client/bedrock.rs') diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index dd94a41..6d9379d 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,7 +1,7 @@ use super::claude::{claude_build_body, claude_extract_completion}; use super::{ - catch_error, generate_prompt, BedrockClient, Client, CompletionStats, ExtraConfig, Model, - ModelConfig, PromptFormat, PromptType, ReplyHandler, SendData, LLAMA2_PROMPT_FORMAT, + catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model, + ModelConfig, PromptFormat, PromptType, SendData, SseHandler, LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, }; @@ -45,7 +45,7 @@ impl Client for BedrockClient { &self, client: &ReqwestClient, data: SendData, - ) -> Result<(String, CompletionStats)> { + ) -> Result<(String, CompletionDetails)> { let model_category = ModelCategory::from_str(&self.model.name)?; let builder = self.request_builder(client, data, &model_category)?; send_message(builder, &model_category).await @@ -54,7 +54,7 @@ impl Client for BedrockClient { async fn send_message_streaming_inner( &self, client: &ReqwestClient, - handler: &mut ReplyHandler, + handler: &mut SseHandler, data: SendData, ) -> Result<()> { let model_category = ModelCategory::from_str(&self.model.name)?; @@ -132,7 +132,7 @@ impl BedrockClient { async fn send_message( builder: RequestBuilder, model_category: &ModelCategory, -) -> Result<(String, CompletionStats)> { +) -> Result<(String, CompletionDetails)> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -150,7 +150,7 @@ async fn send_message( async fn send_message_streaming( builder: RequestBuilder, - handler: &mut ReplyHandler, + handler: &mut SseHandler, model_category: &ModelCategory, ) -> Result<()> { let res = builder.send().await?; @@ -275,23 +275,23 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result { Ok(body) } -fn llama_extract_completion(data: &Value) -> Result<(String, CompletionStats)> { +fn llama_extract_completion(data: &Value) -> Result<(String, CompletionDetails)> { let text = data["generation"] .as_str() .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - let stats = CompletionStats { + let details = CompletionDetails { id: None, input_tokens: data["prompt_token_count"].as_u64(), output_tokens: data["generation_token_count"].as_u64(), }; - Ok((text.to_string(), stats)) + Ok((text.to_string(), details)) } -fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionStats)> { +fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionDetails)> { let text = data["outputs"][0]["text"] .as_str() .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - Ok((text.to_string(), CompletionStats::default())) + Ok((text.to_string(), CompletionDetails::default())) } #[derive(Debug, Clone, Copy, PartialEq, Eq)] -- cgit v1.2.3