use base64::Engine;
use rand::Rng;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct AuthConfig {
pub authentication_hook_server: Option<String>,
pub auth_token_idle_timeout: std::time::Duration,
}
impl Default for AuthConfig {
fn default() -> Self {
Self {
authentication_hook_server: None,
auth_token_idle_timeout: std::time::Duration::from_secs(300), }
}
}
#[derive(Clone, Default)]
pub struct AuthTokenTracker {
tokens: Arc<Mutex<HashMap<Arc<str>, Instant>>>,
}
impl AuthTokenTracker {
pub fn register_token(&self, token: Arc<str>) {
let mut tokens = self.tokens.lock().unwrap();
tokens.insert(token, Instant::now());
}
pub fn is_valid(
&self,
token: &Option<Arc<str>>,
config: &AuthConfig,
) -> bool {
let Some(token) = token else {
return config.authentication_hook_server.is_none();
};
let mut tokens = self.tokens.lock().unwrap();
if let Some(last_used) = tokens.get_mut(token.as_ref()) {
let now = Instant::now();
if now.duration_since(*last_used) < config.auth_token_idle_timeout {
*last_used = now; return true;
}
tokens.remove(token.as_ref());
}
false
}
pub fn prune_expired(&self, config: &AuthConfig) {
let mut tokens = self.tokens.lock().unwrap();
let now = Instant::now();
tokens.retain(|_, last_used| {
now.duration_since(*last_used) < config.auth_token_idle_timeout
});
}
}
#[derive(Debug)]
pub enum AuthenticateError {
Unauthorized,
HookServerError(Box<dyn std::error::Error + Send + Sync>),
OtherError(Box<dyn std::error::Error + Send + Sync>),
}
pub async fn process_authenticate(
config: &AuthConfig,
token_tracker: &AuthTokenTracker,
auth_failures: opentelemetry::metrics::Counter<u64>,
auth_bytes: bytes::Bytes,
) -> Result<Arc<str>, AuthenticateError> {
if let Some(hook_server_url) = &config.authentication_hook_server {
match call_hook_server(hook_server_url, auth_bytes).await {
Ok(token) => {
token_tracker.register_token(token.clone());
Ok(token)
}
Err(err) => {
auth_failures.add(1, &[]);
Err(err)
}
}
} else {
let mut token_bytes = [0u8; 32];
rand::thread_rng().fill(&mut token_bytes);
let token = Arc::<str>::from(
base64::prelude::BASE64_URL_SAFE_NO_PAD.encode(token_bytes),
);
token_tracker.register_token(token.clone());
Ok(token)
}
}
async fn call_hook_server(
base_url: &str,
auth_bytes: bytes::Bytes,
) -> Result<Arc<str>, AuthenticateError> {
let url = format!("{}/authenticate", base_url.trim_end_matches('/'));
let response = tokio::task::spawn_blocking(move || {
let agent: ureq::Agent = ureq::Agent::config_builder()
.timeout_global(Some(Duration::from_secs(10)))
.build()
.into();
agent
.put(&url)
.header("Content-Type", "application/octet-stream")
.send(&auth_bytes[..])
})
.await
.map_err(|e| AuthenticateError::OtherError(Box::new(e)))?;
match response {
Ok(response) => {
let body = response
.into_body()
.read_to_string()
.map_err(|e| AuthenticateError::OtherError(Box::new(e)))?;
let json: serde_json::Value = serde_json::from_str(&body)
.map_err(|e| AuthenticateError::OtherError(Box::new(e)))?;
let token = json["authToken"].as_str().ok_or_else(|| {
AuthenticateError::OtherError(
"Missing authToken in response".into(),
)
})?;
Ok(Arc::<str>::from(token))
}
Err(ureq::Error::StatusCode(401)) => {
Err(AuthenticateError::Unauthorized)
}
Err(e) => Err(AuthenticateError::HookServerError(Box::new(e))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_expiration() {
let config = AuthConfig {
authentication_hook_server: Some(
"http://example.com/auth".to_string(),
),
auth_token_idle_timeout: std::time::Duration::from_millis(200),
};
let tracker = AuthTokenTracker::default();
let token: Arc<str> = Arc::from("test-token");
tracker.register_token(token.clone());
assert!(tracker.is_valid(&Some(token.clone()), &config));
std::thread::sleep(std::time::Duration::from_millis(300));
assert!(!tracker.is_valid(&Some(token.clone()), &config));
}
#[test]
fn test_token_lifetime_extension() {
let config = AuthConfig {
authentication_hook_server: Some(
"http://example.com/auth".to_string(),
),
auth_token_idle_timeout: std::time::Duration::from_millis(500),
};
let tracker = AuthTokenTracker::default();
let token: Arc<str> = Arc::from("test-token");
tracker.register_token(token.clone());
std::thread::sleep(std::time::Duration::from_millis(200));
assert!(tracker.is_valid(&Some(token.clone()), &config));
std::thread::sleep(std::time::Duration::from_millis(200));
assert!(tracker.is_valid(&Some(token.clone()), &config));
std::thread::sleep(std::time::Duration::from_millis(600));
assert!(!tracker.is_valid(&Some(token.clone()), &config));
}
#[test]
fn test_no_auth_configured() {
let config = AuthConfig {
authentication_hook_server: None,
auth_token_idle_timeout: std::time::Duration::from_secs(300),
};
let tracker = AuthTokenTracker::default();
assert!(tracker.is_valid(&None, &config));
assert!(!tracker.is_valid(&Some(Arc::from("any-token")), &config));
let token: Arc<str> = Arc::from("test-token");
tracker.register_token(token.clone());
assert!(tracker.is_valid(&Some(token), &config));
}
#[test]
fn test_prune_expired() {
let config = AuthConfig {
authentication_hook_server: Some(
"http://example.com/auth".to_string(),
),
auth_token_idle_timeout: std::time::Duration::from_millis(100),
};
let tracker = AuthTokenTracker::default();
tracker.register_token(Arc::from("token1"));
tracker.register_token(Arc::from("token2"));
tracker.register_token(Arc::from("token3"));
std::thread::sleep(std::time::Duration::from_millis(150));
tracker.prune_expired(&config);
assert!(!tracker.is_valid(&Some(Arc::from("token1")), &config));
assert!(!tracker.is_valid(&Some(Arc::from("token2")), &config));
assert!(!tracker.is_valid(&Some(Arc::from("token3")), &config));
}
}