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	#[allow(clippy::should_implement_trait)]
51	pub fn from_str(s: &str) -> crate::Result<Self> {
52		Ok(serde_json::from_str(s)?)
53	}
54
55	pub fn from_file<P: AsRef<Path>>(path: P) -> crate::Result<Self> {
56		let json = std::fs::read_to_string(&path)?;
57		Ok(serde_json::from_str(&json)?)
58	}
59
60	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<()> {
65		let json = serde_json::to_string(&self)?;
66		std::fs::write(path, json)?;
67		Ok(())
68	}
69
70	pub fn to_public_set(&self) -> crate::Result<KeySet> {
71		Ok(KeySet {
72			keys: self
73				.keys
74				.iter()
75				.map(|key| key.as_ref().to_public().map(Arc::new))
76				.collect::<Result<Vec<Arc<Key>>, _>>()?,
77		})
78	}
79
80	/// Find the key with the given key ID.
81	pub fn find_key(&self, kid: &str) -> Option<Arc<Key>> {
82		self.keys
83			.iter()
84			.find(|k| k.kid.as_ref().is_some_and(|k| k.encode() == kid))
85			.cloned()
86	}
87
88	/// Find the first key that permits the operation and carries any key material it needs.
89	pub fn find_supported_key(&self, operation: &KeyOperation) -> Option<Arc<Key>> {
90		self.keys
91			.iter()
92			.find(|key| {
93				key.operations.contains(operation) && (*operation != KeyOperation::Sign || key.has_signing_material())
94			})
95			.cloned()
96	}
97
98	/// Sign the claims with the first key in the set that supports signing.
99	pub fn sign(&self, payload: &Claims) -> crate::Result<String> {
100		let key = self
101			.find_supported_key(&KeyOperation::Sign)
102			.ok_or(KeyError::NoSigningKey)?;
103		key.sign(payload)
104	}
105
106	#[doc(hidden)]
107	#[deprecated(note = "renamed to KeySet::sign")]
108	pub fn encode(&self, payload: &Claims) -> crate::Result<String> {
109		self.sign(payload)
110	}
111
112	#[doc(hidden)]
113	#[deprecated(note = "renamed to KeySet::verify")]
114	pub fn decode(&self, token: &str) -> crate::Result<Claims> {
115		self.verify(token)
116	}
117
118	/// Verify a token with the key matching its `kid` header, returning its claims.
119	///
120	/// A token without a `kid` is accepted only when the set holds exactly one key.
121	pub fn verify(&self, token: &str) -> crate::Result<Claims> {
122		let header = jsonwebtoken::decode_header(token)?;
123
124		let key = match header.kid {
125			Some(kid) => self
126				.find_key(kid.as_str())
127				.ok_or_else(|| crate::Error::from(KeyError::KeyNotFound(kid))),
128			None => {
129				// If we only have one key we can use it without a kid
130				if self.keys.len() == 1 {
131					Ok(self.keys[0].clone())
132				} else {
133					Err(KeyError::MissingKid.into())
134				}
135			}
136		}?;
137
138		key.verify(token)
139	}
140}
141
142#[cfg(feature = "jwks-loader")]
143pub async fn load_keys(jwks_uri: &str) -> crate::Result<KeySet> {
144	let client = reqwest::Client::builder().timeout(Duration::from_secs(10)).build()?;
145
146	let jwks_json = client.get(jwks_uri).send().await?.error_for_status()?.text().await?;
147
148	KeySet::from_str(&jwks_json)
149}
150
151#[cfg(test)]
152mod tests {
153	use super::*;
154	use crate::Algorithm;
155	use std::time::{Duration, SystemTime};
156
157	fn create_test_claims() -> Claims {
158		Claims {
159			root: "test-path".to_string(),
160			publish: vec!["test-pub".into()],
161			subscribe: vec!["test-sub".into()],
162			expires: Some(SystemTime::now() + Duration::from_secs(3600)),
163			issued: Some(SystemTime::now()),
164		}
165	}
166
167	fn create_test_key(kid: Option<&str>) -> Key {
168		let kid = kid.map(|s| crate::KeyId::decode(s).unwrap());
169		Key::generate(Algorithm::ES256, kid).expect("failed to generate key")
170	}
171
172	#[test]
173	fn test_keyset_from_str_valid() {
174		let json = r#"{"keys":[{"kty":"oct","k":"2AJvfDJMVfWe9WMRPJP-4zCGN8F62LOy3dUr--rogR8","alg":"HS256","key_ops":["verify","sign"],"kid":"1"}]}"#;
175		let set = KeySet::from_str(json);
176		assert!(set.is_ok());
177		let set = set.unwrap();
178		assert_eq!(set.keys.len(), 1);
179		assert_eq!(set.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
180		assert!(set.find_key("1").is_some());
181	}
182
183	#[test]
184	fn test_keyset_from_str_invalid_json() {
185		let result = KeySet::from_str("invalid json");
186		assert!(result.is_err());
187	}
188
189	#[test]
190	fn test_keyset_from_str_empty() {
191		let json = r#"{"keys":[]}"#;
192		let set = KeySet::from_str(json).unwrap();
193		assert!(set.keys.is_empty());
194	}
195
196	#[test]
197	fn test_keyset_to_str() {
198		let key = create_test_key(Some("1"));
199		let set = KeySet {
200			keys: vec![Arc::new(key)],
201		};
202
203		let json = set.to_str().unwrap();
204		assert!(json.contains("\"keys\""));
205		assert!(json.contains("\"kid\":\"1\""));
206	}
207
208	#[test]
209	fn test_keyset_serde_round_trip() {
210		let key1 = create_test_key(Some("1"));
211		let key2 = create_test_key(Some("2"));
212		let set = KeySet {
213			keys: vec![Arc::new(key1), Arc::new(key2)],
214		};
215
216		let json = set.to_str().unwrap();
217		let deserialized = KeySet::from_str(&json).unwrap();
218
219		assert_eq!(deserialized.keys.len(), 2);
220		assert!(deserialized.find_key("1").is_some());
221		assert!(deserialized.find_key("2").is_some());
222	}
223
224	#[test]
225	fn test_find_key_success() {
226		let key = create_test_key(Some("my-key"));
227		let set = KeySet {
228			keys: vec![Arc::new(key)],
229		};
230
231		let found = set.find_key("my-key");
232		assert!(found.is_some());
233		assert_eq!(found.unwrap().kid.as_ref().map(|k| k.encode()), Some("my-key"));
234	}
235
236	#[test]
237	fn test_find_key_missing() {
238		let key = create_test_key(Some("my-key"));
239		let set = KeySet {
240			keys: vec![Arc::new(key)],
241		};
242
243		let found = set.find_key("other-key");
244		assert!(found.is_none());
245	}
246
247	#[test]
248	fn test_find_key_no_kid() {
249		let key = create_test_key(None);
250		let set = KeySet {
251			keys: vec![Arc::new(key)],
252		};
253
254		let found = set.find_key("any-key");
255		assert!(found.is_none());
256	}
257
258	#[test]
259	fn test_find_supported_key() {
260		let sign_key = create_test_key(Some("sign")).with_operations([KeyOperation::Sign]);
261		let verify_key = create_test_key(Some("verify")).with_operations([KeyOperation::Verify]);
262
263		let set = KeySet {
264			keys: vec![Arc::new(sign_key), Arc::new(verify_key)],
265		};
266
267		let found_sign = set.find_supported_key(&KeyOperation::Sign);
268		assert!(found_sign.is_some());
269		assert_eq!(found_sign.unwrap().kid.as_ref().map(|k| k.encode()), Some("sign"));
270
271		let found_verify = set.find_supported_key(&KeyOperation::Verify);
272		assert!(found_verify.is_some());
273		assert_eq!(found_verify.unwrap().kid.as_ref().map(|k| k.encode()), Some("verify"));
274	}
275
276	#[test]
277	fn test_sign_skips_public_only_key() {
278		let private = create_test_key(Some("active"));
279		let public = create_test_key(Some("old"))
280			.to_public()
281			.unwrap()
282			.with_operations([KeyOperation::Sign, KeyOperation::Verify]);
283
284		let set = KeySet {
285			keys: vec![Arc::new(public), Arc::new(private)],
286		};
287
288		let token = set.sign(&create_test_claims()).unwrap();
289		let header = jsonwebtoken::decode_header(&token).unwrap();
290		assert_eq!(header.kid.as_deref(), Some("active"));
291	}
292
293	#[test]
294	fn test_to_public_set() {
295		// Use asymmetric key (ES256) so we can separate public/private
296		let key = create_test_key(Some("1"));
297
298		let set = KeySet {
299			keys: vec![Arc::new(key)],
300		};
301
302		let public_set = set.to_public_set().expect("failed to convert to public set");
303		assert_eq!(public_set.keys.len(), 1);
304
305		let public_key = &public_set.keys[0];
306		assert_eq!(public_key.kid.as_ref().map(|k| k.encode()), Some("1"));
307		assert!(public_key.operations.contains(&KeyOperation::Verify));
308		assert!(!public_key.operations.contains(&KeyOperation::Sign));
309	}
310
311	#[test]
312	fn test_to_public_set_fails_for_symmetric() {
313		let key = Key::generate(Algorithm::HS256, Some(crate::KeyId::decode("sym").unwrap())).unwrap();
314		let set = KeySet {
315			keys: vec![Arc::new(key)],
316		};
317
318		let result = set.to_public_set();
319		assert!(result.is_err());
320	}
321
322	#[test]
323	fn test_encode_success() {
324		let key = create_test_key(Some("1"));
325		let set = KeySet {
326			keys: vec![Arc::new(key)],
327		};
328		let claims = create_test_claims();
329
330		let token = set.sign(&claims).unwrap();
331		assert!(!token.is_empty());
332	}
333
334	#[test]
335	fn test_encode_no_signing_key() {
336		let key = create_test_key(Some("1")).with_operations([KeyOperation::Verify]);
337		let set = KeySet {
338			keys: vec![Arc::new(key)],
339		};
340		let claims = create_test_claims();
341
342		let result = set.sign(&claims);
343		assert!(result.is_err());
344		assert!(result.unwrap_err().to_string().contains("cannot find signing key"));
345	}
346
347	#[test]
348	fn test_decode_success_with_kid() {
349		let key = create_test_key(Some("1"));
350		let set = KeySet {
351			keys: vec![Arc::new(key)],
352		};
353		let claims = create_test_claims();
354
355		let token = set.sign(&claims).unwrap();
356		let decoded = set.verify(&token).unwrap();
357
358		assert_eq!(decoded.root, claims.root);
359	}
360
361	#[test]
362	fn test_decode_success_single_key_no_kid() {
363		// Create a key without KID
364		let key = create_test_key(None);
365		let claims = create_test_claims();
366
367		// Encode using the key directly
368		let token = key.sign(&claims).unwrap();
369
370		let set = KeySet {
371			keys: vec![Arc::new(key)],
372		};
373
374		// Decode using the set
375		let decoded = set.verify(&token).unwrap();
376		assert_eq!(decoded.root, claims.root);
377	}
378
379	#[test]
380	fn test_decode_fail_multiple_keys_no_kid() {
381		let key1 = create_test_key(None);
382		let key2 = create_test_key(None);
383
384		let set = KeySet {
385			keys: vec![Arc::new(key1), Arc::new(key2)],
386		};
387
388		let claims = create_test_claims();
389		// Encode with one of the keys directly
390		let token = set.keys[0].sign(&claims).unwrap();
391
392		let result = set.verify(&token);
393		assert!(result.is_err());
394		assert!(result.unwrap_err().to_string().contains("missing kid"));
395	}
396
397	#[test]
398	fn test_decode_fail_unknown_kid() {
399		let key1 = create_test_key(Some("1"));
400		let key2 = create_test_key(Some("2"));
401
402		let set1 = KeySet {
403			keys: vec![Arc::new(key1)],
404		};
405		let set2 = KeySet {
406			keys: vec![Arc::new(key2)],
407		};
408
409		let claims = create_test_claims();
410		let token = set1.sign(&claims).unwrap();
411
412		let result = set2.verify(&token);
413		assert!(result.is_err());
414		assert!(result.unwrap_err().to_string().contains("cannot find key with kid 1"));
415	}
416
417	#[test]
418	fn test_file_io() {
419		let key = create_test_key(Some("1"));
420		let set = KeySet {
421			keys: vec![Arc::new(key)],
422		};
423
424		let dir = std::env::temp_dir();
425		// Use a random-ish name to avoid collisions
426		let filename = format!(
427			"test_keyset_{}.json",
428			SystemTime::now()
429				.duration_since(SystemTime::UNIX_EPOCH)
430				.unwrap()
431				.as_nanos()
432		);
433		let path = dir.join(filename);
434
435		set.to_file(&path).expect("failed to write to file");
436
437		let loaded = KeySet::from_file(&path).expect("failed to read from file");
438		assert_eq!(loaded.keys.len(), 1);
439		assert_eq!(loaded.keys[0].kid.as_ref().map(|k| k.encode()), Some("1"));
440
441		// Clean up
442		let _ = std::fs::remove_file(path);
443	}
444}