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/qianwen.rs | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) (limited to 'src/client/qianwen.rs') 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 { -- cgit v1.2.3