Skip to main content

json_web_tolkien/
jwt.rs

1mod encoded_field;
2
3use std::fmt;
4use std::marker::PhantomData;
5
6use base64ct::{Base64UrlUnpadded, Encoding};
7use serde::de::IntoDeserializer;
8use serde::de::value::SeqDeserializer;
9use serde::{Deserialize, Serialize, de, ser};
10
11use crate::error::Error;
12use crate::ser::Serializer;
13
14use encoded_field::EncodedField;
15
16/// The JSON Web Token
17/// This uses generics, for the claims and header, all that is required is that they are structs which implement [Deserialize](https://serde.rs/).
18/// See: [What is a JSON Web Token](https://www.jwt.io/introduction#what-is-json-web-token)
19#[derive(PartialEq)]
20#[cfg_attr(feature = "debug", derive(Debug))]
21pub struct Jwt<C, H> {
22    claims: C,
23    header: H,
24}
25
26impl<C, H> Jwt<C, H>
27where
28    C: de::DeserializeOwned + Sized,
29    H: de::DeserializeOwned + Sized,
30{
31    pub fn new(claims: C, header: H) -> Self {
32        Self { claims, header }
33    }
34
35    pub fn claims(&self) -> &C {
36        &self.claims
37    }
38
39    pub fn header(&self) -> &H {
40        &self.header
41    }
42}
43
44impl<Claims, Header> Jwt<Claims, Header>
45where
46    Claims: Serialize,
47    Header: Serialize,
48{
49    pub fn to_string(&self) -> Result<String, Error> {
50        let mut serializer = Serializer {
51            output: String::new(),
52        };
53
54        self.serialize(&mut serializer)?;
55
56        Ok(serializer.output)
57    }
58}
59
60impl<'de, C, H> Deserialize<'de> for Jwt<C, H>
61where
62    C: de::DeserializeOwned,
63    H: de::DeserializeOwned,
64{
65    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
66    where
67        D: de::Deserializer<'de>,
68    {
69        struct Visitor<C, H>(PhantomData<(C, H)>);
70
71        impl<'de, C, H> de::Visitor<'de> for Visitor<C, H>
72        where
73            C: de::DeserializeOwned,
74            H: de::DeserializeOwned,
75        {
76            type Value = Jwt<C, H>;
77
78            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
79                formatter.write_str("Expecting base64 encoded strings joined with a '.'")
80            }
81
82            fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
83            where
84                A: de::SeqAccess<'de>,
85            {
86                let header: H = seq
87                    .next_element_seed(EncodedField(PhantomData))?
88                    .ok_or_else(|| de::Error::invalid_length(0, &self))?;
89                let claims: C = seq
90                    .next_element_seed(EncodedField(PhantomData))?
91                    .ok_or_else(|| de::Error::invalid_length(1, &self))?;
92
93                Ok(Jwt { header, claims })
94            }
95
96            fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
97            where
98                E: de::Error,
99            {
100                let parts = v.split('.');
101
102                let de = SeqDeserializer::new(parts);
103
104                self.visit_seq(de)
105            }
106        }
107
108        let visitor = Visitor(PhantomData);
109
110        deserializer.deserialize_str(visitor)
111    }
112}
113
114impl<C, H> Serialize for Jwt<C, H>
115where
116    C: Serialize,
117    H: Serialize,
118{
119    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
120    where
121        S: ser::Serializer,
122    {
123        // NOTE: Its possible there may be some optimization writing a base64 serializer
124        // so I can then have serde transcode from json to base64 without allocating as
125        // much.
126
127        let ser_claims: String = {
128            let ser_claims: Vec<u8> =
129                serde_json::to_vec(&self.claims).map_err(ser::Error::custom)?;
130
131            Base64UrlUnpadded::encode_string(&ser_claims)
132        };
133
134        let ser_header: String = {
135            let ser_header: Vec<u8> =
136                serde_json::to_vec(&self.header).map_err(ser::Error::custom)?;
137
138            Base64UrlUnpadded::encode_string(&ser_header)
139        };
140
141        serializer.collect_str(&format_args!("{}.{}", ser_header, ser_claims))
142    }
143}
144
145impl<C, H> TryFrom<&str> for Jwt<C, H>
146where
147    C: de::DeserializeOwned,
148    H: de::DeserializeOwned,
149{
150    type Error = Error;
151
152    fn try_from(value: &str) -> Result<Self, Self::Error> {
153        let de = value.into_deserializer();
154
155        Self::deserialize(de)
156    }
157}
158
159#[cfg(test)]
160mod test {
161    use serde_test::{Token, assert_tokens};
162    use std::time;
163
164    use super::*;
165    use serde::{Deserialize, Serialize};
166
167    use crate::numeric_date::NumericDate;
168
169    #[derive(Debug, Deserialize, PartialEq, Serialize)]
170    struct TestHeader {
171        alg: Box<str>,
172        typ: Box<str>,
173    }
174
175    #[derive(Debug, Deserialize, PartialEq, Serialize)]
176    struct TestClaims {
177        #[serde(alias = "sub")]
178        user_id: Box<str>,
179        name: Box<str>,
180        admin: bool,
181        #[serde(alias = "iat", with = "NumericDate")]
182        issued_at: time::Duration,
183    }
184
185    #[cfg(feature = "debug")]
186    #[test]
187    fn test_ser_de() {
188        let header = TestHeader {
189            alg: Box::from("HS256"),
190            typ: Box::from("JWT"),
191        };
192        let claims = TestClaims {
193            admin: false,
194            issued_at: time::Duration::new(1516239022, 0),
195            name: Box::from("John Doe"),
196            user_id: Box::from("1234567890"),
197        };
198
199        let jwt = Jwt { header, claims };
200
201        assert_tokens(
202            &jwt,
203            &[Token::Str(
204                "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJ1c2VyX2lkIjoiMTIzNDU2Nzg5MCIsIm5hbWUiOiJKb2huIERvZSIsImFkbWluIjpmYWxzZSwiaXNzdWVkX2F0IjoxNTE2MjM5MDIyfQ",
205            )],
206        );
207    }
208}