From a193710a7f88de12ef11c5a5c66a5591150fe247 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 25 Apr 2024 14:03:16 +0800 Subject: refactor: extract common catch_error (#437) --- src/client/claude.rs | 25 ++++++++------------- src/client/cohere.rs | 14 ++---------- src/client/common.rs | 44 +++++++++++++++++++++++++++++++++++- src/client/ernie.rs | 60 ++++++++++++++++---------------------------------- src/client/gemini.rs | 8 +++---- src/client/ollama.rs | 12 ++-------- src/client/openai.rs | 16 +++----------- src/client/qianwen.rs | 16 ++++---------- src/client/vertexai.rs | 46 ++++++++++---------------------------- 9 files changed, 98 insertions(+), 143 deletions(-) (limited to 'src/client') 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 { 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 { +pub async fn claude_send_message(builder: RequestBuilder) -> Result { 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 { 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 { } 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 { 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); @@ -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(mut stream: S, mut handle: F) -> Result<()> where S: Stream> + 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> = 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 { 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 { 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 { 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 { 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) -> 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 { 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 { +pub async fn gemini_send_message(builder: RequestBuilder) -> Result { 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, @@ -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, -- cgit v1.2.3