diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-06 09:08:49 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-06 09:08:49 +0800 |
| commit | 7c6f75a139128e77f6c0e16ee176192dfef75599 (patch) | |
| tree | 4bddb77185d420eee329ebac3f5a3e34de5cf512 /src/client/vertexai.rs | |
| parent | 9b283024b47f57ad7bbf032fa015eab42f8162a9 (diff) | |
| download | aichat-7c6f75a139128e77f6c0e16ee176192dfef75599.tar.gz | |
refactor: unified access token management (#486)
Diffstat (limited to 'src/client/vertexai.rs')
| -rw-r--r-- | src/client/vertexai.rs | 27 |
1 files changed, 14 insertions, 13 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index db04f38..9907a7f 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,3 +1,4 @@ +use super::access_token::*; use super::{ catch_error, json_stream, message::*, patch_system_message, Client, CompletionDetails, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SseHandler, @@ -12,8 +13,6 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::path::PathBuf; -static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0); // safe under linear operation - #[derive(Debug, Clone, Deserialize, Default)] pub struct VertexAIConfig { pub name: Option<String>, @@ -39,6 +38,7 @@ impl VertexAIClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { let project_id = self.get_project_id()?; let location = self.get_location()?; + let access_token = get_access_token(self.name())?; let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers"); @@ -52,10 +52,7 @@ impl VertexAIClient { debug!("VertexAI Request: {url} {body}"); - let builder = client - .post(url) - .bearer_auth(unsafe { &ACCESS_TOKEN.0 }) - .json(&body); + let builder = client.post(url).bearer_auth(access_token).json(&body); Ok(builder) } @@ -70,7 +67,7 @@ impl Client for VertexAIClient { client: &ReqwestClient, data: SendData, ) -> Result<(String, CompletionDetails)> { - prepare_access_token(client, &self.config.adc_file).await?; + prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?; let builder = self.request_builder(client, data)?; gemini_send_message(builder).await } @@ -81,7 +78,7 @@ impl Client for VertexAIClient { handler: &mut SseHandler, data: SendData, ) -> Result<()> { - prepare_access_token(client, &self.config.adc_file).await?; + prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?; let builder = self.request_builder(client, data)?; gemini_send_message_streaming(builder, handler).await } @@ -217,20 +214,24 @@ pub(crate) fn gemini_build_body( Ok(body) } -async fn prepare_access_token(client: &reqwest::Client, adc_file: &Option<String>) -> Result<()> { - if unsafe { ACCESS_TOKEN.0.is_empty() || Utc::now().timestamp() > ACCESS_TOKEN.1 } { - let (token, expires_in) = fetch_gcloud_access_token(client, adc_file) +pub async fn prepare_gcloud_access_token( + client: &reqwest::Client, + client_name: &str, + adc_file: &Option<String>, +) -> Result<()> { + if !is_valid_access_token(client_name) { + let (token, expires_in) = fetch_access_token(client, adc_file) .await .with_context(|| "Failed to fetch access token")?; let expires_at = Utc::now() + Duration::try_seconds(expires_in) .ok_or_else(|| anyhow!("Failed to parse expires_in of access_token"))?; - unsafe { ACCESS_TOKEN = (token, expires_at.timestamp()) }; + set_access_token(client_name, token, expires_at.timestamp()) } Ok(()) } -pub async fn fetch_gcloud_access_token( +async fn fetch_access_token( client: &reqwest::Client, file: &Option<String>, ) -> Result<(String, i64)> { |
