Skip to main content

backbone_integrations/infrastructure/http/
endpoint_guard.rs

1//! The outbound endpoint guard for the one OAuth generation (hand-authored,
2//! user-owned).
3//!
4//! Why this file exists: the upstream code this module replaces guarded its
5//! outbound OAuth calls with a debug `assert` over a host allowlist — an
6//! assertion the runtime strips under optimization, and whose Microsoft
7//! endpoints were overridable at runtime with no validation at all. The guard
8//! here is a real one, and it fails closed at every layer:
9//!
10//! 1. **Registry (compile time)** — [`ProviderRegistry`] carries the adapter
11//!    data for every supported provider: endpoint URLs, default scopes, the
12//!    PKCE capability, and the host allowlist for that provider family.
13//!    There is no code path that constructs a provider URL from anything a
14//!    request supplied.
15//! 2. **Validation (config-load time)** — every endpoint override arriving
16//!    from configuration passes [`validate_endpoint`] before the module using
17//!    it may build: https-only, host suffix-matched against the provider's
18//!    allowlist, no userinfo component, no IP-literal host, no port other
19//!    than 443, a non-root path, no query, no fragment. ANY violation is an
20//!    [`InvalidEndpoint`] and the module builder refuses to build — a bad
21//!    override yields no module, not a degraded one.
22//! 3. **Re-validation (request time)** — a [`ValidatedEndpoint`] can only be
23//!    constructed through validation, and the transport re-runs the full rule
24//!    set ([`ValidatedEndpoint::revalidate`]) before every call, so an
25//!    endpoint that reached the transport by internal drift is refused there
26//!    too. Fail closed at build time AND at request time.
27//! 4. **Resolution guard (transport time)** — the reqwest client follows no
28//!    redirects ([`reqwest::redirect::Policy::none`]) and resolves the host
29//!    before connecting, refusing loopback/private/link-local/unique-local
30//!    addresses ([`assert_public_resolution`]) — a basic DNS-rebinding
31//!    closure on top of the allowlist.
32//!
33//! The type system carries the guarantee: the transport accepts ONLY
34//! [`ValidatedEndpoint`] values — "URL that passed the guard" is the sole
35//! currency — and the field inside is private, so no caller can smuggle an
36//! unvalidated URL into an outbound call.
37//!
38//! [`OAuthTransport`] is the port the OAuth flow and the refresh scheduler
39//! share; tests inject a fake that records every URL it was handed.
40
41use std::collections::BTreeMap;
42use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
43use std::time::Duration;
44
45use chrono::{DateTime, Utc};
46use reqwest::Url;
47
48// ─────────────────────────────────────────────────────────────────────────────
49// Provider registry — compile-time adapter data
50// ─────────────────────────────────────────────────────────────────────────────
51
52/// Provider key constants (the `OAuthProvider` enum's values as plain strings,
53/// so infrastructure stays independent of the generated entity enum).
54pub 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
59/// The host allowlist for the Google provider family. A configured endpoint
60/// for `gmail` / `google_calendar` may live ONLY on these hosts (exact match
61/// or a subdomain of a listed host).
62pub const GOOGLE_HOSTS: &[&str] = &[
63    "accounts.google.com",
64    "oauth2.googleapis.com",
65    "openidconnect.googleapis.com",
66    "www.googleapis.com",
67];
68
69/// The host allowlist for the Microsoft provider family (`outlook` /
70/// `microsoft_calendar`).
71pub const MICROSOFT_HOSTS: &[&str] = &[
72    "login.microsoftonline.com",
73    "graph.microsoft.com",
74    "outlook.office365.com",
75];
76
77/// Which of a provider's three endpoints a URL is being validated as.
78#[derive(Debug, Clone, Copy, PartialEq, Eq)]
79pub enum EndpointKey {
80    Authorize,
81    Token,
82    Userinfo,
83}
84
85impl EndpointKey {
86    /// The key as it appears in configuration (`oauth.endpoints.<provider>.<key>`).
87    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/// One provider's adapter data: the endpoints, default scopes, PKCE
103/// capability, and host allowlist. Purely compile-time — per-deployment
104/// overrides arrive through [`EndpointOverrides`] and are validated against
105/// the allowlist here before use.
106#[derive(Debug, Clone)]
107pub struct ProviderAdapter {
108    /// Provider key (`"gmail"` …) — matches the account row's provider value.
109    pub provider: &'static str,
110    /// The authorization endpoint the browser is sent to.
111    pub authorize_endpoint: &'static str,
112    /// The token endpoint (code exchange AND refresh grant).
113    pub token_endpoint: &'static str,
114    /// The identity endpoint for post-exchange account verification.
115    pub userinfo_endpoint: &'static str,
116    /// Default scopes requested when the caller passes none.
117    pub default_scopes: &'static str,
118    /// Whether the provider supports PKCE S256 (verifier minted per
119    /// authorization, challenge sent, verifier presented at exchange).
120    pub pkce_s256: bool,
121    /// Hosts this provider's endpoints may live on (the SSRF allowlist).
122    pub host_allowlist: &'static [&'static str],
123}
124
125impl ProviderAdapter {
126    /// Validate a candidate URL as one of this provider's endpoints — the
127    /// single rule set every override and every registry value passes.
128    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/// The compile-time provider registry. [`ProviderRegistry::with_builtin`]
174/// carries the four OAuth providers; composition may add adapters, never
175/// remove the validation contract.
176#[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    /// The four built-in providers (gmail / outlook / google_calendar /
187    /// microsoft_calendar).
188    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    /// Register an additional adapter (composition extension point).
202    pub fn register(&mut self, adapter: ProviderAdapter) {
203        self.adapters.insert(adapter.provider, adapter);
204    }
205
206    /// Look up one provider's adapter data.
207    pub fn lookup(&self, provider: &str) -> Option<&ProviderAdapter> {
208        self.adapters.get(provider)
209    }
210
211    /// All registered provider keys, sorted.
212    pub fn providers(&self) -> Vec<&'static str> {
213        self.adapters.keys().copied().collect()
214    }
215}
216
217// ─────────────────────────────────────────────────────────────────────────────
218// Endpoint validation — the guard
219// ─────────────────────────────────────────────────────────────────────────────
220
221/// A configuration-time endpoint rejection. Fail closed: the caller building
222/// the module refuses to proceed.
223#[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
232/// Validate one endpoint URL against a provider's host allowlist. The full
233/// rule set, in order:
234///
235/// - parses as an absolute URL;
236/// - scheme is exactly `https`;
237/// - host is present, carries no userinfo, and is NOT an IP literal;
238/// - host suffix-matches an allowlist entry (exact host or subdomain of a
239///   listed host — a listed apex covers its subdomains, nothing else);
240/// - port is absent (default) or explicitly 443;
241/// - the path is non-empty and not just `/`;
242/// - no query string and no fragment.
243///
244/// Every rule is deny-by-default: anything the rule set does not recognize as
245/// allowed is rejected.
246pub 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    // IPv6 literals arrive bracketed ([::1]) — strip the brackets before the
270    // IP parse so both families are caught.
271    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/// A URL that passed the full rule set. Constructible only through
301/// [`validate_endpoint`] (the inner URL is private); the transport accepts
302/// nothing else, and re-runs the rules before every call.
303#[derive(Debug, Clone)]
304pub struct ValidatedEndpoint {
305    provider_key: String,
306    key: EndpointKey,
307    url: Url,
308}
309
310impl ValidatedEndpoint {
311    /// The validated URL (for building the outbound request).
312    pub fn url(&self) -> &Url {
313        &self.url
314    }
315
316    /// The URL as a string.
317    pub fn as_str(&self) -> &str {
318        self.url.as_str()
319    }
320
321    /// The host this endpoint was validated against the allowlist for.
322    pub fn host(&self) -> &str {
323        self.url.host_str().unwrap_or_default()
324    }
325
326    /// Re-run the full rule set against the stored URL (request-time
327    /// fail-closed). The stored allowlist binding is re-derived from the
328    /// provider key via the registry, so a tampered or drifted value is
329    /// refused here exactly as it would be at config-load time.
330    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/// Per-provider endpoint overrides from configuration. Empty by default — no
348/// override means the registry value. An override is a candidate, never a
349/// truth: it is validated before use, and a bad one fails the build.
350#[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/// The `oauth.endpoints` configuration section: provider key → overrides.
358#[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/// A provider's three endpoints, resolved and validated. Produced by
368/// [`ValidatedEndpoints::resolve`] — registry values by default, validated
369/// overrides where present. The sole endpoint currency the OAuth flow and the
370/// refresh scheduler deal in.
371#[derive(Debug, Clone)]
372pub struct ValidatedEndpoints {
373    pub authorize: ValidatedEndpoint,
374    pub token: ValidatedEndpoint,
375    pub userinfo: ValidatedEndpoint,
376}
377
378impl ValidatedEndpoints {
379    /// Resolve one provider's endpoints: registry defaults unless a
380    /// configuration override is present; EVERY value (registry or override)
381    /// passes [`validate_endpoint`]. Any rejection propagates — the caller
382    /// building the module refuses to build.
383    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// ─────────────────────────────────────────────────────────────────────────────
413// OAuth client configuration (non-secret id + secret handed by composition)
414// ─────────────────────────────────────────────────────────────────────────────
415
416/// One provider's OAuth client registration. The client id is not a secret
417/// (it travels in the authorize URL); the client secret is — redacted in
418/// Debug, zeroized on drop.
419#[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/// The `oauth.clients` configuration section: provider key → client config.
427#[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// ─────────────────────────────────────────────────────────────────────────────
437// The transport port
438// ─────────────────────────────────────────────────────────────────────────────
439
440/// An OAuth token-endpoint request form (code exchange or refresh grant),
441/// serialized as `application/x-www-form-urlencoded`. Debug is redacted — the
442/// code, refresh token, and client secret must never drift into a log line.
443#[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/// A provider token response. `expires_in` missing means the provider claims
477/// no expiry — per the honest-lifetime rule such a response is UNSTOREABLE
478/// and callers refuse it (a "permanent" token is not a value the flow will
479/// store). Debug is redacted.
480#[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    /// The honest expiry of this response: `now + expires_in`. `None` when
510    /// the provider returned no `expires_in` — the unstoreable case.
511    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/// The verified identity of the token's subject, as read server-side (the
519/// userinfo endpoint, or the id_token the token endpoint returned — never a
520/// browser-side claim). `audience` and `nonce` carry the id_token's values so
521/// the flow can enforce the audience check and the nonce binding.
522#[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/// Why an outbound OAuth call failed.
537///
538/// - [`InvalidGrant`](TransportFailureKind::InvalidGrant) — the provider
539///   answered 400 `invalid_grant`: the refresh token is dead; the caller
540///   expires the account (the user must reconnect), never retries.
541/// - [`Provider`](TransportFailureKind::Provider) — the provider refused or
542///   returned an unusable shape. Zero writes; the next tick retries.
543/// - [`Network`](TransportFailureKind::Network) — transport-level failure.
544///   Zero writes; the next tick retries.
545/// - [`EndpointGuard`](TransportFailureKind::EndpointGuard) — the guard
546///   refused (revalidation or resolution). Zero writes, zero HTTP calls.
547#[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/// The outbound transport port: token exchange and server-side identity
577/// fetch. Both methods take ONLY validated endpoints — an unvalidated URL
578/// cannot reach the network through this trait.
579#[async_trait::async_trait]
580pub trait OAuthTransport: Send + Sync {
581    /// POST the form to the token endpoint (code exchange or refresh grant).
582    async fn exchange(
583        &self,
584        endpoint: &ValidatedEndpoint,
585        form: &TokenRequestForm,
586    ) -> Result<TokenResponse, TransportFailure>;
587
588    /// GET the identity endpoint with the access token (server-side
589    /// verification read).
590    async fn fetch_identity(
591        &self,
592        endpoint: &ValidatedEndpoint,
593        access_token: &str,
594    ) -> Result<IdentityClaims, TransportFailure>;
595}
596
597// ─────────────────────────────────────────────────────────────────────────────
598// DNS resolution guard
599// ─────────────────────────────────────────────────────────────────────────────
600
601/// Whether an address is acceptable as an outbound OAuth target: not
602/// loopback, not private (RFC1918 + carrier-grade NAT), not link-local, not
603/// unique-local (fc00::/7), not unspecified, not multicast — for IPv4, IPv6,
604/// and IPv4-mapped IPv6 alike.
605pub 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        // Carrier-grade NAT 100.64.0.0/10 — not globally routable.
626        || (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        // Link-local fe80::/10.
635        || (seg[0] & 0xffc0) == 0xfe80
636        // Unique-local fc00::/7.
637        || (seg[0] & 0xfe00) == 0xfc00)
638}
639
640/// Resolve `host` and refuse any answer in a forbidden range — the basic
641/// DNS-rebinding closure. The allowlist remains the primary control (an
642/// attacker cannot control DNS for an allowlisted provider host); this check
643/// bounds the residual window where a name that passed the allowlist resolves
644/// somewhere it should not. Returns the resolved addresses on success.
645pub 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
671// ─────────────────────────────────────────────────────────────────────────────
672// The reqwest implementation
673// ─────────────────────────────────────────────────────────────────────────────
674
675/// Default connect budget for one outbound OAuth call.
676pub const OAUTH_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
677/// Default total budget for one outbound OAuth call.
678pub const OAUTH_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
679
680/// The production [`OAuthTransport`]: a reqwest client that follows no
681/// redirects, carries explicit timeouts, re-validates the endpoint before
682/// every call, and resolves the host through the private-range guard before
683/// connecting.
684///
685/// The no-redirect policy is load-bearing: an allowlisted endpoint answering
686/// `3xx Location: http://attacker.example/` must not be followed — the
687/// response is surfaced as a provider error instead.
688pub struct ReqwestOAuthTransport {
689    client: reqwest::Client,
690    registry: ProviderRegistry,
691    redirect_policy_none: bool,
692}
693
694impl ReqwestOAuthTransport {
695    /// Build with the guard postures (no redirects, timeouts) over the
696    /// built-in provider registry. This is the production constructor.
697    pub fn new() -> Result<Self, TransportFailure> {
698        Self::with_registry(ProviderRegistry::with_builtin())
699    }
700
701    /// Build with the guard postures over a caller-supplied registry
702    /// (composition that registered extra adapters).
703    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    /// Wrap an externally built client (test injection). The no-redirect
715    /// posture is reported as `false` unless the caller asserts it — the
716    /// caller takes responsibility for the posture it injects.
717    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    /// Whether this transport refuses redirects (the guard posture). Asserted
722    /// by the guard probes for the production constructor.
723    pub fn redirect_policy_is_none(&self) -> bool {
724        self.redirect_policy_none
725    }
726
727    /// The pre-call guard: re-run the full endpoint rule set against this
728    /// transport's registry, then resolve the host through the private-range
729    /// check. Fail closed — a refusal means zero bytes hit the network.
730    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        // A provider refusal: 400 invalid_grant is the dead-refresh-token
783        // signal the refresh path acts on; everything else is a plain
784        // provider error. The body's error code is provider metadata (never
785        // token material).
786        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        // A subdomain of a listed host is allowlisted; the suffix is
887        // dot-anchored, so a host merely ENDING in the allowlist string
888        // (eviloauth2.googleapis.com) does not match.
889        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        // Explicit 443 and a real path pass.
897        assert!(adapter
898            .validate(EndpointKey::Token, "https://oauth2.googleapis.com:443/token")
899            .is_ok());
900        // Microsoft hosts are NOT valid for a Google provider.
901        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(&registry, PROVIDER_GMAIL, &none).expect("defaults resolve");
911        assert_eq!(gmail.token.as_str(), "https://oauth2.googleapis.com/token");
912
913        // A same-host path override passes and is used.
914        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(&registry, PROVIDER_GMAIL, &good).is_ok());
924
925        // An evil override fails the resolve — the module cannot be built.
926        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(&registry, PROVIDER_GMAIL, &evil).expect_err("evil override");
933        assert_eq!(err.reason, "host is not on the provider allowlist");
934
935        // An unknown provider fails closed.
936        assert!(ValidatedEndpoints::resolve(&registry, "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(&registry).is_ok(), "a clean endpoint revalidates");
948
949        // Tamper with the stored URL (only reachable inside the module — the
950        // field is private): revalidation must refuse it at request time.
951        endpoint.url = "https://127.0.0.1/token".parse().unwrap();
952        let err = endpoint.revalidate(&registry).expect_err("tampered endpoint");
953        assert_eq!(err.reason, "IP-literal host forbidden");
954
955        // A provider removed from the registry after validation also fails.
956        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(&registry).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        // localhost resolves via the host table — no network needed.
993        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        // Non-secret fields remain diagnostic.
1031        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        // A provider claiming no expiry — the unstoreable case.
1047        let permanent = TokenResponse { expires_in: None, ..with_expiry };
1048        assert_eq!(permanent.expires_at(now), None);
1049        // A non-positive expiry is not a lifetime either.
1050        let zero = TokenResponse { expires_in: Some(0), ..permanent };
1051        assert_eq!(zero.expires_at(now), None);
1052    }
1053}