use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use wecomx_transport::{
Endpoint, HttpTransportBackend, RequestOptions, Transport, TransportBackend,
};
use crate::bootstrap::{BindSource, fetch_auth};
use crate::bot::BotCredential;
use crate::credentials::{CredentialStore, Credentials};
use crate::error::AuthError;
use crate::gateway;
pub type TokenFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, AuthError>> + Send + 'a>>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AccessToken(String);
impl AccessToken {
pub fn new(token: impl Into<String>) -> Self {
Self(token.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn into_inner(self) -> String {
self.0
}
}
impl std::fmt::Display for AccessToken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl serde::Serialize for AccessToken {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str("***")
}
}
pub trait TokenProvider: Send + Sync {
fn access_token<'a>(&'a self) -> TokenFuture<'a, Option<AccessToken>>;
fn refresh<'a>(
&'a self,
stale_token: Option<&'a str>,
options: Option<RequestOptions>,
) -> TokenFuture<'a, AccessToken>;
}
#[derive(Clone)]
pub struct BotGatewayTokenProvider {
store: Arc<dyn CredentialStore>,
auth_endpoint: Endpoint,
backend: Arc<dyn TransportBackend>,
bind_source: BindSource,
bot: Option<BotCredential>,
}
impl std::fmt::Debug for BotGatewayTokenProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BotGatewayTokenProvider")
.field("bind_source", &self.bind_source)
.finish_non_exhaustive()
}
}
impl BotGatewayTokenProvider {
pub fn new(
bot_id: impl Into<String>,
secret: impl Into<String>,
store: Arc<dyn CredentialStore>,
) -> Self {
Self::from_store(store).with_bot(BotCredential::new(bot_id.into(), secret.into()))
}
pub fn from_store(store: Arc<dyn CredentialStore>) -> Self {
Self {
store,
auth_endpoint: gateway::auth_endpoint(gateway::DEFAULT_AUTH_ENDPOINT),
backend: Arc::new(HttpTransportBackend::default()),
bind_source: BindSource::Interactive,
bot: None,
}
}
#[must_use]
pub fn with_auth_endpoint(mut self, endpoint: Endpoint) -> Self {
self.auth_endpoint = endpoint;
self
}
#[must_use]
pub fn with_backend(mut self, backend: Arc<dyn TransportBackend>) -> Self {
self.backend = backend;
self
}
#[must_use]
pub fn with_bot(mut self, bot: BotCredential) -> Self {
self.bot = Some(bot);
self
}
#[must_use]
pub fn with_bind_source(mut self, bind_source: BindSource) -> Self {
self.bind_source = bind_source;
self
}
fn resolve_bot(&self, creds: Option<&Credentials>) -> Option<BotCredential> {
self.bot
.clone()
.or_else(|| creds.and_then(|c| c.bot.clone()))
}
}
impl TokenProvider for BotGatewayTokenProvider {
fn access_token<'a>(&'a self) -> TokenFuture<'a, Option<AccessToken>> {
Box::pin(async move {
Ok(self
.store
.load()?
.and_then(|c| c.token)
.filter(|t| !t.is_empty())
.map(AccessToken::new))
})
}
fn refresh<'a>(
&'a self,
stale_token: Option<&'a str>,
options: Option<RequestOptions>,
) -> TokenFuture<'a, AccessToken> {
Box::pin(async move {
let creds = self.store.load()?;
if let Some(stored) = creds.as_ref().and_then(|c| c.token.as_deref())
&& Some(stored) != stale_token
{
tracing::debug!("token already refreshed by a concurrent request, reusing it");
return Ok(AccessToken::new(stored));
}
let bot = self.resolve_bot(creds.as_ref()).ok_or_else(|| {
AuthError::MissingCredentials("无 bot 凭据,无法静默刷新 token".into())
})?;
let mut options = options.unwrap_or_default();
options.headers_mut().remove(reqwest::header::AUTHORIZATION);
let transport = Transport::new(self.backend.clone(), options);
let resp = fetch_auth(&transport, &bot, self.bind_source, &self.auth_endpoint).await?;
let token = match resp.token.clone().filter(|t| !t.is_empty()) {
Some(token) => token,
None => {
return Err(AuthError::from(wecomx_transport::Error::Parse {
message: "token 刷新响应缺少访问令牌".to_string(),
endpoint: wecomx_transport::EndpointHttpExt::full_url(&self.auth_endpoint),
body: Box::new(serde_json::to_value(&resp).unwrap_or_default()),
source: None,
}));
}
};
let mut creds = creds.unwrap_or_default();
creds.bot = creds.bot.or(Some(bot));
creds.token = Some(token.clone());
self.store.save(&creds)?;
tracing::info!("access token refreshed (853004) and persisted");
Ok(AccessToken::new(token))
})
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::bot::BotCredential;
use crate::credentials::MemoryCredentialStore;
use crate::gateway::auth_endpoint;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[tokio::test]
async fn access_token_reads_store() {
let store = MemoryCredentialStore::new(Credentials {
bot: Some(BotCredential::new("b".into(), "s".into())),
token: Some("tok-1".into()),
})
.shared();
let provider = BotGatewayTokenProvider::from_store(store);
let token = provider.access_token().await.unwrap();
assert_eq!(token.as_ref().map(AccessToken::as_str), Some("tok-1"));
let empty = BotGatewayTokenProvider::from_store(MemoryCredentialStore::default().shared());
assert!(empty.access_token().await.unwrap().is_none());
}
#[tokio::test]
async fn refresh_fetches_and_persists() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/get_cli_config"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"errcode": 0, "errmsg": "ok", "token": "tok-new",
})))
.expect(1)
.mount(&server)
.await;
let store = MemoryCredentialStore::new(Credentials {
bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
token: None,
})
.shared();
let provider = BotGatewayTokenProvider::from_store(store.clone())
.with_auth_endpoint(auth_endpoint(&format!("{}/get_cli_config", server.uri())));
let token = provider.refresh(Some("tok-old"), None).await.unwrap();
assert_eq!(token.as_str(), "tok-new");
assert_eq!(
store.load().unwrap().unwrap().token.as_deref(),
Some("tok-new")
);
server.verify().await;
}
#[tokio::test]
async fn refresh_reuses_concurrently_refreshed_token() {
let store = MemoryCredentialStore::new(Credentials {
bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
token: Some("tok-fresh".into()),
})
.shared();
let provider = BotGatewayTokenProvider::from_store(store)
.with_auth_endpoint(auth_endpoint("http://localhost/get_cli_config"));
let token = provider.refresh(Some("tok-stale"), None).await.unwrap();
assert_eq!(token.as_str(), "tok-fresh");
}
#[tokio::test]
async fn refresh_without_bot_fails() {
let store = MemoryCredentialStore::new(Credentials {
bot: None,
token: None,
})
.shared();
let provider = BotGatewayTokenProvider::from_store(store);
let err = provider.refresh(None, None).await.unwrap_err();
assert!(matches!(err, AuthError::MissingCredentials(_)));
}
#[tokio::test]
async fn refresh_strips_authorization_header() {
struct NoAuthorization;
impl wiremock::Match for NoAuthorization {
fn matches(&self, request: &wiremock::Request) -> bool {
request.headers.get("authorization").is_none()
}
}
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/get_cli_config"))
.and(NoAuthorization)
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"errcode": 0, "token": "tok-2"})),
)
.expect(1)
.mount(&server)
.await;
let store = MemoryCredentialStore::new(Credentials {
bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
token: None,
})
.shared();
let provider = BotGatewayTokenProvider::from_store(store)
.with_auth_endpoint(auth_endpoint(&format!("{}/get_cli_config", server.uri())));
let mut options = RequestOptions::default();
options.headers_mut().insert(
reqwest::header::AUTHORIZATION,
reqwest::header::HeaderValue::from_static("Bearer stale"),
);
let token = provider
.refresh(Some("stale"), Some(options))
.await
.unwrap();
assert_eq!(token.as_str(), "tok-2");
server.verify().await;
}
#[test]
fn debug_does_not_leak_secret() {
let provider = BotGatewayTokenProvider::new(
"bot1",
"super-secret",
MemoryCredentialStore::default().shared(),
);
let dbg = format!("{provider:?}");
assert!(!dbg.contains("super-secret"), "secret 泄露: {dbg}");
}
}