summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs41
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>,