use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ClusterAuthError {
#[error("Unknown node: {0}")]
UnknownNode(String),
#[error("Token expired at unix timestamp {expired_at}")]
TokenExpired {
expired_at: u64,
},
#[error("Invalid token signature")]
InvalidSignature,
#[error("Token has been revoked: {0}")]
RevokedToken(String),
#[error("Token serialization error: {0}")]
SerializationError(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeIdentity {
pub node_id: String,
pub display_name: String,
pub registered_at: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ClusterNodeToken {
pub node_id: String,
pub issued_at: u64,
pub expires_at: u64,
pub jti: String,
pub signature: String,
}
#[derive(Debug, Clone)]
pub struct ClusterAuthConfig {
pub token_ttl_seconds: u64,
pub max_clock_skew_seconds: u64,
}
impl Default for ClusterAuthConfig {
fn default() -> Self {
Self {
token_ttl_seconds: 300,
max_clock_skew_seconds: 30,
}
}
}
pub struct ClusterAuthManager {
config: ClusterAuthConfig,
nodes: HashMap<String, NodeIdentity>,
secret: Vec<u8>,
revoked_jtis: HashSet<String>,
}
impl ClusterAuthManager {
pub fn new(config: ClusterAuthConfig, secret: &[u8]) -> Self {
Self {
config,
nodes: HashMap::new(),
secret: secret.to_vec(),
revoked_jtis: HashSet::new(),
}
}
pub fn register_node(&mut self, node: NodeIdentity) -> Result<(), ClusterAuthError> {
self.nodes.insert(node.node_id.clone(), node);
Ok(())
}
pub fn remove_node(&mut self, node_id: &str) -> Option<NodeIdentity> {
self.nodes.remove(node_id)
}
pub fn known_nodes(&self) -> Vec<&NodeIdentity> {
self.nodes.values().collect()
}
pub fn issue_token(
&self,
node_id: &str,
now: u64,
) -> Result<ClusterNodeToken, ClusterAuthError> {
if !self.nodes.contains_key(node_id) {
return Err(ClusterAuthError::UnknownNode(node_id.to_string()));
}
let issued_at = now;
let expires_at = now + self.config.token_ttl_seconds;
let jti = generate_jti();
let signature = self.compute_signature(node_id, issued_at, expires_at, &jti)?;
Ok(ClusterNodeToken {
node_id: node_id.to_string(),
issued_at,
expires_at,
jti,
signature,
})
}
pub fn verify_token(
&self,
token: &ClusterNodeToken,
now: u64,
) -> Result<&NodeIdentity, ClusterAuthError> {
let expected_sig = self.compute_signature(
&token.node_id,
token.issued_at,
token.expires_at,
&token.jti,
)?;
if !constant_time_eq(&expected_sig, &token.signature) {
return Err(ClusterAuthError::InvalidSignature);
}
if self.revoked_jtis.contains(&token.jti) {
return Err(ClusterAuthError::RevokedToken(token.jti.clone()));
}
let effective_expiry = token
.expires_at
.saturating_add(self.config.max_clock_skew_seconds);
if now > effective_expiry {
return Err(ClusterAuthError::TokenExpired {
expired_at: token.expires_at,
});
}
self.nodes
.get(&token.node_id)
.ok_or_else(|| ClusterAuthError::UnknownNode(token.node_id.clone()))
}
pub fn revoke_token(&mut self, jti: &str) {
self.revoked_jtis.insert(jti.to_string());
}
pub fn rotate_secret(&mut self, new_secret: &[u8]) {
self.secret = new_secret.to_vec();
}
pub fn encode_token(token: &ClusterNodeToken) -> String {
let payload = ClusterNodePayload {
node_id: token.node_id.clone(),
issued_at: token.issued_at,
expires_at: token.expires_at,
jti: token.jti.clone(),
};
let json = serde_json::to_string(&payload).unwrap_or_default();
let encoded_payload = URL_SAFE_NO_PAD.encode(json.as_bytes());
format!("{}.{}", encoded_payload, token.signature)
}
pub fn decode_token(bearer: &str) -> Result<ClusterNodeToken, ClusterAuthError> {
let (encoded_payload, signature) = bearer
.rsplit_once('.')
.ok_or_else(|| ClusterAuthError::SerializationError("Missing '.' separator".into()))?;
let json_bytes = URL_SAFE_NO_PAD
.decode(encoded_payload)
.map_err(|e| ClusterAuthError::SerializationError(format!("Base64 decode: {e}")))?;
let json_str = std::str::from_utf8(&json_bytes)
.map_err(|e| ClusterAuthError::SerializationError(format!("UTF-8 decode: {e}")))?;
let payload: ClusterNodePayload = serde_json::from_str(json_str)
.map_err(|e| ClusterAuthError::SerializationError(format!("JSON parse: {e}")))?;
Ok(ClusterNodeToken {
node_id: payload.node_id,
issued_at: payload.issued_at,
expires_at: payload.expires_at,
jti: payload.jti,
signature: signature.to_string(),
})
}
fn compute_signature(
&self,
node_id: &str,
issued_at: u64,
expires_at: u64,
jti: &str,
) -> Result<String, ClusterAuthError> {
let message = format!("{node_id}|{issued_at}|{expires_at}|{jti}");
let tag = oxicrypto_mac::hmac_sha256_to_vec(&self.secret, message.as_bytes())
.map_err(|e| ClusterAuthError::SerializationError(format!("HMAC error: {e:?}")))?;
Ok(hex::encode(tag))
}
}
#[derive(Debug, Serialize, Deserialize)]
struct ClusterNodePayload {
node_id: String,
issued_at: u64,
expires_at: u64,
jti: String,
}
fn generate_jti() -> String {
use scirs2_core::random::SecureRandom;
let mut secure = SecureRandom::new();
let bytes = secure.random_bytes(16);
hex::encode(&bytes)
}
fn constant_time_eq(a: &str, b: &str) -> bool {
let a = a.as_bytes();
let b = b.as_bytes();
if a.len() != b.len() {
return false;
}
let diff = a
.iter()
.zip(b.iter())
.fold(0u8, |acc, (x, y)| acc | (x ^ y));
diff == 0
}
#[cfg(test)]
mod tests {
use super::*;
fn test_secret() -> Vec<u8> {
b"test-cluster-secret-32-bytes-long!!".to_vec()
}
fn test_config() -> ClusterAuthConfig {
ClusterAuthConfig {
token_ttl_seconds: 300,
max_clock_skew_seconds: 30,
}
}
fn test_node(suffix: &str) -> NodeIdentity {
NodeIdentity {
node_id: format!("node-{suffix}"),
display_name: format!("Test Node {suffix}"),
registered_at: 1_700_000_000,
}
}
fn manager_with_node(suffix: &str) -> (ClusterAuthManager, String) {
let mut mgr = ClusterAuthManager::new(test_config(), &test_secret());
let node = test_node(suffix);
let node_id = node.node_id.clone();
mgr.register_node(node).unwrap();
(mgr, node_id)
}
#[test]
fn test_register_and_issue_token() {
let (mgr, node_id) = manager_with_node("a1b2c3");
let now: u64 = 1_700_100_000;
let token = mgr.issue_token(&node_id, now).expect("issue_token");
assert_eq!(token.node_id, node_id);
assert_eq!(token.issued_at, now);
assert_eq!(token.expires_at, now + 300);
assert!(!token.jti.is_empty());
assert!(!token.signature.is_empty());
let identity = mgr.verify_token(&token, now).expect("verify_token");
assert_eq!(identity.node_id, node_id);
}
#[test]
fn test_token_expiry() {
let (mgr, node_id) = manager_with_node("exp");
let now: u64 = 1_700_200_000;
let token = mgr.issue_token(&node_id, now).expect("issue_token");
let future = now + 300 + 30 + 1; let err = mgr
.verify_token(&token, future)
.expect_err("should be expired");
assert!(
matches!(err, ClusterAuthError::TokenExpired { expired_at } if expired_at == now + 300),
"expected TokenExpired, got {err:?}"
);
}
#[test]
fn test_invalid_signature() {
let (mgr, node_id) = manager_with_node("sig");
let now: u64 = 1_700_300_000;
let mut token = mgr.issue_token(&node_id, now).expect("issue_token");
let last = token.signature.pop().unwrap();
let flipped = if last == '0' { '1' } else { '0' };
token.signature.push(flipped);
let err = mgr
.verify_token(&token, now)
.expect_err("tampered token should fail");
assert_eq!(err, ClusterAuthError::InvalidSignature);
}
#[test]
fn test_unknown_node_issue() {
let mgr = ClusterAuthManager::new(test_config(), &test_secret());
let err = mgr
.issue_token("node-ghost", 1_700_000_000)
.expect_err("unknown node should fail");
assert!(
matches!(err, ClusterAuthError::UnknownNode(ref id) if id == "node-ghost"),
"unexpected error: {err:?}"
);
}
#[test]
fn test_unknown_node_verify() {
let (issuing_mgr, node_id) = manager_with_node("ghost2");
let now: u64 = 1_700_000_000;
let token = issuing_mgr.issue_token(&node_id, now).unwrap();
let verifying_mgr = ClusterAuthManager::new(test_config(), &test_secret());
let err = verifying_mgr
.verify_token(&token, now)
.expect_err("unregistered node should fail");
assert!(
matches!(err, ClusterAuthError::UnknownNode(_)),
"unexpected error: {err:?}"
);
}
#[test]
fn test_revoke_token() {
let (mut mgr, node_id) = manager_with_node("rev");
let now: u64 = 1_700_400_000;
let token = mgr.issue_token(&node_id, now).unwrap();
let jti = token.jti.clone();
mgr.revoke_token(&jti);
let err = mgr
.verify_token(&token, now)
.expect_err("revoked token should fail");
assert!(
matches!(err, ClusterAuthError::RevokedToken(ref j) if j == &jti),
"unexpected error: {err:?}"
);
}
#[test]
fn test_rotate_secret() {
let (mut mgr, node_id) = manager_with_node("rot");
let now: u64 = 1_700_500_000;
let token = mgr.issue_token(&node_id, now).unwrap();
mgr.verify_token(&token, now)
.expect("should verify before rotation");
let new_secret = b"new-completely-different-secret!!";
mgr.rotate_secret(new_secret);
let err = mgr
.verify_token(&token, now)
.expect_err("old token should be invalid after rotation");
assert_eq!(err, ClusterAuthError::InvalidSignature);
}
#[test]
fn test_encode_decode_roundtrip() {
let (mgr, node_id) = manager_with_node("enc");
let now: u64 = 1_700_600_000;
let original = mgr.issue_token(&node_id, now).unwrap();
let bearer = ClusterAuthManager::encode_token(&original);
let decoded = ClusterAuthManager::decode_token(&bearer).expect("decode_token");
assert_eq!(decoded, original);
}
#[test]
fn test_clock_skew_allowed() {
let (mgr, node_id) = manager_with_node("skew");
let now: u64 = 1_700_700_000;
let token = mgr.issue_token(&node_id, now).unwrap();
let skewed_now = now + 300 + 30 - 1; mgr.verify_token(&token, skewed_now)
.expect("should be accepted within skew window");
let over_skew = now + 300 + 30 + 1;
let err = mgr
.verify_token(&token, over_skew)
.expect_err("should be rejected beyond skew window");
assert!(matches!(err, ClusterAuthError::TokenExpired { .. }));
}
#[test]
fn test_known_nodes_list() {
let mut mgr = ClusterAuthManager::new(test_config(), &test_secret());
mgr.register_node(test_node("n1")).unwrap();
mgr.register_node(test_node("n2")).unwrap();
mgr.register_node(test_node("n3")).unwrap();
let nodes = mgr.known_nodes();
assert_eq!(nodes.len(), 3);
}
#[test]
fn test_remove_node() {
let (mut mgr, node_id) = manager_with_node("rm");
let now: u64 = 1_700_800_000;
let token = mgr.issue_token(&node_id, now).unwrap();
let removed = mgr.remove_node(&node_id);
assert!(removed.is_some());
assert!(!mgr.known_nodes().iter().any(|n| n.node_id == node_id));
let err = mgr
.verify_token(&token, now)
.expect_err("removed node token should fail");
assert!(
matches!(err, ClusterAuthError::UnknownNode(_)),
"unexpected error: {err:?}"
);
}
#[test]
fn test_config_defaults() {
let cfg = ClusterAuthConfig::default();
assert_eq!(cfg.token_ttl_seconds, 300);
assert_eq!(cfg.max_clock_skew_seconds, 30);
}
}