From e7fa6c5a208347b0e5aa779b3ea477f5f1fe41c6 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 30 Apr 2024 00:25:35 +0000 Subject: refactor: improve code quality --- src/client/azure_openai.rs | 4 +--- src/client/bedrock.rs | 46 ++++++++--------------------------------- src/client/claude.rs | 6 ++---- src/client/cloudflare.rs | 4 +--- src/client/cohere.rs | 4 +--- src/client/ernie.rs | 4 +--- src/client/gemini.rs | 4 +--- src/client/ollama.rs | 4 +--- src/client/openai.rs | 4 +--- src/client/openai_compatible.rs | 6 +++--- src/client/qianwen.rs | 9 ++++---- src/client/replicate.rs | 5 ++--- src/client/vertexai.rs | 4 +--- src/config/input.rs | 9 ++++---- src/utils/crypto.rs | 35 +++++++++++++++++++++++++++++++ src/utils/mod.rs | 10 ++------- 16 files changed, 69 insertions(+), 89 deletions(-) create mode 100644 src/utils/crypto.rs 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 { - let mut mac = Hmac::::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 { - 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::>() - .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 { let data = serde_json::from_slice::(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 { .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}; diff --git a/src/config/input.rs b/src/config/input.rs index eaa567a..b7b896c 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -2,10 +2,9 @@ use super::role::Role; use super::session::Session; use crate::client::{ImageUrl, MessageContent, MessageContentPart, ModelCapabilities}; -use crate::utils::sha256sum; +use crate::utils::{base64_encode, sha256}; use anyhow::{bail, Context, Result}; -use base64::{self, engine::general_purpose::STANDARD, Engine}; use fancy_regex::Regex; use lazy_static::lazy_static; use mime_guess::from_path; @@ -56,7 +55,7 @@ impl Input { if is_image { let data_url = read_media_to_data_url(&file_path) .with_context(|| format!("Unable to read media file '{file_item}'"))?; - data_urls.insert(sha256sum(&data_url), file_path.display().to_string()); + data_urls.insert(sha256(&data_url), file_path.display().to_string()); medias.push(data_url) } else { let text = read_file(&file_path) @@ -211,7 +210,7 @@ impl InputContext { pub fn resolve_data_url(data_urls: &HashMap, data_url: String) -> String { if data_url.starts_with("data:") { - let hash = sha256sum(&data_url); + let hash = sha256(&data_url); if let Some(path) = data_urls.get(&hash) { return path.to_string(); } @@ -251,7 +250,7 @@ fn read_media_to_data_url>(image_path: P) -> Result { let mut buffer = Vec::new(); file.read_to_end(&mut buffer)?; - let encoded_image = STANDARD.encode(buffer); + let encoded_image = base64_encode(buffer); let data_url = format!("data:{};base64,{}", mime_type, encoded_image); Ok(data_url) diff --git a/src/utils/crypto.rs b/src/utils/crypto.rs new file mode 100644 index 0000000..9f0ee82 --- /dev/null +++ b/src/utils/crypto.rs @@ -0,0 +1,35 @@ +use base64::{engine::general_purpose::STANDARD, Engine}; +use hmac::{Hmac, Mac}; +use sha2::{Digest, Sha256}; + +pub fn sha256(input: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(input); + format!("{:x}", hasher.finalize()) +} + +pub fn hmac_sha256(key: &[u8], msg: &str) -> Vec { + let mut mac = Hmac::::new_from_slice(key).expect("HMAC can take key of any size"); + mac.update(msg.as_bytes()); + mac.finalize().into_bytes().to_vec() +} + +pub fn hex_encode(bytes: &[u8]) -> String { + bytes + .iter() + .fold(String::new(), |acc, b| acc + &format!("{:02x}", b)) +} + +pub fn encode_uri(uri: &str) -> String { + uri.split('/') + .map(|v| urlencoding::encode(v)) + .collect::>() + .join("/") +} + +pub fn base64_encode>(input: T) -> String { + STANDARD.encode(input) +} +pub fn base64_decode>(input: T) -> Result, base64::DecodeError> { + STANDARD.decode(input) +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 771baf3..86d84d8 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1,5 +1,6 @@ mod abort_signal; mod clipboard; +mod crypto; mod prompt_input; mod render_prompt; mod spinner; @@ -7,6 +8,7 @@ mod tiktoken; pub use self::abort_signal::{create_abort_signal, AbortSignal}; pub use self::clipboard::set_text; +pub use self::crypto::*; pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::spinner::run_spinner; @@ -14,7 +16,6 @@ pub use self::tiktoken::cl100k_base_singleton; use fancy_regex::Regex; use lazy_static::lazy_static; -use sha2::{Digest, Sha256}; use std::env; use std::process::Command; @@ -82,13 +83,6 @@ pub fn light_theme_from_colorfgbg(colorfgbg: &str) -> Option { Some(light) } -pub fn sha256sum(input: &str) -> String { - let mut hasher = Sha256::new(); - hasher.update(input); - let result = hasher.finalize(); - format!("{:x}", result) -} - pub fn detect_os() -> String { let os = env::consts::OS; if os == "linux" { -- cgit v1.2.3