1use std::future::Future;
9use std::pin::Pin;
10use std::sync::Arc;
11
12use wecomx_transport::{
13 Endpoint, HttpTransportBackend, RequestOptions, Transport, TransportBackend,
14};
15
16use crate::bootstrap::{BindSource, fetch_auth};
17use crate::bot::BotCredential;
18use crate::credentials::{CredentialStore, Credentials};
19use crate::error::AuthError;
20use crate::gateway;
21
22pub type TokenFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, AuthError>> + Send + 'a>>;
24
25#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct AccessToken(String);
28
29impl AccessToken {
30 pub fn new(token: impl Into<String>) -> Self {
31 Self(token.into())
32 }
33
34 pub fn as_str(&self) -> &str {
35 &self.0
36 }
37
38 pub fn into_inner(self) -> String {
39 self.0
40 }
41}
42
43impl std::fmt::Display for AccessToken {
44 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45 f.write_str(&self.0)
46 }
47}
48
49impl serde::Serialize for AccessToken {
51 fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
52 serializer.serialize_str("***")
53 }
54}
55
56pub trait TokenProvider: Send + Sync {
62 fn access_token<'a>(&'a self) -> TokenFuture<'a, Option<AccessToken>>;
64
65 fn refresh<'a>(
75 &'a self,
76 stale_token: Option<&'a str>,
77 options: Option<RequestOptions>,
78 ) -> TokenFuture<'a, AccessToken>;
79}
80
81#[derive(Clone)]
90pub struct BotGatewayTokenProvider {
91 store: Arc<dyn CredentialStore>,
92 auth_endpoint: Endpoint,
93 backend: Arc<dyn TransportBackend>,
96 bind_source: BindSource,
97 bot: Option<BotCredential>,
99}
100
101impl std::fmt::Debug for BotGatewayTokenProvider {
102 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
103 f.debug_struct("BotGatewayTokenProvider")
105 .field("bind_source", &self.bind_source)
106 .finish_non_exhaustive()
107 }
108}
109
110impl BotGatewayTokenProvider {
111 pub fn new(
113 bot_id: impl Into<String>,
114 secret: impl Into<String>,
115 store: Arc<dyn CredentialStore>,
116 ) -> Self {
117 Self::from_store(store).with_bot(BotCredential::new(bot_id.into(), secret.into()))
118 }
119
120 pub fn from_store(store: Arc<dyn CredentialStore>) -> Self {
123 Self {
124 store,
125 auth_endpoint: gateway::auth_endpoint(gateway::DEFAULT_AUTH_ENDPOINT),
126 backend: Arc::new(HttpTransportBackend::default()),
127 bind_source: BindSource::Interactive,
128 bot: None,
129 }
130 }
131
132 #[must_use]
134 pub fn with_auth_endpoint(mut self, endpoint: Endpoint) -> Self {
135 self.auth_endpoint = endpoint;
136 self
137 }
138
139 #[must_use]
141 pub fn with_backend(mut self, backend: Arc<dyn TransportBackend>) -> Self {
142 self.backend = backend;
143 self
144 }
145
146 #[must_use]
148 pub fn with_bot(mut self, bot: BotCredential) -> Self {
149 self.bot = Some(bot);
150 self
151 }
152
153 #[must_use]
155 pub fn with_bind_source(mut self, bind_source: BindSource) -> Self {
156 self.bind_source = bind_source;
157 self
158 }
159
160 fn resolve_bot(&self, creds: Option<&Credentials>) -> Option<BotCredential> {
162 self.bot
163 .clone()
164 .or_else(|| creds.and_then(|c| c.bot.clone()))
165 }
166}
167
168impl TokenProvider for BotGatewayTokenProvider {
169 fn access_token<'a>(&'a self) -> TokenFuture<'a, Option<AccessToken>> {
170 Box::pin(async move {
171 Ok(self
172 .store
173 .load()?
174 .and_then(|c| c.token)
175 .filter(|t| !t.is_empty())
176 .map(AccessToken::new))
177 })
178 }
179
180 fn refresh<'a>(
181 &'a self,
182 stale_token: Option<&'a str>,
183 options: Option<RequestOptions>,
184 ) -> TokenFuture<'a, AccessToken> {
185 Box::pin(async move {
186 let creds = self.store.load()?;
187
188 if let Some(stored) = creds.as_ref().and_then(|c| c.token.as_deref())
190 && Some(stored) != stale_token
191 {
192 tracing::debug!("token already refreshed by a concurrent request, reusing it");
193 return Ok(AccessToken::new(stored));
194 }
195
196 let bot = self.resolve_bot(creds.as_ref()).ok_or_else(|| {
197 AuthError::MissingCredentials("无 bot 凭据,无法静默刷新 token".into())
198 })?;
199
200 let mut options = options.unwrap_or_default();
204 options.headers_mut().remove(reqwest::header::AUTHORIZATION);
205 let transport = Transport::new(self.backend.clone(), options);
206
207 let resp = fetch_auth(&transport, &bot, self.bind_source, &self.auth_endpoint).await?;
208 let token = match resp.token.clone().filter(|t| !t.is_empty()) {
209 Some(token) => token,
210 None => {
211 return Err(AuthError::from(wecomx_transport::Error::Parse {
212 message: "token 刷新响应缺少访问令牌".to_string(),
213 endpoint: wecomx_transport::EndpointHttpExt::full_url(&self.auth_endpoint),
214 body: Box::new(serde_json::to_value(&resp).unwrap_or_default()),
215 source: None,
216 }));
217 }
218 };
219
220 let mut creds = creds.unwrap_or_default();
222 creds.bot = creds.bot.or(Some(bot));
223 creds.token = Some(token.clone());
224 self.store.save(&creds)?;
225 tracing::info!("access token refreshed (853004) and persisted");
226
227 Ok(AccessToken::new(token))
228 })
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use serde_json::json;
247
248 use super::*;
249 use crate::bot::BotCredential;
250 use crate::credentials::MemoryCredentialStore;
251 use crate::gateway::auth_endpoint;
252 use wiremock::matchers::{method, path};
253 use wiremock::{Mock, MockServer, ResponseTemplate};
254
255 #[tokio::test]
257 async fn access_token_reads_store() {
258 let store = MemoryCredentialStore::new(Credentials {
259 bot: Some(BotCredential::new("b".into(), "s".into())),
260 token: Some("tok-1".into()),
261 })
262 .shared();
263 let provider = BotGatewayTokenProvider::from_store(store);
264 let token = provider.access_token().await.unwrap();
265 assert_eq!(token.as_ref().map(AccessToken::as_str), Some("tok-1"));
266
267 let empty = BotGatewayTokenProvider::from_store(MemoryCredentialStore::default().shared());
268 assert!(empty.access_token().await.unwrap().is_none());
269 }
270
271 #[tokio::test]
273 async fn refresh_fetches_and_persists() {
274 let server = MockServer::start().await;
275 Mock::given(method("POST"))
276 .and(path("/get_cli_config"))
277 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
278 "errcode": 0, "errmsg": "ok", "token": "tok-new",
279 })))
280 .expect(1)
281 .mount(&server)
282 .await;
283
284 let store = MemoryCredentialStore::new(Credentials {
285 bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
286 token: None,
287 })
288 .shared();
289 let provider = BotGatewayTokenProvider::from_store(store.clone())
290 .with_auth_endpoint(auth_endpoint(&format!("{}/get_cli_config", server.uri())));
291
292 let token = provider.refresh(Some("tok-old"), None).await.unwrap();
293 assert_eq!(token.as_str(), "tok-new");
294 assert_eq!(
295 store.load().unwrap().unwrap().token.as_deref(),
296 Some("tok-new")
297 );
298 server.verify().await;
299 }
300
301 #[tokio::test]
303 async fn refresh_reuses_concurrently_refreshed_token() {
304 let store = MemoryCredentialStore::new(Credentials {
306 bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
307 token: Some("tok-fresh".into()),
308 })
309 .shared();
310 let provider = BotGatewayTokenProvider::from_store(store)
311 .with_auth_endpoint(auth_endpoint("http://localhost/get_cli_config"));
312
313 let token = provider.refresh(Some("tok-stale"), None).await.unwrap();
314 assert_eq!(token.as_str(), "tok-fresh");
315 }
316
317 #[tokio::test]
319 async fn refresh_without_bot_fails() {
320 let store = MemoryCredentialStore::new(Credentials {
321 bot: None,
322 token: None,
323 })
324 .shared();
325 let provider = BotGatewayTokenProvider::from_store(store);
326 let err = provider.refresh(None, None).await.unwrap_err();
327 assert!(matches!(err, AuthError::MissingCredentials(_)));
328 }
329
330 #[tokio::test]
332 async fn refresh_strips_authorization_header() {
333 struct NoAuthorization;
334 impl wiremock::Match for NoAuthorization {
335 fn matches(&self, request: &wiremock::Request) -> bool {
336 request.headers.get("authorization").is_none()
337 }
338 }
339
340 let server = MockServer::start().await;
341 Mock::given(method("POST"))
342 .and(path("/get_cli_config"))
343 .and(NoAuthorization)
344 .respond_with(
345 ResponseTemplate::new(200).set_body_json(json!({"errcode": 0, "token": "tok-2"})),
346 )
347 .expect(1)
348 .mount(&server)
349 .await;
350
351 let store = MemoryCredentialStore::new(Credentials {
352 bot: Some(BotCredential::new("bot1".into(), "secret1".into())),
353 token: None,
354 })
355 .shared();
356 let provider = BotGatewayTokenProvider::from_store(store)
357 .with_auth_endpoint(auth_endpoint(&format!("{}/get_cli_config", server.uri())));
358
359 let mut options = RequestOptions::default();
360 options.headers_mut().insert(
361 reqwest::header::AUTHORIZATION,
362 reqwest::header::HeaderValue::from_static("Bearer stale"),
363 );
364 let token = provider
365 .refresh(Some("stale"), Some(options))
366 .await
367 .unwrap();
368 assert_eq!(token.as_str(), "tok-2");
369 server.verify().await;
370 }
371
372 #[test]
374 fn debug_does_not_leak_secret() {
375 let provider = BotGatewayTokenProvider::new(
376 "bot1",
377 "super-secret",
378 MemoryCredentialStore::default().shared(),
379 );
380 let dbg = format!("{provider:?}");
381 assert!(!dbg.contains("super-secret"), "secret 泄露: {dbg}");
382 }
383}