#![forbid(unsafe_code)]
#![deny(missing_docs)]
#![deny(missing_debug_implementations)]
#![deny(unreachable_pub)]
#![deny(unused_qualifications)]
#![deny(rust_2018_idioms)]
#![warn(clippy::all)]
use axum::{
BoxError,
extract::{ConnectInfo, FromRequestParts, Request},
http::{Method, StatusCode, header, request::Parts},
response::Response,
};
use base64::Engine as _;
use cookie::time::{Duration, OffsetDateTime, PrimitiveDateTime};
pub use cookie::{Key, SameSite};
use futures_util::future::BoxFuture;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use std::net::{IpAddr, SocketAddr};
use std::task::{Context, Poll};
use tower_layer::Layer;
use tower_service::Service;
type HmacSha256 = Hmac<Sha256>;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CidrBlock {
addr: IpAddr,
prefix_len: u8,
}
impl CidrBlock {
pub fn addr(&self) -> IpAddr {
self.addr
}
pub fn prefix_len(&self) -> u8 {
self.prefix_len
}
pub fn contains(&self, ip: &IpAddr) -> bool {
match (self.addr, ip) {
(IpAddr::V4(net), IpAddr::V4(ip)) => {
let mask = if self.prefix_len == 0 {
0
} else {
u32::MAX << (32 - self.prefix_len)
};
net.to_bits() & mask == ip.to_bits() & mask
}
(IpAddr::V6(net), IpAddr::V6(ip)) => {
let mask = if self.prefix_len == 0 {
0
} else {
u128::MAX << (128 - self.prefix_len)
};
net.to_bits() & mask == ip.to_bits() & mask
}
_ => false,
}
}
}
impl std::fmt::Display for CidrBlock {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}/{}", self.addr, self.prefix_len)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CidrParseError;
impl std::fmt::Display for CidrParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("invalid CIDR block (expected `address/prefix`)")
}
}
impl std::error::Error for CidrParseError {}
impl std::str::FromStr for CidrBlock {
type Err = CidrParseError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let (addr_str, prefix_str) = s.split_once('/').ok_or(CidrParseError)?;
let addr: IpAddr = addr_str.parse().map_err(|_| CidrParseError)?;
let prefix_len: u8 = prefix_str.parse().map_err(|_| CidrParseError)?;
let max_prefix = if addr.is_ipv4() { 32 } else { 128 };
if prefix_len > max_prefix {
return Err(CidrParseError);
}
Ok(CidrBlock { addr, prefix_len })
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for CidrBlock {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.collect_str(self)
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for CidrBlock {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let s = std::borrow::Cow::<'de, str>::deserialize(deserializer)?;
s.parse().map_err(serde::de::Error::custom)
}
}
const TOKEN_KEY_INFO: &[u8] = b"axum-token-auth/token-mac/v1";
const TOKEN_B64: base64::engine::general_purpose::GeneralPurpose =
base64::engine::general_purpose::URL_SAFE_NO_PAD;
const TOKEN_VERSION: u8 = 1;
fn token_mac_key(key: &Key) -> [u8; 32] {
let mut kdf =
HmacSha256::new_from_slice(key.master()).expect("HMAC accepts keys of any length");
kdf.update(TOKEN_KEY_INFO);
let out = kdf.finalize().into_bytes();
let mut subkey = [0u8; 32];
subkey.copy_from_slice(&out);
subkey
}
fn token_mac(key: &Key, version: u8, expiry_unix: i64) -> HmacSha256 {
let mut mac =
HmacSha256::new_from_slice(&token_mac_key(key)).expect("HMAC accepts keys of any length");
mac.update(&[version]);
mac.update(&expiry_unix.to_le_bytes());
mac
}
fn sign_token(key: &Key, expiry: OffsetDateTime) -> String {
let expiry_unix = expiry.unix_timestamp();
let mac = token_mac(key, TOKEN_VERSION, expiry_unix)
.finalize()
.into_bytes();
let mut buf = Vec::with_capacity(1 + 8 + mac.len());
buf.push(TOKEN_VERSION);
buf.extend_from_slice(&expiry_unix.to_le_bytes());
buf.extend_from_slice(&mac);
TOKEN_B64.encode(buf)
}
fn verify_token(key: &Key, token: &str, now: OffsetDateTime) -> bool {
let Ok(buf) = TOKEN_B64.decode(token) else {
return false;
};
let Some((&version, rest)) = buf.split_first() else {
return false;
};
if version != TOKEN_VERSION {
return false;
}
let Some((expiry_bytes, mac_bytes)) = rest.split_first_chunk::<8>() else {
return false;
};
let expiry_unix = i64::from_le_bytes(*expiry_bytes);
if token_mac(key, version, expiry_unix)
.verify_slice(mac_bytes)
.is_err()
{
return false;
}
match OffsetDateTime::from_unix_timestamp(expiry_unix) {
Ok(expiry) => now < expiry,
Err(_) => false,
}
}
fn saturating_expiry(now: OffsetDateTime, ttl: std::time::Duration) -> OffsetDateTime {
Duration::try_from(ttl)
.ok()
.and_then(|ttl| now.checked_add(ttl))
.unwrap_or_else(|| PrimitiveDateTime::MAX.assume_utc())
}
fn parse_session_cookie(value: &str) -> Option<(SessionKey, Option<OffsetDateTime>)> {
let (uuid_str, expiry) = match value.split_once('.') {
Some((uuid_str, expiry_str)) => {
let secs: i64 = expiry_str.parse().ok()?;
(
uuid_str,
Some(OffsetDateTime::from_unix_timestamp(secs).ok()?),
)
}
None => (value, None),
};
let uuid = uuid::Uuid::parse_str(uuid_str).ok()?;
Some((SessionKey(uuid), expiry))
}
#[derive(thiserror::Error, Debug)]
#[error("one or more validation errors")]
pub struct ValidationErrors(Vec<String>);
impl ValidationErrors {
pub fn errors(&self) -> impl Iterator<Item = &str> {
self.0.iter().map(String::as_str)
}
}
#[derive(Debug, Clone, Eq, PartialEq, Hash)]
pub struct SessionKey(pub uuid::Uuid);
impl SessionKey {
pub fn is_present(&self) {}
}
impl Default for SessionKey {
fn default() -> Self {
SessionKey(uuid::Uuid::new_v4())
}
}
#[derive(Clone, Debug)]
pub struct TokenConfig {
pub name: String,
}
impl TokenConfig {
pub fn new(name: &str) -> Self {
Self { name: name.into() }
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct AuthConfig<'a> {
pub cookie_name: &'a str,
pub persistent_secret: Key,
pub token_config: Option<TokenConfig>,
pub session_expires: Option<std::time::Duration>,
pub cookie_secure: bool,
pub cookie_http_only: bool,
pub cookie_same_site: Option<SameSite>,
pub trusted_networks: Vec<CidrBlock>,
pub strip_token_redirect: bool,
}
impl Default for AuthConfig<'_> {
fn default() -> Self {
Self {
cookie_name: env!["CARGO_PKG_NAME"],
persistent_secret: Key::generate(),
token_config: None,
session_expires: None,
cookie_secure: false,
cookie_http_only: true,
cookie_same_site: Some(SameSite::Strict),
trusted_networks: Vec::new(),
strip_token_redirect: true,
}
}
}
impl AuthConfig<'_> {
pub fn new(persistent_secret: Key) -> Self {
Self {
persistent_secret,
..Self::default()
}
}
pub fn into_layer(self) -> AuthLayer {
let access_info = AccessInfo::new(self);
AuthLayer { access_info }
}
pub fn generate_token(&self, ttl: std::time::Duration) -> String {
generate_token(&self.persistent_secret, ttl)
}
}
pub fn generate_token(secret: &Key, ttl: std::time::Duration) -> String {
let expiry = saturating_expiry(OffsetDateTime::now_utc(), ttl);
sign_token(secret, expiry)
}
enum SessionAction {
Issue(SessionKey),
Renew(SessionKey),
Keep(SessionKey),
}
impl SessionAction {
fn session_key(&self) -> &SessionKey {
match self {
SessionAction::Issue(sk) | SessionAction::Renew(sk) | SessionAction::Keep(sk) => sk,
}
}
}
#[derive(Clone, Debug)]
struct AccessInfo {
cookie_name: String,
token_config: Option<TokenConfig>,
session_expires: Option<std::time::Duration>,
cookie_secure: bool,
cookie_http_only: bool,
cookie_same_site: Option<SameSite>,
trusted_networks: Vec<CidrBlock>,
strip_token_redirect: bool,
key: Key,
}
impl AccessInfo {
fn new(cfg: AuthConfig<'_>) -> Self {
let AuthConfig {
cookie_name,
persistent_secret,
token_config,
session_expires,
cookie_secure,
cookie_http_only,
cookie_same_site,
trusted_networks,
strip_token_redirect,
} = cfg;
let key = persistent_secret;
Self {
cookie_name: cookie_name.into(),
token_config,
key,
session_expires,
cookie_secure,
cookie_http_only,
cookie_same_site,
trusted_networks,
strip_token_redirect,
}
}
fn is_trusted_client(&self, req: &Request) -> bool {
if self.trusted_networks.is_empty() {
return false;
}
let Some(ConnectInfo(peer)) = req.extensions().get::<ConnectInfo<SocketAddr>>() else {
return false;
};
let ip: IpAddr = peer.ip();
self.trusted_networks.iter().any(|net| net.contains(&ip))
}
fn check_token(&self, req: &Request, now: OffsetDateTime) -> bool {
if self.is_trusted_client(req) {
return true;
}
let Some(token_config) = self.token_config.as_ref() else {
return true;
};
let query = req.uri().query().unwrap_or("");
for (key, value) in url::form_urlencoded::parse(query.as_bytes()) {
if key == token_config.name.as_str() && verify_token(&self.key, &value, now) {
return true;
}
}
false
}
fn token_strip_redirect_location(&self, req: &Request) -> Option<String> {
if !self.strip_token_redirect || req.method() != Method::GET {
return None;
}
let token_name = self.token_config.as_ref()?.name.as_str();
let accepts_html = req
.headers()
.get(header::ACCEPT)
.and_then(|v| v.to_str().ok())
.map(|accept| accept.contains("text/html"))
.unwrap_or(false);
if !accepts_html {
return None;
}
let uri = req.uri();
let query = uri.query()?;
let mut kept = Vec::new();
let mut had_token = false;
for pair in query.split('&') {
let name = pair.split('=').next().unwrap_or("");
if name == token_name {
had_token = true;
} else if !pair.is_empty() {
kept.push(pair);
}
}
if !had_token {
return None;
}
let path = uri.path();
Some(if kept.is_empty() {
path.to_string()
} else {
format!("{path}?{}", kept.join("&"))
})
}
fn authenticate(
&self,
req: &Request,
existing: Option<(SessionKey, Option<OffsetDateTime>)>,
now: OffsetDateTime,
) -> Result<SessionAction, ValidationErrors> {
let valid_session =
existing.filter(|&(_, expiry)| expiry.is_none_or(|expiry| now < expiry));
match valid_session {
Some((session_key, expiry)) => {
if self.should_renew(expiry, now) {
Ok(SessionAction::Renew(session_key))
} else {
Ok(SessionAction::Keep(session_key))
}
}
None => {
if self.check_token(req, now) {
Ok(SessionAction::Issue(SessionKey::default()))
} else {
Err(ValidationErrors(vec![
"No (valid) token in uri and no (valid) session.".into(),
]))
}
}
}
}
fn should_renew(&self, expiry: Option<OffsetDateTime>, now: OffsetDateTime) -> bool {
let Some(ttl) = self.session_expires else {
return false;
};
match expiry {
None => true,
Some(expiry) => {
let ttl = Duration::try_from(ttl).unwrap_or(Duration::ZERO);
(expiry - now) * 2 < ttl
}
}
}
fn build_cookie_value(
&self,
session_key: &SessionKey,
now: OffsetDateTime,
) -> (String, Option<OffsetDateTime>) {
match self.session_expires {
Some(ttl) => {
let expiry = saturating_expiry(now, ttl);
(
format!(
"{}.{}",
session_key.0.as_hyphenated(),
expiry.unix_timestamp()
),
Some(expiry),
)
}
None => (format!("{}", session_key.0.as_hyphenated()), None),
}
}
}
impl<S> FromRequestParts<S> for SessionKey
where
S: Send + Sync,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
if let Some(session_key) = parts.extensions.remove::<SessionKey>() {
Ok(session_key.clone())
} else {
Err((StatusCode::UNAUTHORIZED, "(valid) session key is missing"))
}
}
}
#[derive(Clone, Debug)]
pub struct AuthLayer {
access_info: AccessInfo,
}
impl<S> Layer<S> for AuthLayer {
type Service = tower_cookies::CookieManager<AuthMiddleware<S>>;
fn layer(&self, inner: S) -> Self::Service {
let auth_middleware = AuthMiddleware {
inner,
access_info: self.access_info.clone(),
};
tower_cookies::CookieManager::new(auth_middleware)
}
}
#[derive(Clone, Debug)]
pub struct AuthMiddleware<S> {
inner: S,
access_info: AccessInfo,
}
impl<S> Service<Request> for AuthMiddleware<S>
where
S: Service<Request, Response = Response> + Send + 'static,
S::Error: Into<BoxError>,
S::Future: Send + 'static,
{
type Response = S::Response;
type Error = BoxError;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
match self.inner.poll_ready(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(r) => Poll::Ready(r.map_err(Into::into)),
}
}
fn call(&mut self, mut request: Request) -> Self::Future {
let Some(cookies) = request
.extensions()
.get::<tower_cookies::Cookies>()
.cloned()
else {
tracing::error!("missing cookies request extension");
return Box::pin(std::future::ready(Err(Box::new(ValidationErrors(vec![
"missing cookies request extension".into(),
])) as BoxError)));
};
let signed = cookies.signed(&self.access_info.key);
let mut redirect_location: Option<String> = None;
let err_info = {
let now = OffsetDateTime::now_utc();
let existing = signed
.get(&self.access_info.cookie_name)
.and_then(|received_cookie| parse_session_cookie(received_cookie.value()));
match self.access_info.authenticate(&request, existing, now) {
Ok(action) => {
let session_key = action.session_key().clone();
request.extensions_mut().insert(session_key.clone());
if matches!(action, SessionAction::Issue(_) | SessionAction::Renew(_)) {
let (value, expires) =
self.access_info.build_cookie_value(&session_key, now);
let mut set_cookie =
tower_cookies::Cookie::new(self.access_info.cookie_name.clone(), value);
set_cookie.set_secure(self.access_info.cookie_secure);
set_cookie.set_http_only(self.access_info.cookie_http_only);
set_cookie.set_same_site(self.access_info.cookie_same_site);
if let Some(expires) = expires {
set_cookie.set_expires(expires);
}
signed.add(set_cookie);
}
redirect_location = self.access_info.token_strip_redirect_location(&request);
None
}
Err(val_err) => Some(val_err),
}
};
if let Some(val_err) = err_info {
return Box::pin(std::future::ready(Err(val_err.into())));
}
if let Some(location) = redirect_location {
let response = Response::builder()
.status(StatusCode::SEE_OTHER)
.header(header::LOCATION, location)
.body(axum::body::Body::empty())
.expect("building a redirect response cannot fail");
return Box::pin(std::future::ready(Ok(response)));
}
let fut = self.inner.call(request);
Box::pin(async move {
let response: Response = fut.await.map_err(|e| e.into())?;
Ok(response)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use anyhow::Result;
use axum::body::Body;
use cookie::Cookie;
use http::{Request, StatusCode};
use std::convert::Infallible;
use tower::{ServiceBuilder, ServiceExt};
async fn handler(_: Request<Body>) -> std::result::Result<Response<Body>, Infallible> {
Ok(Response::new(Body::empty()))
}
fn get_cfg() -> AuthConfig<'static> {
AuthConfig {
cookie_name: "auth",
persistent_secret: Key::generate(),
token_config: Some(TokenConfig::new("token")),
session_expires: None,
..Default::default()
}
}
fn valid_token_uri(cfg: &AuthConfig<'_>) -> String {
let name = &cfg.token_config.as_ref().unwrap().name;
let token = cfg.generate_token(std::time::Duration::from_secs(300));
format!("http://example.com/path?{name}={token}")
}
#[tokio::test]
async fn fail_without_token_or_cookie() -> Result<()> {
let auth_layer = get_cfg().into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder().body(Body::empty())?;
let res = svc.oneshot(req).await;
assert!(
!res.err()
.unwrap()
.downcast::<ValidationErrors>()
.unwrap()
.errors()
.collect::<Vec<_>>()
.is_empty()
);
Ok(())
}
async fn get_second_response(
cfg: AuthConfig<'_>,
req: Request<Body>,
) -> Result<Response<Body>> {
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let res = svc.clone().oneshot(req).await.unwrap();
let cookie = {
let set_cookie: Vec<_> = res.headers().get_all(header::SET_COOKIE).iter().collect();
assert_eq!(set_cookie.len(), 1);
Cookie::parse(set_cookie[0].to_str()?.to_string())?
};
let req2 = Request::builder()
.header(header::COOKIE, cookie.stripped().to_string())
.body(Body::empty())
.unwrap();
let res2 = svc.oneshot(req2).await.unwrap();
Ok(res2)
}
#[tokio::test]
async fn set_cookie_with_trusted_socket() -> Result<()> {
let mut cfg = get_cfg();
cfg.token_config = None;
let uri = "http://example.com/path";
let req = Request::builder().uri(uri).body(Body::empty()).unwrap();
let res2 = get_second_response(cfg, req).await?;
assert_eq!(res2.status(), StatusCode::OK);
Ok(())
}
#[tokio::test]
async fn set_cookie_with_valid_token() -> Result<()> {
let cfg = get_cfg();
let uri = valid_token_uri(&cfg);
let req = Request::builder().uri(uri).body(Body::empty()).unwrap();
let res2 = get_second_response(cfg, req).await?;
assert_eq!(res2.status(), StatusCode::OK);
Ok(())
}
#[tokio::test]
async fn legacy_bare_uuid_cookie_is_accepted() -> Result<()> {
let key = Key::generate();
let mut cfg = get_cfg();
cfg.persistent_secret = key.clone();
cfg.session_expires = Some(std::time::Duration::from_secs(60 * 60 * 24 * 400));
let cookie_name = cfg.cookie_name.to_string();
let legacy_value = format!("{}", uuid::Uuid::new_v4().as_hyphenated());
let mut jar = cookie::CookieJar::new();
jar.signed_mut(&key)
.add(Cookie::new(cookie_name.clone(), legacy_value));
let signed = jar.get(&cookie_name).unwrap().stripped().to_string();
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder()
.uri("http://example.com/path")
.header(header::COOKIE, signed)
.body(Body::empty())
.unwrap();
let res = svc.oneshot(req).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
Ok(())
}
#[tokio::test]
async fn issued_cookie_is_httponly_and_samesite_strict() -> Result<()> {
let mut cfg = get_cfg();
cfg.token_config = None;
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder()
.uri("http://example.com/path")
.body(Body::empty())
.unwrap();
let res = svc.oneshot(req).await.unwrap();
let set_cookie = res.headers().get(header::SET_COOKIE).unwrap().to_str()?;
let cookie = Cookie::parse(set_cookie.to_string())?;
assert_eq!(cookie.http_only(), Some(true));
assert_eq!(cookie.same_site(), Some(SameSite::Strict));
Ok(())
}
#[tokio::test]
async fn cookie_attributes_are_configurable() -> Result<()> {
let mut cfg = get_cfg();
cfg.token_config = None;
cfg.cookie_secure = true;
cfg.cookie_http_only = false;
cfg.cookie_same_site = Some(SameSite::Lax);
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder()
.uri("http://example.com/path")
.body(Body::empty())
.unwrap();
let res = svc.oneshot(req).await.unwrap();
let set_cookie = res.headers().get(header::SET_COOKIE).unwrap().to_str()?;
let cookie = Cookie::parse(set_cookie.to_string())?;
assert_eq!(cookie.secure(), Some(true));
assert_eq!(cookie.http_only(), None);
assert_eq!(cookie.same_site(), Some(SameSite::Lax));
Ok(())
}
#[tokio::test]
async fn reject_token_with_wrong_secret() -> Result<()> {
let other = get_cfg();
let uri = valid_token_uri(&other);
let cfg = get_cfg();
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder().uri(uri).body(Body::empty()).unwrap();
let res = svc.oneshot(req).await;
assert!(res.is_err());
Ok(())
}
#[test]
fn cidr_block_parse_and_contains() {
let net: CidrBlock = "100.64.0.0/10".parse().unwrap();
assert!(net.contains(&"100.64.0.1".parse().unwrap()));
assert!(net.contains(&"100.127.255.255".parse().unwrap()));
assert!(!net.contains(&"100.128.0.0".parse().unwrap()));
assert!(!net.contains(&"10.0.0.1".parse().unwrap()));
assert!(!net.contains(&"::1".parse().unwrap()));
assert!(
"0.0.0.0/0"
.parse::<CidrBlock>()
.unwrap()
.contains(&"8.8.8.8".parse().unwrap())
);
let host: CidrBlock = "192.168.1.5/32".parse().unwrap();
assert!(host.contains(&"192.168.1.5".parse().unwrap()));
assert!(!host.contains(&"192.168.1.6".parse().unwrap()));
let v6: CidrBlock = "fd00::/8".parse().unwrap();
assert!(v6.contains(&"fd00::1".parse().unwrap()));
assert!(!v6.contains(&"fe00::1".parse().unwrap()));
assert!("100.64.0.0".parse::<CidrBlock>().is_err());
assert!("100.64.0.0/33".parse::<CidrBlock>().is_err());
assert!("fd00::/129".parse::<CidrBlock>().is_err());
assert!("nonsense/8".parse::<CidrBlock>().is_err());
assert_eq!(net.addr(), "100.64.0.0".parse::<IpAddr>().unwrap());
assert_eq!(net.prefix_len(), 10);
}
#[cfg(feature = "serde")]
#[test]
fn cidr_block_serde_roundtrip() {
let net: CidrBlock = "100.64.0.0/10".parse().unwrap();
let json = serde_json::to_string(&net).unwrap();
assert_eq!(json, "\"100.64.0.0/10\"");
assert_eq!(serde_json::from_str::<CidrBlock>(&json).unwrap(), net);
assert!(serde_json::from_str::<CidrBlock>("\"nonsense\"").is_err());
}
#[test]
fn token_roundtrip_signature_and_expiry() {
let key = Key::generate();
let now = OffsetDateTime::now_utc();
let token = sign_token(&key, now + Duration::minutes(5));
assert!(verify_token(&key, &token, now));
assert!(!verify_token(&key, &token, now + Duration::minutes(6)));
assert!(!verify_token(&key, &format!("{token}x"), now));
assert!(!verify_token(&Key::generate(), &token, now));
assert!(!verify_token(&key, "not base64!!", now));
let mut bytes = TOKEN_B64.decode(&token).unwrap();
bytes[0] = bytes[0].wrapping_add(1);
assert!(!verify_token(&key, &TOKEN_B64.encode(bytes), now));
}
#[test]
fn expiry_saturates_instead_of_panicking() {
let now = OffsetDateTime::from_unix_timestamp(1_700_000_000).unwrap();
let normal = saturating_expiry(now, std::time::Duration::from_secs(60));
assert_eq!(normal, now + Duration::seconds(60));
let huge = saturating_expiry(now, std::time::Duration::from_secs(u64::MAX));
assert_eq!(huge, PrimitiveDateTime::MAX.assume_utc());
}
#[test]
fn session_cookie_parsing() {
let sk = SessionKey::default();
let expiry = OffsetDateTime::from_unix_timestamp(1_900_000_000).unwrap();
let bare = format!("{}", sk.0.as_hyphenated());
assert_eq!(parse_session_cookie(&bare), Some((sk.clone(), None)));
let with_exp = format!("{}.{}", sk.0.as_hyphenated(), expiry.unix_timestamp());
assert_eq!(parse_session_cookie(&with_exp), Some((sk, Some(expiry))));
assert_eq!(parse_session_cookie("nonsense"), None);
assert_eq!(parse_session_cookie(""), None);
}
#[test]
fn renews_past_halfway_point() {
let mut cfg = get_cfg();
cfg.session_expires = Some(std::time::Duration::from_secs(100));
let access_info = AccessInfo::new(cfg);
let now = OffsetDateTime::now_utc();
assert!(!access_info.should_renew(Some(now + Duration::seconds(60)), now));
assert!(access_info.should_renew(Some(now + Duration::seconds(40)), now));
assert!(access_info.should_renew(None, now));
}
#[tokio::test]
async fn trusted_network_skips_token() -> Result<()> {
use axum::extract::ConnectInfo;
use std::net::SocketAddr;
let mut cfg = get_cfg(); cfg.trusted_networks = vec!["100.64.0.0/10".parse().unwrap()];
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let mut trusted = Request::builder()
.uri("http://example.com/path")
.body(Body::empty())
.unwrap();
trusted.extensions_mut().insert(ConnectInfo(
"100.100.1.2:5555".parse::<SocketAddr>().unwrap(),
));
assert_eq!(
svc.clone().oneshot(trusted).await.unwrap().status(),
StatusCode::OK
);
let mut untrusted = Request::builder()
.uri("http://example.com/path")
.body(Body::empty())
.unwrap();
untrusted.extensions_mut().insert(ConnectInfo(
"192.168.1.2:5555".parse::<SocketAddr>().unwrap(),
));
assert!(svc.oneshot(untrusted).await.is_err());
Ok(())
}
#[tokio::test]
async fn browser_token_auth_redirects_without_token() -> Result<()> {
let cfg = get_cfg();
let uri = format!("{}&keep=1", valid_token_uri(&cfg));
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder()
.uri(uri)
.header(header::ACCEPT, "text/html,application/xhtml+xml")
.body(Body::empty())
.unwrap();
let res = svc.oneshot(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SEE_OTHER);
let location = res.headers().get(header::LOCATION).unwrap().to_str()?;
assert_eq!(location, "/path?keep=1");
assert!(res.headers().contains_key(header::SET_COOKIE));
Ok(())
}
#[tokio::test]
async fn programmatic_token_auth_is_not_redirected() -> Result<()> {
let cfg = get_cfg();
let uri = valid_token_uri(&cfg);
let auth_layer = cfg.into_layer();
let svc = ServiceBuilder::new().layer(auth_layer).service_fn(handler);
let req = Request::builder()
.uri(uri)
.header(header::ACCEPT, "*/*")
.body(Body::empty())
.unwrap();
let res = svc.oneshot(req).await.unwrap();
assert_eq!(res.status(), StatusCode::OK);
Ok(())
}
}