summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-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
-rw-r--r--src/config/input.rs9
-rw-r--r--src/utils/crypto.rs35
-rw-r--r--src/utils/mod.rs10
16 files changed, 69 insertions, 89 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};
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<String, String>, 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<P: AsRef<Path>>(image_path: P) -> Result<String> {
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<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()
+}
+
+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::<Vec<_>>()
+ .join("/")
+}
+
+pub fn base64_encode<T: AsRef<[u8]>>(input: T) -> String {
+ STANDARD.encode(input)
+}
+pub fn base64_decode<T: AsRef<[u8]>>(input: T) -> Result<Vec<u8>, 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<bool> {
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" {