Skip to main content

moq_auth/
set.rs

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/// JWK Set to spec <https://datatracker.ietf.org/doc/html/rfc7517#section-5>
8#[derive(Default, Clone)]
9pub struct KeySet {
10	/// Vec of an arbitrary number of Json Web Keys
11	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		// Serialize as a struct with a `keys` field
20		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		// Deserialize into a temporary Vec<Key>
34		#[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	/// Parse a set from a JSON string.
48	#[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	/// Load a set from a JSON file.
54	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	/// Encode the set as JSON.
60	pub fn to_str(&self) -> crate::Result<String> {
61		Ok(serde_json::to_string(&self)?)
62	}
63
64	/// Write the set to a file as JSON.
65	///
66	/// A set holding any private key is written owner-only (mode `0600` on Unix), including when
67	/// it overwrites a file that was more permissive.
68	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	/// Find the key with the given key ID.
86	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	/// Find the first key that permits the operation and carries any key material it needs.
94	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	/// Sign the claims with the first key in the set that supports signing.
102	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	/// Verify a token with the key matching its `kid` header, returning its claims.
110	///
111	/// A token without a `kid` is accepted only when the set holds exactly one key.
112	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 we only have one key we can use it without a kid
121				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		// Use asymmetric key (ES256) so we can separate public/private
282		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		// Create a key without KID
350		let key = create_test_key(None);
351		let claims = create_test_claims();
352
353		// Encode using the key directly
354		let token = key.sign(&claims).unwrap();
355
356		let set = KeySet {
357			keys: vec![Arc::new(key)],
358		};
359
360		// Decode using the set
361		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		// Encode with one of the keys directly
376		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		// Use a random-ish name to avoid collisions
412		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		// Clean up
428		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}