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// ---------------------------------------------------------------------------