1use derive_builder::Builder;
2use jsonwebtoken;
3use serde::{Deserialize, Serialize, de::DeserializeOwned};
4use thiserror::Error;
5
6#[derive(Debug, Serialize, Deserialize, Clone, Builder)]
8#[builder(pattern = "owned", setter(into), build_fn(error = "TokenError"))]
9pub struct TokenClaims<T>
10where
11 T: Serialize + DeserializeOwned,
12{
13 #[cfg_attr(
14 feature = "msgpack",
15 serde(
16 serialize_with = "serde_helpers::serialize_msgpack",
17 deserialize_with = "serde_helpers::deserialize_msgpack"
18 )
19 )]
20 #[cfg_attr(
21 not(feature = "msgpack"),
22 serde(
23 serialize_with = "serde_helpers::serialize_json",
24 deserialize_with = "serde_helpers::deserialize_json"
25 )
26 )]
27 pub sub: Subject<T>, pub exp: TimeStamp, pub iat: TimeStamp, pub typ: String, pub iss: String, pub aud: String, pub jti: JWTID, }
35
36impl<T> TokenClaims<T>
37where
38 T: Serialize + DeserializeOwned,
39{
40 pub fn sub(&self) -> &Subject<T> {
42 &self.sub
43 }
44 pub fn exp(&self) -> TimeStamp {
46 self.exp
47 }
48 pub fn iat(&self) -> TimeStamp {
50 self.iat
51 }
52 pub fn iss(&self) -> &str {
54 &self.iss
55 }
56 pub fn typ(&self) -> &str {
58 &self.typ
59 }
60 pub fn aud(&self) -> &str {
62 &self.aud
63 }
64 pub fn jti(&self) -> &JWTID {
66 &self.jti
67 }
68}
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
72#[serde(transparent)]
73pub struct Subject<T>(pub T);
74
75impl<T> Subject<T> {
76 pub fn new(sub: T) -> Self {
77 Self(sub)
78 }
79 pub fn value(&self) -> &T {
80 &self.0
81 }
82 pub fn into_inner(self) -> T {
83 self.0
84 }
85}
86
87impl<T> std::ops::Deref for Subject<T> {
88 type Target = T;
89 fn deref(&self) -> &Self::Target {
90 &self.0
91 }
92}
93impl<T> std::ops::DerefMut for Subject<T> {
94 fn deref_mut(&mut self) -> &mut Self::Target {
95 &mut self.0
96 }
97}
98
99#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
101#[serde(transparent)]
102pub struct TimeStamp(i64);
103
104impl TimeStamp {
105 pub fn now() -> i64 {
107 chrono::Utc::now().timestamp()
108 }
109 pub fn from_now(seconds: i64) -> Self {
110 TimeStamp(Self::now() + seconds)
111 }
112 pub fn from_i64(timestamp: i64) -> Self {
114 TimeStamp(timestamp)
115 }
116 pub fn is_expired(&self) -> bool {
118 self.0 < TimeStamp::now()
119 }
120 pub fn left_till(&self) -> i64 {
121 self.0 - TimeStamp::now()
122 }
123 pub fn extend(&mut self, seconds: i64) {
124 self.0 += seconds;
125 }
126 pub fn to_i64(&self) -> i64 {
127 self.0
128 }
129}
130
131impl From<i64> for TimeStamp {
132 fn from(value: i64) -> Self {
133 TimeStamp(value)
134 }
135}
136
137impl<T> From<T> for Subject<T>
138where
139 T: Serialize + DeserializeOwned,
140{
141 fn from(value: T) -> Self {
142 Subject::new(value)
143 }
144}
145
146#[derive(Error, Debug)]
148pub enum TokenError {
149 #[error("JWT error: {0}")]
150 JwtError(#[from] jsonwebtoken::errors::Error),
151 #[error("Invalid token format")]
152 InvalidTokenFormat,
153}
154
155impl From<derive_builder::UninitializedFieldError> for TokenError {
156 fn from(_: derive_builder::UninitializedFieldError) -> Self {
157 TokenError::InvalidTokenFormat
158 }
159}
160
161#[derive(Debug, Clone, Serialize, Deserialize)]
163#[serde(transparent)]
164pub struct JWTID(String);
165
166impl JWTID {
167 pub fn new() -> Self {
169 JWTID(uuid::Uuid::new_v4().to_string())
170 }
171 pub fn from_string(id: &str) -> Self {
173 JWTID(id.to_string())
174 }
175 pub fn to_string(&self) -> String {
177 self.0.clone()
178 }
179}
180impl std::fmt::Display for JWTID {
181 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
182 write!(f, "{}", self.0)
183 }
184}
185#[cfg(feature = "msgpack")]
186mod serde_helpers {
187 use base64::{Engine as _, engine::general_purpose};
188 use rmp_serde::{from_slice, to_vec};
189 use serde::Deserialize;
190 use serde::{Serialize, de::DeserializeOwned};
191
192 pub fn serialize_msgpack<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
193 where
194 T: Serialize,
195 S: serde::Serializer,
196 {
197 let bytes = to_vec(value).map_err(serde::ser::Error::custom)?;
198 let b64 = general_purpose::STANDARD.encode(bytes);
199 serializer.serialize_str(&b64)
200 }
201
202 pub fn deserialize_msgpack<'de, T, D>(deserializer: D) -> Result<T, D::Error>
203 where
204 T: DeserializeOwned,
205 D: serde::Deserializer<'de>,
206 {
207 let s = String::deserialize(deserializer)?;
208 let bytes = general_purpose::STANDARD
209 .decode(&s)
210 .map_err(serde::de::Error::custom)?;
211 from_slice(&bytes).map_err(serde::de::Error::custom)
212 }
213
214 pub fn serialize_json<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
215 where
216 T: Serialize,
217 S: serde::Serializer,
218 {
219 let json_str = serde_json::to_string(value).map_err(serde::ser::Error::custom)?;
220 serializer.serialize_str(&json_str)
221 }
222
223 pub fn deserialize_json<'de, T, D>(deserializer: D) -> Result<T, D::Error>
224 where
225 T: DeserializeOwned,
226 D: serde::Deserializer<'de>,
227 {
228 let s = String::deserialize(deserializer)?;
229 serde_json::from_str(&s).map_err(serde::de::Error::custom)
230 }
231}
232
233#[cfg(not(feature = "msgpack"))]
234mod serde_helpers {
235 use serde::Deserialize;
236 use serde::{Serialize, de::DeserializeOwned};
237
238 pub fn serialize_json<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
239 where
240 T: Serialize,
241 S: serde::Serializer,
242 {
243 let json_str = serde_json::to_string(value).map_err(serde::ser::Error::custom)?;
244 serializer.serialize_str(&json_str)
245 }
246
247 pub fn deserialize_json<'de, T, D>(deserializer: D) -> Result<T, D::Error>
248 where
249 T: DeserializeOwned,
250 D: serde::Deserializer<'de>,
251 {
252 let s = String::deserialize(deserializer)?;
253 serde_json::from_str(&s).map_err(serde::de::Error::custom)
254 }
255}
256#[cfg(test)]
257mod tests {
258 use super::*;
259 use serde::{Deserialize, Serialize};
260 #[derive(Debug, Serialize, Deserialize, PartialEq, Eq, Clone)]
261 struct CustomClaims {
262 username: String,
263 admin: bool,
264 }
265
266 #[test]
267 fn test_token_encode_decode_custom_claims() {
268 let claims = TokenClaimsBuilder::<CustomClaims>::default()
269 .sub(Subject::new(CustomClaims {
270 username: "alice".to_string(),
271 admin: true,
272 }))
273 .exp(TimeStamp::from_now(3600))
274 .iat(TimeStamp::from_now(0))
275 .typ("access".to_string())
276 .iss("issuer".to_string())
277 .aud("audience".to_string())
278 .jti(JWTID::new())
279 .build()
280 .unwrap();
281
282 let secret = b"supersecretkey";
283 let token = jsonwebtoken::encode(
284 &jsonwebtoken::Header::default(),
285 &claims,
286 &jsonwebtoken::EncodingKey::from_secret(secret),
287 )
288 .unwrap();
289 let mut validation = jsonwebtoken::Validation::default();
290 validation.set_audience(&["audience"]);
291 let decoded_claims: TokenClaims<CustomClaims> =
292 jsonwebtoken::decode::<TokenClaims<CustomClaims>>(
293 &token,
294 &jsonwebtoken::DecodingKey::from_secret(secret),
295 &validation,
296 )
297 .unwrap()
298 .claims;
299
300 assert_eq!(decoded_claims.sub().value(), claims.sub().value());
301 assert_eq!(decoded_claims.typ(), claims.typ());
302 assert_eq!(decoded_claims.iss(), claims.iss());
303 assert_eq!(decoded_claims.aud(), claims.aud());
304 assert_eq!(decoded_claims.jti().to_string(), claims.jti().to_string());
305 }
306
307 #[test]
308 fn test_token_encode_decode_primitive_claims() {
309 let claims = TokenClaimsBuilder::<u32>::default()
310 .sub(Subject::new(42u32))
311 .exp(TimeStamp::from_now(3600))
312 .iat(TimeStamp::from_now(0))
313 .typ("number".to_string())
314 .iss("issuer".to_string())
315 .aud("audience".to_string())
316 .jti(JWTID::new())
317 .build()
318 .unwrap();
319
320 let secret = b"anothersecret";
321 let token = jsonwebtoken::encode(
322 &jsonwebtoken::Header::default(),
323 &claims,
324 &jsonwebtoken::EncodingKey::from_secret(secret),
325 )
326 .unwrap();
327
328 let mut validation = jsonwebtoken::Validation::default();
329 validation.set_audience(&["audience"]);
330 let decoded_claims: TokenClaims<u32> = jsonwebtoken::decode::<TokenClaims<u32>>(
331 &token,
332 &jsonwebtoken::DecodingKey::from_secret(secret),
333 &validation,
334 )
335 .unwrap()
336 .claims;
337
338 assert_eq!(decoded_claims.sub().value(), claims.sub().value());
339 assert_eq!(decoded_claims.typ(), claims.typ());
340 assert_eq!(decoded_claims.iss(), claims.iss());
341 assert_eq!(decoded_claims.aud(), claims.aud());
342 assert_eq!(decoded_claims.jti().to_string(), claims.jti().to_string());
343 }
344}