diff options
Diffstat (limited to 'src/client/vertexai.rs')
| -rw-r--r-- | src/client/vertexai.rs | 41 |
1 files changed, 22 insertions, 19 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 66bb098..d9c1019 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -108,7 +108,7 @@ pub(crate) async fn send_message(builder: RequestBuilder) -> Result<String> { let status = res.status(); let data: Value = res.json().await?; if status != 200 { - check_error(&data)?; + catch_error(&data, status.as_u16())?; } let output = extract_text(&data)?; Ok(output.to_string()) @@ -119,9 +119,10 @@ pub(crate) async fn send_message_streaming( handler: &mut ReplyHandler, ) -> Result<()> { let res = builder.send().await?; - if res.status() != 200 { + let status = res.status(); + if status != 200 { let data: Value = res.json().await?; - check_error(&data)?; + catch_error(&data, status.as_u16())?; } else { let handle = |value: &str| -> Result<()> { let value: Value = serde_json::from_str(value)?; @@ -149,22 +150,6 @@ fn extract_text(data: &Value) -> Result<&str> { } } -fn check_error(data: &Value) -> Result<()> { - 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!("{status}: {message}") - } else { - bail!("Error {}", data); - } -} - pub(crate) fn build_body( data: SendData, model: &Model, @@ -241,6 +226,24 @@ 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>, |
