Skip to main content

stano_security/
jwt.rs

1use crate::Claims;
2use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
3use serde::{Deserialize, Serialize};
4use thiserror::Error;
5
6/// ES256 EC key material and expiration used by [`encode_jwt`]/[`decode_jwt`].
7#[derive(Debug, Clone)]
8pub struct JwtConfig {
9    /// PEM-encoded EC private key, used to sign tokens.
10    pub private_key_pem: String,
11    /// PEM-encoded EC public key, used to verify tokens.
12    pub public_key_pem: String,
13    /// Token lifetime in seconds, applied when computing `exp`.
14    pub expiration_seconds: u64,
15}
16
17/// Errors returned by [`encode_jwt`]/[`decode_jwt`].
18#[derive(Debug, Error)]
19pub enum JwtError {
20    /// Token signing failed.
21    #[error("Failed to encode JWT: {0}")]
22    EncodingFailed(String),
23
24    /// Token verification/parsing failed.
25    #[error("Failed to decode JWT: {0}")]
26    DecodingFailed(String),
27
28    /// The configured PEM key could not be parsed.
29    #[error("Invalid key format: {0}")]
30    InvalidKey(String),
31
32    /// The token's `exp` claim is in the past.
33    #[error("Token expired")]
34    TokenExpired,
35
36    /// The token failed validation for a reason other than expiration.
37    #[error("Invalid token: {0}")]
38    InvalidToken(String),
39}
40
41/// Encode a Claims struct into a JWT token.
42pub fn encode_jwt<E>(claims: &Claims<E>, config: &JwtConfig) -> Result<String, JwtError>
43where
44    E: Serialize,
45{
46    let encoding_key = EncodingKey::from_ec_pem(config.private_key_pem.as_bytes())
47        .map_err(|e| JwtError::InvalidKey(e.to_string()))?;
48
49    encode(&Header::new(Algorithm::ES256), claims, &encoding_key)
50        .map_err(|e| JwtError::EncodingFailed(e.to_string()))
51}
52
53/// Decode and verify a JWT token, returning the Claims.
54pub fn decode_jwt<E>(token: &str, config: &JwtConfig) -> Result<Claims<E>, JwtError>
55where
56    E: for<'de> Deserialize<'de>,
57{
58    let decoding_key = DecodingKey::from_ec_pem(config.public_key_pem.as_bytes())
59        .map_err(|e| JwtError::InvalidKey(e.to_string()))?;
60
61    let validation = Validation::new(Algorithm::ES256);
62
63    decode::<Claims<E>>(token, &decoding_key, &validation)
64        .map(|token_data| token_data.claims)
65        .map_err(|e| {
66            if e.kind() == &jsonwebtoken::errors::ErrorKind::ExpiredSignature {
67                JwtError::TokenExpired
68            } else {
69                JwtError::DecodingFailed(e.to_string())
70            }
71        })
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77
78    // Test-only ES256 (P-256) EC keypair, generated solely for these unit tests.
79    const PRIVATE_KEY_PEM: &str = "-----BEGIN PRIVATE KEY-----
80MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQgtgbDmCbWzH1rPZlb
81qucYzcKQppWx4YxRh0TfnEd0wd6hRANCAATbjOo4G431D+jMHWgoGXaW/vr20Qxn
82QuoeHrU++Hh7LgqOwXbpqEmKfJa5Os5GQfdQ579fyDqZ/MepnZz2ijhz
83-----END PRIVATE KEY-----";
84
85    const PUBLIC_KEY_PEM: &str = "-----BEGIN PUBLIC KEY-----
86MFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE24zqOBuN9Q/ozB1oKBl2lv769tEM
87Z0LqHh61Pvh4ey4KjsF26ahJinyWuTrORkH3UOe/X8g6mfzHqZ2c9oo4cw==
88-----END PUBLIC KEY-----";
89
90    // A second, unrelated keypair used to exercise signature-mismatch failures.
91    const OTHER_PUBLIC_KEY_PEM: &str = "-----BEGIN PUBLIC KEY-----
92MFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAEQYaZ+hmOmyIcf6OlLbdfrdRDIQVP
93WvgpcJQZdAq9Q3dsB0xGIC4Ea8ps7xzypEj0W6wXZ/zgKyK9NSmDMtgzPg==
94-----END PUBLIC KEY-----";
95
96    fn config() -> JwtConfig {
97        JwtConfig {
98            private_key_pem: PRIVATE_KEY_PEM.to_string(),
99            public_key_pem: PUBLIC_KEY_PEM.to_string(),
100            expiration_seconds: 3600,
101        }
102    }
103
104    fn now() -> usize {
105        std::time::SystemTime::now()
106            .duration_since(std::time::UNIX_EPOCH)
107            .unwrap()
108            .as_secs() as usize
109    }
110
111    fn claims(exp: usize) -> Claims<()> {
112        Claims {
113            sub: "user-1".to_string(),
114            session_id: "session-1".to_string(),
115            exp,
116            ext: (),
117        }
118    }
119
120    #[test]
121    fn test_encode_decode_round_trip_success() {
122        let config = config();
123        let original = claims(now() + config.expiration_seconds as usize);
124        let token = encode_jwt(&original, &config).expect("encode should succeed");
125        let decoded: Claims<()> = decode_jwt(&token, &config).expect("decode should succeed");
126        assert_eq!(decoded.sub, original.sub);
127        assert_eq!(decoded.session_id, original.session_id);
128        assert_eq!(decoded.exp, original.exp);
129    }
130
131    #[test]
132    fn test_decode_expired_token_returns_token_expired() {
133        let config = config();
134        let expired = claims(now().saturating_sub(3600));
135        let token = encode_jwt(&expired, &config).expect("encode should succeed");
136        let result: Result<Claims<()>, JwtError> = decode_jwt(&token, &config);
137        assert!(matches!(result, Err(JwtError::TokenExpired)));
138    }
139
140    #[test]
141    fn test_decode_malformed_token_string_returns_decoding_failed() {
142        let config = config();
143        let result: Result<Claims<()>, JwtError> = decode_jwt("not-a-valid-jwt", &config);
144        assert!(matches!(result, Err(JwtError::DecodingFailed(_))));
145    }
146
147    #[test]
148    fn test_decode_with_mismatched_public_key_returns_decoding_failed() {
149        let signing_config = config();
150        let original = claims(now() + 3600);
151        let token = encode_jwt(&original, &signing_config).expect("encode should succeed");
152
153        let mut verifying_config = signing_config.clone();
154        verifying_config.public_key_pem = OTHER_PUBLIC_KEY_PEM.to_string();
155
156        let result: Result<Claims<()>, JwtError> = decode_jwt(&token, &verifying_config);
157        assert!(matches!(result, Err(JwtError::DecodingFailed(_))));
158    }
159
160    #[test]
161    fn test_encode_with_invalid_pem_returns_invalid_key() {
162        let mut config = config();
163        config.private_key_pem = "not a pem".to_string();
164        let original = claims(now() + 3600);
165        let result = encode_jwt(&original, &config);
166        assert!(matches!(result, Err(JwtError::InvalidKey(_))));
167    }
168
169    #[test]
170    fn test_decode_with_invalid_pem_returns_invalid_key() {
171        let mut config = config();
172        config.public_key_pem = "not a pem".to_string();
173        let result: Result<Claims<()>, JwtError> = decode_jwt("irrelevant.token.value", &config);
174        assert!(matches!(result, Err(JwtError::InvalidKey(_))));
175    }
176
177    #[derive(Debug, PartialEq, Clone, Serialize, Deserialize)]
178    struct CustomExt {
179        email: String,
180        role: String,
181    }
182
183    #[test]
184    fn test_encode_decode_round_trip_with_flatten_ext_struct() {
185        let config = config();
186        let original = Claims {
187            sub: "user-2".to_string(),
188            session_id: "session-2".to_string(),
189            exp: now() + 3600,
190            ext: CustomExt {
191                email: "user@example.com".to_string(),
192                role: "admin".to_string(),
193            },
194        };
195        let token = encode_jwt(&original, &config).expect("encode should succeed");
196        let decoded: Claims<CustomExt> =
197            decode_jwt(&token, &config).expect("decode should succeed");
198        assert_eq!(decoded.ext, original.ext);
199        assert_eq!(decoded.sub, original.sub);
200    }
201}