summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.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/vertexai.rs
parent4db9b309803796bc5f996d0b3713344eb44207ec (diff)
downloadaichat-a193710a7f88de12ef11c5a5c66a5591150fe247.tar.gz
refactor: extract common catch_error (#437)
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs46
1 files changed, 12 insertions, 34 deletions
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<String> {
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<String> {
+pub async fn gemini_send_message(builder: RequestBuilder) -> Result<String> {
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<String>,
@@ -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<String>,