Skip to main content

fraiseql_auth/
provider.rs

1//! OAuth 2.0 / OIDC provider trait and core data types.
2use std::fmt;
3
4use async_trait::async_trait;
5use serde::{Deserialize, Serialize};
6
7use crate::error::{AuthError, Result};
8
9/// User information retrieved from OAuth provider
10#[derive(Debug, Clone, Serialize, Deserialize)]
11pub struct UserInfo {
12    /// Unique user identifier from provider
13    pub id:             String,
14    /// User's email address.
15    ///
16    /// `None` when the provider omits an email claim (e.g. a GitHub account with a
17    /// private email). An empty/whitespace-only claim is normalized to `None` so it
18    /// can never be used as an account-linking key — see the [`crate::account_linking`]
19    /// module, which keys email-less identities on `(provider, provider_id)` instead.
20    pub email:          Option<String>,
21    /// Whether the provider asserts the email address is verified.
22    ///
23    /// Cross-provider account linking requires this to be `true`; an unverified or
24    /// absent `email_verified` claim is treated as `false` (fail-closed).
25    pub email_verified: bool,
26    /// User's display name (optional)
27    pub name:           Option<String>,
28    /// User's profile picture URL (optional)
29    pub picture:        Option<String>,
30    /// Raw claims from provider (for custom fields)
31    pub raw_claims:     serde_json::Value,
32}
33
34/// Token response from OAuth provider
35#[derive(Debug, Clone, Serialize, Deserialize)]
36pub struct TokenResponse {
37    /// Access token (short-lived)
38    pub access_token:  String,
39    /// Refresh token if provider supports it
40    pub refresh_token: Option<String>,
41    /// Token expiration in seconds
42    pub expires_in:    u64,
43    /// Token type (typically "Bearer")
44    pub token_type:    String,
45}
46
47/// OAuth 2.0 / OIDC provider trait
48///
49/// Implement this trait to add support for custom OAuth providers.
50// Reason: used as dyn Trait (Arc<dyn OAuthProvider>, Box<dyn OAuthProvider>); async_trait ensures
51// Send bounds and dyn-compatibility async_trait: dyn-dispatch required; remove when RTN + Send is
52// stable (RFC 3425)
53#[async_trait]
54pub trait OAuthProvider: Send + Sync + fmt::Debug {
55    /// Provider name for logging/debugging
56    fn name(&self) -> &str;
57
58    /// Generate authorization URL for user to visit
59    ///
60    /// # Arguments
61    /// * `state` - CSRF protection state (should be cryptographically random)
62    fn authorization_url(&self, state: &str) -> String;
63
64    /// Exchange authorization code for tokens
65    ///
66    /// # Arguments
67    /// * `code` - Authorization code from provider
68    ///
69    /// # Returns
70    /// Token response with access_token and optional refresh_token
71    async fn exchange_code(&self, code: &str) -> Result<TokenResponse>;
72
73    /// Get user information using access token
74    ///
75    /// # Arguments
76    /// * `access_token` - The access token to use for API call
77    ///
78    /// # Returns
79    /// UserInfo with user details from provider
80    async fn user_info(&self, access_token: &str) -> Result<UserInfo>;
81
82    /// Refresh the access token (optional, default returns error)
83    ///
84    /// # Arguments
85    /// * `refresh_token` - The refresh token
86    ///
87    /// # Returns
88    /// New TokenResponse if provider supports refresh
89    async fn refresh_token(&self, _refresh_token: &str) -> Result<TokenResponse> {
90        Err(AuthError::OAuthError {
91            message: format!("{} does not support token refresh", self.name()),
92        })
93    }
94
95    /// Revoke a token (optional, default is no-op)
96    ///
97    /// # Arguments
98    /// * `token` - Token to revoke
99    async fn revoke_token(&self, _token: &str) -> Result<()> {
100        Ok(())
101    }
102}
103
104/// PKCE (Proof Key for Public Clients) helper
105///
106/// Used to prevent authorization code interception attacks
107#[derive(Debug, Clone)]
108pub struct PkceChallenge {
109    /// Generated code verifier (cryptographically random)
110    pub verifier:  String,
111    /// Code challenge (SHA256 hash of verifier)
112    pub challenge: String,
113}
114
115impl PkceChallenge {
116    /// Generate a new PKCE challenge.
117    ///
118    /// # Errors
119    ///
120    /// Returns [`AuthError::PkceError`] if the generated verifier fails RFC 7636
121    /// length or character-set constraints (essentially never in practice).
122    pub fn generate() -> Result<Self> {
123        use sha2::{Digest, Sha256};
124
125        let verifier = generate_pkce_verifier()?;
126
127        let mut hasher = Sha256::new();
128        hasher.update(verifier.as_bytes());
129        let challenge_bytes = hasher.finalize();
130        let challenge = base64_url_encode(&challenge_bytes);
131
132        Ok(Self {
133            verifier,
134            challenge,
135        })
136    }
137
138    /// Validate a verifier against a challenge.
139    ///
140    /// Uses constant-time equality to prevent timing attacks on the PKCE challenge
141    /// (L-pkce-triplication: this path previously used `==`, drifting from the
142    /// constant-time `oauth::pkce::PkceChallenge::verify`).
143    #[must_use]
144    pub fn validate(&self, verifier: &str) -> bool {
145        use sha2::{Digest, Sha256};
146        use subtle::ConstantTimeEq as _;
147
148        let mut hasher = Sha256::new();
149        hasher.update(verifier.as_bytes());
150        let hash = hasher.finalize();
151        let encoded = base64_url_encode(&hash);
152
153        encoded.as_bytes().ct_eq(self.challenge.as_bytes()).into()
154    }
155}
156
157/// Generate a PKCE verifier (43-128 characters of unreserved characters)
158///
159/// # SECURITY
160///
161/// This uses `rand::rng()` which is cryptographically secure on all major platforms.
162/// It generates a 128-character random string using only unreserved characters as per RFC 7636.
163///
164/// The generated verifier meets these requirements:
165/// - Length: exactly 128 characters (within 43-128 range)
166/// - Characters: only unreserved ASCII characters: [A-Z a-z 0-9 - . _ ~]
167/// - Randomness: cryptographically secure pseudorandom generation
168/// - No padding: can be used directly in PKCE challenge
169///
170/// # Errors
171///
172/// Returns error if:
173/// - Random number generation fails (extremely rare)
174/// - Generated verifier is invalid (should never happen given the constraints)
175///
176/// # Implementation Notes
177///
178/// We use a fixed 128-character length (maximum allowed by RFC 7636) for:
179/// 1. Maximum security: more entropy means harder to guess
180/// 2. Consistency: predictable length for tests and monitoring
181/// 3. Compatibility: all OAuth providers support 128-char verifiers
182fn generate_pkce_verifier() -> Result<String> {
183    use rand::Rng;
184
185    const CHARSET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~";
186    const VERIFIER_LENGTH: usize = 128; // Maximum allowed by RFC 7636
187    const MIN_VERIFIER_LENGTH: usize = 43; // Minimum allowed by RFC 7636
188
189    // SECURITY: rand::rng() is backed by OS-level entropy for PKCE verifiers.
190    let mut rng = rand::rng();
191    let verifier: String = (0..VERIFIER_LENGTH)
192        .map(|_| {
193            let idx = rng.random_range(0..CHARSET.len());
194            CHARSET[idx] as char
195        })
196        .collect();
197
198    // Validate the generated verifier meets RFC 7636 requirements
199    if verifier.len() < MIN_VERIFIER_LENGTH {
200        return Err(AuthError::PkceError {
201            message: format!(
202                "Generated PKCE verifier too short: {} < {} chars",
203                verifier.len(),
204                MIN_VERIFIER_LENGTH
205            ),
206        });
207    }
208
209    if verifier.len() > 128 {
210        return Err(AuthError::PkceError {
211            message: format!("Generated PKCE verifier too long: {} > 128 chars", verifier.len()),
212        });
213    }
214
215    // Verify all characters are from the allowed charset
216    let allowed_chars: std::collections::HashSet<char> =
217        "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-._~"
218            .chars()
219            .collect();
220
221    for (i, c) in verifier.chars().enumerate() {
222        if !allowed_chars.contains(&c) {
223            return Err(AuthError::PkceError {
224                message: format!(
225                    "Generated PKCE verifier contains invalid character '{}' at position {}",
226                    c, i
227                ),
228            });
229        }
230    }
231
232    Ok(verifier)
233}
234
235/// URL-safe base64 encoding for PKCE
236pub(crate) fn base64_url_encode(bytes: &[u8]) -> String {
237    use base64::Engine;
238    base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes)
239}