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}