Skip to main content

fraiseql_auth/
pkce.rs

1//! PKCE state store — RFC 7636 Proof Key for Code Exchange.
2//!
3//! Stores `(code_verifier, redirect_uri)` under a random internal key while
4//! the OAuth2 authorization round-trip is in flight.  The token sent to the
5//! OIDC provider in the `?state=` query parameter is either:
6//! - the raw internal key (no encryption configured), or
7//! - `encrypt(internal_key)` (when [`crate::state_encryption::StateEncryptionService`] is
8//!   attached).
9//!
10//! State lifecycle:
11//! - `create_state(redirect_uri)` → `internal_key = random 32 bytes (base64url)` → `outbound_token
12//!   = encrypt(internal_key)` (or `internal_key` if no encryption) → `store.insert(internal_key,
13//!   {verifier, redirect_uri, ttl})` → returns `(outbound_token, verifier)`
14//! - `consume_state(outbound_token)` → `internal_key = decrypt(outbound_token)` (or
15//!   `outbound_token` if no encryption) → `entry = store.remove(internal_key)?` (
16//!   [`PkceError::StateNotFound`] if absent) → if `entry.elapsed > entry.ttl` →
17//!   [`PkceError::StateExpired`] → returns `{verifier, redirect_uri}`
18//!
19//! Backends:
20//! - **InMemory** — `DashMap`, single-process, per-replica
21//! - **Redis** — distributed, multi-replica (requires the `redis-pkce` Cargo feature)
22
23use std::{sync::Arc, time::Duration};
24
25use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
26use dashmap::DashMap;
27use rand::RngCore as _;
28use sha2::{Digest, Sha256};
29use thiserror::Error;
30
31use crate::state_encryption::StateEncryptionService;
32
33// ---------------------------------------------------------------------------
34// Constants
35// ---------------------------------------------------------------------------
36
37/// Maximum number of in-flight PKCE state entries allowed in the in-memory
38/// store at any one time.
39///
40/// This cap prevents the DashMap from growing without bound under load or
41/// during a DoS attack where an adversary initiates many OAuth flows without
42/// completing them. New inserts beyond this limit are rejected with
43/// [`PkceError::StoreFull`].
44///
45/// 10 000 entries corresponds to approximately 10 000 concurrent in-flight
46/// authorization requests, which is more than sufficient for any realistic
47/// single-node deployment. Multi-replica deployments should use the Redis
48/// backend instead.
49const MAX_PKCE_ENTRIES: usize = 10_000;
50
51// ---------------------------------------------------------------------------
52// Error type
53// ---------------------------------------------------------------------------
54
55/// Errors returned by [`PkceStateStore::consume_state`] and
56/// [`PkceStateStore::create_state`].
57#[derive(Debug, Error)]
58#[non_exhaustive]
59pub enum PkceError {
60    /// The state token was not found — either never issued, already consumed,
61    /// or (when encryption is on) tampered/decryption failed.
62    ///
63    /// Clients receive the same message for unknown and tampered tokens to
64    /// avoid leaking information about the store.
65    #[error(
66        "state not found — the authorization flow may have already been completed or the state is invalid"
67    )]
68    StateNotFound,
69
70    /// The state token was found but its TTL has elapsed.
71    ///
72    /// Distinct from [`PkceError::StateNotFound`] so that clients can show
73    /// a useful "please restart the authorization flow" message rather than
74    /// a generic invalid-state error.
75    #[error("state expired — please restart the authorization flow")]
76    StateExpired,
77
78    /// The in-memory store has reached `MAX_PKCE_ENTRIES` entries.
79    ///
80    /// This prevents unbounded memory growth under load or during a DoS
81    /// attack. Callers should return HTTP 429 to the client and invite them
82    /// to retry after a short delay.
83    #[error("PKCE state store is full — too many concurrent authorization flows")]
84    StoreFull,
85}
86
87// ---------------------------------------------------------------------------
88// Public consumed-state value
89// ---------------------------------------------------------------------------
90
91/// The data recovered after consuming a valid PKCE state token.
92#[derive(Debug)]
93pub struct ConsumedPkceState {
94    /// The `code_verifier` generated during `create_state`, needed for the
95    /// PKCE code exchange at `/token`.
96    pub verifier:     String,
97    /// The `redirect_uri` the client specified at `/auth/start`.
98    pub redirect_uri: String,
99}
100
101// ---------------------------------------------------------------------------
102// InMemoryPkceStateStore
103// ---------------------------------------------------------------------------
104
105struct PkceEntry {
106    verifier:     String,
107    redirect_uri: String,
108    /// Creation time as a Tokio instant so `tokio::time::pause()` +
109    /// `tokio::time::advance()` can control TTL expiry in tests.
110    created_at:   tokio::time::Instant,
111    ttl:          Duration,
112}
113
114/// In-memory PKCE state store backed by a [`DashMap`].
115///
116/// State is per-process: lost on restart, not shared across replicas.
117/// For multi-replica deployments, use `RedisPkceStateStore` instead
118/// (requires the `redis-pkce` Cargo feature).
119pub struct InMemoryPkceStateStore {
120    state_ttl_secs: u64,
121    entries:        DashMap<String, PkceEntry>,
122    encryptor:      Option<Arc<StateEncryptionService>>,
123    /// Maximum number of in-flight entries; defaults to `MAX_PKCE_ENTRIES`.
124    /// Overridable in tests via [`InMemoryPkceStateStore::with_max_entries`].
125    max_entries:    usize,
126}
127
128impl InMemoryPkceStateStore {
129    fn new(state_ttl_secs: u64, encryptor: Option<Arc<StateEncryptionService>>) -> Self {
130        Self {
131            state_ttl_secs,
132            entries: DashMap::new(),
133            encryptor,
134            max_entries: MAX_PKCE_ENTRIES,
135        }
136    }
137
138    /// Create a store with a custom entry cap — for testing only.
139    #[cfg(test)]
140    fn with_max_entries(
141        state_ttl_secs: u64,
142        encryptor: Option<Arc<StateEncryptionService>>,
143        max_entries: usize,
144    ) -> Self {
145        Self {
146            state_ttl_secs,
147            entries: DashMap::new(),
148            encryptor,
149            max_entries,
150        }
151    }
152
153    fn create_state_sync(&self, redirect_uri: &str) -> Result<(String, String), anyhow::Error> {
154        // SECURITY: reject inserts when the map is at capacity to prevent
155        // unbounded memory growth under DoS.  Callers translate this into
156        // HTTP 429 and invite the client to retry.
157        if self.entries.len() >= self.max_entries {
158            return Err(PkceError::StoreFull.into());
159        }
160
161        // code_verifier — RFC 7636 §4.1: 43–128 chars, [A-Za-z0-9\-._~]
162        let mut verifier_bytes = [0u8; 32];
163        rand::rng().fill_bytes(&mut verifier_bytes);
164        let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
165
166        // internal_key — separate from verifier so outbound token cannot reveal it
167        let mut key_bytes = [0u8; 32];
168        rand::rng().fill_bytes(&mut key_bytes);
169        let internal_key = URL_SAFE_NO_PAD.encode(key_bytes);
170
171        self.entries.insert(
172            internal_key.clone(),
173            PkceEntry {
174                verifier:     verifier.clone(),
175                redirect_uri: redirect_uri.to_owned(),
176                created_at:   tokio::time::Instant::now(),
177                ttl:          Duration::from_secs(self.state_ttl_secs),
178            },
179        );
180
181        let outbound_token = match &self.encryptor {
182            Some(enc) => enc.encrypt(internal_key.as_bytes())?,
183            None => internal_key,
184        };
185
186        Ok((outbound_token, verifier))
187    }
188
189    /// Remove all expired entries.
190    ///
191    /// Call this from a background task on a fixed interval to reclaim memory.
192    pub fn purge_expired(&self) {
193        self.entries.retain(|_, e| e.created_at.elapsed() <= e.ttl);
194    }
195
196    fn consume_state_sync(&self, outbound_token: &str) -> Result<ConsumedPkceState, PkceError> {
197        let internal_key = match &self.encryptor {
198            Some(enc) => {
199                let bytes = enc.decrypt(outbound_token).map_err(|_| PkceError::StateNotFound)?;
200                String::from_utf8(bytes).map_err(|_| PkceError::StateNotFound)?
201            },
202            None => outbound_token.to_owned(),
203        };
204
205        let (_, entry) = self.entries.remove(&internal_key).ok_or(PkceError::StateNotFound)?;
206
207        if entry.created_at.elapsed() > entry.ttl {
208            return Err(PkceError::StateExpired);
209        }
210
211        Ok(ConsumedPkceState {
212            verifier:     entry.verifier,
213            redirect_uri: entry.redirect_uri,
214        })
215    }
216
217    fn cleanup_expired_sync(&self) {
218        self.purge_expired();
219    }
220
221    fn len_sync(&self) -> usize {
222        self.entries.len()
223    }
224}
225
226// ---------------------------------------------------------------------------
227// Redis backend
228// ---------------------------------------------------------------------------
229
230/// Cumulative count of Redis PKCE store errors (unreachable Redis, etc.).
231///
232/// Exposed via `/metrics` as `fraiseql_pkce_redis_errors_total`.
233#[cfg(feature = "redis-pkce")]
234pub static REDIS_PKCE_ERRORS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
235
236/// Return the total number of Redis PKCE errors observed so far.
237#[cfg(feature = "redis-pkce")]
238pub fn redis_pkce_error_count_total() -> u64 {
239    REDIS_PKCE_ERRORS.load(std::sync::atomic::Ordering::Relaxed)
240}
241
242/// Redis-backed PKCE state store for distributed, multi-replica deployments.
243///
244/// Stores PKCE state tokens in Redis with TTL, enabling auth flows to be
245/// completed on any replica. Uses `GETDEL` for atomic one-shot consumption
246/// — a state token cannot be reused even under concurrent requests.
247///
248/// Key format:   `fraiseql:pkce:{internal_key}`
249/// Value format: `{"verifier":"...","redirect_uri":"..."}`
250#[cfg(feature = "redis-pkce")]
251pub struct RedisPkceStateStore {
252    pool:           redis::aio::ConnectionManager,
253    state_ttl_secs: u64,
254    encryptor:      Option<Arc<StateEncryptionService>>,
255}
256
257#[cfg(feature = "redis-pkce")]
258impl RedisPkceStateStore {
259    /// Connect to Redis and prepare the PKCE state store.
260    ///
261    /// # Errors
262    ///
263    /// Returns an error if the URL is invalid or the initial connection fails.
264    pub async fn new(
265        url: &str,
266        state_ttl_secs: u64,
267        encryptor: Option<Arc<StateEncryptionService>>,
268    ) -> Result<Self, redis::RedisError> {
269        let client = redis::Client::open(url)?;
270        let pool = redis::aio::ConnectionManager::new(client).await?;
271        Ok(Self {
272            pool,
273            state_ttl_secs,
274            encryptor,
275        })
276    }
277
278    async fn create_state_impl(
279        &self,
280        redirect_uri: &str,
281    ) -> Result<(String, String), anyhow::Error> {
282        // code_verifier (RFC 7636 §4.1)
283        let mut verifier_bytes = [0u8; 32];
284        rand::rng().fill_bytes(&mut verifier_bytes);
285        let verifier = URL_SAFE_NO_PAD.encode(verifier_bytes);
286
287        // Opaque internal key — separate from verifier
288        let mut key_bytes = [0u8; 32];
289        rand::rng().fill_bytes(&mut key_bytes);
290        let internal_key = URL_SAFE_NO_PAD.encode(key_bytes);
291
292        // Serialize state to JSON and store with TTL
293        let redis_key = format!("fraiseql:pkce:{internal_key}");
294        let value = serde_json::json!({
295            "verifier":     verifier,
296            "redirect_uri": redirect_uri,
297        })
298        .to_string();
299
300        let mut conn = self.pool.clone();
301        redis::cmd("SET")
302            .arg(&redis_key)
303            .arg(&value)
304            .arg("EX")
305            .arg(self.state_ttl_secs)
306            .query_async::<()>(&mut conn)
307            .await?;
308
309        let outbound_token = match &self.encryptor {
310            Some(enc) => enc.encrypt(internal_key.as_bytes())?,
311            None => internal_key,
312        };
313
314        Ok((outbound_token, verifier))
315    }
316
317    async fn consume_state_impl(
318        &self,
319        outbound_token: &str,
320    ) -> Result<ConsumedPkceState, PkceError> {
321        #[derive(serde::Deserialize)]
322        struct StoredEntry {
323            verifier:     String,
324            redirect_uri: String,
325        }
326
327        // Recover internal key from outbound token
328        let internal_key = match &self.encryptor {
329            Some(enc) => {
330                let bytes = enc.decrypt(outbound_token).map_err(|_| PkceError::StateNotFound)?;
331                String::from_utf8(bytes).map_err(|_| PkceError::StateNotFound)?
332            },
333            None => outbound_token.to_owned(),
334        };
335
336        let redis_key = format!("fraiseql:pkce:{internal_key}");
337        let mut conn = self.pool.clone();
338
339        // GETDEL — atomically retrieve and delete in a single round-trip.
340        // This guarantees one-shot consumption: no concurrent request can
341        // reuse the same state token, even without application-level locking.
342        let raw: Option<String> = redis::cmd("GETDEL")
343            .arg(&redis_key)
344            .query_async(&mut conn)
345            .await
346            .map_err(|_| PkceError::StateNotFound)?;
347
348        let json = raw.ok_or(PkceError::StateNotFound)?;
349
350        let entry: StoredEntry =
351            serde_json::from_str(&json).map_err(|_| PkceError::StateNotFound)?;
352
353        // Note: TTL expiry is handled by Redis — expired entries are absent.
354        // The Redis backend therefore never returns `PkceError::StateExpired`;
355        // callers receive `StateNotFound` for both absent and expired tokens.
356        Ok(ConsumedPkceState {
357            verifier:     entry.verifier,
358            redirect_uri: entry.redirect_uri,
359        })
360    }
361}
362
363// ---------------------------------------------------------------------------
364// PkceStateStore — unified public interface
365// ---------------------------------------------------------------------------
366
367/// PKCE state store that dispatches to an in-memory or Redis backend.
368///
369/// # Backends
370///
371/// - **InMemory** (default): per-process DashMap. Safe for single-replica deployments. State is
372///   lost on restart.
373///
374/// - **Redis** (requires `redis-pkce` Cargo feature): distributed, shared across all replicas.
375///   Required for multi-instance Kubernetes / ECS / fly.io deployments where `/auth/start` and
376///   `/auth/callback` may hit different nodes.
377///
378/// # Multi-replica requirement
379///
380/// Set `FRAISEQL_REQUIRE_REDIS=1` in your deployment environment to make
381/// FraiseQL refuse to start without a Redis-backed PKCE store. This is the
382/// recommended pattern for production Kubernetes deployments.
383#[non_exhaustive]
384pub enum PkceStateStore {
385    /// Single-node DashMap-backed store.
386    InMemory(InMemoryPkceStateStore),
387    /// Distributed Redis-backed store (requires `redis-pkce` Cargo feature).
388    #[cfg(feature = "redis-pkce")]
389    Redis(RedisPkceStateStore),
390}
391
392impl PkceStateStore {
393    /// Create an in-memory PKCE state store (single-replica deployments).
394    #[must_use]
395    pub fn new(state_ttl_secs: u64, encryptor: Option<Arc<StateEncryptionService>>) -> Self {
396        Self::InMemory(InMemoryPkceStateStore::new(state_ttl_secs, encryptor))
397    }
398
399    /// Create an in-memory PKCE state store with a custom entry cap.
400    ///
401    /// For testing only — use [`Self::new`] in production.
402    #[cfg(test)]
403    pub(crate) fn new_capped(
404        state_ttl_secs: u64,
405        encryptor: Option<Arc<StateEncryptionService>>,
406        max_entries: usize,
407    ) -> Self {
408        Self::InMemory(InMemoryPkceStateStore::with_max_entries(
409            state_ttl_secs,
410            encryptor,
411            max_entries,
412        ))
413    }
414
415    /// Create a Redis-backed distributed PKCE state store.
416    ///
417    /// # Errors
418    ///
419    /// Returns an error if the Redis URL is invalid or the connection fails.
420    #[cfg(feature = "redis-pkce")]
421    pub async fn new_redis(
422        url: &str,
423        state_ttl_secs: u64,
424        encryptor: Option<Arc<StateEncryptionService>>,
425    ) -> Result<Self, redis::RedisError> {
426        let inner = RedisPkceStateStore::new(url, state_ttl_secs, encryptor).await?;
427        Ok(Self::Redis(inner))
428    }
429
430    /// Returns `true` when backed by the in-memory DashMap store.
431    ///
432    /// Used by the `FRAISEQL_REQUIRE_REDIS` startup check.
433    #[must_use]
434    pub const fn is_in_memory(&self) -> bool {
435        matches!(self, Self::InMemory(_))
436    }
437
438    /// Generate an authorization-code verifier and reserve a state slot.
439    ///
440    /// Returns `(outbound_token, code_verifier)`:
441    /// - `outbound_token` goes in the OIDC `?state=` query parameter.
442    /// - `code_verifier` is passed to [`Self::s256_challenge`] and stored until the callback
443    ///   arrives.
444    ///
445    /// # Errors
446    ///
447    /// Returns an error if encryption fails (effectively never with a valid
448    /// key) or the Redis backend is unreachable.
449    pub async fn create_state(
450        &self,
451        redirect_uri: &str,
452    ) -> Result<(String, String), anyhow::Error> {
453        match self {
454            Self::InMemory(s) => s.create_state_sync(redirect_uri),
455            #[cfg(feature = "redis-pkce")]
456            Self::Redis(s) => s.create_state_impl(redirect_uri).await,
457        }
458    }
459
460    /// Consume a state token, atomically removing it from the store.
461    ///
462    /// Returns [`PkceError::StateNotFound`] for:
463    /// - tokens that were never issued,
464    /// - tokens that have already been consumed (one-time use), and
465    /// - tokens that fail decryption (tampered or from a different key).
466    ///
467    /// Returns [`PkceError::StateExpired`] when the in-memory token is valid
468    /// but its TTL has elapsed. The Redis backend returns `StateNotFound` for
469    /// expired tokens (Redis TTL handles expiry).
470    ///
471    /// # Errors
472    ///
473    /// Returns `PkceError::StateNotFound` if the token is unknown, already consumed,
474    /// or fails decryption. Returns `PkceError::StateExpired` if the token's TTL has elapsed.
475    pub async fn consume_state(
476        &self,
477        outbound_token: &str,
478    ) -> Result<ConsumedPkceState, PkceError> {
479        match self {
480            Self::InMemory(s) => s.consume_state_sync(outbound_token),
481            #[cfg(feature = "redis-pkce")]
482            Self::Redis(s) => s.consume_state_impl(outbound_token).await,
483        }
484    }
485
486    /// Compute the S256 code challenge for a given verifier.
487    ///
488    /// Per RFC 7636 §4.2:
489    /// `code_challenge = BASE64URL(SHA256(ASCII(code_verifier)))`
490    /// (no padding).
491    #[must_use]
492    pub fn s256_challenge(verifier: &str) -> String {
493        URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()))
494    }
495
496    /// Remove expired entries from the in-memory store.
497    ///
498    /// No-op for the Redis backend — Redis TTL handles expiry automatically.
499    /// Call from a background task on a fixed interval for the in-memory
500    /// backend to reclaim memory and free capacity below `MAX_PKCE_ENTRIES`.
501    pub async fn cleanup_expired(&self) {
502        match self {
503            Self::InMemory(s) => s.cleanup_expired_sync(),
504            #[cfg(feature = "redis-pkce")]
505            Self::Redis(_) => {}, // Redis TTL handles expiry
506        }
507    }
508
509    /// Synchronously remove all expired entries from the in-memory store.
510    ///
511    /// Identical to [`Self::cleanup_expired`] but callable from synchronous
512    /// contexts (e.g. maintenance hooks, benchmarks). No-op for the Redis
513    /// backend.
514    pub fn purge_expired(&self) {
515        match self {
516            Self::InMemory(s) => s.purge_expired(),
517            #[cfg(feature = "redis-pkce")]
518            Self::Redis(_) => {}, // Redis TTL handles expiry
519        }
520    }
521
522    /// Number of entries currently in the store.
523    ///
524    /// Returns 0 for the Redis backend — Redis state is not enumerable locally.
525    #[must_use]
526    pub fn len(&self) -> usize {
527        match self {
528            Self::InMemory(s) => s.len_sync(),
529            #[cfg(feature = "redis-pkce")]
530            Self::Redis(_) => 0,
531        }
532    }
533
534    /// Returns `true` when the in-memory store contains no entries.
535    ///
536    /// Always returns `true` for the Redis backend.
537    #[must_use]
538    pub fn is_empty(&self) -> bool {
539        self.len() == 0
540    }
541}
542
543// ---------------------------------------------------------------------------
544// Unit tests
545// ---------------------------------------------------------------------------