use std::borrow::Cow;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, RwLock};
use wecomx_auth::{AuthError, RequireAuth, SuppressAuth, TokenProvider};
use wecomx_transport::{
Endpoint, HttpRequestPayload, RequestOptions, TransportBackend, TransportResponse,
};
pub const TOKEN_EXPIRED_ERRCODE: i64 = 853004;
#[derive(Clone)]
pub struct WecomBackend {
inner: Arc<dyn TransportBackend>,
provider: Option<Arc<dyn TokenProvider>>,
token: Arc<RwLock<Option<String>>>,
refresh_lock: Arc<tokio::sync::Mutex<()>>,
bin_name: Arc<str>,
}
impl std::fmt::Debug for WecomBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WecomBackend")
.field("backend", &self.inner.name())
.finish_non_exhaustive()
}
}
impl WecomBackend {
pub fn new(
inner: Arc<dyn TransportBackend>,
provider: Option<Arc<dyn TokenProvider>>,
token: Option<String>,
) -> Self {
Self {
inner,
provider,
token: Arc::new(RwLock::new(token)),
refresh_lock: Arc::new(tokio::sync::Mutex::new(())),
bin_name: Arc::from(crate::DEFAULT_BIN_NAME),
}
}
#[must_use]
pub fn with_bin_name(mut self, bin_name: impl Into<String>) -> Self {
self.bin_name = Arc::from(bin_name.into().as_str());
self
}
pub fn cached_token(&self) -> Option<String> {
self.token.read().unwrap_or_else(|e| e.into_inner()).clone()
}
fn store_token(&self, token: &str) {
*self.token.write().unwrap_or_else(|e| e.into_inner()) = Some(token.to_owned());
}
async fn refresh_token(
&self,
stale_token: Option<&str>,
options: RequestOptions,
) -> Result<String, AuthError> {
let Some(provider) = &self.provider else {
return Err(AuthError::MissingCredentials(format!(
"无 bot 凭据,无法静默刷新 token,请重新运行 `{}` auth init",
self.bin_name
)));
};
let _guard = self.refresh_lock.lock().await;
let token = provider.refresh(stale_token, Some(options)).await?;
self.store_token(token.as_str());
tracing::info!("access token refreshed (853004) and cached");
Ok(token.into_inner())
}
}
impl TransportBackend for WecomBackend {
fn execute<'a>(
&'a self,
endpoint: Cow<'a, Endpoint>,
payload: HttpRequestPayload,
options: RequestOptions,
) -> Pin<
Box<
dyn Future<Output = std::result::Result<TransportResponse, wecomx_transport::Error>>
+ Send
+ 'a,
>,
> {
Box::pin(async move {
let replay_payload = payload.clone();
let mut options = options;
let sent_token = if endpoint.as_ref().get::<SuppressAuth>().is_some() {
None
} else {
let token = self.cached_token();
if endpoint.as_ref().get::<RequireAuth>().is_some() && token.is_none() {
tracing::debug!("endpoint requires auth but no token available");
return Err(AuthError::MissingCredentials(format!(
"该请求需要授权,请先运行 `{}` auth init 登录",
self.bin_name
))
.into());
}
token
.clone()
.inspect(|token| set_bearer_token(&mut options, token))
};
let err = match self
.inner
.execute(endpoint.clone(), payload, options.clone())
.await
{
Ok(resp) => return Ok(resp),
Err(err) => err,
};
if !is_token_expired(&err) {
return Err(err);
}
if sent_token.is_none() {
tracing::warn!("token expired but no token was sent");
return Err(err);
}
if self.provider.is_none() {
tracing::warn!("missing token provider, cannot refresh token");
return Err(err);
}
tracing::info!("token expired (853004), attempting silent refresh");
match self
.refresh_token(sent_token.as_deref(), options.clone())
.await
{
Ok(token) => {
tracing::info!("token refreshed, retrying the original request");
set_bearer_token(&mut options, &token);
self.inner.execute(endpoint, replay_payload, options).await
}
Err(refresh_err) => {
tracing::warn!(error = %refresh_err, "token refresh failed, returning the original error");
Err(err)
}
}
})
}
fn name(&self) -> &str {
self.inner.name()
}
}
pub fn is_token_expired(err: &wecomx_transport::Error) -> bool {
matches!(
err,
wecomx_transport::Error::Api {
code: Some(TOKEN_EXPIRED_ERRCODE),
..
}
)
}
pub fn set_bearer_token(options: &mut RequestOptions, token: &str) {
let Ok(mut value) = reqwest::header::HeaderValue::from_str(&format!("Bearer {token}")) else {
return;
};
value.set_sensitive(true);
options
.wire
.headers
.insert(reqwest::header::AUTHORIZATION, value);
}
#[cfg(test)]
mod tests {
use assert_json_diff::assert_json_eq;
use serde_json::json;
use std::sync::Arc;
use wiremock::matchers::{method, path};
use wiremock::{Match, Mock, MockServer, Request, ResponseTemplate};
use wecomx_auth::{BotCredential, Credentials, MemoryCredentialStore};
use wecomx_transport::{HttpTransportBackend, RequestOptions, Transport};
use super::*;
fn auth_ep(url: &str) -> wecomx_transport::Endpoint {
wecomx_auth::auth_endpoint(url)
}
struct NoAuthorization;
impl Match for NoAuthorization {
fn matches(&self, request: &Request) -> bool {
request.headers.get("authorization").is_none()
}
}
fn wrapped_transport(
base_url: &str,
bot: Option<BotCredential>,
token: Option<&str>,
auth_endpoint_url: &str,
) -> Transport {
let store = MemoryCredentialStore::new(Credentials { bot, token: None }).shared();
let provider = wecomx_auth::BotGatewayTokenProvider::from_store(store)
.with_auth_endpoint(auth_ep(auth_endpoint_url));
HttpTransportBackend::builder()
.base_url(base_url)
.build()
.expect("valid")
.wrap_backend(|backend| {
Arc::new(WecomBackend::new(
backend,
Some(Arc::new(provider)),
token.map(str::to_owned),
))
})
}
fn ep(base: &str, path_str: &str) -> wecomx_transport::Endpoint {
wecomx_transport::Endpoint::new()
.with(wecomx_transport::HttpEndpoint::new(path_str).with_service(base))
}
fn api_error(code: Option<i64>) -> wecomx_transport::Error {
wecomx_transport::Error::Api {
message: "err".into(),
action: "test".into(),
code,
body: Box::new(serde_json::Value::Null),
}
}
#[test]
fn token_expired_errcode_matches() {
assert!(is_token_expired(&api_error(Some(TOKEN_EXPIRED_ERRCODE))));
}
#[test]
fn other_errors_do_not_match() {
assert!(!is_token_expired(&api_error(Some(40001))));
assert!(!is_token_expired(&api_error(None)));
assert!(!is_token_expired(&wecomx_transport::Error::Other(
"x".into()
)));
}
#[test]
fn set_bearer_token_marks_sensitive() {
let mut options = RequestOptions::default();
set_bearer_token(&mut options, "tok-1");
let value = options
.wire
.headers
.get(reqwest::header::AUTHORIZATION)
.unwrap();
assert_eq!(value.to_str().unwrap(), "Bearer tok-1");
assert!(value.is_sensitive(), "token 头应标记敏感");
}
#[test]
fn set_bearer_token_overwrites() {
let mut options = RequestOptions::default();
set_bearer_token(&mut options, "old");
set_bearer_token(&mut options, "new");
let value = options
.wire
.headers
.get(reqwest::header::AUTHORIZATION)
.unwrap();
assert_eq!(value.to_str().unwrap(), "Bearer new");
}
#[test]
fn wrap_backend_decorates_in_place() {
let transport = HttpTransportBackend::builder()
.base_url("http://localhost")
.build()
.expect("valid");
let transport = transport.wrap_backend(|backend| {
Arc::new(WecomBackend::new(backend, None, Some("tok-1".into())))
});
assert_eq!(transport.name(), "http");
}
#[test]
fn debug_does_not_leak_secrets() {
let backend = WecomBackend::new(
Arc::new(HttpTransportBackend::default()),
None,
Some("cached-token".into()),
);
let dbg = format!("{backend:?}");
assert!(!dbg.contains("cached-token"), "token 泄露: {dbg}");
}
#[test]
fn no_provider_token_cached() {
let backend = WecomBackend::new(
Arc::new(HttpTransportBackend::default()),
None,
Some("cached-token".into()),
);
assert!(backend.provider.is_none());
assert_eq!(backend.cached_token().as_deref(), Some("cached-token"));
}
#[tokio::test]
async fn injects_auth_when_require_auth_and_token_available() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/auth"))
.and(wiremock::matchers::header("authorization", "Bearer tok-x"))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"result": "{\"ok\":true}"})),
)
.expect(1)
.mount(&server)
.await;
let transport = wrapped_transport(&server.uri(), None, Some("tok-x"), "http://unused");
let endpoint = ep(&server.uri(), "/auth").with(RequireAuth);
let v = transport
.invoke(&endpoint, json!({}))
.await
.unwrap()
.into_result()
.unwrap();
assert_json_eq!(v, json!({"ok": true}));
server.verify().await;
}
#[tokio::test]
async fn rejects_require_auth_without_token() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/auth"))
.respond_with(ResponseTemplate::new(200))
.expect(0)
.mount(&server)
.await;
let transport = wrapped_transport(&server.uri(), None, None, "http://unused");
let endpoint = ep(&server.uri(), "/auth").with(RequireAuth);
let err = transport.invoke(&endpoint, json!({})).await.unwrap_err();
match err {
wecomx_transport::Error::Other(e) => {
let inner = e.downcast_ref::<AuthError>();
assert!(
inner.is_some_and(|e| matches!(e, AuthError::MissingCredentials(_))),
"expected AuthError::MissingCredentials, got {inner:?}"
);
}
other => panic!("expected Error::Other(AuthError), got {other:?}"),
}
server.verify().await;
}
#[tokio::test]
async fn injects_auth_on_endpoint_without_require_auth() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/open"))
.and(wiremock::matchers::header("authorization", "Bearer tok-x"))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"result": "{\"ok\":true}"})),
)
.expect(1)
.mount(&server)
.await;
let transport = wrapped_transport(&server.uri(), None, Some("tok-x"), "http://unused");
let endpoint = ep(&server.uri(), "/open");
let v = transport
.invoke(&endpoint, json!({}))
.await
.unwrap()
.into_result()
.unwrap();
assert_json_eq!(v, json!({"ok": true}));
server.verify().await;
}
#[tokio::test]
async fn no_token_no_require_auth_omits_auth_header() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/open"))
.and(NoAuthorization)
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"result": "{\"ok\":true}"})),
)
.expect(1)
.mount(&server)
.await;
let transport = wrapped_transport(&server.uri(), None, None, "http://unused");
let endpoint = ep(&server.uri(), "/open");
let v = transport
.invoke(&endpoint, json!({}))
.await
.unwrap()
.into_result()
.unwrap();
assert_json_eq!(v, json!({"ok": true}));
server.verify().await;
}
#[tokio::test]
async fn flat_envelope_bootstrap_suppresses_auth() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/bootstrap"))
.and(NoAuthorization)
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"errcode": 0, "token": "t1"})),
)
.expect(1)
.mount(&server)
.await;
let transport = wrapped_transport(&server.uri(), None, Some("old-token"), "http://unused");
let endpoint = wecomx_transport::Endpoint::new().with(
wecomx_transport::HttpEndpoint::new("/bootstrap")
.with_service(server.uri())
.with_res_envelope(wecomx_auth::FlatRes),
);
let endpoint = endpoint.with(SuppressAuth);
let v = transport
.invoke(&endpoint, json!({}))
.await
.unwrap()
.into_result()
.unwrap();
assert_json_eq!(v, json!({"token": "t1"}));
server.verify().await;
}
#[tokio::test]
async fn refresh_reuses_execute_options_without_stale_auth() {
let server = MockServer::start().await;
let auth_url = format!("{}/bootstrap", server.uri());
Mock::given(method("POST"))
.and(path("/api"))
.and(wiremock::matchers::header(
"authorization",
"Bearer tok-old",
))
.and(wiremock::matchers::header("x-run-scope", "run-1"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"error": {"code": 853004, "message": "token expired"}
})))
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/bootstrap"))
.and(wiremock::matchers::header("x-run-scope", "run-1"))
.and(NoAuthorization)
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"errcode": 0, "token": "tok-new"})),
)
.expect(1)
.mount(&server)
.await;
Mock::given(method("POST"))
.and(path("/api"))
.and(wiremock::matchers::header(
"authorization",
"Bearer tok-new",
))
.and(wiremock::matchers::header("x-run-scope", "run-1"))
.respond_with(
ResponseTemplate::new(200).set_body_json(json!({"result": "{\"ok\":true}"})),
)
.expect(1)
.mount(&server)
.await;
let transport = wrapped_transport(
&server.uri(),
Some(BotCredential::new("bot1".into(), "secret1".into())),
Some("tok-old"),
&auth_url,
);
let endpoint = ep(&server.uri(), "/api").with(RequireAuth);
let v = transport
.invoke(&endpoint, json!({}))
.header("x-run-scope", "run-1")
.await
.unwrap()
.into_result()
.unwrap();
assert_json_eq!(v, json!({"ok": true}));
server.verify().await;
}
}