1use base64::Engine;
44use parking_lot::Mutex;
45use rand::rngs::OsRng;
46use rand::RngCore;
47use sha2::{Digest, Sha256};
48use std::collections::VecDeque;
49use std::sync::Arc;
50use thiserror::Error;
51
52#[derive(Debug, Error)]
58pub enum OAuth2Error {
59 #[error("OAuth2 字段缺失: {0}")]
61 MissingField(String),
62 #[error("OAuth2 授权失败: {0}")]
64 AuthFailed(String),
65 #[error("OAuth2 token 交换失败: {0}")]
67 TokenExchangeFailed(String),
68 #[error("OAuth2 获取用户信息失败: {0}")]
70 UserInfoFailed(String),
71 #[error("OAuth2 HTTP 传输失败: {0}")]
73 HttpTransport(String),
74 #[error("OAuth2 序列化失败: {0}")]
76 Serialize(String),
77}
78
79#[derive(Debug, Clone, Copy, PartialEq, Eq)]
87pub enum PkceMethod {
88 S256,
90}
91
92impl std::fmt::Display for PkceMethod {
93 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
94 match self {
95 PkceMethod::S256 => write!(f, "S256"),
96 }
97 }
98}
99
100pub struct PkceParams {
104 pub code_verifier: String,
106 pub code_challenge: String,
108 pub method: PkceMethod,
110}
111
112impl std::fmt::Debug for PkceParams {
113 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
114 f.debug_struct("PkceParams")
115 .field(
116 "code_verifier",
117 &format!("<redacted, len={}>", self.code_verifier.len()),
118 )
119 .field("code_challenge", &self.code_challenge)
120 .field("method", &self.method)
121 .finish()
122 }
123}
124
125impl Clone for PkceParams {
126 fn clone(&self) -> Self {
127 Self {
128 code_verifier: self.code_verifier.clone(),
129 code_challenge: self.code_challenge.clone(),
130 method: self.method,
131 }
132 }
133}
134
135#[derive(Clone)]
171pub struct OAuth2Config {
172 pub client_id: String,
174 pub client_secret: String,
176 pub redirect_url: String,
178 pub auth_url: String,
180 pub token_url: String,
182 pub user_url: Option<String>,
184 pub scopes: Vec<String>,
186 pub extra_params: Vec<(String, String)>,
188 pub pkce_enabled: bool,
190 pub device_auth_url: Option<String>,
192}
193
194impl std::fmt::Debug for OAuth2Config {
195 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
196 f.debug_struct("OAuth2Config")
197 .field("client_id", &self.client_id)
198 .field("client_secret", &"***")
199 .field("redirect_url", &self.redirect_url)
200 .field("auth_url", &self.auth_url)
201 .field("token_url", &self.token_url)
202 .field("user_url", &self.user_url)
203 .field("scopes", &self.scopes)
204 .field("extra_params", &self.extra_params)
205 .field("pkce_enabled", &self.pkce_enabled)
206 .field("device_auth_url", &self.device_auth_url)
207 .finish()
208 }
209}
210
211impl OAuth2Config {
212 #[allow(clippy::too_many_arguments)]
222 pub fn new(
223 client_id: impl Into<String>,
224 client_secret: impl Into<String>,
225 redirect_url: impl Into<String>,
226 auth_url: impl Into<String>,
227 token_url: impl Into<String>,
228 ) -> Self {
229 Self {
230 client_id: client_id.into(),
231 client_secret: client_secret.into(),
232 redirect_url: redirect_url.into(),
233 auth_url: auth_url.into(),
234 token_url: token_url.into(),
235 user_url: None,
236 scopes: Vec::new(),
237 extra_params: Vec::new(),
238 pkce_enabled: false,
239 device_auth_url: None,
240 }
241 }
242
243 pub fn with_user_url(mut self, user_url: impl Into<String>) -> Self {
245 self.user_url = Some(user_url.into());
246 self
247 }
248
249 pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
251 self.scopes = scopes;
252 self
253 }
254
255 pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
257 self.scopes.push(scope.into());
258 self
259 }
260
261 pub fn with_extra_params(mut self, params: Vec<(String, String)>) -> Self {
263 self.extra_params = params;
264 self
265 }
266
267 pub fn with_extra_param(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
269 self.extra_params.push((key.into(), value.into()));
270 self
271 }
272
273 pub fn with_pkce(mut self, enabled: bool) -> Self {
278 self.pkce_enabled = enabled;
279 self
280 }
281
282 pub fn with_device_auth_url(mut self, url: impl Into<String>) -> Self {
284 self.device_auth_url = Some(url.into());
285 self
286 }
287
288 pub fn generate_state() -> String {
292 let mut bytes = [0u8; 16];
293 OsRng.fill_bytes(&mut bytes);
294 hex::encode(bytes)
295 }
296
297 pub fn generate_pkce_pair() -> PkceParams {
303 let mut bytes = [0u8; 32];
304 OsRng.fill_bytes(&mut bytes);
305 let code_verifier = hex::encode(bytes);
306 let mut hasher = Sha256::new();
307 hasher.update(code_verifier.as_bytes());
308 let digest = hasher.finalize();
309 let code_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
310 PkceParams {
311 code_verifier,
312 code_challenge,
313 method: PkceMethod::S256,
314 }
315 }
316
317 pub fn validate(&self) -> Result<(), OAuth2Error> {
322 if self.client_id.is_empty() {
323 return Err(OAuth2Error::MissingField("client_id".into()));
324 }
325 if self.client_secret.is_empty() {
326 return Err(OAuth2Error::MissingField("client_secret".into()));
327 }
328 if self.redirect_url.is_empty() {
329 return Err(OAuth2Error::MissingField("redirect_url".into()));
330 }
331 if self.auth_url.is_empty() {
332 return Err(OAuth2Error::MissingField("auth_url".into()));
333 }
334 if self.token_url.is_empty() {
335 return Err(OAuth2Error::MissingField("token_url".into()));
336 }
337 Ok(())
338 }
339}
340
341#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
364pub struct SocialiteUser {
365 pub id: String,
367 pub nickname: Option<String>,
369 pub name: Option<String>,
371 pub email: Option<String>,
373 pub avatar: Option<String>,
375 pub raw: serde_json::Value,
377 #[serde(skip_serializing)]
379 pub access_token: Option<String>,
380 #[serde(skip_serializing)]
382 pub refresh_token: Option<String>,
383 pub expires_in: Option<i64>,
385}
386
387#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
396pub struct TokenResponse {
397 #[serde(skip_serializing)]
399 pub access_token: String,
400 pub token_type: Option<String>,
402 pub expires_in: Option<i64>,
404 pub scope: Option<String>,
406 #[serde(skip_serializing)]
408 pub refresh_token: Option<String>,
409}
410
411#[derive(Debug, Clone)]
417pub struct OAuth2AuditEvent {
418 pub client_id: String,
420 pub grant_type: String,
422 pub result: String,
424 pub timestamp: i64,
426 pub alert_code: Option<String>,
428 pub message: Option<String>,
430}
431
432pub trait OAuth2AuditLogger: Send + Sync {
436 fn log_event(&self, event: &OAuth2AuditEvent);
438}
439
440pub trait OAuth2Provider: Send + Sync {
453 fn redirect_url(&self, state: &str) -> String;
464
465 fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error>;
475}
476
477pub trait OAuth2HttpTransport: Send + Sync {
493 fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error>;
504}
505
506#[derive(Debug, Default)]
530pub struct MemoryOAuth2HttpTransport {
531 requests: Mutex<Vec<(String, String)>>,
533 responses: Mutex<VecDeque<String>>,
535}
536
537impl MemoryOAuth2HttpTransport {
538 pub fn new() -> Self {
540 Self::default()
541 }
542
543 pub fn push_response(&self, response: impl Into<String>) {
545 self.responses.lock().push_back(response.into());
546 }
547
548 pub fn count(&self) -> usize {
550 self.requests.lock().len()
551 }
552
553 pub fn all(&self) -> Vec<(String, String)> {
555 self.requests.lock().clone()
556 }
557
558 pub fn last(&self) -> Option<(String, String)> {
560 self.requests.lock().last().cloned()
561 }
562
563 pub fn clear(&self) {
565 self.requests.lock().clear();
566 self.responses.lock().clear();
567 }
568}
569
570impl OAuth2HttpTransport for MemoryOAuth2HttpTransport {
571 fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error> {
572 self.requests
573 .lock()
574 .push((url.to_string(), body.to_string()));
575 let mut responses = self.responses.lock();
576 match responses.pop_front() {
577 Some(resp) => Ok(resp),
578 None => Ok(String::new()),
579 }
580 }
581}
582
583pub struct GenericOAuth2Provider {
624 config: OAuth2Config,
626 transport: Arc<dyn OAuth2HttpTransport>,
628 audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
630 #[cfg(feature = "redis-store")]
632 token_store: Option<Arc<dyn crate::oauth_store::OAuth2TokenStore>>,
633}
634
635impl GenericOAuth2Provider {
636 pub fn new(config: OAuth2Config, transport: Arc<dyn OAuth2HttpTransport>) -> Self {
643 Self {
644 config,
645 transport,
646 audit_logger: None,
647 #[cfg(feature = "redis-store")]
648 token_store: None,
649 }
650 }
651
652 pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
654 self.audit_logger = Some(logger);
655 self
656 }
657
658 #[cfg(feature = "redis-store")]
662 pub fn with_token_store(
663 mut self,
664 store: Arc<dyn crate::oauth_store::OAuth2TokenStore>,
665 ) -> Self {
666 self.token_store = Some(store);
667 self
668 }
669
670 fn log_audit(
672 &self,
673 grant_type: &str,
674 result: &str,
675 alert_code: Option<&str>,
676 message: Option<&str>,
677 ) {
678 if let Some(logger) = &self.audit_logger {
679 let event = OAuth2AuditEvent {
680 client_id: self.config.client_id.clone(),
681 grant_type: grant_type.to_string(),
682 result: result.to_string(),
683 timestamp: chrono::Utc::now().timestamp(),
684 alert_code: alert_code.map(|s| s.to_string()),
685 message: message.map(|s| s.to_string()),
686 };
687 logger.log_event(&event);
688 }
689 }
690
691 pub fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse, OAuth2Error> {
703 if refresh_token.is_empty() {
704 self.log_audit("refresh_token", "failure", None, Some("refresh_token 为空"));
705 return Err(OAuth2Error::AuthFailed("refresh_token 不能为空".into()));
706 }
707
708 let body = serde_json::json!({
709 "grant_type": "refresh_token",
710 "refresh_token": refresh_token,
711 "client_id": self.config.client_id,
712 "client_secret": self.config.client_secret,
713 });
714 let body_str =
715 serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
716
717 let response = self
718 .transport
719 .post_json(&self.config.token_url, &body_str)
720 .map_err(|err| {
721 self.log_audit("refresh_token", "failure", None, Some(&err.to_string()));
722 OAuth2Error::HttpTransport(err.to_string())
723 })?;
724
725 if response.is_empty() {
726 self.log_audit("refresh_token", "failure", None, Some("空响应"));
727 return Err(OAuth2Error::TokenExchangeFailed("token 响应为空".into()));
728 }
729
730 let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
731 self.log_audit(
732 "refresh_token",
733 "failure",
734 None,
735 Some(&format!("JSON 解析失败: {err}")),
736 );
737 OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
738 })?;
739
740 let access_token = json
741 .get("access_token")
742 .and_then(|v| v.as_str())
743 .ok_or_else(|| {
744 self.log_audit(
745 "refresh_token",
746 "failure",
747 None,
748 Some("响应缺少 access_token"),
749 );
750 OAuth2Error::TokenExchangeFailed("token 响应缺少 access_token 字段".into())
751 })?
752 .to_string();
753
754 let token_response = TokenResponse {
755 access_token,
756 token_type: json
757 .get("token_type")
758 .and_then(|v| v.as_str())
759 .map(|s| s.to_string()),
760 expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
761 scope: json
762 .get("scope")
763 .and_then(|v| v.as_str())
764 .map(|s| s.to_string()),
765 refresh_token: json
766 .get("refresh_token")
767 .and_then(|v| v.as_str())
768 .map(|s| s.to_string()),
769 };
770
771 self.log_audit("refresh_token", "success", None, None);
772 Ok(token_response)
773 }
774
775 fn build_redirect_url(&self, state: &str, pkce: Option<&PkceParams>) -> String {
788 let mut params: Vec<(String, String)> = vec![
789 ("client_id".into(), self.config.client_id.clone()),
790 ("redirect_uri".into(), self.config.redirect_url.clone()),
791 ("response_type".into(), "code".into()),
792 ("state".into(), state.to_string()),
793 ];
794
795 if !self.config.scopes.is_empty() {
796 params.push(("scope".into(), self.config.scopes.join(" ")));
797 }
798
799 if let Some(pkce) = pkce {
800 params.push(("code_challenge".into(), pkce.code_challenge.clone()));
801 params.push(("code_challenge_method".into(), pkce.method.to_string()));
802 }
803
804 for (key, value) in &self.config.extra_params {
805 params.push((key.clone(), value.clone()));
806 }
807
808 let query = params
809 .iter()
810 .map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
811 .collect::<Vec<_>>()
812 .join("&");
813
814 let separator = if self.config.auth_url.contains('?') {
815 "&"
816 } else {
817 "?"
818 };
819 format!("{}{}{}", self.config.auth_url, separator, query)
820 }
821
822 fn exchange_token(
838 &self,
839 code: &str,
840 code_verifier: Option<&str>,
841 ) -> Result<serde_json::Value, OAuth2Error> {
842 let mut body = serde_json::json!({
843 "grant_type": "authorization_code",
844 "code": code,
845 "client_id": self.config.client_id,
846 "client_secret": self.config.client_secret,
847 "redirect_uri": self.config.redirect_url,
848 });
849
850 if let Some(verifier) = code_verifier {
851 body["code_verifier"] = serde_json::Value::String(verifier.to_string());
852 }
853
854 let body_str =
855 serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
856
857 let response = self
858 .transport
859 .post_json(&self.config.token_url, &body_str)
860 .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
861
862 if response.is_empty() {
863 return Ok(serde_json::Value::Null);
864 }
865
866 serde_json::from_str(&response)
867 .map_err(|err| OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}")))
868 }
869
870 fn fetch_user_info(
875 &self,
876 access_token: &str,
877 token_json: &serde_json::Value,
878 ) -> Result<serde_json::Value, OAuth2Error> {
879 let user_url = self
880 .config
881 .user_url
882 .as_ref()
883 .ok_or_else(|| OAuth2Error::UserInfoFailed("user_url 未配置".into()))?;
884
885 let body = serde_json::json!({
886 "access_token": access_token,
887 "openid": token_json.get("openid").cloned().unwrap_or(serde_json::Value::Null),
888 });
889 let body_str =
890 serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
891
892 let response = self
893 .transport
894 .post_json(user_url, &body_str)
895 .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
896
897 if response.is_empty() {
898 return Ok(serde_json::Value::Null);
899 }
900
901 serde_json::from_str(&response)
902 .map_err(|err| OAuth2Error::UserInfoFailed(format!("解析用户信息响应失败: {err}")))
903 }
904
905 fn extract_user_fields(user_json: &serde_json::Value) -> SocialiteUser {
914 let id = user_json
915 .get("id")
916 .or_else(|| user_json.get("openid"))
917 .or_else(|| user_json.get("user_id"))
918 .and_then(extract_string)
919 .unwrap_or_default();
920
921 let nickname = user_json
922 .get("nickname")
923 .or_else(|| user_json.get("nick_name"))
924 .and_then(extract_string);
925
926 let name = user_json
927 .get("name")
928 .or_else(|| user_json.get("username"))
929 .and_then(extract_string);
930
931 let email = user_json.get("email").and_then(extract_string);
932
933 let avatar = user_json
934 .get("avatar")
935 .or_else(|| user_json.get("figureurl_qq_1"))
936 .or_else(|| user_json.get("figureurl"))
937 .or_else(|| user_json.get("headimgurl"))
938 .and_then(extract_string);
939
940 SocialiteUser {
941 id,
942 nickname,
943 name,
944 email,
945 avatar,
946 raw: user_json.clone(),
947 access_token: None,
948 refresh_token: None,
949 expires_in: None,
950 }
951 }
952}
953
954impl OAuth2Provider for GenericOAuth2Provider {
955 fn redirect_url(&self, state: &str) -> String {
956 self.build_redirect_url(state, None)
957 }
958
959 fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
960 self.user_from_token_with_pkce(code, None)
961 }
962}
963
964impl GenericOAuth2Provider {
965 pub fn redirect_url_with_pkce(&self, state: &str, pkce: &PkceParams) -> String {
969 self.build_redirect_url(state, Some(pkce))
970 }
971
972 pub fn user_from_token_with_pkce(
976 &self,
977 code: &str,
978 code_verifier: Option<&str>,
979 ) -> Result<SocialiteUser, OAuth2Error> {
980 self.config.validate()?;
982
983 if code.is_empty() {
985 return Err(OAuth2Error::AuthFailed("授权码不能为空".into()));
986 }
987
988 let token_json = self.exchange_token(code, code_verifier)?;
990
991 let access_token = token_json
993 .get("access_token")
994 .and_then(|value| value.as_str())
995 .ok_or_else(|| {
996 self.log_audit(
997 "authorization_code",
998 "failure",
999 None,
1000 Some("token 响应缺少 access_token"),
1001 );
1002 OAuth2Error::TokenExchangeFailed(format!(
1003 "token 响应缺少 access_token 字段: {token_json}"
1004 ))
1005 })?
1006 .to_string();
1007
1008 let refresh_token = token_json
1009 .get("refresh_token")
1010 .and_then(|value| value.as_str())
1011 .map(|value| value.to_string());
1012
1013 let expires_in = token_json
1014 .get("expires_in")
1015 .and_then(|value| value.as_i64());
1016
1017 let (access_token, refresh_token, expires_in) = if let Some(exp) = expires_in {
1019 if exp <= 0 {
1020 if let Some(ref rt) = refresh_token {
1021 match self.refresh_token(rt) {
1022 Ok(new_token) => {
1023 let new_access = new_token.access_token;
1024 let new_refresh = new_token.refresh_token.or(refresh_token.clone());
1025 let new_exp = new_token.expires_in;
1026 (new_access, new_refresh, new_exp)
1027 }
1028 Err(_) => {
1029 (access_token, refresh_token, expires_in)
1031 }
1032 }
1033 } else {
1034 (access_token, refresh_token, expires_in)
1035 }
1036 } else {
1037 (access_token, refresh_token, expires_in)
1038 }
1039 } else {
1040 (access_token, refresh_token, expires_in)
1041 };
1042
1043 let mut user = if self.config.user_url.is_some() {
1045 let user_json = self.fetch_user_info(&access_token, &token_json)?;
1046 Self::extract_user_fields(&user_json)
1047 } else {
1048 SocialiteUser::default()
1049 };
1050
1051 user.access_token = Some(access_token);
1052 user.refresh_token = refresh_token;
1053 user.expires_in = expires_in;
1054
1055 #[cfg(feature = "redis-store")]
1057 if let Some(store) = &self.token_store {
1058 let store = store.clone();
1059 let client_id = self.config.client_id.clone();
1060 let token_to_store = TokenResponse {
1061 access_token: user.access_token.clone().unwrap_or_default(),
1062 token_type: None,
1063 expires_in: user.expires_in,
1064 scope: None,
1065 refresh_token: user.refresh_token.clone(),
1066 };
1067 tokio::task::spawn(async move {
1068 if let Err(err) = store.store_token(&client_id, &token_to_store).await {
1069 tracing::warn!(
1070 error = %err,
1071 client_id = %client_id,
1072 "OAUTH2_TOKEN_STORE_FAILED: token 存储失败(best-effort,不影响主流程)"
1073 );
1074 }
1075 });
1076 }
1077
1078 self.log_audit("authorization_code", "success", None, None);
1079 Ok(user)
1080 }
1081}
1082
1083pub struct ImplicitOAuth2Provider {
1098 config: OAuth2Config,
1100 audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
1102}
1103
1104impl ImplicitOAuth2Provider {
1105 pub fn new(config: OAuth2Config) -> Self {
1107 Self {
1108 config,
1109 audit_logger: None,
1110 }
1111 }
1112
1113 pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
1115 self.audit_logger = Some(logger);
1116 self
1117 }
1118
1119 fn log_audit(
1121 &self,
1122 grant_type: &str,
1123 result: &str,
1124 alert_code: Option<&str>,
1125 message: Option<&str>,
1126 ) {
1127 if let Some(logger) = &self.audit_logger {
1128 let event = OAuth2AuditEvent {
1129 client_id: self.config.client_id.clone(),
1130 grant_type: grant_type.to_string(),
1131 result: result.to_string(),
1132 timestamp: chrono::Utc::now().timestamp(),
1133 alert_code: alert_code.map(|s| s.to_string()),
1134 message: message.map(|s| s.to_string()),
1135 };
1136 logger.log_event(&event);
1137 }
1138 }
1139
1140 pub fn redirect_url(&self, state: &str) -> String {
1151 self.redirect_url_with_pkce(state, None)
1152 }
1153
1154 pub fn redirect_url_with_pkce(&self, state: &str, pkce: Option<&PkceParams>) -> String {
1156 let mut params: Vec<(String, String)> = vec![
1157 ("client_id".into(), self.config.client_id.clone()),
1158 ("redirect_uri".into(), self.config.redirect_url.clone()),
1159 ("response_type".into(), "token".into()),
1160 ("state".into(), state.to_string()),
1161 ];
1162
1163 if !self.config.scopes.is_empty() {
1164 params.push(("scope".into(), self.config.scopes.join(" ")));
1165 }
1166
1167 if let Some(pkce) = pkce {
1168 params.push(("code_challenge".into(), pkce.code_challenge.clone()));
1169 params.push(("code_challenge_method".into(), pkce.method.to_string()));
1170 }
1171
1172 for (key, value) in &self.config.extra_params {
1173 params.push((key.clone(), value.clone()));
1174 }
1175
1176 let query = params
1177 .iter()
1178 .map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
1179 .collect::<Vec<_>>()
1180 .join("&");
1181
1182 let separator = if self.config.auth_url.contains('?') {
1183 "&"
1184 } else {
1185 "?"
1186 };
1187 format!("{}{}{}", self.config.auth_url, separator, query)
1188 }
1189
1190 pub fn parse_fragment(
1201 &self,
1202 fragment: &str,
1203 expected_state: &str,
1204 ) -> Result<TokenResponse, OAuth2Error> {
1205 self.log_audit(
1207 "implicit",
1208 "success",
1209 Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED"),
1210 Some("implicit 流程 token 经 URL fragment 暴露"),
1211 );
1212
1213 if fragment.is_empty() {
1214 self.log_audit("implicit", "failure", None, Some("fragment 为空"));
1215 return Err(OAuth2Error::TokenExchangeFailed("fragment 为空".into()));
1216 }
1217
1218 let params: std::collections::HashMap<&str, &str> = fragment
1220 .split('&')
1221 .filter_map(|pair| {
1222 let (key, value) = pair.split_once('=')?;
1223 Some((key, value))
1224 })
1225 .collect();
1226
1227 let state = params.get("state").copied().unwrap_or("");
1229 if state != expected_state {
1230 self.log_audit(
1231 "implicit",
1232 "failure",
1233 Some("OAUTH2_CSRF_STATE_MISMATCH"),
1234 Some(&format!(
1235 "state 不匹配: expected={expected_state}, actual={state}"
1236 )),
1237 );
1238 return Err(OAuth2Error::AuthFailed("CSRF state mismatch".into()));
1239 }
1240
1241 let access_token = params.get("access_token").copied().ok_or_else(|| {
1243 self.log_audit(
1244 "implicit",
1245 "failure",
1246 None,
1247 Some("fragment 无 access_token"),
1248 );
1249 OAuth2Error::TokenExchangeFailed("fragment 中缺少 access_token".into())
1250 })?;
1251
1252 Ok(TokenResponse {
1253 access_token: access_token.to_string(),
1254 token_type: params.get("token_type").map(|s| s.to_string()),
1255 expires_in: params.get("expires_in").and_then(|s| s.parse().ok()),
1256 scope: params.get("scope").map(|s| s.to_string()),
1257 refresh_token: None,
1259 })
1260 }
1261}
1262
1263impl OAuth2Provider for ImplicitOAuth2Provider {
1264 fn redirect_url(&self, state: &str) -> String {
1265 self.redirect_url(state)
1266 }
1267
1268 fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
1269 Ok(SocialiteUser {
1272 access_token: Some(code.to_string()),
1273 ..Default::default()
1274 })
1275 }
1276}
1277
1278#[cfg(feature = "device-code")]
1284pub mod device_code {
1285 use super::*;
1286 use async_trait::async_trait;
1287
1288 #[async_trait]
1292 pub trait AsyncOAuth2HttpTransport: Send + Sync {
1293 async fn post_form(
1304 &self,
1305 url: &str,
1306 params: &[(&str, &str)],
1307 ) -> Result<String, OAuth2Error>;
1308 }
1309
1310 #[derive(Debug, Clone)]
1312 pub struct DeviceCodeResponse {
1313 pub device_code: String,
1315 pub user_code: String,
1317 pub verification_uri: String,
1319 pub expires_in: i64,
1321 pub interval: i64,
1323 }
1324
1325 pub struct DeviceCodeOAuth2Provider {
1327 config: OAuth2Config,
1329 transport: Arc<dyn AsyncOAuth2HttpTransport>,
1331 audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
1333 }
1334
1335 impl DeviceCodeOAuth2Provider {
1336 pub fn new(config: OAuth2Config, transport: Arc<dyn AsyncOAuth2HttpTransport>) -> Self {
1338 Self {
1339 config,
1340 transport,
1341 audit_logger: None,
1342 }
1343 }
1344
1345 pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
1347 self.audit_logger = Some(logger);
1348 self
1349 }
1350
1351 fn log_audit(
1353 &self,
1354 grant_type: &str,
1355 result: &str,
1356 alert_code: Option<&str>,
1357 message: Option<&str>,
1358 ) {
1359 if let Some(logger) = &self.audit_logger {
1360 let event = OAuth2AuditEvent {
1361 client_id: self.config.client_id.clone(),
1362 grant_type: grant_type.to_string(),
1363 result: result.to_string(),
1364 timestamp: chrono::Utc::now().timestamp(),
1365 alert_code: alert_code.map(|s| s.to_string()),
1366 message: message.map(|s| s.to_string()),
1367 };
1368 logger.log_event(&event);
1369 }
1370 }
1371
1372 pub async fn request_device_code(
1378 &self,
1379 scope: &[String],
1380 ) -> Result<DeviceCodeResponse, OAuth2Error> {
1381 let device_auth_url = self.config.device_auth_url.as_ref().ok_or_else(|| {
1382 self.log_audit(
1383 "device_code",
1384 "failure",
1385 None,
1386 Some("device_auth_url 未配置"),
1387 );
1388 OAuth2Error::MissingField("device_auth_url".into())
1389 })?;
1390
1391 let scope_str = scope.join(" ");
1392 let params: Vec<(&str, &str)> = vec![
1393 ("client_id", self.config.client_id.as_str()),
1394 ("scope", scope_str.as_str()),
1395 ];
1396
1397 let response = self
1398 .transport
1399 .post_form(device_auth_url, ¶ms)
1400 .await
1401 .map_err(|err| {
1402 self.log_audit("device_code", "failure", None, Some(&err.to_string()));
1403 OAuth2Error::HttpTransport(err.to_string())
1404 })?;
1405
1406 let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
1407 self.log_audit(
1408 "device_code",
1409 "failure",
1410 None,
1411 Some(&format!("JSON 解析失败: {err}")),
1412 );
1413 OAuth2Error::TokenExchangeFailed(format!("解析 device code 响应失败: {err}"))
1414 })?;
1415
1416 let device_code = json
1417 .get("device_code")
1418 .and_then(|v| v.as_str())
1419 .ok_or_else(|| {
1420 OAuth2Error::TokenExchangeFailed("device code 响应缺少 device_code 字段".into())
1421 })?
1422 .to_string();
1423
1424 let user_code = json
1425 .get("user_code")
1426 .and_then(|v| v.as_str())
1427 .ok_or_else(|| {
1428 OAuth2Error::TokenExchangeFailed("device code 响应缺少 user_code 字段".into())
1429 })?
1430 .to_string();
1431
1432 let verification_uri = json
1433 .get("verification_uri")
1434 .and_then(|v| v.as_str())
1435 .ok_or_else(|| {
1436 OAuth2Error::TokenExchangeFailed(
1437 "device code 响应缺少 verification_uri 字段".into(),
1438 )
1439 })?
1440 .to_string();
1441
1442 let expires_in = json
1443 .get("expires_in")
1444 .and_then(|v| v.as_i64())
1445 .unwrap_or(600);
1446 let interval = json.get("interval").and_then(|v| v.as_i64()).unwrap_or(5);
1447
1448 self.log_audit("device_code", "success", None, None);
1449 Ok(DeviceCodeResponse {
1450 device_code,
1451 user_code,
1452 verification_uri,
1453 expires_in,
1454 interval,
1455 })
1456 }
1457
1458 pub async fn poll_for_token(
1474 &self,
1475 device_code: &str,
1476 mut interval: i64,
1477 expires_in: i64,
1478 ) -> Result<TokenResponse, OAuth2Error> {
1479 let start = std::time::Instant::now();
1480 let expires_duration = std::time::Duration::from_secs(expires_in.max(0) as u64);
1481
1482 loop {
1483 if start.elapsed() >= expires_duration {
1485 self.log_audit(
1486 "device_code",
1487 "failure",
1488 Some("OAUTH2_DEVICE_CODE_EXPIRED"),
1489 Some("设备码已过期"),
1490 );
1491 return Err(OAuth2Error::AuthFailed(
1492 "OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
1493 ));
1494 }
1495
1496 let params: Vec<(&str, &str)> = vec![
1498 ("grant_type", "device_code"),
1499 ("device_code", device_code),
1500 ("client_id", self.config.client_id.as_str()),
1501 ];
1502
1503 let response = self
1504 .transport
1505 .post_form(&self.config.token_url, ¶ms)
1506 .await
1507 .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
1508
1509 let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
1510 OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
1511 })?;
1512
1513 if let Some(access_token) = json.get("access_token").and_then(|v| v.as_str()) {
1515 let token_response = TokenResponse {
1516 access_token: access_token.to_string(),
1517 token_type: json
1518 .get("token_type")
1519 .and_then(|v| v.as_str())
1520 .map(|s| s.to_string()),
1521 expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
1522 scope: json
1523 .get("scope")
1524 .and_then(|v| v.as_str())
1525 .map(|s| s.to_string()),
1526 refresh_token: json
1527 .get("refresh_token")
1528 .and_then(|v| v.as_str())
1529 .map(|s| s.to_string()),
1530 };
1531 self.log_audit("device_code", "success", None, None);
1532 return Ok(token_response);
1533 }
1534
1535 let error = json.get("error").and_then(|v| v.as_str()).unwrap_or("");
1537
1538 match error {
1539 "authorization_pending" => {
1540 }
1542 "slow_down" => {
1543 interval = (interval + 5).min(60);
1545 }
1546 "access_denied" => {
1547 self.log_audit(
1548 "device_code",
1549 "failure",
1550 Some("OAUTH2_ACCESS_DENIED"),
1551 Some("用户拒绝授权"),
1552 );
1553 return Err(OAuth2Error::AuthFailed(
1554 "OAUTH2_ACCESS_DENIED: 用户拒绝授权".into(),
1555 ));
1556 }
1557 "expired_token" => {
1558 self.log_audit(
1559 "device_code",
1560 "failure",
1561 Some("OAUTH2_DEVICE_CODE_EXPIRED"),
1562 Some("设备码已过期"),
1563 );
1564 return Err(OAuth2Error::AuthFailed(
1565 "OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
1566 ));
1567 }
1568 _ => {
1569 return Err(OAuth2Error::TokenExchangeFailed(format!(
1570 "未知错误: {error}"
1571 )));
1572 }
1573 }
1574
1575 if interval > 0 {
1577 tokio::time::sleep(std::time::Duration::from_secs(interval as u64)).await;
1578 }
1579 }
1580 }
1581 }
1582
1583 #[derive(Default)]
1589 pub struct MemoryAsyncOAuth2HttpTransport {
1590 requests: Mutex<Vec<(String, String)>>,
1591 responses: Mutex<VecDeque<String>>,
1592 }
1593
1594 impl MemoryAsyncOAuth2HttpTransport {
1595 pub fn new() -> Self {
1597 Self::default()
1598 }
1599
1600 pub fn push_response(&self, response: impl Into<String>) {
1602 self.responses.lock().push_back(response.into());
1603 }
1604
1605 pub fn count(&self) -> usize {
1607 self.requests.lock().len()
1608 }
1609 }
1610
1611 #[async_trait]
1612 impl AsyncOAuth2HttpTransport for MemoryAsyncOAuth2HttpTransport {
1613 async fn post_form(
1614 &self,
1615 url: &str,
1616 params: &[(&str, &str)],
1617 ) -> Result<String, OAuth2Error> {
1618 let body = params
1619 .iter()
1620 .map(|(k, v)| format!("{k}={v}"))
1621 .collect::<Vec<_>>()
1622 .join("&");
1623 self.requests.lock().push((url.to_string(), body));
1624 let mut responses = self.responses.lock();
1625 match responses.pop_front() {
1626 Some(resp) => Ok(resp),
1627 None => Ok(String::new()),
1628 }
1629 }
1630 }
1631
1632 #[cfg(test)]
1637 mod tests {
1638 use super::*;
1639
1640 #[tokio::test]
1642 async fn test_device_code_request() {
1643 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1644 transport.push_response(
1645 r#"{"device_code":"dc123","user_code":"UC-ABCD","verification_uri":"https://provider.com/device","expires_in":600,"interval":5}"#,
1646 );
1647
1648 let config = OAuth2Config::new(
1649 "client123",
1650 "secret456",
1651 "https://example.com/callback",
1652 "https://provider.com/authorize",
1653 "https://provider.com/token",
1654 )
1655 .with_device_auth_url("https://provider.com/device_authorize");
1656 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1657
1658 let resp = provider
1659 .request_device_code(&["read".into(), "write".into()])
1660 .await
1661 .expect("request_device_code 失败");
1662
1663 assert_eq!(resp.device_code, "dc123");
1664 assert_eq!(resp.user_code, "UC-ABCD");
1665 assert_eq!(resp.verification_uri, "https://provider.com/device");
1666 assert_eq!(resp.expires_in, 600);
1667 assert_eq!(resp.interval, 5);
1668 }
1669
1670 #[tokio::test]
1672 async fn test_device_code_poll_pending() {
1673 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1674 transport.push_response(r#"{"error":"authorization_pending"}"#);
1676 transport.push_response(
1678 r#"{"access_token":"token123","token_type":"Bearer","expires_in":3600}"#,
1679 );
1680
1681 let config = OAuth2Config::new(
1682 "client123",
1683 "secret456",
1684 "https://example.com/callback",
1685 "https://provider.com/authorize",
1686 "https://provider.com/token",
1687 );
1688 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1689
1690 let token = provider
1691 .poll_for_token("dc123", 0, 600)
1692 .await
1693 .expect("poll_for_token 失败");
1694
1695 assert_eq!(token.access_token, "token123");
1696 assert_eq!(token.token_type.as_deref(), Some("Bearer"));
1697 }
1698
1699 #[tokio::test]
1701 async fn test_device_code_poll_slow_down() {
1702 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1703 transport.push_response(r#"{"error":"slow_down"}"#);
1705 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
1707
1708 let config = OAuth2Config::new(
1709 "client123",
1710 "secret456",
1711 "https://example.com/callback",
1712 "https://provider.com/authorize",
1713 "https://provider.com/token",
1714 );
1715 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1716
1717 let token = provider
1718 .poll_for_token("dc123", 0, 600)
1719 .await
1720 .expect("poll_for_token 失败");
1721
1722 assert_eq!(token.access_token, "token123");
1723 }
1724
1725 #[tokio::test]
1727 async fn test_device_code_expired() {
1728 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1729 transport.push_response(r#"{"error":"authorization_pending"}"#);
1731
1732 let config = OAuth2Config::new(
1733 "client123",
1734 "secret456",
1735 "https://example.com/callback",
1736 "https://provider.com/authorize",
1737 "https://provider.com/token",
1738 );
1739 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1740
1741 let err = provider.poll_for_token("dc123", 0, 0).await.unwrap_err();
1742 assert!(
1743 err.to_string().contains("OAUTH2_DEVICE_CODE_EXPIRED"),
1744 "应返回设备码过期错误: {err}"
1745 );
1746 }
1747
1748 #[tokio::test]
1750 async fn test_device_code_access_denied() {
1751 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1752 transport.push_response(r#"{"error":"access_denied"}"#);
1753
1754 let config = OAuth2Config::new(
1755 "client123",
1756 "secret456",
1757 "https://example.com/callback",
1758 "https://provider.com/authorize",
1759 "https://provider.com/token",
1760 );
1761 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1762
1763 let err = provider.poll_for_token("dc123", 5, 600).await.unwrap_err();
1764 assert!(
1765 err.to_string().contains("OAUTH2_ACCESS_DENIED"),
1766 "应返回 access_denied 错误: {err}"
1767 );
1768 }
1769
1770 #[tokio::test]
1772 async fn test_device_code_no_auth_url() {
1773 let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1774 let config = OAuth2Config::new(
1775 "client123",
1776 "secret456",
1777 "https://example.com/callback",
1778 "https://provider.com/authorize",
1779 "https://provider.com/token",
1780 );
1781 let provider = DeviceCodeOAuth2Provider::new(config, transport);
1783
1784 let err = provider.request_device_code(&[]).await.unwrap_err();
1785 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "device_auth_url"));
1786 }
1787 }
1788}
1789
1790fn extract_string(value: &serde_json::Value) -> Option<String> {
1798 match value {
1799 serde_json::Value::String(string) => Some(string.clone()),
1800 serde_json::Value::Number(number) => number.as_i64().map(|number| number.to_string()),
1801 _ => None,
1802 }
1803}
1804
1805fn percent_encode(input: &str) -> String {
1810 let mut output = String::with_capacity(input.len());
1811 for byte in input.as_bytes() {
1812 if matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~') {
1813 output.push(*byte as char);
1814 } else {
1815 output.push_str(&format!("%{byte:02X}"));
1816 }
1817 }
1818 output
1819}
1820
1821#[cfg(test)]
1826mod tests {
1827 use super::*;
1828 use proptest::{prop_assert, prop_assert_eq};
1829
1830 #[test]
1836 fn test_oauth2_config_builder() {
1837 let config = OAuth2Config::new(
1838 "client123",
1839 "secret456",
1840 "https://example.com/callback",
1841 "https://provider.com/authorize",
1842 "https://provider.com/token",
1843 )
1844 .with_user_url("https://provider.com/user/info")
1845 .with_scopes(vec!["scope1".into(), "scope2".into()])
1846 .with_extra_param("foo", "bar");
1847
1848 assert_eq!(config.client_id, "client123");
1849 assert_eq!(config.client_secret, "secret456");
1850 assert_eq!(config.redirect_url, "https://example.com/callback");
1851 assert_eq!(config.auth_url, "https://provider.com/authorize");
1852 assert_eq!(config.token_url, "https://provider.com/token");
1853 assert_eq!(
1854 config.user_url.as_deref(),
1855 Some("https://provider.com/user/info")
1856 );
1857 assert_eq!(config.scopes, vec!["scope1", "scope2"]);
1858 assert_eq!(config.extra_params, vec![("foo".into(), "bar".into())]);
1859 }
1860
1861 #[test]
1863 fn test_oauth2_config_minimal() {
1864 let config = OAuth2Config::new(
1865 "client123",
1866 "secret456",
1867 "https://example.com/callback",
1868 "https://provider.com/authorize",
1869 "https://provider.com/token",
1870 );
1871
1872 assert_eq!(config.client_id, "client123");
1873 assert_eq!(config.client_secret, "secret456");
1874 assert_eq!(config.redirect_url, "https://example.com/callback");
1875 assert_eq!(config.auth_url, "https://provider.com/authorize");
1876 assert_eq!(config.token_url, "https://provider.com/token");
1877 assert!(config.user_url.is_none());
1878 assert!(config.scopes.is_empty());
1879 assert!(config.extra_params.is_empty());
1880
1881 assert!(config.validate().is_ok());
1883 }
1884
1885 #[test]
1887 fn test_oauth2_config_with_scope_chained() {
1888 let config = OAuth2Config::new(
1889 "id",
1890 "secret",
1891 "https://example.com/callback",
1892 "https://provider.com/authorize",
1893 "https://provider.com/token",
1894 )
1895 .with_scope("get_user_info")
1896 .with_scope("get_unionid");
1897
1898 assert_eq!(config.scopes, vec!["get_user_info", "get_unionid"]);
1899 }
1900
1901 #[test]
1903 fn test_oauth2_config_with_extra_params() {
1904 let config = OAuth2Config::new(
1905 "id",
1906 "secret",
1907 "https://example.com/callback",
1908 "https://provider.com/authorize",
1909 "https://provider.com/token",
1910 )
1911 .with_extra_param("a", "1")
1912 .with_extra_param("b", "2")
1913 .with_extra_params(vec![("x".into(), "10".into())]);
1914
1915 assert_eq!(config.extra_params, vec![("x".into(), "10".into())]);
1916 }
1917
1918 #[test]
1920 fn test_oauth2_config_validate_empty_fields() {
1921 let config = OAuth2Config::new(
1923 "",
1924 "secret",
1925 "https://example.com/callback",
1926 "https://provider.com/authorize",
1927 "https://provider.com/token",
1928 );
1929 let err = config.validate().unwrap_err();
1930 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
1931
1932 let config = OAuth2Config::new(
1934 "id",
1935 "",
1936 "https://example.com/callback",
1937 "https://provider.com/authorize",
1938 "https://provider.com/token",
1939 );
1940 let err = config.validate().unwrap_err();
1941 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_secret"));
1942
1943 let config = OAuth2Config::new(
1945 "id",
1946 "secret",
1947 "",
1948 "https://provider.com/authorize",
1949 "https://provider.com/token",
1950 );
1951 let err = config.validate().unwrap_err();
1952 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "redirect_url"));
1953
1954 let config = OAuth2Config::new(
1956 "id",
1957 "secret",
1958 "https://example.com/callback",
1959 "",
1960 "https://provider.com/token",
1961 );
1962 let err = config.validate().unwrap_err();
1963 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "auth_url"));
1964
1965 let config = OAuth2Config::new(
1967 "id",
1968 "secret",
1969 "https://example.com/callback",
1970 "https://provider.com/authorize",
1971 "",
1972 );
1973 let err = config.validate().unwrap_err();
1974 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
1975 }
1976
1977 #[test]
1983 fn test_socialite_user_default() {
1984 let user = SocialiteUser::default();
1985 assert!(user.id.is_empty());
1986 assert!(user.nickname.is_none());
1987 assert!(user.name.is_none());
1988 assert!(user.email.is_none());
1989 assert!(user.avatar.is_none());
1990 assert!(user.raw.is_null());
1991 assert!(user.access_token.is_none());
1992 assert!(user.refresh_token.is_none());
1993 assert!(user.expires_in.is_none());
1994 }
1995
1996 #[test]
1998 fn test_socialite_user_serialize_deserialize() {
1999 let user = SocialiteUser {
2000 id: "123".into(),
2001 nickname: Some("tester".into()),
2002 name: Some("Test User".into()),
2003 email: Some("test@example.com".into()),
2004 avatar: Some("https://example.com/avatar.png".into()),
2005 raw: serde_json::json!({"key": "value"}),
2006 access_token: Some("token123".into()),
2007 refresh_token: Some("refresh456".into()),
2008 expires_in: Some(3600),
2009 };
2010
2011 let json = serde_json::to_string(&user).expect("序列化失败");
2012
2013 assert!(
2016 !json.contains("access_token"),
2017 "access_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
2018 );
2019 assert!(
2020 !json.contains("refresh_token"),
2021 "refresh_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
2022 );
2023
2024 let parsed: SocialiteUser = serde_json::from_str(&json).expect("反序列化失败");
2025
2026 assert_eq!(parsed.id, "123");
2027 assert_eq!(parsed.nickname.as_deref(), Some("tester"));
2028 assert_eq!(parsed.name.as_deref(), Some("Test User"));
2029 assert_eq!(parsed.email.as_deref(), Some("test@example.com"));
2030 assert_eq!(
2031 parsed.avatar.as_deref(),
2032 Some("https://example.com/avatar.png")
2033 );
2034 assert_eq!(parsed.access_token, None);
2036 assert_eq!(parsed.refresh_token, None);
2037 assert_eq!(parsed.expires_in, Some(3600));
2038 }
2039
2040 #[test]
2046 fn test_redirect_url_contains_required_params() {
2047 let config = OAuth2Config::new(
2048 "client123",
2049 "secret456",
2050 "https://example.com/callback",
2051 "https://provider.com/oauth2.0/authorize",
2052 "https://provider.com/oauth2.0/token",
2053 );
2054 let provider =
2055 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2056
2057 let url = provider.redirect_url("random_state_abc");
2058
2059 assert!(url.starts_with("https://provider.com/oauth2.0/authorize?"));
2060 assert!(url.contains("client_id=client123"));
2061 assert!(url.contains("redirect_uri=https%3A%2F%2Fexample.com%2Fcallback"));
2062 assert!(url.contains("response_type=code"));
2063 assert!(url.contains("state=random_state_abc"));
2064 assert!(!url.contains("scope="));
2066 }
2067
2068 #[test]
2070 fn test_redirect_url_with_scopes() {
2071 let config = OAuth2Config::new(
2072 "client123",
2073 "secret456",
2074 "https://example.com/callback",
2075 "https://provider.com/oauth2.0/authorize",
2076 "https://provider.com/oauth2.0/token",
2077 )
2078 .with_scopes(vec!["get_user_info".into(), "get_unionid".into()]);
2079 let provider =
2080 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2081
2082 let url = provider.redirect_url("state123");
2083
2084 assert!(url.contains("scope=get_user_info%20get_unionid"));
2086 }
2087
2088 #[test]
2090 fn test_redirect_url_with_extra_params() {
2091 let config = OAuth2Config::new(
2092 "client123",
2093 "secret456",
2094 "https://example.com/callback",
2095 "https://provider.com/oauth2.0/authorize",
2096 "https://provider.com/oauth2.0/token",
2097 )
2098 .with_extra_param("foo", "bar")
2099 .with_extra_param("display", "mobile");
2100 let provider =
2101 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2102
2103 let url = provider.redirect_url("state123");
2104
2105 assert!(url.contains("foo=bar"));
2106 assert!(url.contains("display=mobile"));
2107 }
2108
2109 #[test]
2111 fn test_redirect_url_with_existing_query() {
2112 let config = OAuth2Config::new(
2113 "client123",
2114 "secret456",
2115 "https://example.com/callback",
2116 "https://provider.com/authorize?foo=bar",
2117 "https://provider.com/token",
2118 );
2119 let provider =
2120 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2121
2122 let url = provider.redirect_url("state123");
2123
2124 assert!(url.contains("?foo=bar&"));
2126 assert!(url.contains("client_id=client123"));
2127 }
2128
2129 #[test]
2135 fn test_memory_oauth2_http_transport_post_json() {
2136 let transport = MemoryOAuth2HttpTransport::new();
2137 transport.push_response(r#"{"access_token":"token123"}"#);
2138
2139 let response = transport
2140 .post_json("https://example.com/token", r#"{"code":"abc"}"#)
2141 .expect("post_json 失败");
2142
2143 assert_eq!(response, r#"{"access_token":"token123"}"#);
2144 assert_eq!(transport.count(), 1);
2145
2146 let (url, body) = transport.last().expect("应有请求记录");
2147 assert_eq!(url, "https://example.com/token");
2148 assert_eq!(body, r#"{"code":"abc"}"#);
2149 }
2150
2151 #[test]
2153 fn test_memory_oauth2_http_transport_response_queue() {
2154 let transport = MemoryOAuth2HttpTransport::new();
2155 transport.push_response("resp1");
2156 transport.push_response("resp2");
2157
2158 let resp1 = transport
2159 .post_json("url1", "body1")
2160 .expect("第一次调用失败");
2161 let resp2 = transport
2162 .post_json("url2", "body2")
2163 .expect("第二次调用失败");
2164
2165 assert_eq!(resp1, "resp1");
2166 assert_eq!(resp2, "resp2");
2167 assert_eq!(transport.count(), 2);
2168 }
2169
2170 #[test]
2172 fn test_memory_oauth2_http_transport_empty_response() {
2173 let transport = MemoryOAuth2HttpTransport::new();
2174 let response = transport
2176 .post_json("url", "body")
2177 .expect("post_json 不应失败");
2178 assert_eq!(response, "");
2179 }
2180
2181 #[test]
2183 fn test_memory_oauth2_http_transport_clear() {
2184 let transport = MemoryOAuth2HttpTransport::new();
2185 transport.push_response("resp");
2186 transport.post_json("url", "body").expect("调用失败");
2187 assert_eq!(transport.count(), 1);
2188
2189 transport.clear();
2190 assert_eq!(transport.count(), 0);
2191 let response = transport
2193 .post_json("url", "body")
2194 .expect("post_json 不应失败");
2195 assert_eq!(response, "");
2196 }
2197
2198 #[test]
2206 fn test_generic_oauth2_provider_user_from_token() {
2207 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2208 transport.push_response(r#"{"access_token":"token123","refresh_token":"refresh456","expires_in":3600,"openid":"openid_abc"}"#);
2210 transport.push_response(
2211 r#"{"id":"12345","nickname":"test_user","name":"Test","email":"test@example.com","avatar":"https://example.com/avatar.png"}"#,
2212 );
2213
2214 let config = OAuth2Config::new(
2215 "client123",
2216 "secret456",
2217 "https://example.com/callback",
2218 "https://provider.com/authorize",
2219 "https://provider.com/token",
2220 )
2221 .with_user_url("https://provider.com/user/info");
2222 let provider = GenericOAuth2Provider::new(config, transport.clone());
2223
2224 let user = provider
2225 .user_from_token("auth_code_abc")
2226 .expect("user_from_token 失败");
2227
2228 assert_eq!(user.access_token.as_deref(), Some("token123"));
2230 assert_eq!(user.refresh_token.as_deref(), Some("refresh456"));
2231 assert_eq!(user.expires_in, Some(3600));
2232
2233 assert_eq!(user.id, "12345");
2235 assert_eq!(user.nickname.as_deref(), Some("test_user"));
2236 assert_eq!(user.name.as_deref(), Some("Test"));
2237 assert_eq!(user.email.as_deref(), Some("test@example.com"));
2238 assert_eq!(
2239 user.avatar.as_deref(),
2240 Some("https://example.com/avatar.png")
2241 );
2242
2243 assert_eq!(user.raw["id"], "12345");
2245 assert_eq!(user.raw["nickname"], "test_user");
2246
2247 assert_eq!(transport.count(), 2);
2249 }
2250
2251 #[test]
2253 fn test_generic_oauth2_provider_user_from_token_no_user_url() {
2254 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2255 transport.push_response(r#"{"access_token":"token123","expires_in":7200}"#);
2256
2257 let config = OAuth2Config::new(
2258 "client123",
2259 "secret456",
2260 "https://example.com/callback",
2261 "https://provider.com/authorize",
2262 "https://provider.com/token",
2263 );
2264 let provider = GenericOAuth2Provider::new(config, transport.clone());
2266
2267 let user = provider
2268 .user_from_token("auth_code")
2269 .expect("user_from_token 失败");
2270
2271 assert_eq!(user.access_token.as_deref(), Some("token123"));
2272 assert_eq!(user.expires_in, Some(7200));
2273 assert!(user.refresh_token.is_none());
2274 assert!(user.id.is_empty());
2276 assert_eq!(transport.count(), 1);
2278 }
2279
2280 #[test]
2282 fn test_generic_oauth2_provider_missing_code() {
2283 let config = OAuth2Config::new(
2284 "client123",
2285 "secret456",
2286 "https://example.com/callback",
2287 "https://provider.com/authorize",
2288 "https://provider.com/token",
2289 );
2290 let provider =
2291 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2292
2293 let err = provider.user_from_token("").unwrap_err();
2294 assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("授权码")));
2295 }
2296
2297 #[test]
2299 fn test_oauth2_provider_missing_config_fields() {
2300 let config = OAuth2Config::new(
2302 "",
2303 "secret456",
2304 "https://example.com/callback",
2305 "https://provider.com/authorize",
2306 "https://provider.com/token",
2307 );
2308 let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
2309
2310 let err = provider.user_from_token("code").unwrap_err();
2311 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
2312
2313 let config = OAuth2Config::new(
2315 "client123",
2316 "secret456",
2317 "https://example.com/callback",
2318 "https://provider.com/authorize",
2319 "",
2320 );
2321 let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
2322
2323 let err = provider.user_from_token("code").unwrap_err();
2324 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
2325 }
2326
2327 #[test]
2329 fn test_generic_oauth2_provider_token_response_missing_access_token() {
2330 let transport = MemoryOAuth2HttpTransport::new();
2331 transport.push_response(r#"{"error":"invalid_grant"}"#);
2332
2333 let config = OAuth2Config::new(
2334 "client123",
2335 "secret456",
2336 "https://example.com/callback",
2337 "https://provider.com/authorize",
2338 "https://provider.com/token",
2339 );
2340 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2341
2342 let err = provider.user_from_token("code").unwrap_err();
2343 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2344 }
2345
2346 #[test]
2348 fn test_generic_oauth2_provider_token_response_invalid_json() {
2349 let transport = MemoryOAuth2HttpTransport::new();
2350 transport.push_response("not a json");
2351
2352 let config = OAuth2Config::new(
2353 "client123",
2354 "secret456",
2355 "https://example.com/callback",
2356 "https://provider.com/authorize",
2357 "https://provider.com/token",
2358 );
2359 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2360
2361 let err = provider.user_from_token("code").unwrap_err();
2362 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2363 }
2364
2365 #[test]
2367 fn test_generic_oauth2_provider_user_info_field_aliases() {
2368 let transport = MemoryOAuth2HttpTransport::new();
2369 transport.push_response(r#"{"access_token":"token123","openid":"openid_abc"}"#);
2370 transport.push_response(
2371 r#"{"openid":"qq_12345","nickname":"qq_user","figureurl_qq_1":"https://qzapp.qlogo.cn/1.png"}"#,
2372 );
2373
2374 let config = OAuth2Config::new(
2375 "client123",
2376 "secret456",
2377 "https://example.com/callback",
2378 "https://provider.com/authorize",
2379 "https://provider.com/token",
2380 )
2381 .with_user_url("https://provider.com/user/info");
2382 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2383
2384 let user = provider
2385 .user_from_token("code")
2386 .expect("user_from_token 失败");
2387
2388 assert_eq!(user.id, "qq_12345");
2390 assert_eq!(user.nickname.as_deref(), Some("qq_user"));
2391 assert_eq!(user.avatar.as_deref(), Some("https://qzapp.qlogo.cn/1.png"));
2393 }
2394
2395 #[test]
2397 fn test_generic_oauth2_provider_user_id_integer() {
2398 let transport = MemoryOAuth2HttpTransport::new();
2399 transport.push_response(r#"{"access_token":"token123"}"#);
2400 transport.push_response(r#"{"id":12345,"nickname":"github_user"}"#);
2401
2402 let config = OAuth2Config::new(
2403 "client123",
2404 "secret456",
2405 "https://example.com/callback",
2406 "https://provider.com/authorize",
2407 "https://provider.com/token",
2408 )
2409 .with_user_url("https://provider.com/user/info");
2410 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2411
2412 let user = provider
2413 .user_from_token("code")
2414 .expect("user_from_token 失败");
2415
2416 assert_eq!(user.id, "12345");
2417 assert_eq!(user.nickname.as_deref(), Some("github_user"));
2418 }
2419
2420 #[test]
2422 fn test_generic_oauth2_provider_http_transport_failure() {
2423 let config = OAuth2Config::new(
2424 "client123",
2425 "secret456",
2426 "https://example.com/callback",
2427 "https://provider.com/authorize",
2428 "https://provider.com/token",
2429 );
2430 let provider = GenericOAuth2Provider::new(config, Arc::new(FailingTransport));
2431
2432 let err = provider.user_from_token("code").unwrap_err();
2433 assert!(matches!(err, OAuth2Error::HttpTransport(_)));
2434 }
2435
2436 #[test]
2438 fn test_percent_encode() {
2439 assert_eq!(percent_encode("abcXYZ09-._~"), "abcXYZ09-._~");
2441 assert_eq!(percent_encode("a b"), "a%20b");
2443 assert_eq!(percent_encode("/"), "%2F");
2445 assert_eq!(percent_encode(":"), "%3A");
2447 assert_eq!(
2449 percent_encode("https://example.com/path"),
2450 "https%3A%2F%2Fexample.com%2Fpath"
2451 );
2452 assert_eq!(percent_encode("中"), "%E4%B8%AD");
2454 }
2455
2456 #[test]
2458 fn test_extract_string() {
2459 assert_eq!(
2461 extract_string(&serde_json::json!("hello")),
2462 Some("hello".into())
2463 );
2464 assert_eq!(
2466 extract_string(&serde_json::json!(12345)),
2467 Some("12345".into())
2468 );
2469 assert_eq!(extract_string(&serde_json::json!(1.5)), None);
2471 assert_eq!(extract_string(&serde_json::json!(true)), None);
2473 assert_eq!(extract_string(&serde_json::Value::Null), None);
2475 assert_eq!(extract_string(&serde_json::json!({"a": 1})), None);
2477 }
2478
2479 struct FailingTransport;
2485
2486 impl OAuth2HttpTransport for FailingTransport {
2487 fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
2488 Err(OAuth2Error::HttpTransport("connection refused".into()))
2489 }
2490 }
2491
2492 struct MemoryHttpTransport;
2494
2495 impl OAuth2HttpTransport for MemoryHttpTransport {
2496 fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
2497 Ok(String::new())
2498 }
2499 }
2500
2501 #[test]
2507 fn test_state_auto_generate() {
2508 let state1 = OAuth2Config::generate_state();
2509 let state2 = OAuth2Config::generate_state();
2510
2511 assert_eq!(state1.len(), 32, "state 应为 32 字符 hex(16 字节)");
2513 assert_eq!(state2.len(), 32, "state 应为 32 字符 hex(16 字节)");
2514
2515 assert!(
2517 state1.chars().all(|c| c.is_ascii_hexdigit()),
2518 "state 应全为 hex 字符: {state1}"
2519 );
2520 assert!(
2521 state2.chars().all(|c| c.is_ascii_hexdigit()),
2522 "state 应全为 hex 字符: {state2}"
2523 );
2524
2525 assert_ne!(state1, state2, "两次生成的 state 不应相同");
2527 }
2528
2529 #[test]
2531 fn test_pkce_pair_generate() {
2532 let pkce = OAuth2Config::generate_pkce_pair();
2533
2534 assert!(
2536 pkce.code_verifier.len() >= 43 && pkce.code_verifier.len() <= 128,
2537 "code_verifier 长度应在 43-128 之间,实际: {}",
2538 pkce.code_verifier.len()
2539 );
2540 assert_eq!(
2541 pkce.code_verifier.len(),
2542 64,
2543 "code_verifier 应为 64 字符 hex"
2544 );
2545 assert!(
2546 pkce.code_verifier.chars().all(|c| c.is_ascii_hexdigit()),
2547 "code_verifier 应全为 hex 字符"
2548 );
2549
2550 let mut hasher = sha2::Sha256::new();
2552 sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
2553 let digest = sha2::Digest::finalize(hasher);
2554 let expected_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
2555 assert_eq!(
2556 pkce.code_challenge, expected_challenge,
2557 "code_challenge 应等于 base64url(SHA256(code_verifier))"
2558 );
2559
2560 assert_eq!(pkce.method, PkceMethod::S256);
2562 }
2563
2564 #[test]
2566 fn test_authorization_code_with_pkce() {
2567 let config = OAuth2Config::new(
2568 "client123",
2569 "secret456",
2570 "https://example.com/callback",
2571 "https://provider.com/oauth2.0/authorize",
2572 "https://provider.com/oauth2.0/token",
2573 )
2574 .with_pkce(true);
2575 let provider =
2576 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2577
2578 let pkce = OAuth2Config::generate_pkce_pair();
2579 let url = provider.redirect_url_with_pkce("state123", &pkce);
2580
2581 assert!(url.contains("code_challenge="), "URL 应包含 code_challenge");
2582 assert!(
2583 url.contains("code_challenge_method=S256"),
2584 "URL 应包含 code_challenge_method=S256"
2585 );
2586 assert!(url.contains("response_type=code"));
2587 assert!(url.contains("state=state123"));
2588 }
2589
2590 #[test]
2592 fn test_authorization_code_without_pkce() {
2593 let config = OAuth2Config::new(
2594 "client123",
2595 "secret456",
2596 "https://example.com/callback",
2597 "https://provider.com/oauth2.0/authorize",
2598 "https://provider.com/oauth2.0/token",
2599 );
2600 assert!(!config.pkce_enabled);
2602
2603 let provider =
2604 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2605 let url = provider.redirect_url("state123");
2606
2607 assert!(
2608 !url.contains("code_challenge"),
2609 "URL 不应包含 code_challenge(PKCE 未启用)"
2610 );
2611 }
2612
2613 #[test]
2615 fn test_client_secret_not_in_debug() {
2616 let config = OAuth2Config::new(
2617 "client123",
2618 "super_secret_value_456",
2619 "https://example.com/callback",
2620 "https://provider.com/authorize",
2621 "https://provider.com/token",
2622 );
2623
2624 let debug_str = format!("{:?}", config);
2625 assert!(
2626 !debug_str.contains("super_secret_value_456"),
2627 "client_secret 不应出现在 Debug 输出中: {debug_str}"
2628 );
2629 assert!(
2630 debug_str.contains("***"),
2631 "Debug 输出应包含脱敏标记 '***': {debug_str}"
2632 );
2633 }
2634
2635 #[test]
2637 fn test_pkce_params_debug_redacted() {
2638 let pkce = OAuth2Config::generate_pkce_pair();
2639 let debug_str = format!("{:?}", pkce);
2640 assert!(
2641 !debug_str.contains(&pkce.code_verifier),
2642 "code_verifier 明文不应出现在 Debug 输出中: {debug_str}"
2643 );
2644 assert!(
2645 debug_str.contains("redacted"),
2646 "Debug 输出应包含 'redacted' 标记: {debug_str}"
2647 );
2648 }
2649
2650 #[test]
2652 fn test_with_pkce_builder() {
2653 let config = OAuth2Config::new(
2654 "id",
2655 "secret",
2656 "https://example.com/callback",
2657 "https://provider.com/authorize",
2658 "https://provider.com/token",
2659 );
2660 assert!(!config.pkce_enabled, "默认 pkce_enabled 应为 false");
2661
2662 let config = config.with_pkce(true);
2663 assert!(
2664 config.pkce_enabled,
2665 "with_pkce(true) 后 pkce_enabled 应为 true"
2666 );
2667
2668 let config = config.with_pkce(false);
2669 assert!(
2670 !config.pkce_enabled,
2671 "with_pkce(false) 后 pkce_enabled 应为 false"
2672 );
2673 }
2674
2675 #[test]
2677 fn test_with_device_auth_url_builder() {
2678 let config = OAuth2Config::new(
2679 "id",
2680 "secret",
2681 "https://example.com/callback",
2682 "https://provider.com/authorize",
2683 "https://provider.com/token",
2684 );
2685 assert!(
2686 config.device_auth_url.is_none(),
2687 "默认 device_auth_url 应为 None"
2688 );
2689
2690 let config = config.with_device_auth_url("https://provider.com/device_authorize");
2691 assert_eq!(
2692 config.device_auth_url.as_deref(),
2693 Some("https://provider.com/device_authorize"),
2694 );
2695 }
2696
2697 #[test]
2699 fn test_exchange_token_with_pkce_verifier() {
2700 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2701 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2702
2703 let config = OAuth2Config::new(
2704 "client123",
2705 "secret456",
2706 "https://example.com/callback",
2707 "https://provider.com/authorize",
2708 "https://provider.com/token",
2709 )
2710 .with_pkce(true);
2711 let provider = GenericOAuth2Provider::new(config, transport.clone());
2712
2713 let pkce = OAuth2Config::generate_pkce_pair();
2714 let user = provider
2715 .user_from_token_with_pkce("auth_code", Some(&pkce.code_verifier))
2716 .expect("user_from_token_with_pkce 失败");
2717
2718 assert_eq!(user.access_token.as_deref(), Some("token123"));
2719
2720 let (_url, body) = transport.last().expect("应有请求记录");
2722 assert!(
2723 body.contains("code_verifier"),
2724 "token 交换 body 应包含 code_verifier: {body}"
2725 );
2726 assert!(
2727 body.contains(&pkce.code_verifier),
2728 "token 交换 body 应包含 code_verifier 值"
2729 );
2730 }
2731
2732 #[test]
2734 fn test_exchange_token_without_pkce_verifier() {
2735 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2736 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2737
2738 let config = OAuth2Config::new(
2739 "client123",
2740 "secret456",
2741 "https://example.com/callback",
2742 "https://provider.com/authorize",
2743 "https://provider.com/token",
2744 );
2745 let provider = GenericOAuth2Provider::new(config, transport.clone());
2746
2747 let _user = provider
2748 .user_from_token("auth_code")
2749 .expect("user_from_token 失败");
2750
2751 let (_url, body) = transport.last().expect("应有请求记录");
2752 assert!(
2753 !body.contains("code_verifier"),
2754 "token 交换 body 不应包含 code_verifier(PKCE 未启用): {body}"
2755 );
2756 }
2757
2758 #[derive(Default)]
2764 struct MockAuditLogger {
2765 events: Mutex<Vec<OAuth2AuditEvent>>,
2766 }
2767
2768 impl MockAuditLogger {
2769 fn events(&self) -> Vec<OAuth2AuditEvent> {
2770 self.events.lock().clone()
2771 }
2772 }
2773
2774 impl OAuth2AuditLogger for MockAuditLogger {
2775 fn log_event(&self, event: &OAuth2AuditEvent) {
2776 self.events.lock().push(event.clone());
2777 }
2778 }
2779
2780 #[test]
2782 fn test_refresh_token() {
2783 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2784 transport.push_response(
2785 r#"{"access_token":"new_token","token_type":"Bearer","expires_in":7200,"scope":"read"}"#,
2786 );
2787
2788 let config = OAuth2Config::new(
2789 "client123",
2790 "secret456",
2791 "https://example.com/callback",
2792 "https://provider.com/authorize",
2793 "https://provider.com/token",
2794 );
2795 let provider = GenericOAuth2Provider::new(config, transport.clone());
2796
2797 let token_resp = provider
2798 .refresh_token("old_refresh_token")
2799 .expect("refresh_token 失败");
2800
2801 assert_eq!(token_resp.access_token, "new_token");
2802 assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
2803 assert_eq!(token_resp.expires_in, Some(7200));
2804 assert_eq!(token_resp.scope.as_deref(), Some("read"));
2805 }
2806
2807 #[test]
2809 fn test_audit_log_on_token_exchange() {
2810 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2811 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2812
2813 let logger = Arc::new(MockAuditLogger::default());
2814 let config = OAuth2Config::new(
2815 "client123",
2816 "secret456",
2817 "https://example.com/callback",
2818 "https://provider.com/authorize",
2819 "https://provider.com/token",
2820 );
2821 let provider =
2822 GenericOAuth2Provider::new(config, transport).with_audit_logger(logger.clone());
2823
2824 let _user = provider
2825 .user_from_token("auth_code")
2826 .expect("user_from_token 失败");
2827
2828 let events = logger.events();
2829 assert!(
2830 events.iter().any(|e| e.grant_type == "authorization_code"
2831 && e.result == "success"
2832 && e.client_id == "client123"),
2833 "应记录 authorization_code success 事件: {events:?}"
2834 );
2835 }
2836
2837 #[test]
2839 fn test_auto_refresh_on_expired() {
2840 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2841 transport.push_response(
2843 r#"{"access_token":"expired_token","refresh_token":"valid_refresh","expires_in":0}"#,
2844 );
2845 transport.push_response(r#"{"access_token":"refreshed_token","expires_in":3600}"#);
2847
2848 let config = OAuth2Config::new(
2849 "client123",
2850 "secret456",
2851 "https://example.com/callback",
2852 "https://provider.com/authorize",
2853 "https://provider.com/token",
2854 );
2855 let provider = GenericOAuth2Provider::new(config, transport.clone());
2856
2857 let user = provider
2858 .user_from_token("auth_code")
2859 .expect("user_from_token 失败");
2860
2861 assert_eq!(
2863 user.access_token.as_deref(),
2864 Some("refreshed_token"),
2865 "过期 token 应自动刷新"
2866 );
2867 assert_eq!(user.expires_in, Some(3600));
2868 assert_eq!(transport.count(), 2);
2870 }
2871
2872 #[test]
2874 fn test_refresh_token_empty() {
2875 let config = OAuth2Config::new(
2876 "client123",
2877 "secret456",
2878 "https://example.com/callback",
2879 "https://provider.com/authorize",
2880 "https://provider.com/token",
2881 );
2882 let provider =
2883 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2884
2885 let err = provider.refresh_token("").unwrap_err();
2886 assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("refresh_token")));
2887 }
2888
2889 #[test]
2891 fn test_refresh_token_invalid_json() {
2892 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2893 transport.push_response("not a json");
2894
2895 let config = OAuth2Config::new(
2896 "client123",
2897 "secret456",
2898 "https://example.com/callback",
2899 "https://provider.com/authorize",
2900 "https://provider.com/token",
2901 );
2902 let provider = GenericOAuth2Provider::new(config, transport);
2903
2904 let err = provider.refresh_token("valid_refresh").unwrap_err();
2905 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2906 }
2907
2908 #[test]
2910 fn test_token_response_serialize_redacted() {
2911 let token_resp = TokenResponse {
2912 access_token: "secret_access_token".into(),
2913 token_type: Some("Bearer".into()),
2914 expires_in: Some(3600),
2915 scope: Some("read".into()),
2916 refresh_token: Some("secret_refresh_token".into()),
2917 };
2918
2919 let json = serde_json::to_string(&token_resp).expect("序列化失败");
2920 assert!(
2921 !json.contains("secret_access_token"),
2922 "access_token 不应出现在序列化 JSON 中: {json}"
2923 );
2924 assert!(
2925 !json.contains("secret_refresh_token"),
2926 "refresh_token 不应出现在序列化 JSON 中: {json}"
2927 );
2928 }
2929
2930 #[test]
2936 fn test_implicit_redirect_url() {
2937 let config = OAuth2Config::new(
2938 "client123",
2939 "secret456",
2940 "https://example.com/callback",
2941 "https://provider.com/oauth2.0/authorize",
2942 "https://provider.com/oauth2.0/token",
2943 )
2944 .with_scope("profile");
2945 let provider = ImplicitOAuth2Provider::new(config);
2946
2947 let url = provider.redirect_url("state_abc");
2948
2949 assert!(
2950 url.contains("response_type=token"),
2951 "URL 应含 response_type=token"
2952 );
2953 assert!(url.contains("client_id=client123"));
2954 assert!(url.contains("state=state_abc"));
2955 assert!(url.contains("scope=profile"));
2956 }
2957
2958 #[test]
2960 fn test_implicit_parse_fragment() {
2961 let config = OAuth2Config::new(
2962 "client123",
2963 "secret456",
2964 "https://example.com/callback",
2965 "https://provider.com/authorize",
2966 "https://provider.com/token",
2967 );
2968 let provider = ImplicitOAuth2Provider::new(config);
2969
2970 let fragment = "access_token=token123&token_type=Bearer&expires_in=3600&state=mystate";
2971 let token_resp = provider
2972 .parse_fragment(fragment, "mystate")
2973 .expect("parse_fragment 失败");
2974
2975 assert_eq!(token_resp.access_token, "token123");
2976 assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
2977 assert_eq!(token_resp.expires_in, Some(3600));
2978 assert!(
2980 token_resp.refresh_token.is_none(),
2981 "implicit 流程 refresh_token 应为 None"
2982 );
2983 }
2984
2985 #[test]
2987 fn test_implicit_state_mismatch() {
2988 let config = OAuth2Config::new(
2989 "client123",
2990 "secret456",
2991 "https://example.com/callback",
2992 "https://provider.com/authorize",
2993 "https://provider.com/token",
2994 );
2995 let provider = ImplicitOAuth2Provider::new(config);
2996
2997 let fragment = "access_token=token123&state=wrong_state";
2998 let err = provider
2999 .parse_fragment(fragment, "expected_state")
3000 .unwrap_err();
3001 assert!(
3002 matches!(&err, OAuth2Error::AuthFailed(msg) if msg.contains("CSRF state mismatch")),
3003 "state 不匹配应返回 CSRF 错误: {err}"
3004 );
3005 }
3006
3007 #[test]
3009 fn test_implicit_no_refresh_token() {
3010 let config = OAuth2Config::new(
3011 "client123",
3012 "secret456",
3013 "https://example.com/callback",
3014 "https://provider.com/authorize",
3015 "https://provider.com/token",
3016 );
3017 let provider = ImplicitOAuth2Provider::new(config);
3018
3019 let fragment = "access_token=token123&refresh_token=should_be_ignored&state=mystate";
3021 let token_resp = provider
3022 .parse_fragment(fragment, "mystate")
3023 .expect("parse_fragment 失败");
3024
3025 assert!(
3026 token_resp.refresh_token.is_none(),
3027 "implicit 流程 refresh_token 应固定为 None"
3028 );
3029 }
3030
3031 #[test]
3033 fn test_implicit_empty_fragment() {
3034 let config = OAuth2Config::new(
3035 "client123",
3036 "secret456",
3037 "https://example.com/callback",
3038 "https://provider.com/authorize",
3039 "https://provider.com/token",
3040 );
3041 let provider = ImplicitOAuth2Provider::new(config);
3042
3043 let err = provider.parse_fragment("", "state").unwrap_err();
3044 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
3045 }
3046
3047 #[test]
3049 fn test_implicit_no_access_token() {
3050 let config = OAuth2Config::new(
3051 "client123",
3052 "secret456",
3053 "https://example.com/callback",
3054 "https://provider.com/authorize",
3055 "https://provider.com/token",
3056 );
3057 let provider = ImplicitOAuth2Provider::new(config);
3058
3059 let fragment = "token_type=Bearer&state=mystate";
3060 let err = provider.parse_fragment(fragment, "mystate").unwrap_err();
3061 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
3062 }
3063
3064 #[test]
3066 fn test_implicit_empty_scopes() {
3067 let config = OAuth2Config::new(
3068 "client123",
3069 "secret456",
3070 "https://example.com/callback",
3071 "https://provider.com/authorize",
3072 "https://provider.com/token",
3073 );
3074 let provider = ImplicitOAuth2Provider::new(config);
3075
3076 let url = provider.redirect_url("state123");
3077 assert!(
3078 !url.contains("scope="),
3079 "空 scopes 时 URL 不应含 scope 参数"
3080 );
3081 }
3082
3083 #[test]
3085 fn test_implicit_audit_log_token_exposed() {
3086 let config = OAuth2Config::new(
3087 "client123",
3088 "secret456",
3089 "https://example.com/callback",
3090 "https://provider.com/authorize",
3091 "https://provider.com/token",
3092 );
3093 let logger = Arc::new(MockAuditLogger::default());
3094 let provider = ImplicitOAuth2Provider::new(config).with_audit_logger(logger.clone());
3095
3096 let fragment = "access_token=token123&state=mystate";
3097 let _ = provider.parse_fragment(fragment, "mystate");
3098
3099 let events = logger.events();
3100 assert!(
3101 events
3102 .iter()
3103 .any(|e| e.alert_code.as_deref() == Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED")),
3104 "应记录 OAUTH2_IMPLICIT_TOKEN_EXPOSED 告警: {events:?}"
3105 );
3106 }
3107
3108 #[cfg(feature = "redis-store")]
3114 #[tokio::test]
3115 async fn test_token_store_integration() {
3116 use crate::oauth_store::{MemoryOAuth2TokenStore, OAuth2TokenStore};
3117
3118 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
3119 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
3120
3121 let store = Arc::new(MemoryOAuth2TokenStore::new());
3122 let config = OAuth2Config::new(
3123 "client123",
3124 "secret456",
3125 "https://example.com/callback",
3126 "https://provider.com/authorize",
3127 "https://provider.com/token",
3128 );
3129 let provider =
3130 GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
3131
3132 let user = provider
3133 .user_from_token("auth_code")
3134 .expect("user_from_token 失败");
3135 assert_eq!(user.access_token.as_deref(), Some("token123"));
3136
3137 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
3139
3140 let stored = store
3141 .get_token("client123")
3142 .await
3143 .expect("get_token 失败")
3144 .expect("应查到存储的 token");
3145 assert_eq!(stored.access_token, "token123");
3146 }
3147
3148 #[cfg(feature = "redis-store")]
3150 #[tokio::test]
3151 async fn test_token_store_failure_best_effort() {
3152 use crate::oauth_store::MemoryOAuth2TokenStore;
3153
3154 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
3155 transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
3156
3157 let store = Arc::new(MemoryOAuth2TokenStore::new());
3159 let config = OAuth2Config::new(
3160 "client123",
3161 "secret456",
3162 "https://example.com/callback",
3163 "https://provider.com/authorize",
3164 "https://provider.com/token",
3165 );
3166 let provider =
3167 GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
3168
3169 let user = provider
3171 .user_from_token("auth_code")
3172 .expect("user_from_token 应成功(best-effort)");
3173 assert_eq!(user.access_token.as_deref(), Some("token123"));
3174 }
3175
3176 proptest::proptest! {
3182 #[test]
3183 fn proptest_state_unpredictable(_n in 0u32..1000) {
3184 let s1 = OAuth2Config::generate_state();
3185 let s2 = OAuth2Config::generate_state();
3186 prop_assert_eq!(s1.len(), 32);
3187 prop_assert_eq!(s2.len(), 32);
3188 prop_assert!(s1.chars().all(|c| c.is_ascii_hexdigit()));
3189 prop_assert!(s2.chars().all(|c| c.is_ascii_hexdigit()));
3190 }
3191 }
3192
3193 proptest::proptest! {
3195 #[test]
3196 fn proptest_pkce_verifier_length(_n in 0u32..1000) {
3197 let pkce = OAuth2Config::generate_pkce_pair();
3198 prop_assert!(pkce.code_verifier.len() >= 43);
3199 prop_assert!(pkce.code_verifier.len() <= 128);
3200 }
3201 }
3202
3203 proptest::proptest! {
3205 #[test]
3206 fn proptest_pkce_challenge_matches_verifier(_n in 0u32..1000) {
3207 let pkce = OAuth2Config::generate_pkce_pair();
3208 let mut hasher = sha2::Sha256::new();
3209 sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
3210 let digest = sha2::Digest::finalize(hasher);
3211 let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
3212 prop_assert_eq!(pkce.code_challenge, expected);
3213 }
3214 }
3215}