1use std::collections::BTreeMap;
42use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
43use std::time::Duration;
44
45use chrono::{DateTime, Utc};
46use reqwest::Url;
47
48pub const PROVIDER_GMAIL: &str = "gmail";
55pub const PROVIDER_OUTLOOK: &str = "outlook";
56pub const PROVIDER_GOOGLE_CALENDAR: &str = "google_calendar";
57pub const PROVIDER_MICROSOFT_CALENDAR: &str = "microsoft_calendar";
58
59pub const GOOGLE_HOSTS: &[&str] = &[
63 "accounts.google.com",
64 "oauth2.googleapis.com",
65 "openidconnect.googleapis.com",
66 "www.googleapis.com",
67];
68
69pub const MICROSOFT_HOSTS: &[&str] = &[
72 "login.microsoftonline.com",
73 "graph.microsoft.com",
74 "outlook.office365.com",
75];
76
77#[derive(Debug, Clone, Copy, PartialEq, Eq)]
79pub enum EndpointKey {
80 Authorize,
81 Token,
82 Userinfo,
83}
84
85impl EndpointKey {
86 pub fn as_str(&self) -> &'static str {
88 match self {
89 EndpointKey::Authorize => "authorize",
90 EndpointKey::Token => "token",
91 EndpointKey::Userinfo => "userinfo",
92 }
93 }
94}
95
96impl std::fmt::Display for EndpointKey {
97 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
98 f.write_str(self.as_str())
99 }
100}
101
102#[derive(Debug, Clone)]
107pub struct ProviderAdapter {
108 pub provider: &'static str,
110 pub authorize_endpoint: &'static str,
112 pub token_endpoint: &'static str,
114 pub userinfo_endpoint: &'static str,
116 pub default_scopes: &'static str,
118 pub pkce_s256: bool,
121 pub host_allowlist: &'static [&'static str],
123}
124
125impl ProviderAdapter {
126 pub fn validate(&self, key: EndpointKey, url: &str) -> Result<ValidatedEndpoint, InvalidEndpoint> {
129 validate_endpoint(self.provider, key, self.host_allowlist, url)
130 }
131}
132
133const GMAIL_ADAPTER: ProviderAdapter = ProviderAdapter {
134 provider: PROVIDER_GMAIL,
135 authorize_endpoint: "https://accounts.google.com/o/oauth2/v2/auth",
136 token_endpoint: "https://oauth2.googleapis.com/token",
137 userinfo_endpoint: "https://openidconnect.googleapis.com/v1/userinfo",
138 default_scopes: "https://mail.google.com/ openid https://www.googleapis.com/auth/userinfo.email",
139 pkce_s256: true,
140 host_allowlist: GOOGLE_HOSTS,
141};
142
143const GOOGLE_CALENDAR_ADAPTER: ProviderAdapter = ProviderAdapter {
144 provider: PROVIDER_GOOGLE_CALENDAR,
145 authorize_endpoint: "https://accounts.google.com/o/oauth2/v2/auth",
146 token_endpoint: "https://oauth2.googleapis.com/token",
147 userinfo_endpoint: "https://openidconnect.googleapis.com/v1/userinfo",
148 default_scopes: "https://www.googleapis.com/auth/calendar openid https://www.googleapis.com/auth/userinfo.email",
149 pkce_s256: true,
150 host_allowlist: GOOGLE_HOSTS,
151};
152
153const OUTLOOK_ADAPTER: ProviderAdapter = ProviderAdapter {
154 provider: PROVIDER_OUTLOOK,
155 authorize_endpoint: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize",
156 token_endpoint: "https://login.microsoftonline.com/common/oauth2/v2.0/token",
157 userinfo_endpoint: "https://graph.microsoft.com/oidc/userinfo",
158 default_scopes: "https://outlook.office365.com/SMTP.Send https://outlook.office365.com/IMAP.AccessAsUser.All openid email profile",
159 pkce_s256: true,
160 host_allowlist: MICROSOFT_HOSTS,
161};
162
163const MICROSOFT_CALENDAR_ADAPTER: ProviderAdapter = ProviderAdapter {
164 provider: PROVIDER_MICROSOFT_CALENDAR,
165 authorize_endpoint: "https://login.microsoftonline.com/common/oauth2/v2.0/authorize",
166 token_endpoint: "https://login.microsoftonline.com/common/oauth2/v2.0/token",
167 userinfo_endpoint: "https://graph.microsoft.com/oidc/userinfo",
168 default_scopes: "https://graph.microsoft.com/Calendars.ReadWrite openid email profile",
169 pkce_s256: true,
170 host_allowlist: MICROSOFT_HOSTS,
171};
172
173#[derive(Debug, Clone, Default)]
177pub struct ProviderRegistry {
178 adapters: BTreeMap<&'static str, ProviderAdapter>,
179}
180
181impl ProviderRegistry {
182 pub fn new() -> Self {
183 Self { adapters: BTreeMap::new() }
184 }
185
186 pub fn with_builtin() -> Self {
189 let mut r = Self::new();
190 for adapter in [
191 GMAIL_ADAPTER,
192 OUTLOOK_ADAPTER,
193 GOOGLE_CALENDAR_ADAPTER,
194 MICROSOFT_CALENDAR_ADAPTER,
195 ] {
196 r.adapters.insert(adapter.provider, adapter);
197 }
198 r
199 }
200
201 pub fn register(&mut self, adapter: ProviderAdapter) {
203 self.adapters.insert(adapter.provider, adapter);
204 }
205
206 pub fn lookup(&self, provider: &str) -> Option<&ProviderAdapter> {
208 self.adapters.get(provider)
209 }
210
211 pub fn providers(&self) -> Vec<&'static str> {
213 self.adapters.keys().copied().collect()
214 }
215}
216
217#[derive(Debug, Clone, PartialEq, thiserror::Error)]
224#[error("invalid {key} endpoint for {provider} ({url}): {reason}")]
225pub struct InvalidEndpoint {
226 pub provider: String,
227 pub key: EndpointKey,
228 pub url: String,
229 pub reason: &'static str,
230}
231
232pub fn validate_endpoint(
247 provider: &str,
248 key: EndpointKey,
249 allowlist: &[&str],
250 url_str: &str,
251) -> Result<ValidatedEndpoint, InvalidEndpoint> {
252 let reject = |reason: &'static str| InvalidEndpoint {
253 provider: provider.to_string(),
254 key,
255 url: url_str.to_string(),
256 reason,
257 };
258 let url: Url = url_str.parse().map_err(|_| reject("not a parseable URL"))?;
259 if url.scheme() != "https" {
260 return Err(reject("scheme must be https"));
261 }
262 let host = url.host_str().ok_or_else(|| reject("URL carries no host"))?;
263 if host.is_empty() {
264 return Err(reject("URL carries no host"));
265 }
266 if !url.username().is_empty() || url.password().is_some() {
267 return Err(reject("userinfo component forbidden"));
268 }
269 let host_core = host.strip_prefix('[').and_then(|h| h.strip_suffix(']')).unwrap_or(host);
272 if host_core.parse::<IpAddr>().is_ok() {
273 return Err(reject("IP-literal host forbidden"));
274 }
275 let host_lc = host.to_ascii_lowercase();
276 let allowed = allowlist.iter().any(|entry| {
277 let entry = entry.to_ascii_lowercase();
278 host_lc == entry || host_lc.ends_with(&format!(".{entry}"))
279 });
280 if !allowed {
281 return Err(reject("host is not on the provider allowlist"));
282 }
283 match url.port() {
284 None | Some(443) => {}
285 Some(_) => return Err(reject("only port 443 is permitted")),
286 }
287 let path = url.path();
288 if path.is_empty() || path == "/" {
289 return Err(reject("endpoint must carry a non-root path"));
290 }
291 if url.query().is_some() {
292 return Err(reject("query string forbidden in a configured endpoint"));
293 }
294 if url.fragment().is_some() {
295 return Err(reject("fragment forbidden in a configured endpoint"));
296 }
297 Ok(ValidatedEndpoint { provider_key: provider.to_string(), key, url })
298}
299
300#[derive(Debug, Clone)]
304pub struct ValidatedEndpoint {
305 provider_key: String,
306 key: EndpointKey,
307 url: Url,
308}
309
310impl ValidatedEndpoint {
311 pub fn url(&self) -> &Url {
313 &self.url
314 }
315
316 pub fn as_str(&self) -> &str {
318 self.url.as_str()
319 }
320
321 pub fn host(&self) -> &str {
323 self.url.host_str().unwrap_or_default()
324 }
325
326 pub fn revalidate(&self, registry: &ProviderRegistry) -> Result<(), InvalidEndpoint> {
331 let adapter = registry.lookup(&self.provider_key).ok_or_else(|| InvalidEndpoint {
332 provider: self.provider_key.clone(),
333 key: self.key,
334 url: self.url.as_str().to_string(),
335 reason: "provider is not registered",
336 })?;
337 validate_endpoint(
338 adapter.provider,
339 self.key,
340 adapter.host_allowlist,
341 self.url.as_str(),
342 )
343 .map(|_| ())
344 }
345}
346
347#[derive(Debug, Clone, Default, PartialEq, serde::Deserialize)]
351pub struct ProviderEndpointOverride {
352 pub authorize: Option<String>,
353 pub token: Option<String>,
354 pub userinfo: Option<String>,
355}
356
357#[derive(Debug, Clone, Default, PartialEq, serde::Deserialize)]
359pub struct EndpointOverrides(pub BTreeMap<String, ProviderEndpointOverride>);
360
361impl EndpointOverrides {
362 pub fn get(&self, provider: &str) -> ProviderEndpointOverride {
363 self.0.get(provider).cloned().unwrap_or_default()
364 }
365}
366
367#[derive(Debug, Clone)]
372pub struct ValidatedEndpoints {
373 pub authorize: ValidatedEndpoint,
374 pub token: ValidatedEndpoint,
375 pub userinfo: ValidatedEndpoint,
376}
377
378impl ValidatedEndpoints {
379 pub fn resolve(
384 registry: &ProviderRegistry,
385 provider: &str,
386 overrides: &EndpointOverrides,
387 ) -> Result<Self, InvalidEndpoint> {
388 let adapter = registry.lookup(provider).ok_or_else(|| InvalidEndpoint {
389 provider: provider.to_string(),
390 key: EndpointKey::Token,
391 url: String::new(),
392 reason: "provider is not registered",
393 })?;
394 let override_entry = overrides.get(provider);
395 let pick = |key: EndpointKey,
396 registry_value: &'static str,
397 override_value: Option<&String>|
398 -> Result<ValidatedEndpoint, InvalidEndpoint> {
399 match override_value {
400 Some(url) => adapter.validate(key, url),
401 None => adapter.validate(key, registry_value),
402 }
403 };
404 Ok(Self {
405 authorize: pick(EndpointKey::Authorize, adapter.authorize_endpoint, override_entry.authorize.as_ref())?,
406 token: pick(EndpointKey::Token, adapter.token_endpoint, override_entry.token.as_ref())?,
407 userinfo: pick(EndpointKey::Userinfo, adapter.userinfo_endpoint, override_entry.userinfo.as_ref())?,
408 })
409 }
410}
411
412#[derive(Debug, Clone, serde::Deserialize)]
420pub struct OAuthClientConfig {
421 pub client_id: String,
422 #[serde(default)]
423 pub client_secret: Option<String>,
424}
425
426#[derive(Debug, Clone, Default, serde::Deserialize)]
428pub struct OAuthClientConfigs(pub BTreeMap<String, OAuthClientConfig>);
429
430impl OAuthClientConfigs {
431 pub fn get(&self, provider: &str) -> Option<&OAuthClientConfig> {
432 self.0.get(provider)
433 }
434}
435
436#[derive(Clone, Default, serde::Serialize)]
444pub struct TokenRequestForm {
445 pub grant_type: String,
446 #[serde(skip_serializing_if = "Option::is_none")]
447 pub code: Option<String>,
448 #[serde(skip_serializing_if = "Option::is_none")]
449 pub refresh_token: Option<String>,
450 #[serde(skip_serializing_if = "Option::is_none")]
451 pub redirect_uri: Option<String>,
452 #[serde(skip_serializing_if = "Option::is_none")]
453 pub code_verifier: Option<String>,
454 pub client_id: String,
455 #[serde(skip_serializing_if = "Option::is_none")]
456 pub client_secret: Option<String>,
457 #[serde(skip_serializing_if = "Option::is_none")]
458 pub scope: Option<String>,
459}
460
461impl std::fmt::Debug for TokenRequestForm {
462 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
463 f.debug_struct("TokenRequestForm")
464 .field("grant_type", &self.grant_type)
465 .field("code", &self.code.as_ref().map(|_| "[REDACTED]"))
466 .field("refresh_token", &self.refresh_token.as_ref().map(|_| "[REDACTED]"))
467 .field("redirect_uri", &self.redirect_uri)
468 .field("code_verifier", &self.code_verifier.as_ref().map(|_| "[REDACTED]"))
469 .field("client_id", &self.client_id)
470 .field("client_secret", &self.client_secret.as_ref().map(|_| "[REDACTED]"))
471 .field("scope", &self.scope)
472 .finish()
473 }
474}
475
476#[derive(Clone, serde::Deserialize)]
481pub struct TokenResponse {
482 pub access_token: String,
483 #[serde(default)]
484 pub refresh_token: Option<String>,
485 #[serde(default)]
486 pub expires_in: Option<i64>,
487 #[serde(default)]
488 pub scope: Option<String>,
489 #[serde(default)]
490 pub id_token: Option<String>,
491 #[serde(default)]
492 pub token_type: Option<String>,
493}
494
495impl std::fmt::Debug for TokenResponse {
496 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
497 f.debug_struct("TokenResponse")
498 .field("access_token", &"[REDACTED]")
499 .field("refresh_token", &self.refresh_token.as_ref().map(|_| "[REDACTED]"))
500 .field("expires_in", &self.expires_in)
501 .field("scope", &self.scope)
502 .field("id_token", &self.id_token.as_ref().map(|_| "[REDACTED]"))
503 .field("token_type", &self.token_type)
504 .finish()
505 }
506}
507
508impl TokenResponse {
509 pub fn expires_at(&self, now: DateTime<Utc>) -> Option<DateTime<Utc>> {
512 self.expires_in
513 .filter(|secs| *secs > 0)
514 .map(|secs| now + chrono::Duration::seconds(secs))
515 }
516}
517
518#[derive(Debug, Clone, Default, serde::Deserialize)]
523pub struct IdentityClaims {
524 #[serde(default)]
525 pub sub: Option<String>,
526 #[serde(default)]
527 pub email: Option<String>,
528 #[serde(default)]
529 pub email_verified: Option<bool>,
530 #[serde(default, alias = "aud")]
531 pub audience: Option<String>,
532 #[serde(default)]
533 pub nonce: Option<String>,
534}
535
536#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
548pub enum TransportFailureKind {
549 #[error("invalid_grant")]
550 InvalidGrant,
551 #[error("provider refused")]
552 Provider,
553 #[error("network failure")]
554 Network,
555 #[error("endpoint guard refusal")]
556 EndpointGuard,
557}
558
559#[derive(Debug, Clone, thiserror::Error)]
560#[error("oauth transport failure ({kind}): {message}")]
561pub struct TransportFailure {
562 pub kind: TransportFailureKind,
563 pub message: String,
564}
565
566impl TransportFailure {
567 pub fn new(kind: TransportFailureKind, message: impl Into<String>) -> Self {
568 Self { kind, message: message.into() }
569 }
570
571 pub fn is_invalid_grant(&self) -> bool {
572 self.kind == TransportFailureKind::InvalidGrant
573 }
574}
575
576#[async_trait::async_trait]
580pub trait OAuthTransport: Send + Sync {
581 async fn exchange(
583 &self,
584 endpoint: &ValidatedEndpoint,
585 form: &TokenRequestForm,
586 ) -> Result<TokenResponse, TransportFailure>;
587
588 async fn fetch_identity(
591 &self,
592 endpoint: &ValidatedEndpoint,
593 access_token: &str,
594 ) -> Result<IdentityClaims, TransportFailure>;
595}
596
597pub fn ip_is_public(ip: IpAddr) -> bool {
606 match ip {
607 IpAddr::V4(v4) => ipv4_is_public(v4),
608 IpAddr::V6(v6) => {
609 if let Some(mapped) = v6.to_ipv4_mapped() {
610 return ipv4_is_public(mapped);
611 }
612 ipv6_is_public(&v6)
613 }
614 }
615}
616
617fn ipv4_is_public(ip: Ipv4Addr) -> bool {
618 let o = ip.octets();
619 !(ip.is_loopback()
620 || ip.is_private()
621 || ip.is_link_local()
622 || ip.is_unspecified()
623 || ip.is_multicast()
624 || ip.is_broadcast()
625 || (o[0] == 100 && (o[1] & 0b1100_0000) == 0b0100_0000))
627}
628
629fn ipv6_is_public(ip: &Ipv6Addr) -> bool {
630 let seg = ip.segments();
631 !(ip.is_loopback()
632 || ip.is_unspecified()
633 || ip.is_multicast()
634 || (seg[0] & 0xffc0) == 0xfe80
636 || (seg[0] & 0xfe00) == 0xfc00)
638}
639
640pub async fn assert_public_resolution(host: &str) -> Result<Vec<IpAddr>, TransportFailure> {
646 let resolved: Vec<IpAddr> = tokio::net::lookup_host((host, 443))
647 .await
648 .map_err(|e| TransportFailure::new(
649 TransportFailureKind::EndpointGuard,
650 format!("resolving {host}: {e}"),
651 ))?
652 .map(|addr| addr.ip())
653 .collect();
654 if resolved.is_empty() {
655 return Err(TransportFailure::new(
656 TransportFailureKind::EndpointGuard,
657 format!("resolving {host}: no addresses"),
658 ));
659 }
660 for ip in &resolved {
661 if !ip_is_public(*ip) {
662 return Err(TransportFailure::new(
663 TransportFailureKind::EndpointGuard,
664 format!("{host} resolves to forbidden address {ip}"),
665 ));
666 }
667 }
668 Ok(resolved)
669}
670
671pub const OAUTH_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
677pub const OAUTH_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
679
680pub struct ReqwestOAuthTransport {
689 client: reqwest::Client,
690 registry: ProviderRegistry,
691 redirect_policy_none: bool,
692}
693
694impl ReqwestOAuthTransport {
695 pub fn new() -> Result<Self, TransportFailure> {
698 Self::with_registry(ProviderRegistry::with_builtin())
699 }
700
701 pub fn with_registry(registry: ProviderRegistry) -> Result<Self, TransportFailure> {
704 let client = reqwest::Client::builder()
705 .redirect(reqwest::redirect::Policy::none())
706 .connect_timeout(OAUTH_CONNECT_TIMEOUT)
707 .timeout(OAUTH_REQUEST_TIMEOUT)
708 .user_agent("backbone-integrations-oauth/1")
709 .build()
710 .map_err(|e| TransportFailure::new(TransportFailureKind::Network, e.to_string()))?;
711 Ok(Self { client, registry, redirect_policy_none: true })
712 }
713
714 pub fn from_client(client: reqwest::Client, redirect_policy_none: bool) -> Self {
718 Self { client, registry: ProviderRegistry::with_builtin(), redirect_policy_none }
719 }
720
721 pub fn redirect_policy_is_none(&self) -> bool {
724 self.redirect_policy_none
725 }
726
727 async fn guard(&self, endpoint: &ValidatedEndpoint) -> Result<(), TransportFailure> {
731 endpoint
732 .revalidate(&self.registry)
733 .map_err(|e| {
734 TransportFailure::new(
735 TransportFailureKind::EndpointGuard,
736 format!("request-time revalidation refused {e}"),
737 )
738 })?;
739 assert_public_resolution(endpoint.host()).await?;
740 Ok(())
741 }
742}
743
744impl Default for ReqwestOAuthTransport {
745 fn default() -> Self {
746 Self::new().expect("reqwest client construction")
747 }
748}
749
750#[async_trait::async_trait]
751impl OAuthTransport for ReqwestOAuthTransport {
752 async fn exchange(
753 &self,
754 endpoint: &ValidatedEndpoint,
755 form: &TokenRequestForm,
756 ) -> Result<TokenResponse, TransportFailure> {
757 self.guard(endpoint).await?;
758 let response = self
759 .client
760 .post(endpoint.url().clone())
761 .form(form)
762 .send()
763 .await
764 .map_err(|e| TransportFailure::new(TransportFailureKind::Network, e.to_string()))?;
765 let status = response.status();
766 if status.is_success() {
767 let token: TokenResponse = response
768 .json()
769 .await
770 .map_err(|e| TransportFailure::new(
771 TransportFailureKind::Provider,
772 format!("token response was not valid provider JSON: {e}"),
773 ))?;
774 if token.access_token.trim().is_empty() {
775 return Err(TransportFailure::new(
776 TransportFailureKind::Provider,
777 "token response carries no access_token",
778 ));
779 }
780 return Ok(token);
781 }
782 let body = response.text().await.unwrap_or_default();
787 let error_code = serde_json::from_str::<serde_json::Value>(&body)
788 .ok()
789 .and_then(|v| v.get("error").and_then(|e| e.as_str()).map(String::from));
790 if status.as_u16() == 400 && error_code.as_deref() == Some("invalid_grant") {
791 return Err(TransportFailure::new(
792 TransportFailureKind::InvalidGrant,
793 "provider rejected the grant: invalid_grant",
794 ));
795 }
796 Err(TransportFailure::new(
797 TransportFailureKind::Provider,
798 format!("provider answered {status}: {}", error_code.unwrap_or_else(|| body.chars().take(200).collect::<String>())),
799 ))
800 }
801
802 async fn fetch_identity(
803 &self,
804 endpoint: &ValidatedEndpoint,
805 access_token: &str,
806 ) -> Result<IdentityClaims, TransportFailure> {
807 self.guard(endpoint).await?;
808 let response = self
809 .client
810 .get(endpoint.url().clone())
811 .bearer_auth(access_token)
812 .send()
813 .await
814 .map_err(|e| TransportFailure::new(TransportFailureKind::Network, e.to_string()))?;
815 let status = response.status();
816 if status.is_success() {
817 return response.json().await.map_err(|e| {
818 TransportFailure::new(
819 TransportFailureKind::Provider,
820 format!("identity response was not valid provider JSON: {e}"),
821 )
822 });
823 }
824 let body = response.text().await.unwrap_or_default();
825 Err(TransportFailure::new(
826 TransportFailureKind::Provider,
827 format!("provider answered {status}: {}", body.chars().take(200).collect::<String>()),
828 ))
829 }
830}
831
832#[cfg(test)]
833mod tests {
834 use super::*;
835
836 const EVIL: &str = "https://evil.example.com/token";
837
838 #[test]
839 fn registry_endpoints_pass_their_own_allowlist() {
840 let registry = ProviderRegistry::with_builtin();
841 for provider in registry.providers() {
842 let adapter = registry.lookup(provider).expect("adapter");
843 for (key, url) in [
844 (EndpointKey::Authorize, adapter.authorize_endpoint),
845 (EndpointKey::Token, adapter.token_endpoint),
846 (EndpointKey::Userinfo, adapter.userinfo_endpoint),
847 ] {
848 assert!(
849 adapter.validate(key, url).is_ok(),
850 "{provider} {key:?} endpoint {url} must pass its own allowlist",
851 );
852 }
853 }
854 }
855
856 #[test]
857 fn guard_rejects_each_malicious_shape() {
858 let registry = ProviderRegistry::with_builtin();
859 let adapter = registry.lookup(PROVIDER_GMAIL).unwrap();
860 let cases: &[(&str, &'static str)] = &[
861 ("http://oauth2.googleapis.com/token", "scheme must be https"),
862 (EVIL, "host is not on the provider allowlist"),
863 ("https://user@oauth2.googleapis.com/token", "userinfo component forbidden"),
864 ("https://user:pass@oauth2.googleapis.com/token", "userinfo component forbidden"),
865 ("https://127.0.0.1/token", "IP-literal host forbidden"),
866 ("https://[::1]/token", "IP-literal host forbidden"),
867 ("https://oauth2.googleapis.com:8443/token", "only port 443 is permitted"),
868 ("https://oauth2.googleapis.com/", "endpoint must carry a non-root path"),
869 ("https://oauth2.googleapis.com", "endpoint must carry a non-root path"),
870 ("https://oauth2.googleapis.com/token?x=1", "query string forbidden in a configured endpoint"),
871 ("https://oauth2.googleapis.com/token#f", "fragment forbidden in a configured endpoint"),
872 ("not a url", "not a parseable URL"),
873 ];
874 for (url, expected_reason) in cases {
875 let err = adapter
876 .validate(EndpointKey::Token, url)
877 .expect_err(url);
878 assert_eq!(err.reason, *expected_reason, "wrong reason for {url}: {err}");
879 }
880 }
881
882 #[test]
883 fn guard_accepts_subdomains_explicit_port_and_overrides() {
884 let registry = ProviderRegistry::with_builtin();
885 let adapter = registry.lookup(PROVIDER_GMAIL).unwrap();
886 assert!(adapter.validate(EndpointKey::Token, "https://oauth2.googleapis.com/token").is_ok());
890 assert!(adapter
891 .validate(EndpointKey::Token, "https://eviloauth2.googleapis.com/token")
892 .is_err(), "suffix match must be dot-anchored");
893 assert!(adapter
894 .validate(EndpointKey::Token, "https://accounts.google.com.evil.example.com/token")
895 .is_err(), "allowlisted host must be a suffix, not a substring");
896 assert!(adapter
898 .validate(EndpointKey::Token, "https://oauth2.googleapis.com:443/token")
899 .is_ok());
900 assert!(adapter
902 .validate(EndpointKey::Token, "https://login.microsoftonline.com/common/oauth2/v2.0/token")
903 .is_err(), "cross-family host must be refused");
904 }
905
906 #[test]
907 fn resolve_uses_registry_defaults_and_validated_overrides() {
908 let registry = ProviderRegistry::with_builtin();
909 let none = EndpointOverrides::default();
910 let gmail = ValidatedEndpoints::resolve(®istry, PROVIDER_GMAIL, &none).expect("defaults resolve");
911 assert_eq!(gmail.token.as_str(), "https://oauth2.googleapis.com/token");
912
913 let mut map = BTreeMap::new();
915 map.insert(
916 PROVIDER_GMAIL.to_string(),
917 ProviderEndpointOverride {
918 token: Some("https://oauth2.googleapis.com/token".into()),
919 ..Default::default()
920 },
921 );
922 let good = EndpointOverrides(map);
923 assert!(ValidatedEndpoints::resolve(®istry, PROVIDER_GMAIL, &good).is_ok());
924
925 let mut map = BTreeMap::new();
927 map.insert(
928 PROVIDER_GMAIL.to_string(),
929 ProviderEndpointOverride { token: Some(EVIL.into()), ..Default::default() },
930 );
931 let evil = EndpointOverrides(map);
932 let err = ValidatedEndpoints::resolve(®istry, PROVIDER_GMAIL, &evil).expect_err("evil override");
933 assert_eq!(err.reason, "host is not on the provider allowlist");
934
935 assert!(ValidatedEndpoints::resolve(®istry, "not_a_provider", &none).is_err());
937 }
938
939 #[test]
940 fn revalidation_re_runs_the_full_rule_set() {
941 let registry = ProviderRegistry::with_builtin();
942 let mut endpoint = registry
943 .lookup(PROVIDER_OUTLOOK)
944 .unwrap()
945 .validate(EndpointKey::Token, "https://login.microsoftonline.com/common/oauth2/v2.0/token")
946 .expect("valid");
947 assert!(endpoint.revalidate(®istry).is_ok(), "a clean endpoint revalidates");
948
949 endpoint.url = "https://127.0.0.1/token".parse().unwrap();
952 let err = endpoint.revalidate(®istry).expect_err("tampered endpoint");
953 assert_eq!(err.reason, "IP-literal host forbidden");
954
955 let mut endpoint2 = registry
957 .lookup(PROVIDER_GMAIL)
958 .unwrap()
959 .validate(EndpointKey::Token, "https://oauth2.googleapis.com/token")
960 .unwrap();
961 endpoint2.provider_key = "ghost_provider".into();
962 assert!(endpoint2.revalidate(®istry).is_err());
963 }
964
965 #[test]
966 fn public_transport_refuses_redirects_by_construction() {
967 let transport = ReqwestOAuthTransport::new().expect("client builds");
968 assert!(transport.redirect_policy_is_none(), "production client must refuse redirects");
969 }
970
971 #[test]
972 fn private_ranges_are_not_public() {
973 let forbidden = [
974 "127.0.0.1", "10.0.0.1", "172.16.0.1", "192.168.1.254", "169.254.1.1",
975 "100.64.0.1", "0.0.0.0", "255.255.255.255", "224.0.0.1",
976 "::1", "::", "fe80::1", "fc00::1", "fd12:3456::1", "ff02::1",
977 "::ffff:127.0.0.1", "::ffff:10.0.0.1",
978 ];
979 for raw in forbidden {
980 let ip: IpAddr = raw.parse().expect(raw);
981 assert!(!ip_is_public(ip), "{raw} must be refused as an outbound target");
982 }
983 let public = ["8.8.8.8", "34.64.1.1", "2607:f8b0:4005:80a::200e", "2001:4860:4860::8888"];
984 for raw in public {
985 let ip: IpAddr = raw.parse().expect(raw);
986 assert!(ip_is_public(ip), "{raw} is a legitimate outbound target");
987 }
988 }
989
990 #[tokio::test]
991 async fn resolution_guard_refuses_loopback_hosts() {
992 let err = assert_public_resolution("localhost").await.expect_err("loopback must be refused");
994 assert_eq!(err.kind, TransportFailureKind::EndpointGuard);
995 assert!(err.message.contains("forbidden address"), "unexpected message: {err}");
996 let err = assert_public_resolution("127.0.0.1").await.expect_err("IP literal must be refused");
997 assert_eq!(err.kind, TransportFailureKind::EndpointGuard);
998 }
999
1000 #[test]
1001 fn form_and_response_debug_are_redacted() {
1002 let form = TokenRequestForm {
1003 grant_type: "authorization_code".into(),
1004 code: Some("SECRET-CODE".into()),
1005 refresh_token: Some("SECRET-REFRESH".into()),
1006 redirect_uri: Some("https://example.test/cb".into()),
1007 code_verifier: Some("SECRET-VERIFIER".into()),
1008 client_id: "client-1".into(),
1009 client_secret: Some("SECRET-CLIENT-SECRET".into()),
1010 scope: None,
1011 };
1012 let debugged = format!("{form:?}");
1013 for secret in ["SECRET-CODE", "SECRET-REFRESH", "SECRET-VERIFIER", "SECRET-CLIENT-SECRET"] {
1014 assert!(!debugged.contains(secret), "leak in form Debug: {debugged}");
1015 }
1016 assert!(debugged.contains("[REDACTED]"));
1017
1018 let response = TokenResponse {
1019 access_token: "SECRET-ACCESS".into(),
1020 refresh_token: Some("SECRET-REFRESH2".into()),
1021 expires_in: Some(3600),
1022 scope: Some("openid".into()),
1023 id_token: Some("SECRET-ID-TOKEN".into()),
1024 token_type: Some("Bearer".into()),
1025 };
1026 let debugged = format!("{response:?}");
1027 for secret in ["SECRET-ACCESS", "SECRET-REFRESH2", "SECRET-ID-TOKEN"] {
1028 assert!(!debugged.contains(secret), "leak in response Debug: {debugged}");
1029 }
1030 assert!(debugged.contains("3600"));
1032 }
1033
1034 #[test]
1035 fn honest_expiry_derivation() {
1036 let now = Utc::now();
1037 let with_expiry = TokenResponse {
1038 access_token: "a".into(),
1039 refresh_token: None,
1040 expires_in: Some(3600),
1041 scope: None,
1042 id_token: None,
1043 token_type: None,
1044 };
1045 assert_eq!(with_expiry.expires_at(now), Some(now + chrono::Duration::hours(1)));
1046 let permanent = TokenResponse { expires_in: None, ..with_expiry };
1048 assert_eq!(permanent.expires_at(now), None);
1049 let zero = TokenResponse { expires_in: Some(0), ..permanent };
1051 assert_eq!(zero.expires_at(now), None);
1052 }
1053}