1use std::fmt::{Debug, Formatter};
2
3use base64::{Engine, engine::general_purpose::STANDARD};
4use serde::de::DeserializeOwned;
5use zeroize::{Zeroize, ZeroizeOnDrop};
6
7use crate::algorithms::AlgorithmFamily;
8use crate::crypto::{CryptoProvider, JwtVerifier};
9use crate::errors::{ErrorKind, Result, new_error};
10use crate::header::Header;
11use crate::jwk::{AlgorithmParameters, Jwk};
12#[cfg(feature = "use_pem")]
13use crate::pem::decoder::PemEncodedKey;
14use crate::serialization::{DecodedJwtPartClaims, b64_decode};
15use crate::validation::{Validation, validate};
16
17#[derive(Debug)]
19pub struct TokenData<T> {
20 pub header: Header,
22 pub claims: T,
24}
25
26impl<T> Clone for TokenData<T>
27where
28 T: Clone,
29{
30 fn clone(&self) -> Self {
31 Self { header: self.header.clone(), claims: self.claims.clone() }
32 }
33}
34
35macro_rules! expect_two {
38 ($iter:expr) => {{
39 let mut i = $iter;
40 match (i.next(), i.next(), i.next()) {
41 (Some(first), Some(second), None) => (first, second),
42 _ => return Err(new_error(ErrorKind::InvalidToken)),
43 }
44 }};
45}
46
47#[derive(Clone, Zeroize, ZeroizeOnDrop)]
48pub enum DecodingKeyKind {
50 SecretOrDer(Vec<u8>),
52 RsaModulusExponent {
54 n: Vec<u8>,
56 e: Vec<u8>,
58 },
59}
60
61impl Debug for DecodingKeyKind {
62 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
63 match self {
64 Self::SecretOrDer(_) => f.debug_tuple("SecretOrDer").field(&"[redacted]").finish(),
65 Self::RsaModulusExponent { .. } => f
66 .debug_struct("RsaModulusExponent")
67 .field("n", &"[redacted]")
68 .field("e", &"[redacted]")
69 .finish(),
70 }
71 }
72}
73
74#[derive(Clone, Debug, ZeroizeOnDrop, Zeroize)]
77pub struct DecodingKey {
78 #[zeroize(skip)]
79 family: AlgorithmFamily,
80 kind: DecodingKeyKind,
81}
82
83impl DecodingKey {
84 pub fn family(&self) -> AlgorithmFamily {
86 self.family
87 }
88
89 pub fn kind(&self) -> &DecodingKeyKind {
91 &self.kind
92 }
93
94 pub fn from_secret(secret: &[u8]) -> Self {
96 DecodingKey {
97 family: AlgorithmFamily::Hmac,
98 kind: DecodingKeyKind::SecretOrDer(secret.to_vec()),
99 }
100 }
101
102 pub fn from_base64_secret(secret: &str) -> Result<Self> {
104 let out = STANDARD.decode(secret)?;
105 Ok(DecodingKey { family: AlgorithmFamily::Hmac, kind: DecodingKeyKind::SecretOrDer(out) })
106 }
107
108 #[cfg(feature = "use_pem")]
111 pub fn from_rsa_pem(key: &[u8]) -> Result<Self> {
112 let pem_key = PemEncodedKey::new(key)?;
113 let content = pem_key.as_rsa_key()?;
114 Ok(DecodingKey {
115 family: AlgorithmFamily::Rsa,
116 kind: DecodingKeyKind::SecretOrDer(content.to_vec()),
117 })
118 }
119
120 pub fn from_rsa_components(modulus: &str, exponent: &str) -> Result<Self> {
122 let n = b64_decode(modulus)?;
123 let e = b64_decode(exponent)?;
124 Ok(DecodingKey {
125 family: AlgorithmFamily::Rsa,
126 kind: DecodingKeyKind::RsaModulusExponent { n, e },
127 })
128 }
129
130 pub fn from_rsa_raw_components(modulus: &[u8], exponent: &[u8]) -> Self {
132 DecodingKey {
133 family: AlgorithmFamily::Rsa,
134 kind: DecodingKeyKind::RsaModulusExponent { n: modulus.to_vec(), e: exponent.to_vec() },
135 }
136 }
137
138 #[cfg(feature = "use_pem")]
141 pub fn from_ec_pem(key: &[u8]) -> Result<Self> {
142 let pem_key = PemEncodedKey::new(key)?;
143 let content = pem_key.as_ec_public_key()?;
144 Ok(DecodingKey {
145 family: AlgorithmFamily::Ec,
146 kind: DecodingKeyKind::SecretOrDer(content.to_vec()),
147 })
148 }
149
150 pub fn from_ec_components(x: &str, y: &str) -> Result<Self> {
152 let x_cmp = b64_decode(x)?;
153 let y_cmp = b64_decode(y)?;
154
155 let mut public_key = Vec::with_capacity(1 + x.len() + y.len());
156 public_key.push(0x04);
157 public_key.extend_from_slice(&x_cmp);
158 public_key.extend_from_slice(&y_cmp);
159
160 Ok(DecodingKey {
161 family: AlgorithmFamily::Ec,
162 kind: DecodingKeyKind::SecretOrDer(public_key),
163 })
164 }
165
166 #[cfg(feature = "use_pem")]
170 pub fn from_ed_pem(key: &[u8]) -> Result<Self> {
171 let pem_key = PemEncodedKey::new(key)?;
172 let content = pem_key.as_ed_public_key()?;
173 Ok(DecodingKey {
174 family: AlgorithmFamily::Ed,
175 kind: DecodingKeyKind::SecretOrDer(content.to_vec()),
176 })
177 }
178
179 pub fn from_rsa_der(der: &[u8]) -> Self {
181 DecodingKey {
182 family: AlgorithmFamily::Rsa,
183 kind: DecodingKeyKind::SecretOrDer(der.to_vec()),
184 }
185 }
186
187 pub fn from_ec_der(der: &[u8]) -> Self {
189 DecodingKey {
190 family: AlgorithmFamily::Ec,
191 kind: DecodingKeyKind::SecretOrDer(der.to_vec()),
192 }
193 }
194
195 pub fn from_ed_der(der: &[u8]) -> Self {
197 DecodingKey {
198 family: AlgorithmFamily::Ed,
199 kind: DecodingKeyKind::SecretOrDer(der.to_vec()),
200 }
201 }
202
203 pub fn from_ed_components(x: &str) -> Result<Self> {
205 let x_decoded = b64_decode(x)?;
206 Ok(DecodingKey {
207 family: AlgorithmFamily::Ed,
208 kind: DecodingKeyKind::SecretOrDer(x_decoded),
209 })
210 }
211
212 pub fn from_jwk(jwk: &Jwk) -> Result<Self> {
214 match &jwk.algorithm {
215 AlgorithmParameters::RSA(params) => {
216 DecodingKey::from_rsa_components(¶ms.n, ¶ms.e)
217 }
218 AlgorithmParameters::EllipticCurve(params) => {
219 DecodingKey::from_ec_components(¶ms.x, ¶ms.y)
220 }
221 AlgorithmParameters::OctetKeyPair(params) => DecodingKey::from_ed_components(¶ms.x),
222 AlgorithmParameters::OctetKey(params) => {
223 let out = b64_decode(¶ms.value)?;
224 Ok(DecodingKey {
225 family: AlgorithmFamily::Hmac,
226 kind: DecodingKeyKind::SecretOrDer(out),
227 })
228 }
229 AlgorithmParameters::Other(_) => Err(ErrorKind::UnsupportedAlgorithm.into()),
230 }
231 }
232
233 pub fn try_get_as_bytes(&self) -> Result<&[u8]> {
237 match &self.kind {
238 DecodingKeyKind::SecretOrDer(b) => Ok(b),
239 DecodingKeyKind::RsaModulusExponent { .. } => Err(ErrorKind::InvalidKeyFormat.into()),
240 }
241 }
242}
243
244impl TryFrom<&Jwk> for DecodingKey {
245 type Error = crate::errors::Error;
246
247 fn try_from(jwk: &Jwk) -> Result<Self> {
248 Self::from_jwk(jwk)
249 }
250}
251
252pub fn decode<T: DeserializeOwned>(
271 token: impl AsRef<[u8]>,
272 key: &DecodingKey,
273 validation: &Validation,
274) -> Result<TokenData<T>> {
275 let token = token.as_ref();
276 let header = decode_header(token)?;
277
278 if !validation.algorithms.contains(&header.alg) {
279 return Err(new_error(ErrorKind::InvalidAlgorithm));
280 }
281
282 let verifying_provider = (CryptoProvider::get_default().verifier_factory)(&header.alg, key)?;
283
284 let (header, claims) = verify_signature(token, validation, verifying_provider)?;
285
286 let decoded_claims = DecodedJwtPartClaims::from_jwt_part_claims(claims)?;
287 let claims = decoded_claims.deserialize()?;
288 validate(decoded_claims.deserialize()?, validation)?;
289
290 Ok(TokenData { header, claims })
291}
292
293pub fn insecure_decode<T: DeserializeOwned>(token: impl AsRef<[u8]>) -> Result<TokenData<T>> {
297 let token = token.as_ref();
298
299 let (_, message) = expect_two!(token.rsplitn(2, |b| *b == b'.'));
300 let (payload, header) = expect_two!(message.rsplitn(2, |b| *b == b'.'));
301
302 let header = Header::from_encoded(header)?;
303 let claims = DecodedJwtPartClaims::from_jwt_part_claims(payload)?.deserialize()?;
304
305 Ok(TokenData { header, claims })
306}
307
308pub fn insecure_decode_claims<T: DeserializeOwned>(token: impl AsRef<[u8]>) -> Result<T> {
312 let token = token.as_ref();
313
314 let (_, message) = expect_two!(token.rsplitn(2, |b| *b == b'.'));
315 let (payload, _) = expect_two!(message.rsplitn(2, |b| *b == b'.'));
316
317 let claims = DecodedJwtPartClaims::from_jwt_part_claims(payload)?.deserialize()?;
318
319 Ok(claims)
320}
321
322pub fn decode_header(token: impl AsRef<[u8]>) -> Result<Header> {
333 let token = token.as_ref();
334 let (_, message) = expect_two!(token.rsplitn(2, |b| *b == b'.'));
335 let (_, header) = expect_two!(message.rsplitn(2, |b| *b == b'.'));
336 Header::from_encoded(header)
337}
338
339pub(crate) fn verify_signature_body(
340 message: &[u8],
341 signature: &[u8],
342 header: &Header,
343 validation: &Validation,
344 verifying_provider: Box<dyn JwtVerifier>,
345) -> Result<()> {
346 if validation.algorithms.is_empty() {
347 return Err(new_error(ErrorKind::MissingAlgorithm));
348 }
349
350 for alg in &validation.algorithms {
351 if verifying_provider.algorithm().family() != alg.family() {
352 return Err(new_error(ErrorKind::InvalidAlgorithm));
353 }
354 }
355
356 if !validation.algorithms.contains(&header.alg) {
357 return Err(new_error(ErrorKind::InvalidAlgorithm));
358 }
359
360 if verifying_provider.verify(message, &b64_decode(signature)?).is_err() {
361 return Err(new_error(ErrorKind::InvalidSignature));
362 }
363
364 Ok(())
365}
366
367fn verify_signature<'a>(
371 token: &'a [u8],
372 validation: &Validation,
373 verifying_provider: Box<dyn JwtVerifier>,
374) -> Result<(Header, &'a [u8])> {
375 let (signature, message) = expect_two!(token.rsplitn(2, |b| *b == b'.'));
376 let (payload, header) = expect_two!(message.rsplitn(2, |b| *b == b'.'));
377 let header = Header::from_encoded(header)?;
378 verify_signature_body(message, signature, &header, validation, verifying_provider)?;
379
380 Ok((header, payload))
381}