summaryrefslogtreecommitdiffstats
path: root/src/client/bedrock.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-30 17:10:02 +0800
committerGitHub <noreply@github.com>2024-05-30 17:10:02 +0800
commit54a837784c3df7e73c7c95adfef96e68c842984b (patch)
treee090da40436bd3c1220d5ab88071ac33397ce57c /src/client/bedrock.rs
parentfa4bf14e02b1e6361b38014fe6a986738cf06837 (diff)
downloadaichat-54a837784c3df7e73c7c95adfef96e68c842984b.tar.gz
refactor: rename `SendData` to `CompletionData` (#553)
Diffstat (limited to 'src/client/bedrock.rs')
-rw-r--r--src/client/bedrock.rs26
1 files changed, 15 insertions, 11 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 1f13c7f..355565f 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -1,7 +1,7 @@
use super::{
- prompt_format::*, claude::*,
- catch_error, BedrockClient, Client, CompletionOutput, ExtraConfig, Model, ModelData,
- ModelPatches, PromptAction, PromptKind, SendData, SseHandler,
+ catch_error, claude::*, prompt_format::*, BedrockClient, Client, CompletionData,
+ CompletionOutput, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind,
+ SseHandler,
};
use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256};
@@ -41,7 +41,7 @@ impl Client for BedrockClient {
async fn send_message_inner(
&self,
client: &ReqwestClient,
- data: SendData,
+ data: CompletionData,
) -> Result<CompletionOutput> {
let model_category = ModelCategory::from_str(self.model.name())?;
let builder = self.request_builder(client, data, &model_category)?;
@@ -52,7 +52,7 @@ impl Client for BedrockClient {
&self,
client: &ReqwestClient,
handler: &mut SseHandler,
- data: SendData,
+ data: CompletionData,
) -> Result<()> {
let model_category = ModelCategory::from_str(self.model.name())?;
let builder = self.request_builder(client, data, &model_category)?;
@@ -84,7 +84,7 @@ impl BedrockClient {
fn request_builder(
&self,
client: &ReqwestClient,
- data: SendData,
+ data: CompletionData,
model_category: &ModelCategory,
) -> Result<RequestBuilder> {
let access_key_id = self.get_access_key_id()?;
@@ -211,7 +211,11 @@ async fn send_message_streaming(
Ok(())
}
-fn build_body(data: SendData, model: &Model, model_category: &ModelCategory) -> Result<Value> {
+fn build_body(
+ data: CompletionData,
+ model: &Model,
+ model_category: &ModelCategory,
+) -> Result<Value> {
match model_category {
ModelCategory::Anthropic => {
let mut body = claude_build_body(data, model)?;
@@ -227,8 +231,8 @@ fn build_body(data: SendData, model: &Model, model_category: &ModelCategory) ->
}
}
-fn meta_llama_build_body(data: SendData, model: &Model, pt: PromptFormat) -> Result<Value> {
- let SendData {
+fn meta_llama_build_body(data: CompletionData, model: &Model, pt: PromptFormat) -> Result<Value> {
+ let CompletionData {
messages,
temperature,
top_p,
@@ -251,8 +255,8 @@ fn meta_llama_build_body(data: SendData, model: &Model, pt: PromptFormat) -> Res
Ok(body)
}
-fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> {
- let SendData {
+fn mistral_build_body(data: CompletionData, model: &Model) -> Result<Value> {
+ let CompletionData {
messages,
temperature,
top_p,