223 lines
6.8 KiB
Rust
223 lines
6.8 KiB
Rust
// Authentik Service Account Token Provider for Phase 6.6
|
|||
|
|
// Handles OAuth2 client_credentials flow for Memory service
|
||
|
|
// Used by webhook execution + metrics persistence
|
||
|
|
|
||
|
|
use std::sync::{Arc, RwLock};
|
||
|
|
use std::time::{Duration, Instant};
|
||
|
|
use serde::{Deserialize, Serialize};
|
||
|
|
use reqwest::Client;
|
||
|
|
use log::{debug, warn, error};
|
||
|
|
|
||
|
|
#[derive(Clone, Debug)]
|
||
|
|
pub struct AuthentikServiceAccountConfig {
|
||
|
|
pub client_id: String,
|
||
|
|
pub client_secret: String,
|
||
|
|
pub token_endpoint: String,
|
||
|
|
pub cache_ttl_secs: u64,
|
||
|
|
}
|
||
|
|
|
||
|
|
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||
|
|
pub struct TokenResponse {
|
||
|
|
pub access_token: String,
|
||
|
|
pub token_type: String,
|
||
|
|
pub expires_in: u64,
|
||
|
|
}
|
||
|
|
|
||
|
|
#[derive(Clone, Debug)]
|
||
|
|
struct CachedToken {
|
||
|
|
token: String,
|
||
|
|
expires_at: Instant,
|
||
|
|
}
|
||
|
|
|
||
|
|
/// AuthentikServiceAccount provides OAuth2 service account tokens
|
||
|
|
/// Caches tokens with TTL to reduce token endpoint calls
|
||
|
|
pub struct AuthentikServiceAccount {
|
||
|
|
config: AuthentikServiceAccountConfig,
|
||
|
|
client: Client,
|
||
|
|
cached_token: Arc<RwLock<Option<CachedToken>>>,
|
||
|
|
}
|
||
|
|
|
||
|
|
impl AuthentikServiceAccount {
|
||
|
|
pub fn new(config: AuthentikServiceAccountConfig) -> Self {
|
||
|
|
Self {
|
||
|
|
config,
|
||
|
|
client: Client::new(),
|
||
|
|
cached_token: Arc::new(RwLock::new(None)),
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Get valid service account token
|
||
|
|
/// Returns cached token if valid, otherwise fetches new token
|
||
|
|
pub async fn get_token(&self) -> Result<String, String> {
|
||
|
|
// Check cache first
|
||
|
|
{
|
||
|
|
let cache = self.cached_token.read()
|
||
|
|
.map_err(|e| format!("Cache lock failed: {}", e))?;
|
||
|
|
|
||
|
|
if let Some(cached) = cache.as_ref() {
|
||
|
|
if cached.expires_at > Instant::now() {
|
||
|
|
debug!("Using cached Authentik service account token");
|
||
|
|
return Ok(cached.token.clone());
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// Cache miss or expired: fetch new token
|
||
|
|
debug!("Fetching new Authentik service account token");
|
||
|
|
let response = self.fetch_token().await?;
|
||
|
|
|
||
|
|
// Cache the new token
|
||
|
|
let token = response.access_token.clone();
|
||
|
|
let expires_in = response.expires_in.saturating_sub(60); // Refresh 60s before expiry
|
||
|
|
let expires_at = Instant::now() + Duration::from_secs(expires_in);
|
||
|
|
|
||
|
|
{
|
||
|
|
let mut cache = self.cached_token.write()
|
||
|
|
.map_err(|e| format!("Cache lock failed: {}", e))?;
|
||
|
|
*cache = Some(CachedToken { token: token.clone(), expires_at });
|
||
|
|
}
|
||
|
|
|
||
|
|
Ok(token)
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Fetch token from Authentik token endpoint
|
||
|
|
async fn fetch_token(&self) -> Result<TokenResponse, String> {
|
||
|
|
let params = [
|
||
|
|
("grant_type", "client_credentials"),
|
||
|
|
("client_id", &self.config.client_id),
|
||
|
|
("client_secret", &self.config.client_secret),
|
||
|
|
];
|
||
|
|
|
||
|
|
let response = self.client
|
||
|
|
.post(&self.config.token_endpoint)
|
||
|
|
.form(¶ms)
|
||
|
|
.send()
|
||
|
|
.await
|
||
|
|
.map_err(|e| format!("Token request failed: {}", e))?;
|
||
|
|
|
||
|
|
if !response.status().is_success() {
|
||
|
|
let status = response.status();
|
||
|
|
let body = response.text().await.unwrap_or_default();
|
||
|
|
error!("Authentik token endpoint error: {} {}", status, body);
|
||
|
|
return Err(format!("Token endpoint error: {}", status));
|
||
|
|
}
|
||
|
|
|
||
|
|
let token_response: TokenResponse = response
|
||
|
|
.json()
|
||
|
|
.await
|
||
|
|
.map_err(|e| format!("Failed to parse token response: {}", e))?;
|
||
|
|
|
||
|
|
debug!("Successfully fetched token, expires_in: {}s", token_response.expires_in);
|
||
|
|
Ok(token_response)
|
||
|
|
}
|
||
|
|
|
||
|
|
/// Invalidate cached token (force refresh on next request)
|
||
|
|
pub fn invalidate_cache(&self) {
|
||
|
|
if let Ok(mut cache) = self.cached_token.write() {
|
||
|
|
*cache = None;
|
||
|
|
debug!("Invalidated cached Authentik service account token");
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[cfg(test)]
|
||
|
|
mod tests {
|
||
|
|
use super::*;
|
||
|
|
|
||
|
|
fn test_config() -> AuthentikServiceAccountConfig {
|
||
|
|
AuthentikServiceAccountConfig {
|
||
|
|
client_id: "test-client".to_string(),
|
||
|
|
client_secret: "test-secret".to_string(),
|
||
|
|
token_endpoint: "http://localhost:8080/token".to_string(),
|
||
|
|
cache_ttl_secs: 3600,
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_authentik_service_account_new() {
|
||
|
|
let config = test_config();
|
||
|
|
let sa = AuthentikServiceAccount::new(config.clone());
|
||
|
|
assert_eq!(sa.config.client_id, "test-client");
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_cached_token_expiry() {
|
||
|
|
let expired = CachedToken {
|
||
|
|
token: "expired".to_string(),
|
||
|
|
expires_at: Instant::now() - Duration::from_secs(60),
|
||
|
|
};
|
||
|
|
assert!(expired.expires_at < Instant::now());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_cached_token_valid() {
|
||
|
|
let valid = CachedToken {
|
||
|
|
token: "valid".to_string(),
|
||
|
|
expires_at: Instant::now() + Duration::from_secs(3600),
|
||
|
|
};
|
||
|
|
assert!(valid.expires_at > Instant::now());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[tokio::test]
|
||
|
|
async fn test_invalidate_cache() {
|
||
|
|
let config = test_config();
|
||
|
|
let sa = AuthentikServiceAccount::new(config);
|
||
|
|
|
||
|
|
// Pre-populate cache
|
||
|
|
{
|
||
|
|
let mut cache = sa.cached_token.write().unwrap();
|
||
|
|
*cache = Some(CachedToken {
|
||
|
|
token: "test".to_string(),
|
||
|
|
expires_at: Instant::now() + Duration::from_secs(3600),
|
||
|
|
});
|
||
|
|
}
|
||
|
|
|
||
|
|
// Verify cached
|
||
|
|
{
|
||
|
|
let cache = sa.cached_token.read().unwrap();
|
||
|
|
assert!(cache.is_some());
|
||
|
|
}
|
||
|
|
|
||
|
|
// Invalidate
|
||
|
|
sa.invalidate_cache();
|
||
|
|
|
||
|
|
// Verify empty
|
||
|
|
{
|
||
|
|
let cache = sa.cached_token.read().unwrap();
|
||
|
|
assert!(cache.is_none());
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_token_response_parse() {
|
||
|
|
let json = r#"{"access_token": "abc123", "token_type": "Bearer", "expires_in": 3600}"#;
|
||
|
|
let token: TokenResponse = serde_json::from_str(json).unwrap();
|
||
|
|
assert_eq!(token.access_token, "abc123");
|
||
|
|
assert_eq!(token.expires_in, 3600);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_service_account_config() {
|
||
|
|
let config = test_config();
|
||
|
|
assert_eq!(config.client_id, "test-client");
|
||
|
|
assert_eq!(config.client_secret, "test-secret");
|
||
|
|
assert_eq!(config.cache_ttl_secs, 3600);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_cache_ttl_expiry_calculation() {
|
||
|
|
let expires_in = 3600u64;
|
||
|
|
let buffer_secs = 60u64;
|
||
|
|
let final_ttl = expires_in.saturating_sub(buffer_secs);
|
||
|
|
assert_eq!(final_ttl, 3540);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn test_cache_ttl_edge_case_small_expiry() {
|
||
|
|
let expires_in = 30u64;
|
||
|
|
let buffer_secs = 60u64;
|
||
|
|
let final_ttl = expires_in.saturating_sub(buffer_secs);
|
||
|
|
assert_eq!(final_ttl, 0); // saturating_sub prevents underflow
|
||
|
|
}
|
||
|
|
}
|