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 }
151 }
152
153 fn create_test_key(kid: Option<&str>) -> Key {
154 let kid = kid.map(|s| crate::KeyId::decode(s).unwrap());
155 Key::generate(Algorithm::ES256, kid).expect("failed to generate key")
156 }
157
158 #[test]
159 fn test_keyset_from_str_valid() {
160 let json = r#"{"keys":[{"kty":"oct","k":"2AJvfDJMVfWe9WMRPJP-4zCGN8F62LOy3dUr--rogR8","alg":"HS256","key_ops":["verify","sign"],"kid":"1"}]}"#;
161 let set = KeySet::from_str(json);
162 assert!(set.is_ok());
163 let set = set.unwrap();
164 assert_eq!(set.keys.len(), 1);
165 assert_eq!(set.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
166 assert!(set.find_key("1").is_some());
167 }
168
169 #[test]
170 fn test_keyset_from_str_invalid_json() {
171 let result = KeySet::from_str("invalid json");
172 assert!(result.is_err());
173 }
174
175 #[test]
176 fn test_keyset_from_str_empty() {
177 let json = r#"{"keys":[]}"#;
178 let set = KeySet::from_str(json).unwrap();
179 assert!(set.keys.is_empty());
180 }
181
182 #[test]
183 fn test_keyset_to_str() {
184 let key = create_test_key(Some("1"));
185 let set = KeySet {
186 keys: vec![Arc::new(key)],
187 };
188
189 let json = set.to_str().unwrap();
190 assert!(json.contains("\"keys\""));
191 assert!(json.contains("\"kid\":\"1\""));
192 }
193
194 #[test]
195 fn test_keyset_serde_round_trip() {
196 let key1 = create_test_key(Some("1"));
197 let key2 = create_test_key(Some("2"));
198 let set = KeySet {
199 keys: vec![Arc::new(key1), Arc::new(key2)],
200 };
201
202 let json = set.to_str().unwrap();
203 let deserialized = KeySet::from_str(&json).unwrap();
204
205 assert_eq!(deserialized.keys.len(), 2);
206 assert!(deserialized.find_key("1").is_some());
207 assert!(deserialized.find_key("2").is_some());
208 }
209
210 #[test]
211 fn test_find_key_success() {
212 let key = create_test_key(Some("my-key"));
213 let set = KeySet {
214 keys: vec![Arc::new(key)],
215 };
216
217 let found = set.find_key("my-key");
218 assert!(found.is_some());
219 assert_eq!(found.unwrap().kid.as_ref().map(|k| k.encode()), Some("my-key"));
220 }
221
222 #[test]
223 fn test_find_key_missing() {
224 let key = create_test_key(Some("my-key"));
225 let set = KeySet {
226 keys: vec![Arc::new(key)],
227 };
228
229 let found = set.find_key("other-key");
230 assert!(found.is_none());
231 }
232
233 #[test]
234 fn test_find_key_no_kid() {
235 let key = create_test_key(None);
236 let set = KeySet {
237 keys: vec![Arc::new(key)],
238 };
239
240 let found = set.find_key("any-key");
241 assert!(found.is_none());
242 }
243
244 #[test]
245 fn test_find_supported_key() {
246 let sign_key = create_test_key(Some("sign")).with_operations([KeyOperation::Sign]);
247 let verify_key = create_test_key(Some("verify")).with_operations([KeyOperation::Verify]);
248
249 let set = KeySet {
250 keys: vec![Arc::new(sign_key), Arc::new(verify_key)],
251 };
252
253 let found_sign = set.find_supported_key(&KeyOperation::Sign);
254 assert!(found_sign.is_some());
255 assert_eq!(found_sign.unwrap().kid.as_ref().map(|k| k.encode()), Some("sign"));
256
257 let found_verify = set.find_supported_key(&KeyOperation::Verify);
258 assert!(found_verify.is_some());
259 assert_eq!(found_verify.unwrap().kid.as_ref().map(|k| k.encode()), Some("verify"));
260 }
261
262 #[test]
263 fn test_sign_skips_public_only_key() {
264 let private = create_test_key(Some("active"));
265 let public = create_test_key(Some("old"))
266 .to_public()
267 .unwrap()
268 .with_operations([KeyOperation::Sign, KeyOperation::Verify]);
269
270 let set = KeySet {
271 keys: vec![Arc::new(public), Arc::new(private)],
272 };
273
274 let token = set.sign(&create_test_claims()).unwrap();
275 let header = jsonwebtoken::decode_header(&token).unwrap();
276 assert_eq!(header.kid.as_deref(), Some("active"));
277 }
278
279 #[test]
280 fn test_to_public_set() {
281 let key = create_test_key(Some("1"));
283
284 let set = KeySet {
285 keys: vec![Arc::new(key)],
286 };
287
288 let public_set = set.to_public_set().expect("failed to convert to public set");
289 assert_eq!(public_set.keys.len(), 1);
290
291 let public_key = &public_set.keys[0];
292 assert_eq!(public_key.kid.as_ref().map(|k| k.encode()), Some("1"));
293 assert!(public_key.operations.contains(&KeyOperation::Verify));
294 assert!(!public_key.operations.contains(&KeyOperation::Sign));
295 }
296
297 #[test]
298 fn test_to_public_set_fails_for_symmetric() {
299 let key = Key::generate(Algorithm::HS256, Some(crate::KeyId::decode("sym").unwrap())).unwrap();
300 let set = KeySet {
301 keys: vec![Arc::new(key)],
302 };
303
304 let result = set.to_public_set();
305 assert!(result.is_err());
306 }
307
308 #[test]
309 fn test_encode_success() {
310 let key = create_test_key(Some("1"));
311 let set = KeySet {
312 keys: vec![Arc::new(key)],
313 };
314 let claims = create_test_claims();
315
316 let token = set.sign(&claims).unwrap();
317 assert!(!token.is_empty());
318 }
319
320 #[test]
321 fn test_encode_no_signing_key() {
322 let key = create_test_key(Some("1")).with_operations([KeyOperation::Verify]);
323 let set = KeySet {
324 keys: vec![Arc::new(key)],
325 };
326 let claims = create_test_claims();
327
328 let result = set.sign(&claims);
329 assert!(result.is_err());
330 assert!(result.unwrap_err().to_string().contains("cannot find signing key"));
331 }
332
333 #[test]
334 fn test_decode_success_with_kid() {
335 let key = create_test_key(Some("1"));
336 let set = KeySet {
337 keys: vec![Arc::new(key)],
338 };
339 let claims = create_test_claims();
340
341 let token = set.sign(&claims).unwrap();
342 let decoded = set.verify(&token).unwrap();
343
344 assert_eq!(decoded.root, claims.root);
345 }
346
347 #[test]
348 fn test_decode_success_single_key_no_kid() {
349 let key = create_test_key(None);
351 let claims = create_test_claims();
352
353 let token = key.sign(&claims).unwrap();
355
356 let set = KeySet {
357 keys: vec![Arc::new(key)],
358 };
359
360 let decoded = set.verify(&token).unwrap();
362 assert_eq!(decoded.root, claims.root);
363 }
364
365 #[test]
366 fn test_decode_fail_multiple_keys_no_kid() {
367 let key1 = create_test_key(None);
368 let key2 = create_test_key(None);
369
370 let set = KeySet {
371 keys: vec![Arc::new(key1), Arc::new(key2)],
372 };
373
374 let claims = create_test_claims();
375 let token = set.keys[0].sign(&claims).unwrap();
377
378 let result = set.verify(&token);
379 assert!(result.is_err());
380 assert!(result.unwrap_err().to_string().contains("missing kid"));
381 }
382
383 #[test]
384 fn test_decode_fail_unknown_kid() {
385 let key1 = create_test_key(Some("1"));
386 let key2 = create_test_key(Some("2"));
387
388 let set1 = KeySet {
389 keys: vec![Arc::new(key1)],
390 };
391 let set2 = KeySet {
392 keys: vec![Arc::new(key2)],
393 };
394
395 let claims = create_test_claims();
396 let token = set1.sign(&claims).unwrap();
397
398 let result = set2.verify(&token);
399 assert!(result.is_err());
400 assert!(result.unwrap_err().to_string().contains("cannot find key with kid 1"));
401 }
402
403 #[test]
404 fn test_file_io() {
405 let key = create_test_key(Some("1"));
406 let set = KeySet {
407 keys: vec![Arc::new(key)],
408 };
409
410 let dir = std::env::temp_dir();
411 let filename = format!(
413 "test_keyset_{}.json",
414 SystemTime::now()
415 .duration_since(SystemTime::UNIX_EPOCH)
416 .unwrap()
417 .as_nanos()
418 );
419 let path = dir.join(filename);
420
421 set.to_file(&path).expect("failed to write to file");
422
423 let loaded = KeySet::from_file(&path).expect("failed to read from file");
424 assert_eq!(loaded.keys.len(), 1);
425 assert_eq!(loaded.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
426
427 let _ = std::fs::remove_file(path);
429 }
430
431 #[cfg(unix)]
432 #[test]
433 fn test_file_io_permissions() {
434 use std::os::unix::fs::PermissionsExt;
435
436 let unique = SystemTime::now()
437 .duration_since(SystemTime::UNIX_EPOCH)
438 .unwrap()
439 .as_nanos();
440 let path = std::env::temp_dir().join(format!("test_keyset_perms_{unique}.json"));
441
442 let set = KeySet {
443 keys: vec![Arc::new(create_test_key(Some("1")))],
444 };
445 set.to_file(&path).expect("failed to write to file");
446
447 let mode = std::fs::metadata(&path).unwrap().permissions().mode() & 0o777;
448 assert_eq!(mode, 0o600, "a set holding a private key must be owner-only");
449
450 let _ = std::fs::remove_file(path);
451 }
452}