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#[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 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}