Skip to main content

token_claims/
lib.rs

1use derive_builder::Builder;
2use jsonwebtoken;
3use serde::{Deserialize, Serialize, de::DeserializeOwned};
4use thiserror::Error;
5
6/// Represents the claims in a JWT token.
7#[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>, // Subject
28    pub exp: TimeStamp, // Expiration time (Unix timestamp)
29    pub iat: TimeStamp, // Issued at (Unix timestamp)
30    pub typ: String,    // Type
31    pub iss: String,    // Issuer
32    pub aud: String,    // Audience
33    pub jti: JWTID,     // JWT ID
34}
35
36impl<T> TokenClaims<T>
37where
38    T: Serialize + DeserializeOwned,
39{
40    /// Returns the subject of the token.
41    pub fn sub(&self) -> &Subject<T> {
42        &self.sub
43    }
44    /// Returns the expiration time of the token.
45    pub fn exp(&self) -> TimeStamp {
46        self.exp
47    }
48    /// Returns the issued-at time of the token.
49    pub fn iat(&self) -> TimeStamp {
50        self.iat
51    }
52    /// Returns the issuer of the token.
53    pub fn iss(&self) -> &str {
54        &self.iss
55    }
56    /// Returns the type of the token.
57    pub fn typ(&self) -> &str {
58        &self.typ
59    }
60    /// Returns the audience of the token.
61    pub fn aud(&self) -> &str {
62        &self.aud
63    }
64    /// Returns the JWT ID of the token.
65    pub fn jti(&self) -> &JWTID {
66        &self.jti
67    }
68}
69
70/// Represents the subject of a JWT token.
71#[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/// Represents a timestamp in a JWT token.
100#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
101#[serde(transparent)]
102pub struct TimeStamp(i64);
103
104impl TimeStamp {
105    /// Returns the current timestamp.
106    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    /// Creates a `TimeStamp` from a raw Unix timestamp.
113    pub fn from_i64(timestamp: i64) -> Self {
114        TimeStamp(timestamp)
115    }
116    /// Checks if the timestamp is expired.
117    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/// Custom error type for token-related operations.
147#[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/// Represents a unique JWT ID (JTI).
162#[derive(Debug, Clone, Serialize, Deserialize)]
163#[serde(transparent)]
164pub struct JWTID(String);
165
166impl JWTID {
167    /// Creates a new unique JWT ID.
168    pub fn new() -> Self {
169        JWTID(uuid::Uuid::new_v4().to_string())
170    }
171    /// Creates a `JWTID` from a string.
172    pub fn from_string(id: &str) -> Self {
173        JWTID(id.to_string())
174    }
175    /// Converts the `JWTID` to a string.
176    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}