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			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		// Use asymmetric key (ES256) so we can separate public/private
283		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		// Create a key without KID
351		let key = create_test_key(None);
352		let claims = create_test_claims();
353
354		// Encode using the key directly
355		let token = key.sign(&claims).unwrap();
356
357		let set = KeySet {
358			keys: vec![Arc::new(key)],
359		};
360
361		// Decode using the set
362		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		// Encode with one of the keys directly
377		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		// Use a random-ish name to avoid collisions
413		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		// Clean up
429		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}