1use crate::error::KeyError;
2use crate::{Claims, Key, KeyOperation};
3use serde::{Deserialize, Deserializer, Serialize, Serializer};
4use std::path::Path;
5use std::sync::Arc;
6
7#[cfg(feature = "jwks-loader")]
8use std::time::Duration;
9
10#[derive(Default, Clone)]
12pub struct KeySet {
13 pub keys: Vec<Arc<Key>>,
15}
16
17impl Serialize for KeySet {
18 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
19 where
20 S: Serializer,
21 {
22 use serde::ser::SerializeStruct;
24
25 let mut state = serializer.serialize_struct("KeySet", 1)?;
26 state.serialize_field("keys", &self.keys.iter().map(|k| k.as_ref()).collect::<Vec<_>>())?;
27 state.end()
28 }
29}
30
31impl<'de> Deserialize<'de> for KeySet {
32 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
33 where
34 D: Deserializer<'de>,
35 {
36 #[derive(Deserialize)]
38 struct RawKeySet {
39 keys: Vec<Key>,
40 }
41
42 let raw = RawKeySet::deserialize(deserializer)?;
43 Ok(KeySet {
44 keys: raw.keys.into_iter().map(Arc::new).collect(),
45 })
46 }
47}
48
49impl KeySet {
50 #[allow(clippy::should_implement_trait)]
52 pub fn from_str(s: &str) -> crate::Result<Self> {
53 Ok(serde_json::from_str(s)?)
54 }
55
56 pub fn from_file<P: AsRef<Path>>(path: P) -> crate::Result<Self> {
58 let json = std::fs::read_to_string(&path)?;
59 Ok(serde_json::from_str(&json)?)
60 }
61
62 pub fn to_str(&self) -> crate::Result<String> {
64 Ok(serde_json::to_string(&self)?)
65 }
66
67 pub fn to_file<P: AsRef<Path>>(&self, path: P) -> crate::Result<()> {
72 let json = serde_json::to_string(&self)?;
73 let private = self.keys.iter().any(|key| key.is_private());
74 crate::fs::write(path.as_ref(), &json, private)?;
75 Ok(())
76 }
77
78 pub fn to_public_set(&self) -> crate::Result<KeySet> {
79 Ok(KeySet {
80 keys: self
81 .keys
82 .iter()
83 .map(|key| key.as_ref().to_public().map(Arc::new))
84 .collect::<Result<Vec<Arc<Key>>, _>>()?,
85 })
86 }
87
88 pub fn find_key(&self, kid: &str) -> Option<Arc<Key>> {
90 self.keys
91 .iter()
92 .find(|k| k.kid.as_ref().is_some_and(|k| k.encode() == kid))
93 .cloned()
94 }
95
96 pub fn find_supported_key(&self, operation: &KeyOperation) -> Option<Arc<Key>> {
98 self.keys
99 .iter()
100 .find(|key| key.operations.contains(operation) && (*operation != KeyOperation::Sign || key.is_private()))
101 .cloned()
102 }
103
104 pub fn sign(&self, payload: &Claims) -> crate::Result<String> {
106 let key = self
107 .find_supported_key(&KeyOperation::Sign)
108 .ok_or(KeyError::NoSigningKey)?;
109 key.sign(payload)
110 }
111
112 #[doc(hidden)]
113 #[deprecated(note = "renamed to KeySet::sign")]
114 pub fn encode(&self, payload: &Claims) -> crate::Result<String> {
115 self.sign(payload)
116 }
117
118 #[doc(hidden)]
119 #[deprecated(note = "renamed to KeySet::verify")]
120 pub fn decode(&self, token: &str) -> crate::Result<Claims> {
121 self.verify(token)
122 }
123
124 pub fn verify(&self, token: &str) -> crate::Result<Claims> {
128 let header = jsonwebtoken::decode_header(token)?;
129
130 let key = match header.kid {
131 Some(kid) => self
132 .find_key(kid.as_str())
133 .ok_or_else(|| crate::Error::from(KeyError::KeyNotFound(kid))),
134 None => {
135 if self.keys.len() == 1 {
137 Ok(self.keys[0].clone())
138 } else {
139 Err(KeyError::MissingKid.into())
140 }
141 }
142 }?;
143
144 key.verify(token)
145 }
146}
147
148#[cfg(feature = "jwks-loader")]
149pub async fn load_keys(jwks_uri: &str) -> crate::Result<KeySet> {
150 let client = reqwest::Client::builder().timeout(Duration::from_secs(10)).build()?;
151
152 let jwks_json = client.get(jwks_uri).send().await?.error_for_status()?.text().await?;
153
154 KeySet::from_str(&jwks_json)
155}
156
157#[cfg(test)]
158mod tests {
159 use super::*;
160 use crate::Algorithm;
161 use std::time::{Duration, SystemTime};
162
163 fn create_test_claims() -> Claims {
164 Claims {
165 root: "test-path".to_string(),
166 publish: vec!["test-pub".into()],
167 subscribe: vec!["test-sub".into()],
168 expires: Some(SystemTime::now() + Duration::from_secs(3600)),
169 issued: Some(SystemTime::now()),
170 }
171 }
172
173 fn create_test_key(kid: Option<&str>) -> Key {
174 let kid = kid.map(|s| crate::KeyId::decode(s).unwrap());
175 Key::generate(Algorithm::ES256, kid).expect("failed to generate key")
176 }
177
178 #[test]
179 fn test_keyset_from_str_valid() {
180 let json = r#"{"keys":[{"kty":"oct","k":"2AJvfDJMVfWe9WMRPJP-4zCGN8F62LOy3dUr--rogR8","alg":"HS256","key_ops":["verify","sign"],"kid":"1"}]}"#;
181 let set = KeySet::from_str(json);
182 assert!(set.is_ok());
183 let set = set.unwrap();
184 assert_eq!(set.keys.len(), 1);
185 assert_eq!(set.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
186 assert!(set.find_key("1").is_some());
187 }
188
189 #[test]
190 fn test_keyset_from_str_invalid_json() {
191 let result = KeySet::from_str("invalid json");
192 assert!(result.is_err());
193 }
194
195 #[test]
196 fn test_keyset_from_str_empty() {
197 let json = r#"{"keys":[]}"#;
198 let set = KeySet::from_str(json).unwrap();
199 assert!(set.keys.is_empty());
200 }
201
202 #[test]
203 fn test_keyset_to_str() {
204 let key = create_test_key(Some("1"));
205 let set = KeySet {
206 keys: vec![Arc::new(key)],
207 };
208
209 let json = set.to_str().unwrap();
210 assert!(json.contains("\"keys\""));
211 assert!(json.contains("\"kid\":\"1\""));
212 }
213
214 #[test]
215 fn test_keyset_serde_round_trip() {
216 let key1 = create_test_key(Some("1"));
217 let key2 = create_test_key(Some("2"));
218 let set = KeySet {
219 keys: vec![Arc::new(key1), Arc::new(key2)],
220 };
221
222 let json = set.to_str().unwrap();
223 let deserialized = KeySet::from_str(&json).unwrap();
224
225 assert_eq!(deserialized.keys.len(), 2);
226 assert!(deserialized.find_key("1").is_some());
227 assert!(deserialized.find_key("2").is_some());
228 }
229
230 #[test]
231 fn test_find_key_success() {
232 let key = create_test_key(Some("my-key"));
233 let set = KeySet {
234 keys: vec![Arc::new(key)],
235 };
236
237 let found = set.find_key("my-key");
238 assert!(found.is_some());
239 assert_eq!(found.unwrap().kid.as_ref().map(|k| k.encode()), Some("my-key"));
240 }
241
242 #[test]
243 fn test_find_key_missing() {
244 let key = create_test_key(Some("my-key"));
245 let set = KeySet {
246 keys: vec![Arc::new(key)],
247 };
248
249 let found = set.find_key("other-key");
250 assert!(found.is_none());
251 }
252
253 #[test]
254 fn test_find_key_no_kid() {
255 let key = create_test_key(None);
256 let set = KeySet {
257 keys: vec![Arc::new(key)],
258 };
259
260 let found = set.find_key("any-key");
261 assert!(found.is_none());
262 }
263
264 #[test]
265 fn test_find_supported_key() {
266 let sign_key = create_test_key(Some("sign")).with_operations([KeyOperation::Sign]);
267 let verify_key = create_test_key(Some("verify")).with_operations([KeyOperation::Verify]);
268
269 let set = KeySet {
270 keys: vec![Arc::new(sign_key), Arc::new(verify_key)],
271 };
272
273 let found_sign = set.find_supported_key(&KeyOperation::Sign);
274 assert!(found_sign.is_some());
275 assert_eq!(found_sign.unwrap().kid.as_ref().map(|k| k.encode()), Some("sign"));
276
277 let found_verify = set.find_supported_key(&KeyOperation::Verify);
278 assert!(found_verify.is_some());
279 assert_eq!(found_verify.unwrap().kid.as_ref().map(|k| k.encode()), Some("verify"));
280 }
281
282 #[test]
283 fn test_sign_skips_public_only_key() {
284 let private = create_test_key(Some("active"));
285 let public = create_test_key(Some("old"))
286 .to_public()
287 .unwrap()
288 .with_operations([KeyOperation::Sign, KeyOperation::Verify]);
289
290 let set = KeySet {
291 keys: vec![Arc::new(public), Arc::new(private)],
292 };
293
294 let token = set.sign(&create_test_claims()).unwrap();
295 let header = jsonwebtoken::decode_header(&token).unwrap();
296 assert_eq!(header.kid.as_deref(), Some("active"));
297 }
298
299 #[test]
300 fn test_to_public_set() {
301 let key = create_test_key(Some("1"));
303
304 let set = KeySet {
305 keys: vec![Arc::new(key)],
306 };
307
308 let public_set = set.to_public_set().expect("failed to convert to public set");
309 assert_eq!(public_set.keys.len(), 1);
310
311 let public_key = &public_set.keys[0];
312 assert_eq!(public_key.kid.as_ref().map(|k| k.encode()), Some("1"));
313 assert!(public_key.operations.contains(&KeyOperation::Verify));
314 assert!(!public_key.operations.contains(&KeyOperation::Sign));
315 }
316
317 #[test]
318 fn test_to_public_set_fails_for_symmetric() {
319 let key = Key::generate(Algorithm::HS256, Some(crate::KeyId::decode("sym").unwrap())).unwrap();
320 let set = KeySet {
321 keys: vec![Arc::new(key)],
322 };
323
324 let result = set.to_public_set();
325 assert!(result.is_err());
326 }
327
328 #[test]
329 fn test_encode_success() {
330 let key = create_test_key(Some("1"));
331 let set = KeySet {
332 keys: vec![Arc::new(key)],
333 };
334 let claims = create_test_claims();
335
336 let token = set.sign(&claims).unwrap();
337 assert!(!token.is_empty());
338 }
339
340 #[test]
341 fn test_encode_no_signing_key() {
342 let key = create_test_key(Some("1")).with_operations([KeyOperation::Verify]);
343 let set = KeySet {
344 keys: vec![Arc::new(key)],
345 };
346 let claims = create_test_claims();
347
348 let result = set.sign(&claims);
349 assert!(result.is_err());
350 assert!(result.unwrap_err().to_string().contains("cannot find signing key"));
351 }
352
353 #[test]
354 fn test_decode_success_with_kid() {
355 let key = create_test_key(Some("1"));
356 let set = KeySet {
357 keys: vec![Arc::new(key)],
358 };
359 let claims = create_test_claims();
360
361 let token = set.sign(&claims).unwrap();
362 let decoded = set.verify(&token).unwrap();
363
364 assert_eq!(decoded.root, claims.root);
365 }
366
367 #[test]
368 fn test_decode_success_single_key_no_kid() {
369 let key = create_test_key(None);
371 let claims = create_test_claims();
372
373 let token = key.sign(&claims).unwrap();
375
376 let set = KeySet {
377 keys: vec![Arc::new(key)],
378 };
379
380 let decoded = set.verify(&token).unwrap();
382 assert_eq!(decoded.root, claims.root);
383 }
384
385 #[test]
386 fn test_decode_fail_multiple_keys_no_kid() {
387 let key1 = create_test_key(None);
388 let key2 = create_test_key(None);
389
390 let set = KeySet {
391 keys: vec![Arc::new(key1), Arc::new(key2)],
392 };
393
394 let claims = create_test_claims();
395 let token = set.keys[0].sign(&claims).unwrap();
397
398 let result = set.verify(&token);
399 assert!(result.is_err());
400 assert!(result.unwrap_err().to_string().contains("missing kid"));
401 }
402
403 #[test]
404 fn test_decode_fail_unknown_kid() {
405 let key1 = create_test_key(Some("1"));
406 let key2 = create_test_key(Some("2"));
407
408 let set1 = KeySet {
409 keys: vec![Arc::new(key1)],
410 };
411 let set2 = KeySet {
412 keys: vec![Arc::new(key2)],
413 };
414
415 let claims = create_test_claims();
416 let token = set1.sign(&claims).unwrap();
417
418 let result = set2.verify(&token);
419 assert!(result.is_err());
420 assert!(result.unwrap_err().to_string().contains("cannot find key with kid 1"));
421 }
422
423 #[test]
424 fn test_file_io() {
425 let key = create_test_key(Some("1"));
426 let set = KeySet {
427 keys: vec![Arc::new(key)],
428 };
429
430 let dir = std::env::temp_dir();
431 let filename = format!(
433 "test_keyset_{}.json",
434 SystemTime::now()
435 .duration_since(SystemTime::UNIX_EPOCH)
436 .unwrap()
437 .as_nanos()
438 );
439 let path = dir.join(filename);
440
441 set.to_file(&path).expect("failed to write to file");
442
443 let loaded = KeySet::from_file(&path).expect("failed to read from file");
444 assert_eq!(loaded.keys.len(), 1);
445 assert_eq!(loaded.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
446
447 let _ = std::fs::remove_file(path);
449 }
450
451 #[cfg(unix)]
452 #[test]
453 fn test_file_io_permissions() {
454 use std::os::unix::fs::PermissionsExt;
455
456 let unique = SystemTime::now()
457 .duration_since(SystemTime::UNIX_EPOCH)
458 .unwrap()
459 .as_nanos();
460 let path = std::env::temp_dir().join(format!("test_keyset_perms_{unique}.json"));
461
462 let set = KeySet {
463 keys: vec![Arc::new(create_test_key(Some("1")))],
464 };
465 set.to_file(&path).expect("failed to write to file");
466
467 let mode = std::fs::metadata(&path).unwrap().permissions().mode() & 0o777;
468 assert_eq!(mode, 0o600, "a set holding a private key must be owner-only");
469
470 let _ = std::fs::remove_file(path);
471 }
472}