use chrono::{DateTime, Duration, Utc};
use serde::{Deserialize, Serialize};
use crate::ValidationError;
#[derive(Debug, Clone, Copy)]
pub struct TimeOptions {
pub leeway: Duration,
pub current_time: Option<DateTime<Utc>>,
}
impl Default for TimeOptions {
fn default() -> Self {
Self {
leeway: Duration::seconds(60),
current_time: None,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct Empty {}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Claims<T> {
#[serde(
rename = "exp",
default,
skip_serializing_if = "Option::is_none",
with = "self::serde_timestamp"
)]
pub expiration_date: Option<DateTime<Utc>>,
#[serde(
rename = "nbf",
default,
skip_serializing_if = "Option::is_none",
with = "self::serde_timestamp"
)]
pub not_before: Option<DateTime<Utc>>,
#[serde(
rename = "iat",
default,
skip_serializing_if = "Option::is_none",
with = "self::serde_timestamp"
)]
pub issued_at: Option<DateTime<Utc>>,
#[serde(flatten)]
pub custom: T,
}
impl Claims<Empty> {
pub fn empty() -> Self {
Self {
expiration_date: None,
not_before: None,
issued_at: None,
custom: Empty {},
}
}
}
impl<T> Claims<T> {
pub fn new(custom_claims: T) -> Self {
Self {
expiration_date: None,
not_before: None,
issued_at: None,
custom: custom_claims,
}
}
pub fn set_duration(self, duration: Duration) -> Self {
Self {
expiration_date: Some(Utc::now() + duration),
..self
}
}
pub fn set_duration_and_issuance(self, duration: Duration) -> Self {
let issued_at = Utc::now();
Self {
expiration_date: Some(issued_at + duration),
issued_at: Some(issued_at),
..self
}
}
pub fn set_not_before(self, moment: DateTime<Utc>) -> Self {
Self {
not_before: Some(moment),
..self
}
}
pub fn validate_expiration(&self, options: TimeOptions) -> Result<&Self, ValidationError> {
if let Some(expiration) = self.expiration_date {
let current_time = options.current_time.unwrap_or_else(Utc::now);
if current_time > expiration + options.leeway {
Err(ValidationError::Expired)
} else {
Ok(self)
}
} else {
Err(ValidationError::NoClaim)
}
}
pub fn validate_maturity(&self, options: TimeOptions) -> Result<&Self, ValidationError> {
if let Some(not_before) = self.not_before {
let current_time = options.current_time.unwrap_or_else(Utc::now);
if current_time < not_before - options.leeway {
Err(ValidationError::NotMature)
} else {
Ok(self)
}
} else {
Err(ValidationError::NoClaim)
}
}
}
mod serde_timestamp {
use chrono::{offset::TimeZone, DateTime, Utc};
use serde::{
de::{Error as DeError, Visitor},
Deserializer, Serializer,
};
use std::{convert::TryFrom, fmt};
struct TimestampVisitor;
impl<'de> Visitor<'de> for TimestampVisitor {
type Value = DateTime<Utc>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("UTC timestamp")
}
fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E>
where
E: DeError,
{
Ok(Utc.timestamp(value, 0))
}
fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E>
where
E: DeError,
{
let value = i64::try_from(value).map_err(DeError::custom)?;
Ok(Utc.timestamp(value, 0))
}
}
pub fn serialize<S: Serializer>(
time: &Option<DateTime<Utc>>,
serializer: S,
) -> Result<S::Ok, S::Error> {
serializer.serialize_i64(time.unwrap().timestamp())
}
pub fn deserialize<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<DateTime<Utc>>, D::Error> {
deserializer.deserialize_i64(TimestampVisitor).map(Some)
}
}
#[cfg(test)]
mod tests {
use super::*;
use assert_matches::assert_matches;
#[test]
fn empty_claims_can_be_serialized() {
let mut claims = Claims::empty();
assert!(serde_json::to_string(&claims).is_ok());
assert!(serde_cbor::to_vec(&claims).is_ok());
claims.expiration_date = Some(Utc::now());
assert!(serde_json::to_string(&claims).is_ok());
assert!(serde_cbor::to_vec(&claims).is_ok());
claims.not_before = Some(Utc::now());
assert!(serde_json::to_string(&claims).is_ok());
assert!(serde_cbor::to_vec(&claims).is_ok());
}
#[test]
fn expired_claim() {
let mut claims = Claims::empty();
assert_matches!(
claims
.validate_expiration(TimeOptions::default())
.unwrap_err(),
ValidationError::NoClaim
);
claims.expiration_date = Some(Utc::now() - Duration::hours(1));
assert_matches!(
claims
.validate_expiration(TimeOptions::default())
.unwrap_err(),
ValidationError::Expired
);
claims.expiration_date = Some(Utc::now() - Duration::seconds(10));
assert!(claims.validate_expiration(TimeOptions::default()).is_ok());
assert_matches!(
claims
.validate_expiration(TimeOptions {
leeway: Duration::seconds(5),
..Default::default()
})
.unwrap_err(),
ValidationError::Expired
);
}
#[test]
fn immature_claim() {
let mut claims = Claims::empty();
assert_matches!(
claims
.validate_maturity(TimeOptions::default())
.unwrap_err(),
ValidationError::NoClaim
);
claims.not_before = Some(Utc::now() + Duration::hours(1));
assert_matches!(
claims
.validate_maturity(TimeOptions::default())
.unwrap_err(),
ValidationError::NotMature
);
claims.not_before = Some(Utc::now() + Duration::seconds(10));
assert!(claims.validate_maturity(TimeOptions::default()).is_ok());
assert_matches!(
claims
.validate_maturity(TimeOptions {
leeway: Duration::seconds(5),
..Default::default()
})
.unwrap_err(),
ValidationError::NotMature
);
}
}