1use parking_lot::Mutex;
44use std::collections::VecDeque;
45use std::sync::Arc;
46use thiserror::Error;
47
48#[derive(Debug, Error)]
54pub enum OAuth2Error {
55 #[error("OAuth2 字段缺失: {0}")]
57 MissingField(String),
58 #[error("OAuth2 授权失败: {0}")]
60 AuthFailed(String),
61 #[error("OAuth2 token 交换失败: {0}")]
63 TokenExchangeFailed(String),
64 #[error("OAuth2 获取用户信息失败: {0}")]
66 UserInfoFailed(String),
67 #[error("OAuth2 HTTP 传输失败: {0}")]
69 HttpTransport(String),
70 #[error("OAuth2 序列化失败: {0}")]
72 Serialize(String),
73}
74
75#[derive(Debug, Clone)]
111pub struct OAuth2Config {
112 pub client_id: String,
114 pub client_secret: String,
116 pub redirect_url: String,
118 pub auth_url: String,
120 pub token_url: String,
122 pub user_url: Option<String>,
124 pub scopes: Vec<String>,
126 pub extra_params: Vec<(String, String)>,
128}
129
130impl OAuth2Config {
131 #[allow(clippy::too_many_arguments)]
141 pub fn new(
142 client_id: impl Into<String>,
143 client_secret: impl Into<String>,
144 redirect_url: impl Into<String>,
145 auth_url: impl Into<String>,
146 token_url: impl Into<String>,
147 ) -> Self {
148 Self {
149 client_id: client_id.into(),
150 client_secret: client_secret.into(),
151 redirect_url: redirect_url.into(),
152 auth_url: auth_url.into(),
153 token_url: token_url.into(),
154 user_url: None,
155 scopes: Vec::new(),
156 extra_params: Vec::new(),
157 }
158 }
159
160 pub fn with_user_url(mut self, user_url: impl Into<String>) -> Self {
162 self.user_url = Some(user_url.into());
163 self
164 }
165
166 pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
168 self.scopes = scopes;
169 self
170 }
171
172 pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
174 self.scopes.push(scope.into());
175 self
176 }
177
178 pub fn with_extra_params(mut self, params: Vec<(String, String)>) -> Self {
180 self.extra_params = params;
181 self
182 }
183
184 pub fn with_extra_param(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
186 self.extra_params.push((key.into(), value.into()));
187 self
188 }
189
190 pub fn validate(&self) -> Result<(), OAuth2Error> {
195 if self.client_id.is_empty() {
196 return Err(OAuth2Error::MissingField("client_id".into()));
197 }
198 if self.client_secret.is_empty() {
199 return Err(OAuth2Error::MissingField("client_secret".into()));
200 }
201 if self.redirect_url.is_empty() {
202 return Err(OAuth2Error::MissingField("redirect_url".into()));
203 }
204 if self.auth_url.is_empty() {
205 return Err(OAuth2Error::MissingField("auth_url".into()));
206 }
207 if self.token_url.is_empty() {
208 return Err(OAuth2Error::MissingField("token_url".into()));
209 }
210 Ok(())
211 }
212}
213
214#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
237pub struct SocialiteUser {
238 pub id: String,
240 pub nickname: Option<String>,
242 pub name: Option<String>,
244 pub email: Option<String>,
246 pub avatar: Option<String>,
248 pub raw: serde_json::Value,
250 #[serde(skip_serializing)]
252 pub access_token: Option<String>,
253 #[serde(skip_serializing)]
255 pub refresh_token: Option<String>,
256 pub expires_in: Option<i64>,
258}
259
260pub trait OAuth2Provider: Send + Sync {
273 fn redirect_url(&self, state: &str) -> String;
284
285 fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error>;
295}
296
297pub trait OAuth2HttpTransport: Send + Sync {
313 fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error>;
324}
325
326#[derive(Debug, Default)]
350pub struct MemoryOAuth2HttpTransport {
351 requests: Mutex<Vec<(String, String)>>,
353 responses: Mutex<VecDeque<String>>,
355}
356
357impl MemoryOAuth2HttpTransport {
358 pub fn new() -> Self {
360 Self::default()
361 }
362
363 pub fn push_response(&self, response: impl Into<String>) {
365 self.responses.lock().push_back(response.into());
366 }
367
368 pub fn count(&self) -> usize {
370 self.requests.lock().len()
371 }
372
373 pub fn all(&self) -> Vec<(String, String)> {
375 self.requests.lock().clone()
376 }
377
378 pub fn last(&self) -> Option<(String, String)> {
380 self.requests.lock().last().cloned()
381 }
382
383 pub fn clear(&self) {
385 self.requests.lock().clear();
386 self.responses.lock().clear();
387 }
388}
389
390impl OAuth2HttpTransport for MemoryOAuth2HttpTransport {
391 fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error> {
392 self.requests
393 .lock()
394 .push((url.to_string(), body.to_string()));
395 let mut responses = self.responses.lock();
396 match responses.pop_front() {
397 Some(resp) => Ok(resp),
398 None => Ok(String::new()),
399 }
400 }
401}
402
403pub struct GenericOAuth2Provider {
444 config: OAuth2Config,
446 transport: Arc<dyn OAuth2HttpTransport>,
448}
449
450impl GenericOAuth2Provider {
451 pub fn new(config: OAuth2Config, transport: Arc<dyn OAuth2HttpTransport>) -> Self {
458 Self { config, transport }
459 }
460
461 fn build_redirect_url(&self, state: &str) -> String {
473 let mut params: Vec<(String, String)> = vec![
474 ("client_id".into(), self.config.client_id.clone()),
475 ("redirect_uri".into(), self.config.redirect_url.clone()),
476 ("response_type".into(), "code".into()),
477 ("state".into(), state.to_string()),
478 ];
479
480 if !self.config.scopes.is_empty() {
481 params.push(("scope".into(), self.config.scopes.join(" ")));
482 }
483
484 for (key, value) in &self.config.extra_params {
485 params.push((key.clone(), value.clone()));
486 }
487
488 let query = params
489 .iter()
490 .map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
491 .collect::<Vec<_>>()
492 .join("&");
493
494 let separator = if self.config.auth_url.contains('?') {
495 "&"
496 } else {
497 "?"
498 };
499 format!("{}{}{}", self.config.auth_url, separator, query)
500 }
501
502 fn exchange_token(&self, code: &str) -> Result<serde_json::Value, OAuth2Error> {
517 let body = serde_json::json!({
518 "grant_type": "authorization_code",
519 "code": code,
520 "client_id": self.config.client_id,
521 "client_secret": self.config.client_secret,
522 "redirect_uri": self.config.redirect_url,
523 });
524 let body_str =
525 serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
526
527 let response = self
528 .transport
529 .post_json(&self.config.token_url, &body_str)
530 .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
531
532 if response.is_empty() {
533 return Ok(serde_json::Value::Null);
534 }
535
536 serde_json::from_str(&response)
537 .map_err(|err| OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}")))
538 }
539
540 fn fetch_user_info(
545 &self,
546 access_token: &str,
547 token_json: &serde_json::Value,
548 ) -> Result<serde_json::Value, OAuth2Error> {
549 let user_url = self
550 .config
551 .user_url
552 .as_ref()
553 .ok_or_else(|| OAuth2Error::UserInfoFailed("user_url 未配置".into()))?;
554
555 let body = serde_json::json!({
556 "access_token": access_token,
557 "openid": token_json.get("openid").cloned().unwrap_or(serde_json::Value::Null),
558 });
559 let body_str =
560 serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
561
562 let response = self
563 .transport
564 .post_json(user_url, &body_str)
565 .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
566
567 if response.is_empty() {
568 return Ok(serde_json::Value::Null);
569 }
570
571 serde_json::from_str(&response)
572 .map_err(|err| OAuth2Error::UserInfoFailed(format!("解析用户信息响应失败: {err}")))
573 }
574
575 fn extract_user_fields(user_json: &serde_json::Value) -> SocialiteUser {
584 let id = user_json
585 .get("id")
586 .or_else(|| user_json.get("openid"))
587 .or_else(|| user_json.get("user_id"))
588 .and_then(extract_string)
589 .unwrap_or_default();
590
591 let nickname = user_json
592 .get("nickname")
593 .or_else(|| user_json.get("nick_name"))
594 .and_then(extract_string);
595
596 let name = user_json
597 .get("name")
598 .or_else(|| user_json.get("username"))
599 .and_then(extract_string);
600
601 let email = user_json.get("email").and_then(extract_string);
602
603 let avatar = user_json
604 .get("avatar")
605 .or_else(|| user_json.get("figureurl_qq_1"))
606 .or_else(|| user_json.get("figureurl"))
607 .or_else(|| user_json.get("headimgurl"))
608 .and_then(extract_string);
609
610 SocialiteUser {
611 id,
612 nickname,
613 name,
614 email,
615 avatar,
616 raw: user_json.clone(),
617 access_token: None,
618 refresh_token: None,
619 expires_in: None,
620 }
621 }
622}
623
624impl OAuth2Provider for GenericOAuth2Provider {
625 fn redirect_url(&self, state: &str) -> String {
626 self.build_redirect_url(state)
627 }
628
629 fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
630 self.config.validate()?;
632
633 if code.is_empty() {
635 return Err(OAuth2Error::AuthFailed("授权码不能为空".into()));
636 }
637
638 let token_json = self.exchange_token(code)?;
640
641 let access_token = token_json
643 .get("access_token")
644 .and_then(|value| value.as_str())
645 .ok_or_else(|| {
646 OAuth2Error::TokenExchangeFailed(format!(
647 "token 响应缺少 access_token 字段: {token_json}"
648 ))
649 })?
650 .to_string();
651
652 let refresh_token = token_json
653 .get("refresh_token")
654 .and_then(|value| value.as_str())
655 .map(|value| value.to_string());
656
657 let expires_in = token_json
658 .get("expires_in")
659 .and_then(|value| value.as_i64());
660
661 let mut user = if self.config.user_url.is_some() {
663 let user_json = self.fetch_user_info(&access_token, &token_json)?;
664 Self::extract_user_fields(&user_json)
665 } else {
666 SocialiteUser::default()
667 };
668
669 user.access_token = Some(access_token);
670 user.refresh_token = refresh_token;
671 user.expires_in = expires_in;
672
673 Ok(user)
674 }
675}
676
677fn extract_string(value: &serde_json::Value) -> Option<String> {
685 match value {
686 serde_json::Value::String(string) => Some(string.clone()),
687 serde_json::Value::Number(number) => number.as_i64().map(|number| number.to_string()),
688 _ => None,
689 }
690}
691
692fn percent_encode(input: &str) -> String {
697 let mut output = String::with_capacity(input.len());
698 for byte in input.as_bytes() {
699 if matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~') {
700 output.push(*byte as char);
701 } else {
702 output.push_str(&format!("%{byte:02X}"));
703 }
704 }
705 output
706}
707
708#[cfg(test)]
713mod tests {
714 use super::*;
715
716 #[test]
722 fn test_oauth2_config_builder() {
723 let config = OAuth2Config::new(
724 "client123",
725 "secret456",
726 "https://example.com/callback",
727 "https://provider.com/authorize",
728 "https://provider.com/token",
729 )
730 .with_user_url("https://provider.com/user/info")
731 .with_scopes(vec!["scope1".into(), "scope2".into()])
732 .with_extra_param("foo", "bar");
733
734 assert_eq!(config.client_id, "client123");
735 assert_eq!(config.client_secret, "secret456");
736 assert_eq!(config.redirect_url, "https://example.com/callback");
737 assert_eq!(config.auth_url, "https://provider.com/authorize");
738 assert_eq!(config.token_url, "https://provider.com/token");
739 assert_eq!(
740 config.user_url.as_deref(),
741 Some("https://provider.com/user/info")
742 );
743 assert_eq!(config.scopes, vec!["scope1", "scope2"]);
744 assert_eq!(config.extra_params, vec![("foo".into(), "bar".into())]);
745 }
746
747 #[test]
749 fn test_oauth2_config_minimal() {
750 let config = OAuth2Config::new(
751 "client123",
752 "secret456",
753 "https://example.com/callback",
754 "https://provider.com/authorize",
755 "https://provider.com/token",
756 );
757
758 assert_eq!(config.client_id, "client123");
759 assert_eq!(config.client_secret, "secret456");
760 assert_eq!(config.redirect_url, "https://example.com/callback");
761 assert_eq!(config.auth_url, "https://provider.com/authorize");
762 assert_eq!(config.token_url, "https://provider.com/token");
763 assert!(config.user_url.is_none());
764 assert!(config.scopes.is_empty());
765 assert!(config.extra_params.is_empty());
766
767 assert!(config.validate().is_ok());
769 }
770
771 #[test]
773 fn test_oauth2_config_with_scope_chained() {
774 let config = OAuth2Config::new(
775 "id",
776 "secret",
777 "https://example.com/callback",
778 "https://provider.com/authorize",
779 "https://provider.com/token",
780 )
781 .with_scope("get_user_info")
782 .with_scope("get_unionid");
783
784 assert_eq!(config.scopes, vec!["get_user_info", "get_unionid"]);
785 }
786
787 #[test]
789 fn test_oauth2_config_with_extra_params() {
790 let config = OAuth2Config::new(
791 "id",
792 "secret",
793 "https://example.com/callback",
794 "https://provider.com/authorize",
795 "https://provider.com/token",
796 )
797 .with_extra_param("a", "1")
798 .with_extra_param("b", "2")
799 .with_extra_params(vec![("x".into(), "10".into())]);
800
801 assert_eq!(config.extra_params, vec![("x".into(), "10".into())]);
802 }
803
804 #[test]
806 fn test_oauth2_config_validate_empty_fields() {
807 let config = OAuth2Config::new(
809 "",
810 "secret",
811 "https://example.com/callback",
812 "https://provider.com/authorize",
813 "https://provider.com/token",
814 );
815 let err = config.validate().unwrap_err();
816 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
817
818 let config = OAuth2Config::new(
820 "id",
821 "",
822 "https://example.com/callback",
823 "https://provider.com/authorize",
824 "https://provider.com/token",
825 );
826 let err = config.validate().unwrap_err();
827 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_secret"));
828
829 let config = OAuth2Config::new(
831 "id",
832 "secret",
833 "",
834 "https://provider.com/authorize",
835 "https://provider.com/token",
836 );
837 let err = config.validate().unwrap_err();
838 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "redirect_url"));
839
840 let config = OAuth2Config::new(
842 "id",
843 "secret",
844 "https://example.com/callback",
845 "",
846 "https://provider.com/token",
847 );
848 let err = config.validate().unwrap_err();
849 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "auth_url"));
850
851 let config = OAuth2Config::new(
853 "id",
854 "secret",
855 "https://example.com/callback",
856 "https://provider.com/authorize",
857 "",
858 );
859 let err = config.validate().unwrap_err();
860 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
861 }
862
863 #[test]
869 fn test_socialite_user_default() {
870 let user = SocialiteUser::default();
871 assert!(user.id.is_empty());
872 assert!(user.nickname.is_none());
873 assert!(user.name.is_none());
874 assert!(user.email.is_none());
875 assert!(user.avatar.is_none());
876 assert!(user.raw.is_null());
877 assert!(user.access_token.is_none());
878 assert!(user.refresh_token.is_none());
879 assert!(user.expires_in.is_none());
880 }
881
882 #[test]
884 fn test_socialite_user_serialize_deserialize() {
885 let user = SocialiteUser {
886 id: "123".into(),
887 nickname: Some("tester".into()),
888 name: Some("Test User".into()),
889 email: Some("test@example.com".into()),
890 avatar: Some("https://example.com/avatar.png".into()),
891 raw: serde_json::json!({"key": "value"}),
892 access_token: Some("token123".into()),
893 refresh_token: Some("refresh456".into()),
894 expires_in: Some(3600),
895 };
896
897 let json = serde_json::to_string(&user).expect("序列化失败");
898
899 assert!(
902 !json.contains("access_token"),
903 "access_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
904 );
905 assert!(
906 !json.contains("refresh_token"),
907 "refresh_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
908 );
909
910 let parsed: SocialiteUser = serde_json::from_str(&json).expect("反序列化失败");
911
912 assert_eq!(parsed.id, "123");
913 assert_eq!(parsed.nickname.as_deref(), Some("tester"));
914 assert_eq!(parsed.name.as_deref(), Some("Test User"));
915 assert_eq!(parsed.email.as_deref(), Some("test@example.com"));
916 assert_eq!(
917 parsed.avatar.as_deref(),
918 Some("https://example.com/avatar.png")
919 );
920 assert_eq!(parsed.access_token, None);
922 assert_eq!(parsed.refresh_token, None);
923 assert_eq!(parsed.expires_in, Some(3600));
924 }
925
926 #[test]
932 fn test_redirect_url_contains_required_params() {
933 let config = OAuth2Config::new(
934 "client123",
935 "secret456",
936 "https://example.com/callback",
937 "https://provider.com/oauth2.0/authorize",
938 "https://provider.com/oauth2.0/token",
939 );
940 let provider =
941 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
942
943 let url = provider.redirect_url("random_state_abc");
944
945 assert!(url.starts_with("https://provider.com/oauth2.0/authorize?"));
946 assert!(url.contains("client_id=client123"));
947 assert!(url.contains("redirect_uri=https%3A%2F%2Fexample.com%2Fcallback"));
948 assert!(url.contains("response_type=code"));
949 assert!(url.contains("state=random_state_abc"));
950 assert!(!url.contains("scope="));
952 }
953
954 #[test]
956 fn test_redirect_url_with_scopes() {
957 let config = OAuth2Config::new(
958 "client123",
959 "secret456",
960 "https://example.com/callback",
961 "https://provider.com/oauth2.0/authorize",
962 "https://provider.com/oauth2.0/token",
963 )
964 .with_scopes(vec!["get_user_info".into(), "get_unionid".into()]);
965 let provider =
966 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
967
968 let url = provider.redirect_url("state123");
969
970 assert!(url.contains("scope=get_user_info%20get_unionid"));
972 }
973
974 #[test]
976 fn test_redirect_url_with_extra_params() {
977 let config = OAuth2Config::new(
978 "client123",
979 "secret456",
980 "https://example.com/callback",
981 "https://provider.com/oauth2.0/authorize",
982 "https://provider.com/oauth2.0/token",
983 )
984 .with_extra_param("foo", "bar")
985 .with_extra_param("display", "mobile");
986 let provider =
987 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
988
989 let url = provider.redirect_url("state123");
990
991 assert!(url.contains("foo=bar"));
992 assert!(url.contains("display=mobile"));
993 }
994
995 #[test]
997 fn test_redirect_url_with_existing_query() {
998 let config = OAuth2Config::new(
999 "client123",
1000 "secret456",
1001 "https://example.com/callback",
1002 "https://provider.com/authorize?foo=bar",
1003 "https://provider.com/token",
1004 );
1005 let provider =
1006 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
1007
1008 let url = provider.redirect_url("state123");
1009
1010 assert!(url.contains("?foo=bar&"));
1012 assert!(url.contains("client_id=client123"));
1013 }
1014
1015 #[test]
1021 fn test_memory_oauth2_http_transport_post_json() {
1022 let transport = MemoryOAuth2HttpTransport::new();
1023 transport.push_response(r#"{"access_token":"token123"}"#);
1024
1025 let response = transport
1026 .post_json("https://example.com/token", r#"{"code":"abc"}"#)
1027 .expect("post_json 失败");
1028
1029 assert_eq!(response, r#"{"access_token":"token123"}"#);
1030 assert_eq!(transport.count(), 1);
1031
1032 let (url, body) = transport.last().expect("应有请求记录");
1033 assert_eq!(url, "https://example.com/token");
1034 assert_eq!(body, r#"{"code":"abc"}"#);
1035 }
1036
1037 #[test]
1039 fn test_memory_oauth2_http_transport_response_queue() {
1040 let transport = MemoryOAuth2HttpTransport::new();
1041 transport.push_response("resp1");
1042 transport.push_response("resp2");
1043
1044 let resp1 = transport
1045 .post_json("url1", "body1")
1046 .expect("第一次调用失败");
1047 let resp2 = transport
1048 .post_json("url2", "body2")
1049 .expect("第二次调用失败");
1050
1051 assert_eq!(resp1, "resp1");
1052 assert_eq!(resp2, "resp2");
1053 assert_eq!(transport.count(), 2);
1054 }
1055
1056 #[test]
1058 fn test_memory_oauth2_http_transport_empty_response() {
1059 let transport = MemoryOAuth2HttpTransport::new();
1060 let response = transport
1062 .post_json("url", "body")
1063 .expect("post_json 不应失败");
1064 assert_eq!(response, "");
1065 }
1066
1067 #[test]
1069 fn test_memory_oauth2_http_transport_clear() {
1070 let transport = MemoryOAuth2HttpTransport::new();
1071 transport.push_response("resp");
1072 transport.post_json("url", "body").expect("调用失败");
1073 assert_eq!(transport.count(), 1);
1074
1075 transport.clear();
1076 assert_eq!(transport.count(), 0);
1077 let response = transport
1079 .post_json("url", "body")
1080 .expect("post_json 不应失败");
1081 assert_eq!(response, "");
1082 }
1083
1084 #[test]
1092 fn test_generic_oauth2_provider_user_from_token() {
1093 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
1094 transport.push_response(r#"{"access_token":"token123","refresh_token":"refresh456","expires_in":3600,"openid":"openid_abc"}"#);
1096 transport.push_response(
1097 r#"{"id":"12345","nickname":"test_user","name":"Test","email":"test@example.com","avatar":"https://example.com/avatar.png"}"#,
1098 );
1099
1100 let config = OAuth2Config::new(
1101 "client123",
1102 "secret456",
1103 "https://example.com/callback",
1104 "https://provider.com/authorize",
1105 "https://provider.com/token",
1106 )
1107 .with_user_url("https://provider.com/user/info");
1108 let provider = GenericOAuth2Provider::new(config, transport.clone());
1109
1110 let user = provider
1111 .user_from_token("auth_code_abc")
1112 .expect("user_from_token 失败");
1113
1114 assert_eq!(user.access_token.as_deref(), Some("token123"));
1116 assert_eq!(user.refresh_token.as_deref(), Some("refresh456"));
1117 assert_eq!(user.expires_in, Some(3600));
1118
1119 assert_eq!(user.id, "12345");
1121 assert_eq!(user.nickname.as_deref(), Some("test_user"));
1122 assert_eq!(user.name.as_deref(), Some("Test"));
1123 assert_eq!(user.email.as_deref(), Some("test@example.com"));
1124 assert_eq!(
1125 user.avatar.as_deref(),
1126 Some("https://example.com/avatar.png")
1127 );
1128
1129 assert_eq!(user.raw["id"], "12345");
1131 assert_eq!(user.raw["nickname"], "test_user");
1132
1133 assert_eq!(transport.count(), 2);
1135 }
1136
1137 #[test]
1139 fn test_generic_oauth2_provider_user_from_token_no_user_url() {
1140 let transport = Arc::new(MemoryOAuth2HttpTransport::new());
1141 transport.push_response(r#"{"access_token":"token123","expires_in":7200}"#);
1142
1143 let config = OAuth2Config::new(
1144 "client123",
1145 "secret456",
1146 "https://example.com/callback",
1147 "https://provider.com/authorize",
1148 "https://provider.com/token",
1149 );
1150 let provider = GenericOAuth2Provider::new(config, transport.clone());
1152
1153 let user = provider
1154 .user_from_token("auth_code")
1155 .expect("user_from_token 失败");
1156
1157 assert_eq!(user.access_token.as_deref(), Some("token123"));
1158 assert_eq!(user.expires_in, Some(7200));
1159 assert!(user.refresh_token.is_none());
1160 assert!(user.id.is_empty());
1162 assert_eq!(transport.count(), 1);
1164 }
1165
1166 #[test]
1168 fn test_generic_oauth2_provider_missing_code() {
1169 let config = OAuth2Config::new(
1170 "client123",
1171 "secret456",
1172 "https://example.com/callback",
1173 "https://provider.com/authorize",
1174 "https://provider.com/token",
1175 );
1176 let provider =
1177 GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
1178
1179 let err = provider.user_from_token("").unwrap_err();
1180 assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("授权码")));
1181 }
1182
1183 #[test]
1185 fn test_oauth2_provider_missing_config_fields() {
1186 let config = OAuth2Config::new(
1188 "",
1189 "secret456",
1190 "https://example.com/callback",
1191 "https://provider.com/authorize",
1192 "https://provider.com/token",
1193 );
1194 let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
1195
1196 let err = provider.user_from_token("code").unwrap_err();
1197 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
1198
1199 let config = OAuth2Config::new(
1201 "client123",
1202 "secret456",
1203 "https://example.com/callback",
1204 "https://provider.com/authorize",
1205 "",
1206 );
1207 let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
1208
1209 let err = provider.user_from_token("code").unwrap_err();
1210 assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
1211 }
1212
1213 #[test]
1215 fn test_generic_oauth2_provider_token_response_missing_access_token() {
1216 let transport = MemoryOAuth2HttpTransport::new();
1217 transport.push_response(r#"{"error":"invalid_grant"}"#);
1218
1219 let config = OAuth2Config::new(
1220 "client123",
1221 "secret456",
1222 "https://example.com/callback",
1223 "https://provider.com/authorize",
1224 "https://provider.com/token",
1225 );
1226 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
1227
1228 let err = provider.user_from_token("code").unwrap_err();
1229 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
1230 }
1231
1232 #[test]
1234 fn test_generic_oauth2_provider_token_response_invalid_json() {
1235 let transport = MemoryOAuth2HttpTransport::new();
1236 transport.push_response("not a json");
1237
1238 let config = OAuth2Config::new(
1239 "client123",
1240 "secret456",
1241 "https://example.com/callback",
1242 "https://provider.com/authorize",
1243 "https://provider.com/token",
1244 );
1245 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
1246
1247 let err = provider.user_from_token("code").unwrap_err();
1248 assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
1249 }
1250
1251 #[test]
1253 fn test_generic_oauth2_provider_user_info_field_aliases() {
1254 let transport = MemoryOAuth2HttpTransport::new();
1255 transport.push_response(r#"{"access_token":"token123","openid":"openid_abc"}"#);
1256 transport.push_response(
1257 r#"{"openid":"qq_12345","nickname":"qq_user","figureurl_qq_1":"https://qzapp.qlogo.cn/1.png"}"#,
1258 );
1259
1260 let config = OAuth2Config::new(
1261 "client123",
1262 "secret456",
1263 "https://example.com/callback",
1264 "https://provider.com/authorize",
1265 "https://provider.com/token",
1266 )
1267 .with_user_url("https://provider.com/user/info");
1268 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
1269
1270 let user = provider
1271 .user_from_token("code")
1272 .expect("user_from_token 失败");
1273
1274 assert_eq!(user.id, "qq_12345");
1276 assert_eq!(user.nickname.as_deref(), Some("qq_user"));
1277 assert_eq!(user.avatar.as_deref(), Some("https://qzapp.qlogo.cn/1.png"));
1279 }
1280
1281 #[test]
1283 fn test_generic_oauth2_provider_user_id_integer() {
1284 let transport = MemoryOAuth2HttpTransport::new();
1285 transport.push_response(r#"{"access_token":"token123"}"#);
1286 transport.push_response(r#"{"id":12345,"nickname":"github_user"}"#);
1287
1288 let config = OAuth2Config::new(
1289 "client123",
1290 "secret456",
1291 "https://example.com/callback",
1292 "https://provider.com/authorize",
1293 "https://provider.com/token",
1294 )
1295 .with_user_url("https://provider.com/user/info");
1296 let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
1297
1298 let user = provider
1299 .user_from_token("code")
1300 .expect("user_from_token 失败");
1301
1302 assert_eq!(user.id, "12345");
1303 assert_eq!(user.nickname.as_deref(), Some("github_user"));
1304 }
1305
1306 #[test]
1308 fn test_generic_oauth2_provider_http_transport_failure() {
1309 let config = OAuth2Config::new(
1310 "client123",
1311 "secret456",
1312 "https://example.com/callback",
1313 "https://provider.com/authorize",
1314 "https://provider.com/token",
1315 );
1316 let provider = GenericOAuth2Provider::new(config, Arc::new(FailingTransport));
1317
1318 let err = provider.user_from_token("code").unwrap_err();
1319 assert!(matches!(err, OAuth2Error::HttpTransport(_)));
1320 }
1321
1322 #[test]
1324 fn test_percent_encode() {
1325 assert_eq!(percent_encode("abcXYZ09-._~"), "abcXYZ09-._~");
1327 assert_eq!(percent_encode("a b"), "a%20b");
1329 assert_eq!(percent_encode("/"), "%2F");
1331 assert_eq!(percent_encode(":"), "%3A");
1333 assert_eq!(
1335 percent_encode("https://example.com/path"),
1336 "https%3A%2F%2Fexample.com%2Fpath"
1337 );
1338 assert_eq!(percent_encode("中"), "%E4%B8%AD");
1340 }
1341
1342 #[test]
1344 fn test_extract_string() {
1345 assert_eq!(
1347 extract_string(&serde_json::json!("hello")),
1348 Some("hello".into())
1349 );
1350 assert_eq!(
1352 extract_string(&serde_json::json!(12345)),
1353 Some("12345".into())
1354 );
1355 assert_eq!(extract_string(&serde_json::json!(1.5)), None);
1357 assert_eq!(extract_string(&serde_json::json!(true)), None);
1359 assert_eq!(extract_string(&serde_json::Value::Null), None);
1361 assert_eq!(extract_string(&serde_json::json!({"a": 1})), None);
1363 }
1364
1365 struct FailingTransport;
1371
1372 impl OAuth2HttpTransport for FailingTransport {
1373 fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
1374 Err(OAuth2Error::HttpTransport("connection refused".into()))
1375 }
1376 }
1377
1378 struct MemoryHttpTransport;
1380
1381 impl OAuth2HttpTransport for MemoryHttpTransport {
1382 fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
1383 Ok(String::new())
1384 }
1385 }
1386}