summaryrefslogtreecommitdiffstats
path: root/src/client/bedrock.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/bedrock.rs')
-rw-r--r--src/client/bedrock.rs51
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,