use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use chrono::{DateTime, Utc};
use hmac::{Hmac, Mac};
use serde::{Deserialize, Serialize};
use sha2::Sha256;
use crate::api::auth::constant_time_eq;
pub const STOP_ACTION: &str = "stop";
pub const DEFAULT_TTL_SECS: i64 = 30 * 60;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
struct ActionClaims {
session_id: String,
action: String,
exp: i64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ActionTokenError {
Malformed,
BadSignature,
Expired,
WrongAction,
}
fn compute_hmac(secret: &str, data: &[u8]) -> String {
let mut mac =
Hmac::<Sha256>::new_from_slice(secret.as_bytes()).expect("HMAC accepts any key length");
mac.update(data);
hex::encode(mac.finalize().into_bytes())
}
pub fn sign_action_token(
secret: &str,
session_id: &str,
action: &str,
now: DateTime<Utc>,
ttl_secs: i64,
) -> String {
let claims = ActionClaims {
session_id: session_id.to_owned(),
action: action.to_owned(),
exp: now.timestamp() + ttl_secs,
};
let payload = serde_json::to_vec(&claims).unwrap_or_default();
let payload_b64 = URL_SAFE_NO_PAD.encode(payload);
let signature = compute_hmac(secret, payload_b64.as_bytes());
format!("{payload_b64}.{signature}")
}
pub fn verify_action_token(
secret: &str,
token: &str,
expected_action: &str,
now: DateTime<Utc>,
) -> Result<String, ActionTokenError> {
let (payload_b64, signature) = token.split_once('.').ok_or(ActionTokenError::Malformed)?;
let expected_signature = compute_hmac(secret, payload_b64.as_bytes());
if !constant_time_eq(&expected_signature, signature) {
return Err(ActionTokenError::BadSignature);
}
let payload_bytes = URL_SAFE_NO_PAD
.decode(payload_b64)
.map_err(|_| ActionTokenError::Malformed)?;
let claims: ActionClaims =
serde_json::from_slice(&payload_bytes).map_err(|_| ActionTokenError::Malformed)?;
if claims.action != expected_action {
return Err(ActionTokenError::WrongAction);
}
if claims.exp < now.timestamp() {
return Err(ActionTokenError::Expired);
}
Ok(claims.session_id)
}
#[cfg(test)]
mod tests {
use super::*;
fn now() -> DateTime<Utc> {
DateTime::parse_from_rfc3339("2026-07-08T12:00:00Z")
.unwrap()
.with_timezone(&Utc)
}
#[test]
fn test_sign_and_verify_roundtrip() {
let token = sign_action_token("secret", "sess-1", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let session_id = verify_action_token("secret", &token, STOP_ACTION, now()).unwrap();
assert_eq!(session_id, "sess-1");
}
#[test]
fn test_verify_rejects_wrong_secret() {
let token = sign_action_token("secret-a", "sess-1", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let err = verify_action_token("secret-b", &token, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::BadSignature);
}
#[test]
fn test_verify_rejects_expired_token() {
let token = sign_action_token("secret", "sess-1", STOP_ACTION, now(), -1);
let err = verify_action_token("secret", &token, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::Expired);
}
#[test]
fn test_verify_accepts_token_at_exact_expiry_boundary() {
let token = sign_action_token("secret", "sess-1", STOP_ACTION, now(), 0);
assert!(verify_action_token("secret", &token, STOP_ACTION, now()).is_ok());
}
#[test]
fn test_verify_rejects_wrong_action() {
let token = sign_action_token("secret", "sess-1", "purge", now(), DEFAULT_TTL_SECS);
let err = verify_action_token("secret", &token, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::WrongAction);
}
#[test]
fn test_verify_rejects_tampered_signature() {
let token = sign_action_token("secret", "sess-1", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let (payload, _sig) = token.split_once('.').unwrap();
let tampered =
format!("{payload}.0000000000000000000000000000000000000000000000000000000000000000");
let err = verify_action_token("secret", &tampered, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::BadSignature);
}
#[test]
fn test_verify_rejects_tampered_session_id() {
let token_a = sign_action_token("secret", "sess-a", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let token_b = sign_action_token("secret", "sess-b", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let (payload_b, _) = token_b.split_once('.').unwrap();
let (_, signature_a) = token_a.split_once('.').unwrap();
let frankenstein = format!("{payload_b}.{signature_a}");
let err = verify_action_token("secret", &frankenstein, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::BadSignature);
}
#[test]
fn test_verify_rejects_malformed_no_separator() {
let err = verify_action_token("secret", "not-a-token", STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::Malformed);
}
#[test]
fn test_verify_rejects_malformed_payload_not_base64() {
let sig = compute_hmac("secret", b"not-base64!!!");
let token = format!("not-base64!!!.{sig}");
let err = verify_action_token("secret", &token, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::Malformed);
}
#[test]
fn test_verify_rejects_malformed_payload_not_json() {
let payload_b64 = URL_SAFE_NO_PAD.encode(b"not json");
let sig = compute_hmac("secret", payload_b64.as_bytes());
let token = format!("{payload_b64}.{sig}");
let err = verify_action_token("secret", &token, STOP_ACTION, now()).unwrap_err();
assert_eq!(err, ActionTokenError::Malformed);
}
#[test]
fn test_different_sessions_yield_different_tokens() {
let a = sign_action_token("secret", "sess-a", STOP_ACTION, now(), DEFAULT_TTL_SECS);
let b = sign_action_token("secret", "sess-b", STOP_ACTION, now(), DEFAULT_TTL_SECS);
assert_ne!(a, b);
}
#[test]
fn test_action_token_error_debug_clone_copy_eq() {
let e = ActionTokenError::Expired;
let cloned = e;
assert_eq!(e, cloned);
assert!(format!("{e:?}").contains("Expired"));
}
}