summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/claude.rs25
-rw-r--r--src/client/cohere.rs14
-rw-r--r--src/client/common.rs44
-rw-r--r--src/client/ernie.rs60
-rw-r--r--src/client/gemini.rs8
-rw-r--r--src/client/ollama.rs12
-rw-r--r--src/client/openai.rs16
-rw-r--r--src/client/qianwen.rs16
-rw-r--r--src/client/vertexai.rs46
9 files changed, 98 insertions, 143 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 68ab509..68e9567 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -1,6 +1,6 @@
use super::{
- extract_sytem_message, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent,
- MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData,
+ catch_error, extract_sytem_message, ClaudeClient, Client, ExtraConfig, ImageUrl,
+ MessageContent, MessageContentPart, Model, ModelConfig, PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -30,7 +30,7 @@ impl Client for ClaudeClient {
async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
let builder = self.request_builder(client, data)?;
- send_message(builder).await
+ claude_send_message(builder).await
}
async fn send_message_streaming_inner(
@@ -40,7 +40,7 @@ impl Client for ClaudeClient {
data: SendData,
) -> Result<()> {
let builder = self.request_builder(client, data)?;
- send_message_streaming(builder, handler).await
+ claude_send_message_streaming(builder, handler).await
}
}
@@ -79,7 +79,7 @@ impl ClaudeClient {
}
}
-async fn send_message(builder: RequestBuilder) -> Result<String> {
+pub async fn claude_send_message(builder: RequestBuilder) -> Result<String> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -94,7 +94,10 @@ async fn send_message(builder: RequestBuilder) -> Result<String> {
Ok(output.to_string())
}
-async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> {
+pub async fn claude_send_message_streaming(
+ builder: RequestBuilder,
+ handler: &mut ReplyHandler,
+) -> Result<()> {
let mut es = builder.eventsource()?;
while let Some(event) = es.next().await {
match event {
@@ -214,13 +217,3 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> {
}
Ok(body)
}
-
-fn catch_error(data: &Value, status: u16) -> Result<()> {
- debug!("Invalid response, status: {status}, data: {data}");
- if let Some(error) = data["error"].as_object() {
- if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
- bail!("{message} (type: {typ})");
- }
- }
- bail!("Invalid response, status: {status}, data: {data}");
-}
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 828b799..df86843 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,6 +1,6 @@
use super::{
- extract_sytem_message, json_stream, message::*, Client, CohereClient, ExtraConfig, Model,
- ModelConfig, PromptType, ReplyHandler, SendData,
+ catch_error, extract_sytem_message, json_stream, message::*, Client, CohereClient, ExtraConfig,
+ Model, ModelConfig, PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -184,16 +184,6 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-fn catch_error(data: &Value, status: u16) -> Result<()> {
- debug!("Invalid response, status: {status}, data: {data}");
-
- if let Some(message) = data["message"].as_str() {
- bail!("{message}");
- } else {
- bail!("Invalid response, status: {status}, data: {data}");
- }
-}
-
fn extract_text(data: &Value) -> Result<&str> {
match data["text"].as_str() {
Some(text) => Ok(text),
diff --git a/src/client/common.rs b/src/client/common.rs
index e003e7e..1b74de8 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -6,7 +6,7 @@ use crate::{
utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind},
};
-use anyhow::{Context, Result};
+use anyhow::{bail, Context, Result};
use async_trait::async_trait;
use futures_util::{Stream, StreamExt};
use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder};
@@ -224,6 +224,13 @@ macro_rules! list_models_fn {
};
}
+#[macro_export]
+macro_rules! unsupported_model {
+ ($name:expr) => {
+ anyhow::bail!("Unsupported model '{}'", $name)
+ };
+}
+
#[async_trait]
pub trait Client: Sync + Send {
fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>);
@@ -409,6 +416,41 @@ where
Ok(())
}
+pub fn catch_error(data: &Value, status: u16) -> Result<()> {
+ if (200..300).contains(&status) {
+ return Ok(());
+ }
+ debug!("Invalid response, status: {status}, data: {data}");
+ if let Some(error) = data["error"].as_object() {
+ if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
+ bail!("{message} (type: {typ})");
+ }
+ } else if let Some(error) = data[0]["error"].as_object() {
+ if let (Some(status), Some(message)) = (error["status"].as_str(), error["message"].as_str())
+ {
+ bail!("{message} (status: {status})")
+ }
+ } else if let Some(error) = data["error"].as_str() {
+ bail!("{error}");
+ } else if let Some(message) = data["message"].as_str() {
+ bail!("{message}");
+ }
+ bail!("Invalid response, status: {status}, data: {data}");
+}
+
+pub fn maybe_catch_error(data: &Value) -> Result<()> {
+ if let (Some(code), Some(message)) = (data["code"].as_str(), data["message"].as_str()) {
+ debug!("Invalid response: {}", data);
+ bail!("{message} (code: {code})");
+ } else if let (Some(error_code), Some(error_msg)) =
+ (data["error_code"].as_number(), data["error_msg"].as_str())
+ {
+ debug!("Invalid response: {}", data);
+ bail!("{error_msg} (error_code: {error_code})");
+ }
+ Ok(())
+}
+
pub async fn json_stream<S, F>(mut stream: S, mut handle: F) -> Result<()>
where
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 31ec536..22f9362 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,26 +1,24 @@
use super::{
- patch_system_message, Client, ErnieClient, ExtraConfig, Model, ModelConfig, PromptType,
- ReplyHandler, SendData,
+ maybe_catch_error, patch_system_message, Client, ErnieClient, ExtraConfig, Model, ModelConfig,
+ PromptType, ReplyHandler, SendData,
};
use crate::utils::PromptKind;
use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
+use chrono::Utc;
use futures_util::StreamExt;
-use lazy_static::lazy_static;
use reqwest::{Client as ReqwestClient, RequestBuilder};
use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt};
use serde::Deserialize;
use serde_json::{json, Value};
-use std::{env, sync::Mutex};
+use std::env;
const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1";
const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token";
-lazy_static! {
- static ref ACCESS_TOKEN: Mutex<Option<String>> = Mutex::new(None);
-}
+static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0);
#[derive(Debug, Clone, Deserialize, Default)]
pub struct ErnieConfig {
@@ -78,23 +76,17 @@ impl ErnieClient {
let body = build_body(data, &self.model);
let endpoint = match self.model.name.as_str() {
- "ernie-4.0-8k" => "/wenxinworkshop/chat/completions_pro",
- "ernie-3.5-8k" => "/wenxinworkshop/chat/ernie-3.5-8k-0205",
- "ernie-3.5-4k" => "/wenxinworkshop/chat/ernie-3.5-4k-0205",
- "ernie-speed-8k" => "/wenxinworkshop/chat/ernie_speed",
- "ernie-speed-128k" => "/wenxinworkshop/chat/ernie-speed-128k",
- "ernie-lite-8k" => "/wenxinworkshop/chat/ernie-lite-8k",
- "ernie-tiny-8k" => "/wenxinworkshop/chat/ernie-tiny-8k",
- _ => bail!("Miss Model '{}'", self.model.id()),
+ "ernie-4.0-8k" => "completions_pro",
+ "ernie-3.5-8k" => "ernie-3.5-8k-0205",
+ "ernie-3.5-4k" => "ernie-3.5-4k-0205",
+ "ernie-speed-8k" => "ernie_speed",
+ _ => &self.model.name,
};
- let access_token = ACCESS_TOKEN
- .lock()
- .unwrap()
- .clone()
- .ok_or_else(|| anyhow!("Failed to load access token"))?;
-
- let url = format!("{API_BASE}{endpoint}?access_token={access_token}");
+ let url = format!(
+ "{API_BASE}/wenxinworkshop/chat/{endpoint}?access_token={}",
+ unsafe { &ACCESS_TOKEN.0 }
+ );
debug!("Ernie Request: {url} {body}");
@@ -104,7 +96,7 @@ impl ErnieClient {
}
async fn prepare_access_token(&self) -> Result<()> {
- if ACCESS_TOKEN.lock().unwrap().is_none() {
+ if unsafe { ACCESS_TOKEN.0.is_empty() || Utc::now().timestamp() > ACCESS_TOKEN.1 } {
let env_prefix = Self::name(&self.config).to_uppercase();
let api_key = self.config.api_key.clone();
let api_key = api_key
@@ -120,7 +112,7 @@ impl ErnieClient {
let token = fetch_access_token(&client, &api_key, &secret_key)
.await
.with_context(|| "Failed to fetch access token")?;
- *ACCESS_TOKEN.lock().unwrap() = Some(token);
+ unsafe { ACCESS_TOKEN = (token, 86400) };
}
Ok(())
}
@@ -128,7 +120,7 @@ impl ErnieClient {
async fn send_message(builder: RequestBuilder) -> Result<String> {
let data: Value = builder.send().await?.json().await?;
- catch_error(&data)?;
+ maybe_catch_error(&data)?;
let output = data["result"]
.as_str()
@@ -156,8 +148,8 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHand
.map_err(|_| anyhow!("Invalid response header"))?;
if content_type.contains("application/json") {
let data: Value = res.json().await?;
- catch_error(&data)?;
- bail!("Request failed");
+ maybe_catch_error(&data)?;
+ bail!("Invalid response data: {data}");
} else {
let text = res.text().await?;
if let Some(text) = text.strip_prefix("data: ") {
@@ -214,20 +206,6 @@ fn build_body(data: SendData, model: &Model) -> Value {
body
}
-fn catch_error(data: &Value) -> Result<()> {
- if let (Some(error_code), Some(error_msg)) =
- (data["error_code"].as_number(), data["error_msg"].as_str())
- {
- debug!("Invalid response: {}", data);
- let error_code = error_code.as_i64().unwrap_or_default();
- if error_code == 110 {
- *ACCESS_TOKEN.lock().unwrap() = None;
- }
- bail!("{error_msg} (error_code: {error_code})");
- }
- Ok(())
-}
-
async fn fetch_access_token(
client: &reqwest::Client,
api_key: &str,
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index b930d79..05e19a3 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -1,4 +1,4 @@
-use super::vertexai::{build_body, send_message, send_message_streaming};
+use super::vertexai::{gemini_build_body, gemini_send_message, gemini_send_message_streaming};
use super::{
Client, ExtraConfig, GeminiClient, Model, ModelConfig, PromptType, ReplyHandler, SendData,
};
@@ -28,7 +28,7 @@ impl Client for GeminiClient {
async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
let builder = self.request_builder(client, data)?;
- send_message(builder).await
+ gemini_send_message(builder).await
}
async fn send_message_streaming_inner(
@@ -38,7 +38,7 @@ impl Client for GeminiClient {
data: SendData,
) -> Result<()> {
let builder = self.request_builder(client, data)?;
- send_message_streaming(builder, handler).await
+ gemini_send_message_streaming(builder, handler).await
}
}
@@ -67,7 +67,7 @@ impl GeminiClient {
let block_threshold = self.config.block_threshold.clone();
- let body = build_body(data, &self.model, block_threshold)?;
+ let body = gemini_build_body(data, &self.model, block_threshold)?;
let model = &self.model.name;
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 1394ced..f2340c6 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -1,6 +1,6 @@
use super::{
- message::*, Client, ExtraConfig, Model, ModelConfig, OllamaClient, PromptType, ReplyHandler,
- SendData,
+ catch_error, message::*, Client, ExtraConfig, Model, ModelConfig, OllamaClient, PromptType,
+ ReplyHandler, SendData,
};
use crate::utils::PromptKind;
@@ -191,11 +191,3 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-
-fn catch_error(data: &Value, status: u16) -> Result<()> {
- debug!("Invalid response, status: {status}, data: {data}");
- if let Some(error) = data["error"].as_str() {
- bail!("{error}");
- }
- bail!("Invalid response, status: {status}, data: {data}");
-}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 37da878..e5aae24 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,4 +1,6 @@
-use super::{ExtraConfig, Model, ModelConfig, OpenAIClient, PromptType, ReplyHandler, SendData};
+use super::{
+ catch_error, ExtraConfig, Model, ModelConfig, OpenAIClient, PromptType, ReplyHandler, SendData,
+};
use crate::utils::PromptKind;
@@ -154,15 +156,3 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
}
body
}
-
-fn catch_error(data: &Value, status: u16) -> Result<()> {
- debug!("Invalid response, status: {status}, data: {data}");
- if let Some(error) = data["error"].as_object() {
- if let (Some(type_), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
- bail!("{message} (type: {type_})");
- }
- } else if let Some(message) = data["message"].as_str() {
- bail!("{message}");
- }
- bail!("Invalid response, status: {status}, data: {data}");
-}
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 52548d8..76cc712 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,6 +1,6 @@
use super::{
- message::*, Client, ExtraConfig, Model, ModelConfig, PromptType, QianwenClient, ReplyHandler,
- SendData,
+ maybe_catch_error, message::*, Client, ExtraConfig, Model, ModelConfig, PromptType,
+ QianwenClient, ReplyHandler, SendData,
};
use crate::utils::{sha256sum, PromptKind};
@@ -112,7 +112,7 @@ impl QianwenClient {
async fn send_message(builder: RequestBuilder, is_vl: bool) -> Result<String> {
let data: Value = builder.send().await?.json().await?;
- catch_error(&data)?;
+ maybe_catch_error(&data)?;
let output = if is_vl {
data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
@@ -137,7 +137,7 @@ async fn send_message_streaming(
Ok(Event::Open) => {}
Ok(Event::Message(message)) => {
let data: Value = serde_json::from_str(&message.data)?;
- catch_error(&data)?;
+ maybe_catch_error(&data)?;
if is_vl {
if let Some(text) =
data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
@@ -231,14 +231,6 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
Ok((body, has_upload))
}
-fn catch_error(data: &Value) -> Result<()> {
- if let (Some(code), Some(message)) = (data["code"].as_str(), data["message"].as_str()) {
- debug!("Invalid response: {}", data);
- bail!("{message} (code: {code})");
- }
- Ok(())
-}
-
/// Patch messsages, upload embedded images to oss
async fn patch_messages(model: &str, api_key: &str, messages: &mut Vec<Message>) -> Result<()> {
for message in messages {
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 30aba75..47e1a38 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,6 +1,6 @@
use super::{
- json_stream, message::*, patch_system_message, Client, ExtraConfig, Model, ModelConfig,
- PromptType, ReplyHandler, SendData, VertexAIClient,
+ catch_error, json_stream, message::*, patch_system_message, Client, ExtraConfig, Model,
+ ModelConfig, PromptType, ReplyHandler, SendData, VertexAIClient,
};
use crate::utils::PromptKind;
@@ -33,7 +33,7 @@ impl Client for VertexAIClient {
async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> {
self.prepare_access_token().await?;
let builder = self.request_builder(client, data)?;
- send_message(builder).await
+ gemini_send_message(builder).await
}
async fn send_message_streaming_inner(
@@ -44,7 +44,7 @@ impl Client for VertexAIClient {
) -> Result<()> {
self.prepare_access_token().await?;
let builder = self.request_builder(client, data)?;
- send_message_streaming(builder, handler).await
+ gemini_send_message_streaming(builder, handler).await
}
}
@@ -70,14 +70,10 @@ impl VertexAIClient {
true => "streamGenerateContent",
false => "generateContent",
};
+ let url = format!("{api_base}/{}:{}", &self.model.name, func);
let block_threshold = self.config.block_threshold.clone();
-
- let body = build_body(data, &self.model, block_threshold)?;
-
- let model = &self.model.name;
-
- let url = format!("{api_base}/{}:{}", model, func);
+ let body = gemini_build_body(data, &self.model, block_threshold)?;
debug!("VertexAI Request: {url} {body}");
@@ -104,18 +100,18 @@ impl VertexAIClient {
}
}
-pub(crate) async fn send_message(builder: RequestBuilder) -> Result<String> {
+pub async fn gemini_send_message(builder: RequestBuilder) -> Result<String> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
if status != 200 {
catch_error(&data, status.as_u16())?;
}
- let output = extract_text(&data)?;
+ let output = gemini_extract_text(&data)?;
Ok(output.to_string())
}
-pub(crate) async fn send_message_streaming(
+pub async fn gemini_send_message_streaming(
builder: RequestBuilder,
handler: &mut ReplyHandler,
) -> Result<()> {
@@ -127,7 +123,7 @@ pub(crate) async fn send_message_streaming(
} else {
let handle = |value: &str| -> Result<()> {
let value: Value = serde_json::from_str(value)?;
- handler.text(extract_text(&value)?)?;
+ handler.text(gemini_extract_text(&value)?)?;
Ok(())
};
json_stream(res.bytes_stream(), handle).await?;
@@ -135,7 +131,7 @@ pub(crate) async fn send_message_streaming(
Ok(())
}
-fn extract_text(data: &Value) -> Result<&str> {
+fn gemini_extract_text(data: &Value) -> Result<&str> {
match data["candidates"][0]["content"]["parts"][0]["text"].as_str() {
Some(text) => Ok(text),
None => {
@@ -151,7 +147,7 @@ fn extract_text(data: &Value) -> Result<&str> {
}
}
-pub(crate) fn build_body(
+pub(crate) fn gemini_build_body(
data: SendData,
model: &Model,
block_threshold: Option<String>,
@@ -230,24 +226,6 @@ pub(crate) fn build_body(
Ok(body)
}
-fn catch_error(data: &Value, status: u16) -> Result<()> {
- debug!("Invalid response, status: {status}, data: {data}");
-
- if let Some((Some(status), Some(message))) = data[0]["error"].as_object().map(|v| {
- (
- v.get("status").and_then(|v| v.as_str()),
- v.get("message").and_then(|v| v.as_str()),
- )
- }) {
- if status == "UNAUTHENTICATED" {
- unsafe { ACCESS_TOKEN = (String::new(), 0) }
- }
- bail!("{message} (status: {status})")
- } else {
- bail!("Invalid response, status: {status}, data: {data}",);
- }
-}
-
async fn fetch_access_token(
client: &reqwest::Client,
file: &Option<String>,