summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-30 00:25:35 +0000
committersigoden <sigoden@gmail.com>2024-04-30 00:25:35 +0000
commite7fa6c5a208347b0e5aa779b3ea477f5f1fe41c6 (patch)
treed6888543d645d63dbf98363365c14cdda9de1406 /src/client
parenta50b32ca21d1d103c6c5f239f30fe8062a23e1fa (diff)
downloadaichat-e7fa6c5a208347b0e5aa779b3ea477f5f1fe41c6.tar.gz
refactor: improve code quality
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs4
-rw-r--r--src/client/bedrock.rs46
-rw-r--r--src/client/claude.rs6
-rw-r--r--src/client/cloudflare.rs4
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/ernie.rs4
-rw-r--r--src/client/gemini.rs4
-rw-r--r--src/client/ollama.rs4
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/openai_compatible.rs6
-rw-r--r--src/client/qianwen.rs9
-rw-r--r--src/client/replicate.rs5
-rw-r--r--src/client/vertexai.rs4
13 files changed, 28 insertions, 76 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index dd83ae1..005351f 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -1,7 +1,5 @@
use super::openai::openai_build_body;
-use super::{AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptType, SendData};
-
-use crate::utils::PromptKind;
+use super::{AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData};
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 61cb799..a889162 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -1,21 +1,19 @@
use super::claude::{claude_build_body, claude_extract_completion};
use super::{
catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model,
- ModelConfig, PromptFormat, PromptType, SendData, SseHandler, LLAMA2_PROMPT_FORMAT,
+ ModelConfig, PromptFormat, PromptKind, PromptType, SendData, SseHandler, LLAMA2_PROMPT_FORMAT,
LLAMA3_PROMPT_FORMAT,
};
-use crate::utils::PromptKind;
+use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256};
use anyhow::{anyhow, bail, Result};
use async_trait::async_trait;
use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder};
use aws_smithy_eventstream::smithy::parse_response_headers;
-use base64::{engine::general_purpose::STANDARD, Engine};
use bytes::BytesMut;
use chrono::{DateTime, Utc};
use futures_util::StreamExt;
-use hmac::{Hmac, Mac};
use indexmap::IndexMap;
use reqwest::{
header::{HeaderMap, HeaderName, HeaderValue},
@@ -23,7 +21,6 @@ use reqwest::{
};
use serde::Deserialize;
use serde_json::{json, Value};
-use sha2::{Digest, Sha256};
use std::str::FromStr;
#[derive(Debug, Clone, Deserialize)]
@@ -199,7 +196,7 @@ async fn send_message_streaming(
}
}
("exception", _) => {
- let payload = STANDARD.decode(message.payload())?;
+ let payload = base64_decode(message.payload())?;
let data = String::from_utf8_lossy(&payload);
bail!("Invalid response data: {data} (smithy_type: {smithy_type})")
@@ -402,7 +399,7 @@ fn aws_fetch(
region,
&service,
);
- let signature = sign(&signing_key, &string_to_sign);
+ let signature = hmac_sha256(&signing_key, &string_to_sign);
let signature = hex_encode(&signature);
let authorization_header = format!(
@@ -426,41 +423,16 @@ fn aws_fetch(
Ok(request_builder)
}
-fn sha256(data: &str) -> String {
- let mut hasher = Sha256::new();
- hasher.update(data.as_bytes());
- format!("{:x}", hasher.finalize())
-}
-
-fn sign(key: &[u8], msg: &str) -> Vec<u8> {
- let mut mac = Hmac::<Sha256>::new_from_slice(key).expect("HMAC can take key of any size");
- mac.update(msg.as_bytes());
- mac.finalize().into_bytes().to_vec()
-}
-
fn gen_signing_key(key: &str, date_stamp: &str, region: &str, service: &str) -> Vec<u8> {
- let k_date = sign(format!("AWS4{}", key).as_bytes(), date_stamp);
- let k_region = sign(&k_date, region);
- let k_service = sign(&k_region, service);
- sign(&k_service, "aws4_request")
-}
-
-fn hex_encode(bytes: &[u8]) -> String {
- bytes
- .iter()
- .fold(String::new(), |acc, b| acc + &format!("{:02x}", b))
-}
-
-fn encode_uri(uri: &str) -> String {
- uri.split('/')
- .map(|v| urlencoding::encode(v))
- .collect::<Vec<_>>()
- .join("/")
+ let k_date = hmac_sha256(format!("AWS4{}", key).as_bytes(), date_stamp);
+ let k_region = hmac_sha256(&k_date, region);
+ let k_service = hmac_sha256(&k_region, service);
+ hmac_sha256(&k_service, "aws4_request")
}
fn decode_chunk(data: &[u8]) -> Option<Value> {
let data = serde_json::from_slice::<Value>(data).ok()?;
let data = data["bytes"].as_str()?;
- let data = STANDARD.decode(data).ok()?;
+ let data = base64_decode(data).ok()?;
serde_json::from_slice(&data).ok()
}
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 0bd248b..84b4a6c 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,11 +1,9 @@
use super::{
catch_error, extract_system_message, sse_stream, ClaudeClient, CompletionDetails, ExtraConfig,
- ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptType, SendData,
- SsMmessage, SseHandler,
+ ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptKind, PromptType,
+ SendData, SsMmessage, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, bail, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 33520a7..14a5828 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,10 +1,8 @@
use super::{
catch_error, sse_stream, CloudflareClient, CompletionDetails, ExtraConfig, Model, ModelConfig,
- PromptType, SendData, SsMmessage, SseHandler,
+ PromptKind, PromptType, SendData, SsMmessage, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 7f07640..0069c2c 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,10 +1,8 @@
use super::{
catch_error, extract_system_message, json_stream, message::*, CohereClient, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptType, SendData, SseHandler,
+ ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, bail, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 5e600af..a0187a4 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,10 +1,8 @@
use super::{
maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient,
- ExtraConfig, Model, ModelConfig, PromptType, SendData, SsMmessage, SseHandler,
+ ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SsMmessage, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, Context, Result};
use async_trait::async_trait;
use chrono::Utc;
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 7f51739..783b674 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -1,7 +1,5 @@
use super::vertexai::gemini_build_body;
-use super::{ExtraConfig, GeminiClient, Model, ModelConfig, PromptType, SendData};
-
-use crate::utils::PromptKind;
+use super::{ExtraConfig, GeminiClient, Model, ModelConfig, PromptKind, PromptType, SendData};
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index eebf301..a2688a5 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,10 +1,8 @@
use super::{
catch_error, message::*, CompletionDetails, ExtraConfig, Model, ModelConfig, OllamaClient,
- PromptType, SendData, SseHandler,
+ PromptKind, PromptType, SendData, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, bail, Result};
use futures_util::StreamExt;
use reqwest::{Client as ReqwestClient, RequestBuilder};
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 76a9f78..7e8fb87 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,10 +1,8 @@
use super::{
catch_error, sse_stream, CompletionDetails, ExtraConfig, Model, ModelConfig, OpenAIClient,
- PromptType, SendData, SsMmessage, SseHandler,
+ PromptKind, PromptType, SendData, SsMmessage, SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index ba3a751..d7aff2b 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -1,7 +1,7 @@
use super::openai::openai_build_body;
-use super::{ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptType, SendData};
-
-use crate::utils::PromptKind;
+use super::{
+ ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptKind, PromptType, SendData,
+};
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 2f298c3..b1ba093 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,13 +1,12 @@
use super::{
maybe_catch_error, message::*, sse_stream, Client, CompletionDetails, ExtraConfig, Model,
- ModelConfig, PromptType, QianwenClient, SendData, SsMmessage, SseHandler,
+ ModelConfig, PromptKind, PromptType, QianwenClient, SendData, SsMmessage, SseHandler,
};
-use crate::utils::{sha256sum, PromptKind};
+use crate::utils::{base64_decode, sha256};
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
-use base64::{engine::general_purpose::STANDARD, Engine};
use reqwest::{
multipart::{Form, Part},
Client as ReqwestClient, RequestBuilder,
@@ -254,12 +253,12 @@ async fn upload(model: &str, api_key: &str, url: &str) -> Result<String> {
.strip_prefix("data:")
.and_then(|v| v.split_once(";base64,"))
.ok_or_else(|| anyhow!("Invalid image url"))?;
- let mut name = sha256sum(data);
+ let mut name = sha256(data);
if let Some(ext) = mime_type.strip_prefix("image/") {
name.push('.');
name.push_str(ext);
}
- let data = STANDARD.decode(data)?;
+ let data = base64_decode(data)?;
let client = reqwest::Client::new();
let policy: Policy = client
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index e399927..aef992d 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -2,11 +2,10 @@ use std::time::Duration;
use super::{
catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptType, ReplicateClient, SendData, SsMmessage, SseHandler,
+ ExtraConfig, Model, ModelConfig, PromptKind, PromptType, ReplicateClient, SendData, SsMmessage,
+ SseHandler,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, Result};
use async_trait::async_trait;
use reqwest::{Client as ReqwestClient, RequestBuilder};
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index e798451..e801025 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,11 +1,9 @@
use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming};
use super::{
catch_error, json_stream, message::*, patch_system_message, Client, CompletionDetails,
- ExtraConfig, Model, ModelConfig, PromptType, SendData, SseHandler, VertexAIClient,
+ ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SseHandler, VertexAIClient,
};
-use crate::utils::PromptKind;
-
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
use chrono::{Duration, Utc};