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/ernie.rs | |
| parent | 9b283024b47f57ad7bbf032fa015eab42f8162a9 (diff) | |
| download | aichat-7c6f75a139128e77f6c0e16ee176192dfef75599.tar.gz | |
refactor: unified access token management (#486)
Diffstat (limited to 'src/client/ernie.rs')
| -rw-r--r-- | src/client/ernie.rs | 13 |
1 files changed, 6 insertions, 7 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 982edae..d3002ff 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,3 +1,4 @@ +use super::access_token::*; use super::{ maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SsMmessage, SseHandler, @@ -5,7 +6,6 @@ use super::{ use anyhow::{anyhow, Context, Result}; use async_trait::async_trait; -use chrono::Utc; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -14,8 +14,6 @@ use std::env; const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1"; const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token"; -static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0); - #[derive(Debug, Clone, Deserialize, Default)] pub struct ErnieConfig { pub name: Option<String>, @@ -34,11 +32,11 @@ impl ErnieClient { fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { let body = build_body(data, &self.model); + let access_token = get_access_token(self.name())?; let url = format!( - "{API_BASE}/wenxinworkshop/chat/{}?access_token={}", + "{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}", &self.model.name, - unsafe { &ACCESS_TOKEN.0 } ); debug!("Ernie Request: {url} {body}"); @@ -49,7 +47,8 @@ impl ErnieClient { } async fn prepare_access_token(&self) -> Result<()> { - if unsafe { ACCESS_TOKEN.0.is_empty() || Utc::now().timestamp() > ACCESS_TOKEN.1 } { + let client_name = self.name(); + if !is_valid_access_token(client_name) { let env_prefix = Self::name(&self.config).to_uppercase(); let api_key = self.config.api_key.clone(); let api_key = api_key @@ -65,7 +64,7 @@ impl ErnieClient { let token = fetch_access_token(&client, &api_key, &secret_key) .await .with_context(|| "Failed to fetch access token")?; - unsafe { ACCESS_TOKEN = (token, 86400) }; + set_access_token(client_name, token, 86400); } Ok(()) } |
