1use std::collections::BTreeMap;
2use std::result;
3
4use base64::{Engine, engine::general_purpose::STANDARD};
5use serde::de::DeserializeOwned;
6use serde::{Deserialize, Deserializer, Serialize, Serializer};
7
8use crate::algorithms::Algorithm;
9use crate::errors::Result;
10use crate::jwk::Jwk;
11use crate::serialization::b64_decode;
12
13const ZIP_SERIAL_DEFLATE: &str = "DEF";
14const ENC_A128CBC_HS256: &str = "A128CBC-HS256";
15const ENC_A192CBC_HS384: &str = "A192CBC-HS384";
16const ENC_A256CBC_HS512: &str = "A256CBC-HS512";
17const ENC_A128GCM: &str = "A128GCM";
18const ENC_A192GCM: &str = "A192GCM";
19const ENC_A256GCM: &str = "A256GCM";
20
21#[derive(Debug, Clone, PartialEq, Eq, Hash)]
27#[allow(clippy::upper_case_acronyms, non_camel_case_types)]
28pub enum Enc {
29 A128CBC_HS256,
30 A192CBC_HS384,
31 A256CBC_HS512,
32 A128GCM,
33 A192GCM,
34 A256GCM,
35 Other(String),
36}
37
38impl Serialize for Enc {
39 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
40 where
41 S: Serializer,
42 {
43 match self {
44 Enc::A128CBC_HS256 => ENC_A128CBC_HS256,
45 Enc::A192CBC_HS384 => ENC_A192CBC_HS384,
46 Enc::A256CBC_HS512 => ENC_A256CBC_HS512,
47 Enc::A128GCM => ENC_A128GCM,
48 Enc::A192GCM => ENC_A192GCM,
49 Enc::A256GCM => ENC_A256GCM,
50 Enc::Other(v) => v,
51 }
52 .serialize(serializer)
53 }
54}
55
56impl<'de> Deserialize<'de> for Enc {
57 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
58 where
59 D: Deserializer<'de>,
60 {
61 let s = String::deserialize(deserializer)?;
62 match s.as_str() {
63 ENC_A128CBC_HS256 => return Ok(Enc::A128CBC_HS256),
64 ENC_A192CBC_HS384 => return Ok(Enc::A192CBC_HS384),
65 ENC_A256CBC_HS512 => return Ok(Enc::A256CBC_HS512),
66 ENC_A128GCM => return Ok(Enc::A128GCM),
67 ENC_A192GCM => return Ok(Enc::A192GCM),
68 ENC_A256GCM => return Ok(Enc::A256GCM),
69 _ => (),
70 }
71 Ok(Enc::Other(s))
72 }
73}
74
75#[derive(Debug, Clone, PartialEq, Eq, Hash)]
79pub enum Zip {
80 Deflate,
81 Other(String),
82}
83
84impl Serialize for Zip {
85 fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
86 where
87 S: Serializer,
88 {
89 match self {
90 Zip::Deflate => ZIP_SERIAL_DEFLATE,
91 Zip::Other(v) => v,
92 }
93 .serialize(serializer)
94 }
95}
96
97impl<'de> Deserialize<'de> for Zip {
98 fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
99 where
100 D: Deserializer<'de>,
101 {
102 let s = String::deserialize(deserializer)?;
103 match s.as_str() {
104 ZIP_SERIAL_DEFLATE => Ok(Zip::Deflate),
105 _ => Ok(Zip::Other(s)),
106 }
107 }
108}
109
110#[derive(Debug, Clone, Default, Hash, PartialEq, Eq, Serialize, Deserialize)]
112pub struct Extras {
113 #[serde(flatten)]
114 inner: BTreeMap<String, serde_json::Value>,
115}
116
117impl Extras {
118 pub fn get<T>(&self, key: &str) -> Result<Option<T>>
120 where
121 T: DeserializeOwned,
122 {
123 match self.inner.get(key) {
124 Some(value) => {
125 let parsed = serde_json::from_value(value.clone())?;
126 Ok(Some(parsed))
127 }
128 None => Ok(None),
129 }
130 }
131
132 pub fn insert<T>(&mut self, key: impl Into<String>, value: T)
134 where
135 T: Serialize,
136 {
137 let value =
138 serde_json::to_value(value).expect("serializing extra header value must not fail");
139 self.inner.insert(key.into(), value);
140 }
141
142 pub fn inner(&self) -> &BTreeMap<String, serde_json::Value> {
144 &self.inner
145 }
146}
147
148#[derive(Debug, Clone, Hash, PartialEq, Eq, Serialize, Deserialize)]
151pub struct Header {
152 #[serde(skip_serializing_if = "Option::is_none")]
156 pub typ: Option<String>,
157 pub alg: Algorithm,
161 #[serde(skip_serializing_if = "Option::is_none")]
165 pub cty: Option<String>,
166 #[serde(skip_serializing_if = "Option::is_none")]
170 pub jku: Option<String>,
171 #[serde(skip_serializing_if = "Option::is_none")]
175 pub jwk: Option<Jwk>,
176 #[serde(skip_serializing_if = "Option::is_none")]
180 pub kid: Option<String>,
181 #[serde(skip_serializing_if = "Option::is_none")]
185 pub x5u: Option<String>,
186 #[serde(skip_serializing_if = "Option::is_none")]
190 pub x5c: Option<Vec<String>>,
191 #[serde(skip_serializing_if = "Option::is_none")]
195 pub x5t: Option<String>,
196 #[serde(skip_serializing_if = "Option::is_none")]
202 #[serde(rename = "x5t#S256")]
203 pub x5t_s256: Option<String>,
204 #[serde(skip_serializing_if = "Option::is_none")]
208 pub crit: Option<Vec<String>>,
209 #[serde(skip_serializing_if = "Option::is_none")]
211 pub enc: Option<Enc>,
212 #[serde(skip_serializing_if = "Option::is_none")]
214 pub zip: Option<Zip>,
215 #[serde(skip_serializing_if = "Option::is_none")]
219 pub url: Option<String>,
220 #[serde(skip_serializing_if = "Option::is_none")]
224 pub nonce: Option<String>,
225 #[serde(flatten)]
229 pub extras: Extras,
230}
231
232impl Header {
233 pub fn new(algorithm: Algorithm) -> Self {
235 Header {
236 typ: Some("JWT".to_string()),
237 alg: algorithm,
238 cty: None,
239 jku: None,
240 jwk: None,
241 kid: None,
242 x5u: None,
243 x5c: None,
244 x5t: None,
245 x5t_s256: None,
246 crit: None,
247 enc: None,
248 zip: None,
249 url: None,
250 nonce: None,
251 extras: Extras::default(),
252 }
253 }
254
255 pub(crate) fn from_encoded<T: AsRef<[u8]>>(encoded_part: T) -> Result<Self> {
257 let decoded = b64_decode(encoded_part)?;
258 Ok(serde_json::from_slice(&decoded)?)
259 }
260
261 pub fn x5c_der(&self) -> Result<Option<Vec<Vec<u8>>>> {
263 Ok(self
264 .x5c
265 .as_ref()
266 .map(|b64_certs| {
267 b64_certs.iter().map(|x| STANDARD.decode(x)).collect::<result::Result<_, _>>()
268 })
269 .transpose()?)
270 }
271}
272
273impl Default for Header {
274 fn default() -> Self {
276 Header::new(Algorithm::default())
277 }
278}
279
280#[cfg(test)]
281mod tests {
282 use std::hash::{DefaultHasher, Hash, Hasher};
283
284 use crate::{Algorithm, Extras, Header};
285
286 fn hash<T>(value: &T) -> u64
287 where
288 T: Hash,
289 {
290 let mut hasher = DefaultHasher::new();
291 value.hash(&mut hasher);
292 hasher.finish()
293 }
294
295 #[test]
296 fn test_header_extras_hash() {
297 assert_eq!(hash(&Extras::default()), hash(&Extras::default()));
298
299 let mut a = Extras::default();
300 a.insert("foo", "bar");
301 a.insert("answer", 42);
302
303 let mut b = Extras::default();
304 b.insert("answer", 42);
305 b.insert("foo", "bar");
306
307 assert_eq!(a, b);
308 assert_eq!(hash(&a), hash(&b));
309
310 b.insert("more", "values");
311
312 assert_ne!(a, b);
313 assert_ne!(hash(&a), hash(&b));
314 }
315
316 #[test]
317 fn test_header_hash() {
318 assert_eq!(hash(&Header::default()), hash(&Header::default()));
319
320 let mut extras_a = Extras::default();
321 extras_a.insert("foo", "bar");
322 extras_a.insert("answer", 42);
323
324 let mut extras_b = Extras::default();
325 extras_b.insert("answer", 42);
326 extras_b.insert("foo", "bar");
327
328 let mut a = Header::new(Algorithm::HS512);
329 a.extras = extras_a;
330
331 let mut b = Header::new(Algorithm::HS512);
332 b.extras = extras_b.clone();
333
334 assert_eq!(a, b);
335 assert_eq!(hash(&a), hash(&b));
336
337 extras_b.insert("more", "values");
338 b.extras = extras_b;
339
340 assert_ne!(a, b);
341 assert_ne!(hash(&a), hash(&b));
342
343 assert_ne!(
344 hash(&Header { alg: Algorithm::EdDSA, ..Default::default() }),
345 hash(&Header { alg: Algorithm::ES256, ..Default::default() })
346 )
347 }
348}