Skip to main content

fraiseql_auth/
state_encryption.rs

1//! State encryption for PKCE and OAuth state parameter protection.
2//!
3//! Encrypts OAuth `state` (and PKCE) blobs with AEAD ciphers so that the
4//! outbound token sent to the identity provider cannot be deciphered or
5//! tampered with by an attacker who intercepts the redirect.
6//!
7//! Supports two algorithms selectable at runtime:
8//! - [`EncryptionAlgorithm::Chacha20Poly1305`] (default, constant-time in software)
9//! - [`EncryptionAlgorithm::Aes256Gcm`] (hardware-accelerated on modern CPUs)
10
11use std::{fmt, sync::Arc};
12
13// aes_gcm and chacha20poly1305 both re-export the same underlying `aead` traits.
14// We import them once from chacha20poly1305 and reuse for both cipher types.
15use aes_gcm::Aes256Gcm;
16use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
17use chacha20poly1305::{
18    ChaCha20Poly1305, Nonce,
19    aead::{Aead, AeadCore, KeyInit, OsRng, Payload},
20};
21use rand::RngCore as _;
22use serde::{Deserialize, Serialize};
23use zeroize::Zeroizing;
24
25use crate::{AuthError, error::Result};
26
27/// Encrypted state container with nonce
28#[derive(Debug, Clone)]
29pub struct EncryptedState {
30    /// Ciphertext with authentication tag appended
31    pub ciphertext: Vec<u8>,
32    /// 96-bit nonce used for encryption
33    pub nonce:      [u8; 12],
34}
35
36impl EncryptedState {
37    /// Create new encrypted state
38    #[must_use]
39    pub const fn new(ciphertext: Vec<u8>, nonce: [u8; 12]) -> Self {
40        Self { ciphertext, nonce }
41    }
42
43    /// Serialize to bytes for storage
44    /// Format: [12-byte nonce][ciphertext with auth tag]
45    #[must_use]
46    pub fn to_bytes(&self) -> Vec<u8> {
47        let mut bytes = Vec::with_capacity(12 + self.ciphertext.len());
48        bytes.extend_from_slice(&self.nonce);
49        bytes.extend_from_slice(&self.ciphertext);
50        bytes
51    }
52
53    /// Deserialize from bytes.
54    ///
55    /// # Errors
56    ///
57    /// Returns [`AuthError::InvalidState`] if `bytes` is shorter than 12 bytes
58    /// (minimum nonce size).
59    pub fn from_bytes(bytes: &[u8]) -> Result<Self> {
60        if bytes.len() < 12 {
61            return Err(AuthError::InvalidState);
62        }
63
64        let mut nonce = [0u8; 12];
65        nonce.copy_from_slice(&bytes[0..12]);
66        let ciphertext = bytes[12..].to_vec();
67
68        Ok(Self::new(ciphertext, nonce))
69    }
70}
71
72/// State encryption using ChaCha20-Poly1305 AEAD
73///
74/// Provides authenticated encryption for OAuth state parameters.
75/// Uses a fixed encryption key for the deployment lifetime.
76/// Each encryption uses a random nonce for security.
77///
78/// # Security Properties
79/// - **Confidentiality**: State values are encrypted with ChaCha20
80/// - **Authenticity**: Authentication tag prevents tampering detection
81/// - **Replay Prevention**: Random nonce in each encryption
82/// - **Key Isolation**: Separate from signing keys, used only for state
83pub struct StateEncryption {
84    cipher: ChaCha20Poly1305,
85}
86
87impl StateEncryption {
88    /// Create a new state encryption instance
89    ///
90    /// # Arguments
91    /// * `key` - 32-byte encryption key (must be cryptographically random)
92    ///
93    /// # Errors
94    /// Returns error if key is invalid
95    pub fn new(key_bytes: &[u8; 32]) -> Result<Self> {
96        let cipher =
97            ChaCha20Poly1305::new_from_slice(key_bytes).map_err(|_| AuthError::ConfigError {
98                message: "Invalid state encryption key".to_string(),
99            })?;
100
101        Ok(Self { cipher })
102    }
103
104    /// Encrypt a state value
105    ///
106    /// Generates a random 96-bit nonce and encrypts the state using ChaCha20-Poly1305.
107    /// The authentication tag is appended to the ciphertext.
108    ///
109    /// # Arguments
110    /// * `state` - The plaintext state value to encrypt
111    ///
112    /// # Returns
113    /// EncryptedState containing ciphertext and nonce
114    ///
115    /// # Errors
116    /// Returns error if encryption fails (should be rare)
117    pub fn encrypt(&self, state: &str) -> Result<EncryptedState> {
118        // Generate random 96-bit nonce
119        let mut nonce_bytes = [0u8; 12];
120        rand::rng().fill_bytes(&mut nonce_bytes);
121        let nonce = Nonce::from(nonce_bytes);
122
123        // Encrypt with AEAD (includes authentication tag)
124        let ciphertext =
125            self.cipher.encrypt(&nonce, Payload::from(state.as_bytes())).map_err(|_| {
126                AuthError::Internal {
127                    message: "State encryption failed".to_string(),
128                }
129            })?;
130
131        Ok(EncryptedState::new(ciphertext, nonce_bytes))
132    }
133
134    /// Decrypt and verify a state value
135    ///
136    /// Uses the nonce from EncryptedState to decrypt the ciphertext.
137    /// Authentication tag verification is automatic - tampering is detected.
138    ///
139    /// # Arguments
140    /// * `encrypted` - The encrypted state to decrypt
141    ///
142    /// # Returns
143    /// The decrypted plaintext state value
144    ///
145    /// # Errors
146    /// Returns error if:
147    /// - Authentication tag verification fails (tampering detected)
148    /// - Decryption fails
149    /// - Result is not valid UTF-8
150    pub fn decrypt(&self, encrypted: &EncryptedState) -> Result<String> {
151        let nonce = Nonce::from(encrypted.nonce);
152
153        // Decrypt and verify authentication tag
154        let plaintext = self
155            .cipher
156            .decrypt(&nonce, Payload::from(encrypted.ciphertext.as_slice()))
157            .map_err(|_| AuthError::InvalidState)?;
158
159        // Convert bytes to UTF-8 string
160        String::from_utf8(plaintext).map_err(|_| AuthError::InvalidState)
161    }
162
163    /// Encrypt state and serialize to bytes.
164    ///
165    /// # Errors
166    ///
167    /// Returns [`AuthError::Internal`] if AEAD encryption fails (essentially never).
168    pub fn encrypt_to_bytes(&self, state: &str) -> Result<Vec<u8>> {
169        let encrypted = self.encrypt(state)?;
170        Ok(encrypted.to_bytes())
171    }
172
173    /// Decrypt state from serialized bytes.
174    ///
175    /// # Errors
176    ///
177    /// Returns [`AuthError::InvalidState`] if `bytes` is too short, if AEAD
178    /// authentication fails (tampered or wrong key), or if decrypted bytes are
179    /// not valid UTF-8.
180    pub fn decrypt_from_bytes(&self, bytes: &[u8]) -> Result<String> {
181        let encrypted = EncryptedState::from_bytes(bytes)?;
182        self.decrypt(&encrypted)
183    }
184}
185
186/// Generate a cryptographically random encryption key
187#[must_use]
188pub fn generate_state_encryption_key() -> Zeroizing<[u8; 32]> {
189    let mut key = [0u8; 32];
190    rand::rng().fill_bytes(&mut key);
191    Zeroizing::new(key)
192}
193
194// ── StateEncryptionService ────────────────────────────────────────────────────
195//
196// A higher-level service that wraps the low-level `StateEncryption` struct.
197// Differences from `StateEncryption`:
198//   - Supports both ChaCha20-Poly1305 AND AES-256-GCM (runtime-selectable)
199//   - Wire format: URL-safe base64 of `[12-byte nonce || ciphertext || tag]`
200//   - Accepts keys as 64-char hex strings or env-var names
201//   - Can be constructed from the compiled schema JSON
202//   - Key never appears in `Debug` output
203//
204// This is the PKCE state encryption service wired into `Server`.
205
206/// Errors that can occur during decryption by `StateEncryptionService`.
207#[derive(Debug, thiserror::Error)]
208#[non_exhaustive]
209pub enum DecryptionError {
210    /// Ciphertext was tampered with or encrypted with a different key.
211    #[error("authentication failed — ciphertext may be tampered or key is wrong")]
212    AuthenticationFailed,
213    /// Input is malformed (empty, too short, bad base64, etc.).
214    #[error("invalid input: {0}")]
215    InvalidInput(String),
216}
217
218/// Errors that can occur when constructing a `StateEncryptionService` key.
219#[derive(Debug, thiserror::Error)]
220#[non_exhaustive]
221pub enum KeyError {
222    /// Hex string was not 64 characters (32 bytes).
223    #[error("hex key must be 64 chars (32 bytes); got {0} chars")]
224    WrongLength(usize),
225    /// Hex string contained a non-hex character.
226    #[error("invalid hex character in key")]
227    InvalidHex,
228}
229
230/// AEAD algorithm selection for `StateEncryptionService`.
231#[derive(Debug, Clone, Default, Deserialize, Serialize, PartialEq, Eq)]
232#[non_exhaustive]
233pub enum EncryptionAlgorithm {
234    /// ChaCha20-Poly1305 (recommended — constant-time, software-friendly).
235    #[default]
236    #[serde(rename = "chacha20-poly1305")]
237    Chacha20Poly1305,
238    /// AES-256-GCM (hardware-accelerated on modern CPUs).
239    #[serde(rename = "aes-256-gcm")]
240    Aes256Gcm,
241}
242
243impl fmt::Display for EncryptionAlgorithm {
244    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
245        match self {
246            Self::Chacha20Poly1305 => f.write_str("chacha20-poly1305"),
247            Self::Aes256Gcm => f.write_str("aes-256-gcm"),
248        }
249    }
250}
251
252/// Deserialized from `compiled.security.state_encryption`.
253#[derive(Debug, Clone, Deserialize, Serialize)]
254#[serde(default)]
255pub struct StateEncryptionConfig {
256    /// Enable the service; when `false`, `from_compiled_schema` returns `None`.
257    pub enabled:   bool,
258    /// AEAD algorithm to use.
259    pub algorithm: EncryptionAlgorithm,
260    /// Name of the environment variable holding the 64-char hex key.
261    pub key_env:   Option<String>,
262}
263
264impl Default for StateEncryptionConfig {
265    fn default() -> Self {
266        Self {
267            enabled:   false,
268            algorithm: EncryptionAlgorithm::default(),
269            key_env:   Some("STATE_ENCRYPTION_KEY".to_string()),
270        }
271    }
272}
273
274/// AEAD encryption service for OAuth state and PKCE blobs.
275///
276/// Wire format: URL-safe base64 of `[12-byte nonce || ciphertext || 16-byte tag]`.
277///
278/// The 32-byte key is never printed in [`fmt::Debug`] output.
279pub struct StateEncryptionService {
280    algorithm: EncryptionAlgorithm,
281    key:       [u8; 32],
282}
283
284impl fmt::Debug for StateEncryptionService {
285    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
286        f.debug_struct("StateEncryptionService")
287            .field("algorithm", &self.algorithm)
288            .field("key", &"[REDACTED]")
289            .finish()
290    }
291}
292
293impl StateEncryptionService {
294    /// Construct from a raw 32-byte key slice.
295    #[must_use]
296    pub const fn from_raw_key(key: &[u8; 32], algorithm: EncryptionAlgorithm) -> Self {
297        Self {
298            algorithm,
299            key: *key,
300        }
301    }
302
303    /// Construct from a 64-character hex string (= 32 bytes).
304    ///
305    /// # Errors
306    ///
307    /// Returns [`KeyError::WrongLength`] if `hex` is not 64 chars.
308    /// Returns [`KeyError::InvalidHex`] if `hex` contains non-hex chars.
309    pub fn from_hex_key(
310        hex: &str,
311        algorithm: EncryptionAlgorithm,
312    ) -> std::result::Result<Self, KeyError> {
313        if hex.len() != 64 {
314            return Err(KeyError::WrongLength(hex.len()));
315        }
316        let bytes = hex::decode(hex).map_err(|_| KeyError::InvalidHex)?;
317        let mut key = [0u8; 32];
318        key.copy_from_slice(&bytes);
319        Ok(Self { algorithm, key })
320    }
321
322    /// Load the key from an environment variable containing a 64-char hex string.
323    ///
324    /// # Errors
325    ///
326    /// Returns an error if the env var is absent or the value is not valid hex/length.
327    pub fn new_from_env(
328        var: &str,
329        algorithm: EncryptionAlgorithm,
330    ) -> std::result::Result<Self, anyhow::Error> {
331        let hex = std::env::var(var).map_err(|_| anyhow::anyhow!("env var {var} not set"))?;
332        Ok(Self::from_hex_key(&hex, algorithm)?)
333    }
334
335    /// Build from the `security` blob of a compiled schema, if enabled.
336    ///
337    /// Returns `Ok(None)` when the `state_encryption` key is absent or `enabled = false`.
338    ///
339    /// # Errors
340    ///
341    /// Returns `Err` when `enabled = true` but the key environment variable is absent
342    /// or contains an invalid value.  The server must refuse to start in this case.
343    pub fn from_compiled_schema(
344        security_json: &serde_json::Value,
345    ) -> std::result::Result<Option<Arc<Self>>, anyhow::Error> {
346        let cfg: StateEncryptionConfig = match security_json.get("state_encryption") {
347            None | Some(serde_json::Value::Null) => return Ok(None),
348            Some(v) => serde_json::from_value(v.clone())
349                .map_err(|e| anyhow::anyhow!("invalid state_encryption config: {e}"))?,
350        };
351
352        if !cfg.enabled {
353            return Ok(None);
354        }
355
356        let key_env = cfg.key_env.as_deref().unwrap_or("STATE_ENCRYPTION_KEY");
357        Self::new_from_env(key_env, cfg.algorithm)
358            .map(|svc| Some(Arc::new(svc)))
359            .map_err(|e| {
360                anyhow::anyhow!(
361                    "state_encryption enabled but key env var '{}' failed: {e}",
362                    key_env
363                )
364            })
365    }
366
367    /// Encrypt `plaintext` to a URL-safe base64 string.
368    ///
369    /// A fresh random nonce is generated on every call.
370    ///
371    /// # Errors
372    ///
373    /// Returns an error only on internal cipher failure (essentially never).
374    pub fn encrypt(&self, plaintext: &[u8]) -> std::result::Result<String, anyhow::Error> {
375        let combined = match self.algorithm {
376            EncryptionAlgorithm::Chacha20Poly1305 => {
377                let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
378                    .map_err(|_| anyhow::anyhow!("invalid key for ChaCha20-Poly1305"))?;
379                let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
380                let ct = cipher
381                    .encrypt(&nonce, plaintext)
382                    .map_err(|_| anyhow::anyhow!("ChaCha20-Poly1305 encryption failed"))?;
383                let mut out = nonce.to_vec();
384                out.extend_from_slice(&ct);
385                out
386            },
387            EncryptionAlgorithm::Aes256Gcm => {
388                let cipher = Aes256Gcm::new_from_slice(&self.key)
389                    .map_err(|_| anyhow::anyhow!("invalid key for AES-256-GCM"))?;
390                let nonce = Aes256Gcm::generate_nonce(&mut OsRng);
391                let ct = cipher
392                    .encrypt(&nonce, plaintext)
393                    .map_err(|_| anyhow::anyhow!("AES-256-GCM encryption failed"))?;
394                let mut out = nonce.to_vec();
395                out.extend_from_slice(&ct);
396                out
397            },
398        };
399        Ok(URL_SAFE_NO_PAD.encode(&combined))
400    }
401
402    /// Decrypt a URL-safe base64 string produced by [`Self::encrypt`].
403    ///
404    /// # Errors
405    ///
406    /// - [`DecryptionError::InvalidInput`] — empty / too-short / bad base64
407    /// - [`DecryptionError::AuthenticationFailed`] — tampered or wrong-key
408    pub fn decrypt(&self, encoded: &str) -> std::result::Result<Vec<u8>, DecryptionError> {
409        const NONCE_SIZE: usize = 12;
410        if encoded.is_empty() {
411            return Err(DecryptionError::InvalidInput("empty input".into()));
412        }
413        let combined = URL_SAFE_NO_PAD
414            .decode(encoded)
415            .map_err(|_| DecryptionError::InvalidInput("invalid base64".into()))?;
416
417        if combined.len() < NONCE_SIZE {
418            return Err(DecryptionError::InvalidInput(format!(
419                "too short: {} bytes (minimum {NONCE_SIZE})",
420                combined.len()
421            )));
422        }
423        let (nonce_bytes, ct) = combined.split_at(NONCE_SIZE);
424
425        match self.algorithm {
426            EncryptionAlgorithm::Chacha20Poly1305 => {
427                let cipher = ChaCha20Poly1305::new_from_slice(&self.key)
428                    .map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
429                let nonce = chacha20poly1305::Nonce::from_slice(nonce_bytes);
430                cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
431            },
432            EncryptionAlgorithm::Aes256Gcm => {
433                let cipher = Aes256Gcm::new_from_slice(&self.key)
434                    .map_err(|_| DecryptionError::InvalidInput("invalid key".into()))?;
435                let nonce = aes_gcm::Nonce::from_slice(nonce_bytes);
436                cipher.decrypt(nonce, ct).map_err(|_| DecryptionError::AuthenticationFailed)
437            },
438        }
439    }
440}