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#[derive(Default, Clone)]
12pub struct KeySet {
13 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 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 #[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 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 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 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 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 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 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 let key = create_test_key(None);
365 let claims = create_test_claims();
366
367 let token = key.sign(&claims).unwrap();
369
370 let set = KeySet {
371 keys: vec![Arc::new(key)],
372 };
373
374 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 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 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 let _ = std::fs::remove_file(path);
443 }
444}