use parking_lot::Mutex;
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)]
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)>,
}
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(),
}
}
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 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>,
}
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>,
}
impl GenericOAuth2Provider {
pub fn new(config: OAuth2Config, transport: Arc<dyn OAuth2HttpTransport>) -> Self {
Self { config, transport }
}
fn build_redirect_url(&self, state: &str) -> 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(" ")));
}
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) -> Result<serde_json::Value, OAuth2Error> {
let 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,
});
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)
}
fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
self.config.validate()?;
if code.is_empty() {
return Err(OAuth2Error::AuthFailed("授权码不能为空".into()));
}
let token_json = self.exchange_token(code)?;
let access_token = token_json
.get("access_token")
.and_then(|value| value.as_str())
.ok_or_else(|| {
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 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;
Ok(user)
}
}
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::*;
#[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())
}
}
}