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