Skip to main content

moq_token/
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#[cfg(feature = "jwks-loader")]
8use std::time::Duration;
9
10/// JWK Set to spec <https://datatracker.ietf.org/doc/html/rfc7517#section-5>
11#[derive(Default, Clone)]
12pub struct KeySet {
13	/// Vec of an arbitrary number of Json Web Keys
14	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		// Serialize as a struct with a `keys` field
23		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		// Deserialize into a temporary Vec<Key>
37		#[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	/// Parse a set from a JSON string.
51	#[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	/// Load a set from a JSON file.
57	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	/// Encode the set as JSON.
63	pub fn to_str(&self) -> crate::Result<String> {
64		Ok(serde_json::to_string(&self)?)
65	}
66
67	/// Write the set to a file as JSON.
68	///
69	/// A set holding any private key is written owner-only (mode `0600` on Unix), including when
70	/// it overwrites a file that was more permissive.
71	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	/// Find the key with the given key ID.
89	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	/// Find the first key that permits the operation and carries any key material it needs.
97	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	/// Sign the claims with the first key in the set that supports signing.
105	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	/// Verify a token with the key matching its `kid` header, returning its claims.
125	///
126	/// A token without a `kid` is accepted only when the set holds exactly one key.
127	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 we only have one key we can use it without a kid
136				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		// Use asymmetric key (ES256) so we can separate public/private
302		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		// Create a key without KID
370		let key = create_test_key(None);
371		let claims = create_test_claims();
372
373		// Encode using the key directly
374		let token = key.sign(&claims).unwrap();
375
376		let set = KeySet {
377			keys: vec![Arc::new(key)],
378		};
379
380		// Decode using the set
381		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		// Encode with one of the keys directly
396		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		// Use a random-ish name to avoid collisions
432		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		// Clean up
448		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}