summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-06 09:08:49 +0800
committerGitHub <noreply@github.com>2024-05-06 09:08:49 +0800
commit7c6f75a139128e77f6c0e16ee176192dfef75599 (patch)
tree4bddb77185d420eee329ebac3f5a3e34de5cf512 /src/client/ernie.rs
parent9b283024b47f57ad7bbf032fa015eab42f8162a9 (diff)
downloadaichat-7c6f75a139128e77f6c0e16ee176192dfef75599.tar.gz
refactor: unified access token management (#486)
Diffstat (limited to 'src/client/ernie.rs')
-rw-r--r--src/client/ernie.rs13
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(())
}