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