use base64::Engine;
use parking_lot::Mutex;
use rand::rngs::OsRng;
use rand::RngCore;
use sha2::{Digest, Sha256};
use std::collections::VecDeque;
use std::sync::Arc;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum OAuth2Error {
#[error("OAuth2 字段缺失: {0}")]
MissingField(String),
#[error("OAuth2 授权失败: {0}")]
AuthFailed(String),
#[error("OAuth2 token 交换失败: {0}")]
TokenExchangeFailed(String),
#[error("OAuth2 获取用户信息失败: {0}")]
UserInfoFailed(String),
#[error("OAuth2 HTTP 传输失败: {0}")]
HttpTransport(String),
#[error("OAuth2 序列化失败: {0}")]
Serialize(String),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PkceMethod {
S256,
}
impl std::fmt::Display for PkceMethod {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PkceMethod::S256 => write!(f, "S256"),
}
}
}
pub struct PkceParams {
pub code_verifier: String,
pub code_challenge: String,
pub method: PkceMethod,
}
impl std::fmt::Debug for PkceParams {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PkceParams")
.field(
"code_verifier",
&format!("<redacted, len={}>", self.code_verifier.len()),
)
.field("code_challenge", &self.code_challenge)
.field("method", &self.method)
.finish()
}
}
impl Clone for PkceParams {
fn clone(&self) -> Self {
Self {
code_verifier: self.code_verifier.clone(),
code_challenge: self.code_challenge.clone(),
method: self.method,
}
}
}
#[derive(Clone)]
pub struct OAuth2Config {
pub client_id: String,
pub client_secret: String,
pub redirect_url: String,
pub auth_url: String,
pub token_url: String,
pub user_url: Option<String>,
pub scopes: Vec<String>,
pub extra_params: Vec<(String, String)>,
pub pkce_enabled: bool,
pub device_auth_url: Option<String>,
}
impl std::fmt::Debug for OAuth2Config {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuth2Config")
.field("client_id", &self.client_id)
.field("client_secret", &"***")
.field("redirect_url", &self.redirect_url)
.field("auth_url", &self.auth_url)
.field("token_url", &self.token_url)
.field("user_url", &self.user_url)
.field("scopes", &self.scopes)
.field("extra_params", &self.extra_params)
.field("pkce_enabled", &self.pkce_enabled)
.field("device_auth_url", &self.device_auth_url)
.finish()
}
}
impl OAuth2Config {
#[allow(clippy::too_many_arguments)]
pub fn new(
client_id: impl Into<String>,
client_secret: impl Into<String>,
redirect_url: impl Into<String>,
auth_url: impl Into<String>,
token_url: impl Into<String>,
) -> Self {
Self {
client_id: client_id.into(),
client_secret: client_secret.into(),
redirect_url: redirect_url.into(),
auth_url: auth_url.into(),
token_url: token_url.into(),
user_url: None,
scopes: Vec::new(),
extra_params: Vec::new(),
pkce_enabled: false,
device_auth_url: None,
}
}
pub fn with_user_url(mut self, user_url: impl Into<String>) -> Self {
self.user_url = Some(user_url.into());
self
}
pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
self.scopes = scopes;
self
}
pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
self.scopes.push(scope.into());
self
}
pub fn with_extra_params(mut self, params: Vec<(String, String)>) -> Self {
self.extra_params = params;
self
}
pub fn with_extra_param(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.extra_params.push((key.into(), value.into()));
self
}
pub fn with_pkce(mut self, enabled: bool) -> Self {
self.pkce_enabled = enabled;
self
}
pub fn with_device_auth_url(mut self, url: impl Into<String>) -> Self {
self.device_auth_url = Some(url.into());
self
}
pub fn generate_state() -> String {
let mut bytes = [0u8; 16];
OsRng.fill_bytes(&mut bytes);
hex::encode(bytes)
}
pub fn generate_pkce_pair() -> PkceParams {
let mut bytes = [0u8; 32];
OsRng.fill_bytes(&mut bytes);
let code_verifier = hex::encode(bytes);
let mut hasher = Sha256::new();
hasher.update(code_verifier.as_bytes());
let digest = hasher.finalize();
let code_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
PkceParams {
code_verifier,
code_challenge,
method: PkceMethod::S256,
}
}
pub fn validate(&self) -> Result<(), OAuth2Error> {
if self.client_id.is_empty() {
return Err(OAuth2Error::MissingField("client_id".into()));
}
if self.client_secret.is_empty() {
return Err(OAuth2Error::MissingField("client_secret".into()));
}
if self.redirect_url.is_empty() {
return Err(OAuth2Error::MissingField("redirect_url".into()));
}
if self.auth_url.is_empty() {
return Err(OAuth2Error::MissingField("auth_url".into()));
}
if self.token_url.is_empty() {
return Err(OAuth2Error::MissingField("token_url".into()));
}
Ok(())
}
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct SocialiteUser {
pub id: String,
pub nickname: Option<String>,
pub name: Option<String>,
pub email: Option<String>,
pub avatar: Option<String>,
pub raw: serde_json::Value,
#[serde(skip_serializing)]
pub access_token: Option<String>,
#[serde(skip_serializing)]
pub refresh_token: Option<String>,
pub expires_in: Option<i64>,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct TokenResponse {
#[serde(skip_serializing)]
pub access_token: String,
pub token_type: Option<String>,
pub expires_in: Option<i64>,
pub scope: Option<String>,
#[serde(skip_serializing)]
pub refresh_token: Option<String>,
}
#[derive(Debug, Clone)]
pub struct OAuth2AuditEvent {
pub client_id: String,
pub grant_type: String,
pub result: String,
pub timestamp: i64,
pub alert_code: Option<String>,
pub message: Option<String>,
}
pub trait OAuth2AuditLogger: Send + Sync {
fn log_event(&self, event: &OAuth2AuditEvent);
}
pub trait OAuth2Provider: Send + Sync {
fn redirect_url(&self, state: &str) -> String;
fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error>;
}
pub trait OAuth2HttpTransport: Send + Sync {
fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error>;
}
#[derive(Debug, Default)]
pub struct MemoryOAuth2HttpTransport {
requests: Mutex<Vec<(String, String)>>,
responses: Mutex<VecDeque<String>>,
}
impl MemoryOAuth2HttpTransport {
pub fn new() -> Self {
Self::default()
}
pub fn push_response(&self, response: impl Into<String>) {
self.responses.lock().push_back(response.into());
}
pub fn count(&self) -> usize {
self.requests.lock().len()
}
pub fn all(&self) -> Vec<(String, String)> {
self.requests.lock().clone()
}
pub fn last(&self) -> Option<(String, String)> {
self.requests.lock().last().cloned()
}
pub fn clear(&self) {
self.requests.lock().clear();
self.responses.lock().clear();
}
}
impl OAuth2HttpTransport for MemoryOAuth2HttpTransport {
fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error> {
self.requests
.lock()
.push((url.to_string(), body.to_string()));
let mut responses = self.responses.lock();
match responses.pop_front() {
Some(resp) => Ok(resp),
None => Ok(String::new()),
}
}
}
pub struct GenericOAuth2Provider {
config: OAuth2Config,
transport: Arc<dyn OAuth2HttpTransport>,
audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
#[cfg(feature = "redis-store")]
token_store: Option<Arc<dyn crate::oauth_store::OAuth2TokenStore>>,
}
impl GenericOAuth2Provider {
pub fn new(config: OAuth2Config, transport: Arc<dyn OAuth2HttpTransport>) -> Self {
Self {
config,
transport,
audit_logger: None,
#[cfg(feature = "redis-store")]
token_store: None,
}
}
pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
self.audit_logger = Some(logger);
self
}
#[cfg(feature = "redis-store")]
pub fn with_token_store(
mut self,
store: Arc<dyn crate::oauth_store::OAuth2TokenStore>,
) -> Self {
self.token_store = Some(store);
self
}
fn log_audit(
&self,
grant_type: &str,
result: &str,
alert_code: Option<&str>,
message: Option<&str>,
) {
if let Some(logger) = &self.audit_logger {
let event = OAuth2AuditEvent {
client_id: self.config.client_id.clone(),
grant_type: grant_type.to_string(),
result: result.to_string(),
timestamp: chrono::Utc::now().timestamp(),
alert_code: alert_code.map(|s| s.to_string()),
message: message.map(|s| s.to_string()),
};
logger.log_event(&event);
}
}
pub fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse, OAuth2Error> {
if refresh_token.is_empty() {
self.log_audit("refresh_token", "failure", None, Some("refresh_token 为空"));
return Err(OAuth2Error::AuthFailed("refresh_token 不能为空".into()));
}
let body = serde_json::json!({
"grant_type": "refresh_token",
"refresh_token": refresh_token,
"client_id": self.config.client_id,
"client_secret": self.config.client_secret,
});
let body_str =
serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
let response = self
.transport
.post_json(&self.config.token_url, &body_str)
.map_err(|err| {
self.log_audit("refresh_token", "failure", None, Some(&err.to_string()));
OAuth2Error::HttpTransport(err.to_string())
})?;
if response.is_empty() {
self.log_audit("refresh_token", "failure", None, Some("空响应"));
return Err(OAuth2Error::TokenExchangeFailed("token 响应为空".into()));
}
let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
self.log_audit(
"refresh_token",
"failure",
None,
Some(&format!("JSON 解析失败: {err}")),
);
OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
})?;
let access_token = json
.get("access_token")
.and_then(|v| v.as_str())
.ok_or_else(|| {
self.log_audit(
"refresh_token",
"failure",
None,
Some("响应缺少 access_token"),
);
OAuth2Error::TokenExchangeFailed("token 响应缺少 access_token 字段".into())
})?
.to_string();
let token_response = TokenResponse {
access_token,
token_type: json
.get("token_type")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
scope: json
.get("scope")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
refresh_token: json
.get("refresh_token")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
};
self.log_audit("refresh_token", "success", None, None);
Ok(token_response)
}
fn build_redirect_url(&self, state: &str, pkce: Option<&PkceParams>) -> String {
let mut params: Vec<(String, String)> = vec![
("client_id".into(), self.config.client_id.clone()),
("redirect_uri".into(), self.config.redirect_url.clone()),
("response_type".into(), "code".into()),
("state".into(), state.to_string()),
];
if !self.config.scopes.is_empty() {
params.push(("scope".into(), self.config.scopes.join(" ")));
}
if let Some(pkce) = pkce {
params.push(("code_challenge".into(), pkce.code_challenge.clone()));
params.push(("code_challenge_method".into(), pkce.method.to_string()));
}
for (key, value) in &self.config.extra_params {
params.push((key.clone(), value.clone()));
}
let query = params
.iter()
.map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
.collect::<Vec<_>>()
.join("&");
let separator = if self.config.auth_url.contains('?') {
"&"
} else {
"?"
};
format!("{}{}{}", self.config.auth_url, separator, query)
}
fn exchange_token(
&self,
code: &str,
code_verifier: Option<&str>,
) -> Result<serde_json::Value, OAuth2Error> {
let mut body = serde_json::json!({
"grant_type": "authorization_code",
"code": code,
"client_id": self.config.client_id,
"client_secret": self.config.client_secret,
"redirect_uri": self.config.redirect_url,
});
if let Some(verifier) = code_verifier {
body["code_verifier"] = serde_json::Value::String(verifier.to_string());
}
let body_str =
serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
let response = self
.transport
.post_json(&self.config.token_url, &body_str)
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
if response.is_empty() {
return Ok(serde_json::Value::Null);
}
serde_json::from_str(&response)
.map_err(|err| OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}")))
}
fn fetch_user_info(
&self,
access_token: &str,
token_json: &serde_json::Value,
) -> Result<serde_json::Value, OAuth2Error> {
let user_url = self
.config
.user_url
.as_ref()
.ok_or_else(|| OAuth2Error::UserInfoFailed("user_url 未配置".into()))?;
let body = serde_json::json!({
"access_token": access_token,
"openid": token_json.get("openid").cloned().unwrap_or(serde_json::Value::Null),
});
let body_str =
serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
let response = self
.transport
.post_json(user_url, &body_str)
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
if response.is_empty() {
return Ok(serde_json::Value::Null);
}
serde_json::from_str(&response)
.map_err(|err| OAuth2Error::UserInfoFailed(format!("解析用户信息响应失败: {err}")))
}
fn extract_user_fields(user_json: &serde_json::Value) -> SocialiteUser {
let id = user_json
.get("id")
.or_else(|| user_json.get("openid"))
.or_else(|| user_json.get("user_id"))
.and_then(extract_string)
.unwrap_or_default();
let nickname = user_json
.get("nickname")
.or_else(|| user_json.get("nick_name"))
.and_then(extract_string);
let name = user_json
.get("name")
.or_else(|| user_json.get("username"))
.and_then(extract_string);
let email = user_json.get("email").and_then(extract_string);
let avatar = user_json
.get("avatar")
.or_else(|| user_json.get("figureurl_qq_1"))
.or_else(|| user_json.get("figureurl"))
.or_else(|| user_json.get("headimgurl"))
.and_then(extract_string);
SocialiteUser {
id,
nickname,
name,
email,
avatar,
raw: user_json.clone(),
access_token: None,
refresh_token: None,
expires_in: None,
}
}
}
impl OAuth2Provider for GenericOAuth2Provider {
fn redirect_url(&self, state: &str) -> String {
self.build_redirect_url(state, None)
}
fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
self.user_from_token_with_pkce(code, None)
}
}
impl GenericOAuth2Provider {
pub fn redirect_url_with_pkce(&self, state: &str, pkce: &PkceParams) -> String {
self.build_redirect_url(state, Some(pkce))
}
pub fn user_from_token_with_pkce(
&self,
code: &str,
code_verifier: Option<&str>,
) -> Result<SocialiteUser, OAuth2Error> {
self.config.validate()?;
if code.is_empty() {
return Err(OAuth2Error::AuthFailed("授权码不能为空".into()));
}
let token_json = self.exchange_token(code, code_verifier)?;
let access_token = token_json
.get("access_token")
.and_then(|value| value.as_str())
.ok_or_else(|| {
self.log_audit(
"authorization_code",
"failure",
None,
Some("token 响应缺少 access_token"),
);
OAuth2Error::TokenExchangeFailed(format!(
"token 响应缺少 access_token 字段: {token_json}"
))
})?
.to_string();
let refresh_token = token_json
.get("refresh_token")
.and_then(|value| value.as_str())
.map(|value| value.to_string());
let expires_in = token_json
.get("expires_in")
.and_then(|value| value.as_i64());
let (access_token, refresh_token, expires_in) = if let Some(exp) = expires_in {
if exp <= 0 {
if let Some(ref rt) = refresh_token {
match self.refresh_token(rt) {
Ok(new_token) => {
let new_access = new_token.access_token;
let new_refresh = new_token.refresh_token.or(refresh_token.clone());
let new_exp = new_token.expires_in;
(new_access, new_refresh, new_exp)
}
Err(_) => {
(access_token, refresh_token, expires_in)
}
}
} else {
(access_token, refresh_token, expires_in)
}
} else {
(access_token, refresh_token, expires_in)
}
} else {
(access_token, refresh_token, expires_in)
};
let mut user = if self.config.user_url.is_some() {
let user_json = self.fetch_user_info(&access_token, &token_json)?;
Self::extract_user_fields(&user_json)
} else {
SocialiteUser::default()
};
user.access_token = Some(access_token);
user.refresh_token = refresh_token;
user.expires_in = expires_in;
#[cfg(feature = "redis-store")]
if let Some(store) = &self.token_store {
let store = store.clone();
let client_id = self.config.client_id.clone();
let token_to_store = TokenResponse {
access_token: user.access_token.clone().unwrap_or_default(),
token_type: None,
expires_in: user.expires_in,
scope: None,
refresh_token: user.refresh_token.clone(),
};
tokio::task::spawn(async move {
if let Err(err) = store.store_token(&client_id, &token_to_store).await {
tracing::warn!(
error = %err,
client_id = %client_id,
"OAUTH2_TOKEN_STORE_FAILED: token 存储失败(best-effort,不影响主流程)"
);
}
});
}
self.log_audit("authorization_code", "success", None, None);
Ok(user)
}
}
pub struct ImplicitOAuth2Provider {
config: OAuth2Config,
audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
}
impl ImplicitOAuth2Provider {
pub fn new(config: OAuth2Config) -> Self {
Self {
config,
audit_logger: None,
}
}
pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
self.audit_logger = Some(logger);
self
}
fn log_audit(
&self,
grant_type: &str,
result: &str,
alert_code: Option<&str>,
message: Option<&str>,
) {
if let Some(logger) = &self.audit_logger {
let event = OAuth2AuditEvent {
client_id: self.config.client_id.clone(),
grant_type: grant_type.to_string(),
result: result.to_string(),
timestamp: chrono::Utc::now().timestamp(),
alert_code: alert_code.map(|s| s.to_string()),
message: message.map(|s| s.to_string()),
};
logger.log_event(&event);
}
}
pub fn redirect_url(&self, state: &str) -> String {
self.redirect_url_with_pkce(state, None)
}
pub fn redirect_url_with_pkce(&self, state: &str, pkce: Option<&PkceParams>) -> String {
let mut params: Vec<(String, String)> = vec![
("client_id".into(), self.config.client_id.clone()),
("redirect_uri".into(), self.config.redirect_url.clone()),
("response_type".into(), "token".into()),
("state".into(), state.to_string()),
];
if !self.config.scopes.is_empty() {
params.push(("scope".into(), self.config.scopes.join(" ")));
}
if let Some(pkce) = pkce {
params.push(("code_challenge".into(), pkce.code_challenge.clone()));
params.push(("code_challenge_method".into(), pkce.method.to_string()));
}
for (key, value) in &self.config.extra_params {
params.push((key.clone(), value.clone()));
}
let query = params
.iter()
.map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
.collect::<Vec<_>>()
.join("&");
let separator = if self.config.auth_url.contains('?') {
"&"
} else {
"?"
};
format!("{}{}{}", self.config.auth_url, separator, query)
}
pub fn parse_fragment(
&self,
fragment: &str,
expected_state: &str,
) -> Result<TokenResponse, OAuth2Error> {
self.log_audit(
"implicit",
"success",
Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED"),
Some("implicit 流程 token 经 URL fragment 暴露"),
);
if fragment.is_empty() {
self.log_audit("implicit", "failure", None, Some("fragment 为空"));
return Err(OAuth2Error::TokenExchangeFailed("fragment 为空".into()));
}
let params: std::collections::HashMap<&str, &str> = fragment
.split('&')
.filter_map(|pair| {
let (key, value) = pair.split_once('=')?;
Some((key, value))
})
.collect();
let state = params.get("state").copied().unwrap_or("");
if state != expected_state {
self.log_audit(
"implicit",
"failure",
Some("OAUTH2_CSRF_STATE_MISMATCH"),
Some(&format!(
"state 不匹配: expected={expected_state}, actual={state}"
)),
);
return Err(OAuth2Error::AuthFailed("CSRF state mismatch".into()));
}
let access_token = params.get("access_token").copied().ok_or_else(|| {
self.log_audit(
"implicit",
"failure",
None,
Some("fragment 无 access_token"),
);
OAuth2Error::TokenExchangeFailed("fragment 中缺少 access_token".into())
})?;
Ok(TokenResponse {
access_token: access_token.to_string(),
token_type: params.get("token_type").map(|s| s.to_string()),
expires_in: params.get("expires_in").and_then(|s| s.parse().ok()),
scope: params.get("scope").map(|s| s.to_string()),
refresh_token: None,
})
}
}
impl OAuth2Provider for ImplicitOAuth2Provider {
fn redirect_url(&self, state: &str) -> String {
self.redirect_url(state)
}
fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
Ok(SocialiteUser {
access_token: Some(code.to_string()),
..Default::default()
})
}
}
#[cfg(feature = "device-code")]
pub mod device_code {
use super::*;
use async_trait::async_trait;
#[async_trait]
pub trait AsyncOAuth2HttpTransport: Send + Sync {
async fn post_form(
&self,
url: &str,
params: &[(&str, &str)],
) -> Result<String, OAuth2Error>;
}
#[derive(Debug, Clone)]
pub struct DeviceCodeResponse {
pub device_code: String,
pub user_code: String,
pub verification_uri: String,
pub expires_in: i64,
pub interval: i64,
}
pub struct DeviceCodeOAuth2Provider {
config: OAuth2Config,
transport: Arc<dyn AsyncOAuth2HttpTransport>,
audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
}
impl DeviceCodeOAuth2Provider {
pub fn new(config: OAuth2Config, transport: Arc<dyn AsyncOAuth2HttpTransport>) -> Self {
Self {
config,
transport,
audit_logger: None,
}
}
pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
self.audit_logger = Some(logger);
self
}
fn log_audit(
&self,
grant_type: &str,
result: &str,
alert_code: Option<&str>,
message: Option<&str>,
) {
if let Some(logger) = &self.audit_logger {
let event = OAuth2AuditEvent {
client_id: self.config.client_id.clone(),
grant_type: grant_type.to_string(),
result: result.to_string(),
timestamp: chrono::Utc::now().timestamp(),
alert_code: alert_code.map(|s| s.to_string()),
message: message.map(|s| s.to_string()),
};
logger.log_event(&event);
}
}
pub async fn request_device_code(
&self,
scope: &[String],
) -> Result<DeviceCodeResponse, OAuth2Error> {
let device_auth_url = self.config.device_auth_url.as_ref().ok_or_else(|| {
self.log_audit(
"device_code",
"failure",
None,
Some("device_auth_url 未配置"),
);
OAuth2Error::MissingField("device_auth_url".into())
})?;
let scope_str = scope.join(" ");
let params: Vec<(&str, &str)> = vec![
("client_id", self.config.client_id.as_str()),
("scope", scope_str.as_str()),
];
let response = self
.transport
.post_form(device_auth_url, ¶ms)
.await
.map_err(|err| {
self.log_audit("device_code", "failure", None, Some(&err.to_string()));
OAuth2Error::HttpTransport(err.to_string())
})?;
let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
self.log_audit(
"device_code",
"failure",
None,
Some(&format!("JSON 解析失败: {err}")),
);
OAuth2Error::TokenExchangeFailed(format!("解析 device code 响应失败: {err}"))
})?;
let device_code = json
.get("device_code")
.and_then(|v| v.as_str())
.ok_or_else(|| {
OAuth2Error::TokenExchangeFailed("device code 响应缺少 device_code 字段".into())
})?
.to_string();
let user_code = json
.get("user_code")
.and_then(|v| v.as_str())
.ok_or_else(|| {
OAuth2Error::TokenExchangeFailed("device code 响应缺少 user_code 字段".into())
})?
.to_string();
let verification_uri = json
.get("verification_uri")
.and_then(|v| v.as_str())
.ok_or_else(|| {
OAuth2Error::TokenExchangeFailed(
"device code 响应缺少 verification_uri 字段".into(),
)
})?
.to_string();
let expires_in = json
.get("expires_in")
.and_then(|v| v.as_i64())
.unwrap_or(600);
let interval = json.get("interval").and_then(|v| v.as_i64()).unwrap_or(5);
self.log_audit("device_code", "success", None, None);
Ok(DeviceCodeResponse {
device_code,
user_code,
verification_uri,
expires_in,
interval,
})
}
pub async fn poll_for_token(
&self,
device_code: &str,
mut interval: i64,
expires_in: i64,
) -> Result<TokenResponse, OAuth2Error> {
let start = std::time::Instant::now();
let expires_duration = std::time::Duration::from_secs(expires_in.max(0) as u64);
loop {
if start.elapsed() >= expires_duration {
self.log_audit(
"device_code",
"failure",
Some("OAUTH2_DEVICE_CODE_EXPIRED"),
Some("设备码已过期"),
);
return Err(OAuth2Error::AuthFailed(
"OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
));
}
let params: Vec<(&str, &str)> = vec![
("grant_type", "device_code"),
("device_code", device_code),
("client_id", self.config.client_id.as_str()),
];
let response = self
.transport
.post_form(&self.config.token_url, ¶ms)
.await
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
})?;
if let Some(access_token) = json.get("access_token").and_then(|v| v.as_str()) {
let token_response = TokenResponse {
access_token: access_token.to_string(),
token_type: json
.get("token_type")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
scope: json
.get("scope")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
refresh_token: json
.get("refresh_token")
.and_then(|v| v.as_str())
.map(|s| s.to_string()),
};
self.log_audit("device_code", "success", None, None);
return Ok(token_response);
}
let error = json.get("error").and_then(|v| v.as_str()).unwrap_or("");
match error {
"authorization_pending" => {
}
"slow_down" => {
interval = (interval + 5).min(60);
}
"access_denied" => {
self.log_audit(
"device_code",
"failure",
Some("OAUTH2_ACCESS_DENIED"),
Some("用户拒绝授权"),
);
return Err(OAuth2Error::AuthFailed(
"OAUTH2_ACCESS_DENIED: 用户拒绝授权".into(),
));
}
"expired_token" => {
self.log_audit(
"device_code",
"failure",
Some("OAUTH2_DEVICE_CODE_EXPIRED"),
Some("设备码已过期"),
);
return Err(OAuth2Error::AuthFailed(
"OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
));
}
_ => {
return Err(OAuth2Error::TokenExchangeFailed(format!(
"未知错误: {error}"
)));
}
}
if interval > 0 {
tokio::time::sleep(std::time::Duration::from_secs(interval as u64)).await;
}
}
}
}
#[derive(Default)]
pub struct MemoryAsyncOAuth2HttpTransport {
requests: Mutex<Vec<(String, String)>>,
responses: Mutex<VecDeque<String>>,
}
impl MemoryAsyncOAuth2HttpTransport {
pub fn new() -> Self {
Self::default()
}
pub fn push_response(&self, response: impl Into<String>) {
self.responses.lock().push_back(response.into());
}
pub fn count(&self) -> usize {
self.requests.lock().len()
}
}
#[async_trait]
impl AsyncOAuth2HttpTransport for MemoryAsyncOAuth2HttpTransport {
async fn post_form(
&self,
url: &str,
params: &[(&str, &str)],
) -> Result<String, OAuth2Error> {
let body = params
.iter()
.map(|(k, v)| format!("{k}={v}"))
.collect::<Vec<_>>()
.join("&");
self.requests.lock().push((url.to_string(), body));
let mut responses = self.responses.lock();
match responses.pop_front() {
Some(resp) => Ok(resp),
None => Ok(String::new()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_device_code_request() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
transport.push_response(
r#"{"device_code":"dc123","user_code":"UC-ABCD","verification_uri":"https://provider.com/device","expires_in":600,"interval":5}"#,
);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_device_auth_url("https://provider.com/device_authorize");
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let resp = provider
.request_device_code(&["read".into(), "write".into()])
.await
.expect("request_device_code 失败");
assert_eq!(resp.device_code, "dc123");
assert_eq!(resp.user_code, "UC-ABCD");
assert_eq!(resp.verification_uri, "https://provider.com/device");
assert_eq!(resp.expires_in, 600);
assert_eq!(resp.interval, 5);
}
#[tokio::test]
async fn test_device_code_poll_pending() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
transport.push_response(r#"{"error":"authorization_pending"}"#);
transport.push_response(
r#"{"access_token":"token123","token_type":"Bearer","expires_in":3600}"#,
);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let token = provider
.poll_for_token("dc123", 0, 600)
.await
.expect("poll_for_token 失败");
assert_eq!(token.access_token, "token123");
assert_eq!(token.token_type.as_deref(), Some("Bearer"));
}
#[tokio::test]
async fn test_device_code_poll_slow_down() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
transport.push_response(r#"{"error":"slow_down"}"#);
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let token = provider
.poll_for_token("dc123", 0, 600)
.await
.expect("poll_for_token 失败");
assert_eq!(token.access_token, "token123");
}
#[tokio::test]
async fn test_device_code_expired() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
transport.push_response(r#"{"error":"authorization_pending"}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let err = provider.poll_for_token("dc123", 0, 0).await.unwrap_err();
assert!(
err.to_string().contains("OAUTH2_DEVICE_CODE_EXPIRED"),
"应返回设备码过期错误: {err}"
);
}
#[tokio::test]
async fn test_device_code_access_denied() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
transport.push_response(r#"{"error":"access_denied"}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let err = provider.poll_for_token("dc123", 5, 600).await.unwrap_err();
assert!(
err.to_string().contains("OAUTH2_ACCESS_DENIED"),
"应返回 access_denied 错误: {err}"
);
}
#[tokio::test]
async fn test_device_code_no_auth_url() {
let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = DeviceCodeOAuth2Provider::new(config, transport);
let err = provider.request_device_code(&[]).await.unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "device_auth_url"));
}
}
}
fn extract_string(value: &serde_json::Value) -> Option<String> {
match value {
serde_json::Value::String(string) => Some(string.clone()),
serde_json::Value::Number(number) => number.as_i64().map(|number| number.to_string()),
_ => None,
}
}
fn percent_encode(input: &str) -> String {
let mut output = String::with_capacity(input.len());
for byte in input.as_bytes() {
if matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~') {
output.push(*byte as char);
} else {
output.push_str(&format!("%{byte:02X}"));
}
}
output
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::{prop_assert, prop_assert_eq};
#[test]
fn test_oauth2_config_builder() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_user_url("https://provider.com/user/info")
.with_scopes(vec!["scope1".into(), "scope2".into()])
.with_extra_param("foo", "bar");
assert_eq!(config.client_id, "client123");
assert_eq!(config.client_secret, "secret456");
assert_eq!(config.redirect_url, "https://example.com/callback");
assert_eq!(config.auth_url, "https://provider.com/authorize");
assert_eq!(config.token_url, "https://provider.com/token");
assert_eq!(
config.user_url.as_deref(),
Some("https://provider.com/user/info")
);
assert_eq!(config.scopes, vec!["scope1", "scope2"]);
assert_eq!(config.extra_params, vec![("foo".into(), "bar".into())]);
}
#[test]
fn test_oauth2_config_minimal() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
assert_eq!(config.client_id, "client123");
assert_eq!(config.client_secret, "secret456");
assert_eq!(config.redirect_url, "https://example.com/callback");
assert_eq!(config.auth_url, "https://provider.com/authorize");
assert_eq!(config.token_url, "https://provider.com/token");
assert!(config.user_url.is_none());
assert!(config.scopes.is_empty());
assert!(config.extra_params.is_empty());
assert!(config.validate().is_ok());
}
#[test]
fn test_oauth2_config_with_scope_chained() {
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_scope("get_user_info")
.with_scope("get_unionid");
assert_eq!(config.scopes, vec!["get_user_info", "get_unionid"]);
}
#[test]
fn test_oauth2_config_with_extra_params() {
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_extra_param("a", "1")
.with_extra_param("b", "2")
.with_extra_params(vec![("x".into(), "10".into())]);
assert_eq!(config.extra_params, vec![("x".into(), "10".into())]);
}
#[test]
fn test_oauth2_config_validate_empty_fields() {
let config = OAuth2Config::new(
"",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let err = config.validate().unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
let config = OAuth2Config::new(
"id",
"",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let err = config.validate().unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_secret"));
let config = OAuth2Config::new(
"id",
"secret",
"",
"https://provider.com/authorize",
"https://provider.com/token",
);
let err = config.validate().unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "redirect_url"));
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"",
"https://provider.com/token",
);
let err = config.validate().unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "auth_url"));
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"",
);
let err = config.validate().unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
}
#[test]
fn test_socialite_user_default() {
let user = SocialiteUser::default();
assert!(user.id.is_empty());
assert!(user.nickname.is_none());
assert!(user.name.is_none());
assert!(user.email.is_none());
assert!(user.avatar.is_none());
assert!(user.raw.is_null());
assert!(user.access_token.is_none());
assert!(user.refresh_token.is_none());
assert!(user.expires_in.is_none());
}
#[test]
fn test_socialite_user_serialize_deserialize() {
let user = SocialiteUser {
id: "123".into(),
nickname: Some("tester".into()),
name: Some("Test User".into()),
email: Some("test@example.com".into()),
avatar: Some("https://example.com/avatar.png".into()),
raw: serde_json::json!({"key": "value"}),
access_token: Some("token123".into()),
refresh_token: Some("refresh456".into()),
expires_in: Some(3600),
};
let json = serde_json::to_string(&user).expect("序列化失败");
assert!(
!json.contains("access_token"),
"access_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
);
assert!(
!json.contains("refresh_token"),
"refresh_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
);
let parsed: SocialiteUser = serde_json::from_str(&json).expect("反序列化失败");
assert_eq!(parsed.id, "123");
assert_eq!(parsed.nickname.as_deref(), Some("tester"));
assert_eq!(parsed.name.as_deref(), Some("Test User"));
assert_eq!(parsed.email.as_deref(), Some("test@example.com"));
assert_eq!(
parsed.avatar.as_deref(),
Some("https://example.com/avatar.png")
);
assert_eq!(parsed.access_token, None);
assert_eq!(parsed.refresh_token, None);
assert_eq!(parsed.expires_in, Some(3600));
}
#[test]
fn test_redirect_url_contains_required_params() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let url = provider.redirect_url("random_state_abc");
assert!(url.starts_with("https://provider.com/oauth2.0/authorize?"));
assert!(url.contains("client_id=client123"));
assert!(url.contains("redirect_uri=https%3A%2F%2Fexample.com%2Fcallback"));
assert!(url.contains("response_type=code"));
assert!(url.contains("state=random_state_abc"));
assert!(!url.contains("scope="));
}
#[test]
fn test_redirect_url_with_scopes() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
)
.with_scopes(vec!["get_user_info".into(), "get_unionid".into()]);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let url = provider.redirect_url("state123");
assert!(url.contains("scope=get_user_info%20get_unionid"));
}
#[test]
fn test_redirect_url_with_extra_params() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
)
.with_extra_param("foo", "bar")
.with_extra_param("display", "mobile");
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let url = provider.redirect_url("state123");
assert!(url.contains("foo=bar"));
assert!(url.contains("display=mobile"));
}
#[test]
fn test_redirect_url_with_existing_query() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize?foo=bar",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let url = provider.redirect_url("state123");
assert!(url.contains("?foo=bar&"));
assert!(url.contains("client_id=client123"));
}
#[test]
fn test_memory_oauth2_http_transport_post_json() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response(r#"{"access_token":"token123"}"#);
let response = transport
.post_json("https://example.com/token", r#"{"code":"abc"}"#)
.expect("post_json 失败");
assert_eq!(response, r#"{"access_token":"token123"}"#);
assert_eq!(transport.count(), 1);
let (url, body) = transport.last().expect("应有请求记录");
assert_eq!(url, "https://example.com/token");
assert_eq!(body, r#"{"code":"abc"}"#);
}
#[test]
fn test_memory_oauth2_http_transport_response_queue() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response("resp1");
transport.push_response("resp2");
let resp1 = transport
.post_json("url1", "body1")
.expect("第一次调用失败");
let resp2 = transport
.post_json("url2", "body2")
.expect("第二次调用失败");
assert_eq!(resp1, "resp1");
assert_eq!(resp2, "resp2");
assert_eq!(transport.count(), 2);
}
#[test]
fn test_memory_oauth2_http_transport_empty_response() {
let transport = MemoryOAuth2HttpTransport::new();
let response = transport
.post_json("url", "body")
.expect("post_json 不应失败");
assert_eq!(response, "");
}
#[test]
fn test_memory_oauth2_http_transport_clear() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response("resp");
transport.post_json("url", "body").expect("调用失败");
assert_eq!(transport.count(), 1);
transport.clear();
assert_eq!(transport.count(), 0);
let response = transport
.post_json("url", "body")
.expect("post_json 不应失败");
assert_eq!(response, "");
}
#[test]
fn test_generic_oauth2_provider_user_from_token() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","refresh_token":"refresh456","expires_in":3600,"openid":"openid_abc"}"#);
transport.push_response(
r#"{"id":"12345","nickname":"test_user","name":"Test","email":"test@example.com","avatar":"https://example.com/avatar.png"}"#,
);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_user_url("https://provider.com/user/info");
let provider = GenericOAuth2Provider::new(config, transport.clone());
let user = provider
.user_from_token("auth_code_abc")
.expect("user_from_token 失败");
assert_eq!(user.access_token.as_deref(), Some("token123"));
assert_eq!(user.refresh_token.as_deref(), Some("refresh456"));
assert_eq!(user.expires_in, Some(3600));
assert_eq!(user.id, "12345");
assert_eq!(user.nickname.as_deref(), Some("test_user"));
assert_eq!(user.name.as_deref(), Some("Test"));
assert_eq!(user.email.as_deref(), Some("test@example.com"));
assert_eq!(
user.avatar.as_deref(),
Some("https://example.com/avatar.png")
);
assert_eq!(user.raw["id"], "12345");
assert_eq!(user.raw["nickname"], "test_user");
assert_eq!(transport.count(), 2);
}
#[test]
fn test_generic_oauth2_provider_user_from_token_no_user_url() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":7200}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, transport.clone());
let user = provider
.user_from_token("auth_code")
.expect("user_from_token 失败");
assert_eq!(user.access_token.as_deref(), Some("token123"));
assert_eq!(user.expires_in, Some(7200));
assert!(user.refresh_token.is_none());
assert!(user.id.is_empty());
assert_eq!(transport.count(), 1);
}
#[test]
fn test_generic_oauth2_provider_missing_code() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let err = provider.user_from_token("").unwrap_err();
assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("授权码")));
}
#[test]
fn test_oauth2_provider_missing_config_fields() {
let config = OAuth2Config::new(
"",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
let err = provider.user_from_token("code").unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"",
);
let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
let err = provider.user_from_token("code").unwrap_err();
assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
}
#[test]
fn test_generic_oauth2_provider_token_response_missing_access_token() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response(r#"{"error":"invalid_grant"}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
let err = provider.user_from_token("code").unwrap_err();
assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
}
#[test]
fn test_generic_oauth2_provider_token_response_invalid_json() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response("not a json");
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
let err = provider.user_from_token("code").unwrap_err();
assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
}
#[test]
fn test_generic_oauth2_provider_user_info_field_aliases() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response(r#"{"access_token":"token123","openid":"openid_abc"}"#);
transport.push_response(
r#"{"openid":"qq_12345","nickname":"qq_user","figureurl_qq_1":"https://qzapp.qlogo.cn/1.png"}"#,
);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_user_url("https://provider.com/user/info");
let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
let user = provider
.user_from_token("code")
.expect("user_from_token 失败");
assert_eq!(user.id, "qq_12345");
assert_eq!(user.nickname.as_deref(), Some("qq_user"));
assert_eq!(user.avatar.as_deref(), Some("https://qzapp.qlogo.cn/1.png"));
}
#[test]
fn test_generic_oauth2_provider_user_id_integer() {
let transport = MemoryOAuth2HttpTransport::new();
transport.push_response(r#"{"access_token":"token123"}"#);
transport.push_response(r#"{"id":12345,"nickname":"github_user"}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_user_url("https://provider.com/user/info");
let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
let user = provider
.user_from_token("code")
.expect("user_from_token 失败");
assert_eq!(user.id, "12345");
assert_eq!(user.nickname.as_deref(), Some("github_user"));
}
#[test]
fn test_generic_oauth2_provider_http_transport_failure() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, Arc::new(FailingTransport));
let err = provider.user_from_token("code").unwrap_err();
assert!(matches!(err, OAuth2Error::HttpTransport(_)));
}
#[test]
fn test_percent_encode() {
assert_eq!(percent_encode("abcXYZ09-._~"), "abcXYZ09-._~");
assert_eq!(percent_encode("a b"), "a%20b");
assert_eq!(percent_encode("/"), "%2F");
assert_eq!(percent_encode(":"), "%3A");
assert_eq!(
percent_encode("https://example.com/path"),
"https%3A%2F%2Fexample.com%2Fpath"
);
assert_eq!(percent_encode("中"), "%E4%B8%AD");
}
#[test]
fn test_extract_string() {
assert_eq!(
extract_string(&serde_json::json!("hello")),
Some("hello".into())
);
assert_eq!(
extract_string(&serde_json::json!(12345)),
Some("12345".into())
);
assert_eq!(extract_string(&serde_json::json!(1.5)), None);
assert_eq!(extract_string(&serde_json::json!(true)), None);
assert_eq!(extract_string(&serde_json::Value::Null), None);
assert_eq!(extract_string(&serde_json::json!({"a": 1})), None);
}
struct FailingTransport;
impl OAuth2HttpTransport for FailingTransport {
fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
Err(OAuth2Error::HttpTransport("connection refused".into()))
}
}
struct MemoryHttpTransport;
impl OAuth2HttpTransport for MemoryHttpTransport {
fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
Ok(String::new())
}
}
#[test]
fn test_state_auto_generate() {
let state1 = OAuth2Config::generate_state();
let state2 = OAuth2Config::generate_state();
assert_eq!(state1.len(), 32, "state 应为 32 字符 hex(16 字节)");
assert_eq!(state2.len(), 32, "state 应为 32 字符 hex(16 字节)");
assert!(
state1.chars().all(|c| c.is_ascii_hexdigit()),
"state 应全为 hex 字符: {state1}"
);
assert!(
state2.chars().all(|c| c.is_ascii_hexdigit()),
"state 应全为 hex 字符: {state2}"
);
assert_ne!(state1, state2, "两次生成的 state 不应相同");
}
#[test]
fn test_pkce_pair_generate() {
let pkce = OAuth2Config::generate_pkce_pair();
assert!(
pkce.code_verifier.len() >= 43 && pkce.code_verifier.len() <= 128,
"code_verifier 长度应在 43-128 之间,实际: {}",
pkce.code_verifier.len()
);
assert_eq!(
pkce.code_verifier.len(),
64,
"code_verifier 应为 64 字符 hex"
);
assert!(
pkce.code_verifier.chars().all(|c| c.is_ascii_hexdigit()),
"code_verifier 应全为 hex 字符"
);
let mut hasher = sha2::Sha256::new();
sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
let digest = sha2::Digest::finalize(hasher);
let expected_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
assert_eq!(
pkce.code_challenge, expected_challenge,
"code_challenge 应等于 base64url(SHA256(code_verifier))"
);
assert_eq!(pkce.method, PkceMethod::S256);
}
#[test]
fn test_authorization_code_with_pkce() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
)
.with_pkce(true);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let pkce = OAuth2Config::generate_pkce_pair();
let url = provider.redirect_url_with_pkce("state123", &pkce);
assert!(url.contains("code_challenge="), "URL 应包含 code_challenge");
assert!(
url.contains("code_challenge_method=S256"),
"URL 应包含 code_challenge_method=S256"
);
assert!(url.contains("response_type=code"));
assert!(url.contains("state=state123"));
}
#[test]
fn test_authorization_code_without_pkce() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
);
assert!(!config.pkce_enabled);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let url = provider.redirect_url("state123");
assert!(
!url.contains("code_challenge"),
"URL 不应包含 code_challenge(PKCE 未启用)"
);
}
#[test]
fn test_client_secret_not_in_debug() {
let config = OAuth2Config::new(
"client123",
"super_secret_value_456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let debug_str = format!("{:?}", config);
assert!(
!debug_str.contains("super_secret_value_456"),
"client_secret 不应出现在 Debug 输出中: {debug_str}"
);
assert!(
debug_str.contains("***"),
"Debug 输出应包含脱敏标记 '***': {debug_str}"
);
}
#[test]
fn test_pkce_params_debug_redacted() {
let pkce = OAuth2Config::generate_pkce_pair();
let debug_str = format!("{:?}", pkce);
assert!(
!debug_str.contains(&pkce.code_verifier),
"code_verifier 明文不应出现在 Debug 输出中: {debug_str}"
);
assert!(
debug_str.contains("redacted"),
"Debug 输出应包含 'redacted' 标记: {debug_str}"
);
}
#[test]
fn test_with_pkce_builder() {
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
assert!(!config.pkce_enabled, "默认 pkce_enabled 应为 false");
let config = config.with_pkce(true);
assert!(
config.pkce_enabled,
"with_pkce(true) 后 pkce_enabled 应为 true"
);
let config = config.with_pkce(false);
assert!(
!config.pkce_enabled,
"with_pkce(false) 后 pkce_enabled 应为 false"
);
}
#[test]
fn test_with_device_auth_url_builder() {
let config = OAuth2Config::new(
"id",
"secret",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
assert!(
config.device_auth_url.is_none(),
"默认 device_auth_url 应为 None"
);
let config = config.with_device_auth_url("https://provider.com/device_authorize");
assert_eq!(
config.device_auth_url.as_deref(),
Some("https://provider.com/device_authorize"),
);
}
#[test]
fn test_exchange_token_with_pkce_verifier() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
)
.with_pkce(true);
let provider = GenericOAuth2Provider::new(config, transport.clone());
let pkce = OAuth2Config::generate_pkce_pair();
let user = provider
.user_from_token_with_pkce("auth_code", Some(&pkce.code_verifier))
.expect("user_from_token_with_pkce 失败");
assert_eq!(user.access_token.as_deref(), Some("token123"));
let (_url, body) = transport.last().expect("应有请求记录");
assert!(
body.contains("code_verifier"),
"token 交换 body 应包含 code_verifier: {body}"
);
assert!(
body.contains(&pkce.code_verifier),
"token 交换 body 应包含 code_verifier 值"
);
}
#[test]
fn test_exchange_token_without_pkce_verifier() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, transport.clone());
let _user = provider
.user_from_token("auth_code")
.expect("user_from_token 失败");
let (_url, body) = transport.last().expect("应有请求记录");
assert!(
!body.contains("code_verifier"),
"token 交换 body 不应包含 code_verifier(PKCE 未启用): {body}"
);
}
#[derive(Default)]
struct MockAuditLogger {
events: Mutex<Vec<OAuth2AuditEvent>>,
}
impl MockAuditLogger {
fn events(&self) -> Vec<OAuth2AuditEvent> {
self.events.lock().clone()
}
}
impl OAuth2AuditLogger for MockAuditLogger {
fn log_event(&self, event: &OAuth2AuditEvent) {
self.events.lock().push(event.clone());
}
}
#[test]
fn test_refresh_token() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(
r#"{"access_token":"new_token","token_type":"Bearer","expires_in":7200,"scope":"read"}"#,
);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, transport.clone());
let token_resp = provider
.refresh_token("old_refresh_token")
.expect("refresh_token 失败");
assert_eq!(token_resp.access_token, "new_token");
assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
assert_eq!(token_resp.expires_in, Some(7200));
assert_eq!(token_resp.scope.as_deref(), Some("read"));
}
#[test]
fn test_audit_log_on_token_exchange() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let logger = Arc::new(MockAuditLogger::default());
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, transport).with_audit_logger(logger.clone());
let _user = provider
.user_from_token("auth_code")
.expect("user_from_token 失败");
let events = logger.events();
assert!(
events.iter().any(|e| e.grant_type == "authorization_code"
&& e.result == "success"
&& e.client_id == "client123"),
"应记录 authorization_code success 事件: {events:?}"
);
}
#[test]
fn test_auto_refresh_on_expired() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(
r#"{"access_token":"expired_token","refresh_token":"valid_refresh","expires_in":0}"#,
);
transport.push_response(r#"{"access_token":"refreshed_token","expires_in":3600}"#);
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, transport.clone());
let user = provider
.user_from_token("auth_code")
.expect("user_from_token 失败");
assert_eq!(
user.access_token.as_deref(),
Some("refreshed_token"),
"过期 token 应自动刷新"
);
assert_eq!(user.expires_in, Some(3600));
assert_eq!(transport.count(), 2);
}
#[test]
fn test_refresh_token_empty() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
let err = provider.refresh_token("").unwrap_err();
assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("refresh_token")));
}
#[test]
fn test_refresh_token_invalid_json() {
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response("not a json");
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = GenericOAuth2Provider::new(config, transport);
let err = provider.refresh_token("valid_refresh").unwrap_err();
assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
}
#[test]
fn test_token_response_serialize_redacted() {
let token_resp = TokenResponse {
access_token: "secret_access_token".into(),
token_type: Some("Bearer".into()),
expires_in: Some(3600),
scope: Some("read".into()),
refresh_token: Some("secret_refresh_token".into()),
};
let json = serde_json::to_string(&token_resp).expect("序列化失败");
assert!(
!json.contains("secret_access_token"),
"access_token 不应出现在序列化 JSON 中: {json}"
);
assert!(
!json.contains("secret_refresh_token"),
"refresh_token 不应出现在序列化 JSON 中: {json}"
);
}
#[test]
fn test_implicit_redirect_url() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/oauth2.0/authorize",
"https://provider.com/oauth2.0/token",
)
.with_scope("profile");
let provider = ImplicitOAuth2Provider::new(config);
let url = provider.redirect_url("state_abc");
assert!(
url.contains("response_type=token"),
"URL 应含 response_type=token"
);
assert!(url.contains("client_id=client123"));
assert!(url.contains("state=state_abc"));
assert!(url.contains("scope=profile"));
}
#[test]
fn test_implicit_parse_fragment() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let fragment = "access_token=token123&token_type=Bearer&expires_in=3600&state=mystate";
let token_resp = provider
.parse_fragment(fragment, "mystate")
.expect("parse_fragment 失败");
assert_eq!(token_resp.access_token, "token123");
assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
assert_eq!(token_resp.expires_in, Some(3600));
assert!(
token_resp.refresh_token.is_none(),
"implicit 流程 refresh_token 应为 None"
);
}
#[test]
fn test_implicit_state_mismatch() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let fragment = "access_token=token123&state=wrong_state";
let err = provider
.parse_fragment(fragment, "expected_state")
.unwrap_err();
assert!(
matches!(&err, OAuth2Error::AuthFailed(msg) if msg.contains("CSRF state mismatch")),
"state 不匹配应返回 CSRF 错误: {err}"
);
}
#[test]
fn test_implicit_no_refresh_token() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let fragment = "access_token=token123&refresh_token=should_be_ignored&state=mystate";
let token_resp = provider
.parse_fragment(fragment, "mystate")
.expect("parse_fragment 失败");
assert!(
token_resp.refresh_token.is_none(),
"implicit 流程 refresh_token 应固定为 None"
);
}
#[test]
fn test_implicit_empty_fragment() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let err = provider.parse_fragment("", "state").unwrap_err();
assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
}
#[test]
fn test_implicit_no_access_token() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let fragment = "token_type=Bearer&state=mystate";
let err = provider.parse_fragment(fragment, "mystate").unwrap_err();
assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
}
#[test]
fn test_implicit_empty_scopes() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider = ImplicitOAuth2Provider::new(config);
let url = provider.redirect_url("state123");
assert!(
!url.contains("scope="),
"空 scopes 时 URL 不应含 scope 参数"
);
}
#[test]
fn test_implicit_audit_log_token_exposed() {
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let logger = Arc::new(MockAuditLogger::default());
let provider = ImplicitOAuth2Provider::new(config).with_audit_logger(logger.clone());
let fragment = "access_token=token123&state=mystate";
let _ = provider.parse_fragment(fragment, "mystate");
let events = logger.events();
assert!(
events
.iter()
.any(|e| e.alert_code.as_deref() == Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED")),
"应记录 OAUTH2_IMPLICIT_TOKEN_EXPOSED 告警: {events:?}"
);
}
#[cfg(feature = "redis-store")]
#[tokio::test]
async fn test_token_store_integration() {
use crate::oauth_store::{MemoryOAuth2TokenStore, OAuth2TokenStore};
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let store = Arc::new(MemoryOAuth2TokenStore::new());
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
let user = provider
.user_from_token("auth_code")
.expect("user_from_token 失败");
assert_eq!(user.access_token.as_deref(), Some("token123"));
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let stored = store
.get_token("client123")
.await
.expect("get_token 失败")
.expect("应查到存储的 token");
assert_eq!(stored.access_token, "token123");
}
#[cfg(feature = "redis-store")]
#[tokio::test]
async fn test_token_store_failure_best_effort() {
use crate::oauth_store::MemoryOAuth2TokenStore;
let transport = Arc::new(MemoryOAuth2HttpTransport::new());
transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
let store = Arc::new(MemoryOAuth2TokenStore::new());
let config = OAuth2Config::new(
"client123",
"secret456",
"https://example.com/callback",
"https://provider.com/authorize",
"https://provider.com/token",
);
let provider =
GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
let user = provider
.user_from_token("auth_code")
.expect("user_from_token 应成功(best-effort)");
assert_eq!(user.access_token.as_deref(), Some("token123"));
}
proptest::proptest! {
#[test]
fn proptest_state_unpredictable(_n in 0u32..1000) {
let s1 = OAuth2Config::generate_state();
let s2 = OAuth2Config::generate_state();
prop_assert_eq!(s1.len(), 32);
prop_assert_eq!(s2.len(), 32);
prop_assert!(s1.chars().all(|c| c.is_ascii_hexdigit()));
prop_assert!(s2.chars().all(|c| c.is_ascii_hexdigit()));
}
}
proptest::proptest! {
#[test]
fn proptest_pkce_verifier_length(_n in 0u32..1000) {
let pkce = OAuth2Config::generate_pkce_pair();
prop_assert!(pkce.code_verifier.len() >= 43);
prop_assert!(pkce.code_verifier.len() <= 128);
}
}
proptest::proptest! {
#[test]
fn proptest_pkce_challenge_matches_verifier(_n in 0u32..1000) {
let pkce = OAuth2Config::generate_pkce_pair();
let mut hasher = sha2::Sha256::new();
sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
let digest = sha2::Digest::finalize(hasher);
let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
prop_assert_eq!(pkce.code_challenge, expected);
}
}
}