summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-25 14:03:16 +0800
committerGitHub <noreply@github.com>2024-04-25 14:03:16 +0800
commita193710a7f88de12ef11c5a5c66a5591150fe247 (patch)
tree937a64da3dad52706d155ac9191e324fd8648ccc /src/client/common.rs
parent4db9b309803796bc5f996d0b3713344eb44207ec (diff)
downloadaichat-a193710a7f88de12ef11c5a5c66a5591150fe247.tar.gz
refactor: extract common catch_error (#437)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs44
1 files changed, 43 insertions, 1 deletions
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,