summaryrefslogtreecommitdiffstats
path: root/src/client/bedrock.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-29 08:33:17 +0800
committerGitHub <noreply@github.com>2024-04-29 08:33:17 +0800
commit37a0cd08a92f07ef24e39bdab7c8bced5b59c146 (patch)
tree1082ed49fa584a7f8e07763146ab0e73a475b383 /src/client/bedrock.rs
parent865be2bf75bb62b6aeee059f684400b4b9938a15 (diff)
downloadaichat-37a0cd08a92f07ef24e39bdab7c8bced5b59c146.tar.gz
refactor: rename some structs (#457)
Diffstat (limited to 'src/client/bedrock.rs')
-rw-r--r--src/client/bedrock.rs22
1 files changed, 11 insertions, 11 deletions
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<Value> {
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)]