Skip to main content

authkestra_engine/token/
mod.rs

1use crate::auth::{error::AuthError, state::Identity};
2
3use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation};
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6
7#[derive(Debug, Serialize, Deserialize, Clone)]
8pub struct Claims {
9    // Standard OIDC claims
10    pub iss: Option<String>,
11    pub sub: String,
12    pub aud: Option<String>,
13    pub exp: usize,
14    pub iat: usize,
15    pub nbf: Option<usize>,
16    pub jti: Option<String>,
17
18    // Authkestra-specific core fields
19    pub scope: Option<String>,
20    /// Optional identity data for user-centric tokens.
21    /// If None, this is likely a machine-to-machine token.
22    #[serde(skip_serializing_if = "Option::is_none")]
23    pub identity: Option<Identity>,
24
25    // Isolated custom claims
26    #[serde(flatten)]
27    pub extra: HashMap<String, serde_json::Value>,
28}
29
30#[derive(Clone)]
31pub struct TokenManager {
32    encoding_key: EncodingKey,
33    decoding_key: DecodingKey,
34    issuer: Option<String>,
35    kid: Option<String>,
36    alg: Algorithm,
37    public_jwk: Option<crate::token::jwk::Jwk>,
38}
39
40impl TokenManager {
41    /// Creates a TokenManager for symmetric signing (HS256).
42    pub fn new(secret: &[u8], issuer: Option<String>) -> Self {
43        Self {
44            encoding_key: EncodingKey::from_secret(secret),
45            decoding_key: DecodingKey::from_secret(secret),
46            issuer,
47            kid: None,
48            alg: Algorithm::HS256,
49            public_jwk: None,
50        }
51    }
52
53    /// Creates a TokenManager for asymmetric signing (RS256).
54    /// `private_key_pem` must be a valid RSA private key in PEM format.
55    /// OP/external verification should use this path; internal resource servers
56    /// can continue to use `new` (HS256).
57    pub fn new_asymmetric(
58        private_key_pem: &[u8],
59        issuer: Option<String>,
60        kid: Option<String>,
61    ) -> Result<Self, AuthError> {
62        let encoding_key = EncodingKey::from_rsa_pem(private_key_pem)
63            .map_err(|e| AuthError::Token(e.to_string()))?;
64        let decoding_key = DecodingKey::from_rsa_pem(private_key_pem)
65            .map_err(|e| AuthError::Token(e.to_string()))?;
66
67        let pem_str = std::str::from_utf8(private_key_pem)
68            .map_err(|_| AuthError::Token("Invalid PEM UTF-8".into()))?;
69
70        use rsa::pkcs1::DecodeRsaPrivateKey;
71        use rsa::pkcs8::DecodePrivateKey;
72        let rsa_key = rsa::RsaPrivateKey::from_pkcs8_pem(pem_str)
73            .or_else(|_| rsa::RsaPrivateKey::from_pkcs1_pem(pem_str))
74            .map_err(|e| AuthError::Token(format!("Failed to parse RSA key: {}", e)))?;
75
76        use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
77        use rsa::traits::PublicKeyParts;
78
79        let n = URL_SAFE_NO_PAD.encode(rsa_key.n().to_bytes_be());
80        let e = URL_SAFE_NO_PAD.encode(rsa_key.e().to_bytes_be());
81
82        let kid_val = kid.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
83
84        let jwk = crate::token::jwk::Jwk {
85            kid: Some(kid_val.clone()),
86            kty: "RSA".to_string(),
87            alg: Some("RS256".to_string()),
88            n: Some(n),
89            e: Some(e),
90        };
91
92        Ok(Self {
93            encoding_key,
94            decoding_key,
95            issuer,
96            kid: Some(kid_val),
97            alg: Algorithm::RS256,
98            public_jwk: Some(jwk),
99        })
100    }
101
102    pub fn public_jwk(&self) -> Option<crate::token::jwk::Jwk> {
103        self.public_jwk.clone()
104    }
105
106    pub fn with_issuer(mut self, issuer: String) -> Self {
107        self.issuer = Some(issuer);
108        self
109    }
110
111    /// Issues a token for a user identity.
112    pub fn issue_user_token(
113        &self,
114        identity: Identity,
115        expires_in_secs: u64,
116        scope: Option<String>,
117        aud: Option<String>,
118    ) -> Result<String, AuthError> {
119        let now = chrono::Utc::now().timestamp() as usize;
120        let expiration = now + expires_in_secs as usize;
121
122        let claims = Claims {
123            iss: self.issuer.clone(),
124            sub: identity.external_id.clone(),
125            aud,
126            exp: expiration,
127            iat: now,
128            nbf: Some(now),
129            jti: Some(uuid::Uuid::new_v4().to_string()),
130            scope,
131            identity: Some(identity),
132            extra: HashMap::new(),
133        };
134
135        let mut header = Header::new(self.alg);
136        if let Some(ref kid) = self.kid {
137            header.kid = Some(kid.clone());
138        }
139
140        encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
141    }
142
143    /// Issues an OIDC-conformant ID token.
144    pub fn issue_id_token(
145        &self,
146        identity: Identity,
147        client_id: &str,
148        nonce: Option<String>,
149        expires_in_secs: u64,
150    ) -> Result<String, AuthError> {
151        let now = chrono::Utc::now().timestamp() as usize;
152        let expiration = now + expires_in_secs as usize;
153
154        let mut claims = Claims {
155            iss: self.issuer.clone(),
156            sub: identity.external_id.clone(),
157            aud: Some(client_id.to_string()),
158            exp: expiration,
159            iat: now,
160            nbf: Some(now),
161            jti: Some(uuid::Uuid::new_v4().to_string()),
162            scope: None,
163            identity: Some(identity),
164            extra: HashMap::new(),
165        };
166
167        if let Some(n) = nonce {
168            claims
169                .extra
170                .insert("nonce".to_string(), serde_json::Value::String(n));
171        }
172
173        let mut header = Header::new(self.alg);
174        if let Some(ref kid) = self.kid {
175            header.kid = Some(kid.clone());
176        }
177
178        encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
179    }
180
181    /// Issues a machine-to-machine (M2M) token for a client.
182    pub fn issue_client_token(
183        &self,
184        client_id: &str,
185        expires_in_secs: u64,
186        scope: Option<String>,
187        aud: Option<String>,
188    ) -> Result<String, AuthError> {
189        let now = chrono::Utc::now().timestamp() as usize;
190        let expiration = now + expires_in_secs as usize;
191
192        let claims = Claims {
193            iss: self.issuer.clone(),
194            sub: client_id.to_string(),
195            aud,
196            exp: expiration,
197            iat: now,
198            nbf: Some(now),
199            jti: Some(uuid::Uuid::new_v4().to_string()),
200            scope,
201            identity: None,
202            extra: HashMap::new(),
203        };
204
205        let mut header = Header::new(self.alg);
206        if let Some(ref kid) = self.kid {
207            header.kid = Some(kid.clone());
208        }
209
210        encode(&header, &claims, &self.encoding_key).map_err(|e| AuthError::Token(e.to_string()))
211    }
212
213    pub fn validate_token(
214        &self,
215        token: &str,
216        expected_aud: Option<&str>,
217    ) -> Result<Claims, AuthError> {
218        let mut validation = Validation::new(self.alg);
219        if let Some(aud) = expected_aud {
220            validation.set_audience(&[aud]);
221        } else {
222            validation.validate_aud = false;
223        }
224        if let Some(ref iss) = self.issuer {
225            validation.set_issuer(&[iss]);
226        }
227
228        let token_data = decode::<Claims>(token, &self.decoding_key, &validation)
229            .map_err(|e| AuthError::Token(e.to_string()))?;
230
231        Ok(token_data.claims)
232    }
233}
234
235#[cfg(test)]
236mod tests {
237
238    use super::*;
239    use crate::auth::state::Identity;
240    use std::collections::HashMap;
241
242    #[test]
243    fn test_claims_serialization() {
244        let mut extra = HashMap::new();
245        extra.insert(
246            "custom".to_string(),
247            serde_json::Value::String("value".to_string()),
248        );
249
250        let claims = Claims {
251            iss: Some("issuer".to_string()),
252            sub: "user123".to_string(),
253            aud: Some("audience".to_string()),
254            exp: 1000,
255            iat: 500,
256            nbf: Some(500),
257            jti: Some("jti".to_string()),
258            scope: Some("openid profile".to_string()),
259            identity: Some(Identity {
260                provider_id: "google".to_string(),
261                external_id: "user123".to_string(),
262                email: Some("user@example.com".to_string()),
263                username: Some("user".to_string()),
264                attributes: HashMap::new(),
265            }),
266            extra,
267        };
268
269        let serialized = serde_json::to_string(&claims).unwrap();
270        let deserialized: Claims = serde_json::from_str(&serialized).unwrap();
271
272        assert_eq!(deserialized.iss, claims.iss);
273        assert_eq!(deserialized.sub, claims.sub);
274        assert_eq!(deserialized.extra.get("custom").unwrap(), "value");
275    }
276
277    #[test]
278    fn test_token_manager_issuance() {
279        let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
280        let identity = Identity {
281            provider_id: "mock".to_string(),
282            external_id: "user123".to_string(),
283            email: None,
284            username: None,
285            attributes: HashMap::new(),
286        };
287
288        let token = manager
289            .issue_user_token(identity, 3600, None, None)
290            .unwrap();
291        let claims = manager.validate_token(&token, None).unwrap();
292
293        assert_eq!(claims.iss, Some("issuer".to_string()));
294        assert_eq!(claims.sub, "user123");
295        assert!(claims.jti.is_some());
296        assert!(claims.nbf.is_some());
297    }
298
299    #[test]
300    fn test_token_manager_asymmetric_issuance() {
301        let pem = b"-----BEGIN PRIVATE KEY-----
302MIIEvAIBADANBgkqhkiG9w0BAQEFAASCBKYwggSiAgEAAoIBAQDA5hJIcQ+2rxMz
303VM8ZH5WAmguCr0xmNDAdy0IzzsUeFLG7BebB7izOkU36J4t8t5tUaQwrBMnx2Fvt
304VqJjbdE242UDpvWF/8m9zJ2HR5298cbwT5cGMKLB0HWzDMahugs+Bbh2lCgwyLZk
305Tr3Diwxp5SwFew/Wb+Ke9cNG9Hu5IFH3BCuJ839d9hfqisIeYrBPfb52xxckM37R
3067zSGu/eDP/HZAeLkQuptZJW4A3u7xni14u4qyqXDqsHsYFNgJaxMSAwWgBRY6HNu
307TnvBArTXCiVfL+F73B2L6mdYr64g+QS9nK9v97MlJu/E3mSduz54pren4mpCHc9m
308/S2+VjCZAgMBAAECggEAASC9qQbGnL7XuExRDOIn/m4bWx92ehjo0lCTibhpY3LW
309umbSbpfbhmmuSj3CjW9VZsaM3hBTgSjoTX72lbY/eIUXD7c0memUK5pV4XcEIrQw
310AZlPIye6ckx4I7ZGnKasO8FoAel9dd7DXw36AuBK3LBzJwtzkEFsBc0e3/wixqmG
311UJBbbt/+5ya7CxyjuePaQhKtkLD5R6DpvN2XnCYq5nHJNJdvSVg1pOzsTHYIf+Ee
3122Rz42fGsfFKqeEQCcBFRZaGb/ELeP4c6UZdktZAvmHb1p1fursVZc6X9JXmiJ2OJ
313Kv2H2tMKuysP8L0fXFOMgkH2SVt6rcdHkO6xhlhWsQKBgQDqR8rAJeEE5BFoXA8T
314VVW6CLMlW51x4ey7PEGOaYh39dTG2Q+GZQBZ9G+SZk3f5Y85UCACSyc//4qaz/c3
3150nWsegZ+JPyymmuc79wzIAFFvXB7pL6wyn0Ed1P620kOZTtA8iBcXrsuxL+KP7iu
316MXfWmU1QiZpbndILtyDnY+70uwKBgQDSyCljWkydQCaPU+fiAXLxP8CvcJTSSNQD
317mVUlwJ+OpHnU+Alsi1rBavMgUtLlYbFqzH7NmYrLC8Yadq3ZOwLt0VEK0r8qstAL
3187QCDUD2WNuQjpZupRnXuMUl3iXB96i2gb+VQKGuUAJvVWjdIbYa4+Gu+sBMfcDcX
319dBihDLuEuwKBgAgX4tEwfc2Fc3R/eaXZVNTQaB/qQk4k1+C//CPHUYeTXn5gEUE7
320S//PiesszZPmgkQgmHp7zidP1KH0fT3Yb2g97ut8q54f54fMYXcCrAiUusYKsuu4
321kwkMdkI8QRHWPW3I74VBYIYFFfjYqrCZ1OH8+cbGeiagFRmCggh8U0zxAoGAVW3u
3226Ge22Z0gg8LcHsu7jG/sZq7Ygool8/d3fT+e669Z+ak2GJo6hF4WgClRdMqtn72W
323PzpV+ImjFyK2v26dd0n48MwN0v56N/ss1Av3iiRhPtlmR6tZLNspDZvUzhPVvkrb
324xCs9vtSoVEamVWKe0eVNthGjDoDqs0TInq2MavUCgYB6REavSJs/CLkSS7iimjxZ
325G7g5YQi9/p1lXLOEUDiwEmvRr0XTwzzxUsIc535IXhh/ZUYpthenW+qBBzn85pEC
326TowIqciHu5redqlQ8rITA8/AOY98vaDIhppDg1rfpnHHaZHFbXD/keYAEbhBtbvf
327a0QMqKUcs8+YTy5R5K6qtw==
328-----END PRIVATE KEY-----";
329
330        let manager = TokenManager::new_asymmetric(
331            pem,
332            Some("issuer".to_string()),
333            Some("my-kid-123".to_string()),
334        )
335        .unwrap();
336
337        let identity = Identity {
338            provider_id: "mock".to_string(),
339            external_id: "user123".to_string(),
340            email: None,
341            username: None,
342            attributes: HashMap::new(),
343        };
344
345        let token = manager
346            .issue_user_token(identity, 3600, None, None)
347            .unwrap();
348
349        // Decode directly via jsonwebtoken to prove independent verification
350        let jwk = manager.public_jwk().unwrap();
351        assert_eq!(jwk.kid.as_deref(), Some("my-kid-123"));
352
353        let decoding_key = jwk.to_decoding_key().unwrap();
354        let mut validation = jsonwebtoken::Validation::new(jsonwebtoken::Algorithm::RS256);
355        validation.set_issuer(&["issuer"]);
356
357        let token_data =
358            jsonwebtoken::decode::<Claims>(&token, &decoding_key, &validation).unwrap();
359        assert_eq!(token_data.claims.sub, "user123");
360        assert_eq!(token_data.header.kid.as_deref(), Some("my-kid-123"));
361    }
362
363    #[test]
364    fn test_issue_id_token() {
365        let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
366        let identity = Identity {
367            provider_id: "mock".to_string(),
368            external_id: "user123".to_string(),
369            email: None,
370            username: None,
371            attributes: HashMap::new(),
372        };
373
374        let token = manager
375            .issue_id_token(identity, "client-1", Some("nonce123".to_string()), 3600)
376            .unwrap();
377
378        let claims = manager.validate_token(&token, None).unwrap();
379
380        assert_eq!(claims.iss, Some("issuer".to_string()));
381        assert_eq!(claims.sub, "user123");
382        assert_eq!(claims.aud, Some("client-1".to_string()));
383        assert_eq!(claims.extra.get("nonce").unwrap(), "nonce123");
384    }
385    #[test]
386    fn test_token_manager_audience_validation() {
387        let manager = TokenManager::new(b"secret", Some("issuer".to_string()));
388        let identity = Identity {
389            provider_id: "mock".to_string(),
390            external_id: "user123".to_string(),
391            email: None,
392            username: None,
393            attributes: HashMap::new(),
394        };
395
396        // Issue token for "client-1"
397        let token = manager
398            .issue_id_token(identity, "client-1", None, 3600)
399            .unwrap();
400
401        // Validate with correct audience
402        let claims = manager.validate_token(&token, Some("client-1")).unwrap();
403        assert_eq!(claims.aud, Some("client-1".to_string()));
404
405        // Validate with incorrect audience (should fail)
406        let err = manager
407            .validate_token(&token, Some("client-2"))
408            .unwrap_err();
409        assert!(err.to_string().contains("InvalidAudience"));
410    }
411}
412pub mod jwk;