use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
use tracing;
use crate::error::{PepError, Result};
use crate::oidc_client::{OidcClient, TokenResponse};
use crate::token_store::{TokenStore, StoredToken};
#[allow(async_fn_in_trait)]
pub trait TokenProvider: Send + Sync {
async fn get_token(&self) -> Result<String>;
}
#[derive(Clone, Debug)]
struct CachedToken {
token: String,
expires_at: Instant,
}
impl CachedToken {
fn from_response(token_response: &TokenResponse) -> Self {
let expires_in = token_response.expires_in.unwrap_or(900);
let buffer_secs = 30;
let effective_secs = expires_in.saturating_sub(buffer_secs);
Self {
token: token_response.access_token.clone(),
expires_at: Instant::now() + Duration::from_secs(effective_secs),
}
}
fn is_valid(&self) -> bool {
Instant::now() < self.expires_at
}
}
#[derive(Clone)]
pub struct StaticTokenProvider {
token: String,
}
impl std::fmt::Debug for StaticTokenProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StaticTokenProvider")
.field("token", &format!("{}...", &self.token[..self.token.len().min(8)]))
.finish()
}
}
impl StaticTokenProvider {
pub fn new(token: String) -> Self {
Self { token }
}
}
impl TokenProvider for StaticTokenProvider {
async fn get_token(&self) -> Result<String> {
Ok(self.token.clone())
}
}
#[derive(Debug, Clone)]
pub struct ServiceAccountConfig {
pub service_token: String,
pub issuer_url: String,
pub client_id: String,
pub client_secret: Option<String>,
pub audience: String,
pub scope: Option<String>,
}
#[derive(Clone)]
pub struct ServiceAccountTokenProvider {
oidc_client: OidcClient,
config: ServiceAccountConfig,
cache: Arc<RwLock<Option<CachedToken>>>,
}
impl std::fmt::Debug for ServiceAccountTokenProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ServiceAccountTokenProvider")
.field("audience", &self.config.audience)
.finish()
}
}
impl ServiceAccountTokenProvider {
pub fn new(config: ServiceAccountConfig) -> Self {
Self {
oidc_client: OidcClient::new(),
config,
cache: Arc::new(RwLock::new(None)),
}
}
pub fn with_client(oidc_client: OidcClient, config: ServiceAccountConfig) -> Self {
Self {
oidc_client,
config,
cache: Arc::new(RwLock::new(None)),
}
}
async fn exchange(&self) -> Result<CachedToken> {
tracing::debug!(
"Exchanging service account token for audience '{}'",
self.config.audience
);
let response = self
.oidc_client
.exchange_token(
&self.config.issuer_url,
&self.config.client_id,
self.config.client_secret.as_deref(),
&self.config.service_token,
&self.config.audience,
self.config.scope.as_deref(),
)
.await?;
tracing::info!(
"Token exchange successful, expires_in={:?}s",
response.expires_in
);
Ok(CachedToken::from_response(&response))
}
}
impl TokenProvider for ServiceAccountTokenProvider {
async fn get_token(&self) -> Result<String> {
{
let cache = self.cache.read().await;
if let Some(cached) = cache.as_ref() {
if cached.is_valid() {
return Ok(cached.token.clone());
}
}
}
let cached = self.exchange().await?;
let token = cached.token.clone();
{
let mut cache = self.cache.write().await;
*cache = Some(cached);
}
Ok(token)
}
}
#[derive(Debug, Clone)]
pub struct InteractiveConfig {
pub issuer_url: String,
pub client_id: String,
pub client_secret: Option<String>,
pub redirect_uri: String,
pub scope: String,
pub credential_name: String,
}
#[derive(Clone)]
pub struct InteractiveTokenProvider {
oidc_client: OidcClient,
config: InteractiveConfig,
token_store: Arc<dyn TokenStore>,
cache: Arc<RwLock<Option<CachedToken>>>,
}
impl std::fmt::Debug for InteractiveTokenProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("InteractiveTokenProvider")
.field("issuer_url", &self.config.issuer_url)
.field("credential_name", &self.config.credential_name)
.finish()
}
}
impl InteractiveTokenProvider {
pub fn with_store(config: InteractiveConfig, token_store: Arc<dyn TokenStore>) -> Self {
Self {
oidc_client: OidcClient::new(),
config,
token_store,
cache: Arc::new(RwLock::new(None)),
}
}
}
impl TokenProvider for InteractiveTokenProvider {
async fn get_token(&self) -> Result<String> {
{
let cache = self.cache.read().await;
if let Some(cached) = cache.as_ref() {
if cached.is_valid() {
return Ok(cached.token.clone());
}
}
}
let stored = self.token_store.load(&self.config.credential_name)?;
let stored = match stored {
Some(s) => s,
None => {
return Err(PepError::BadRequest(format!(
"Not authenticated. Run: trustee mcp auth {}",
self.config.credential_name
)));
}
};
if !stored.is_expired() {
let cached = CachedToken {
token: stored.access_token.clone(),
expires_at: Instant::now()
+ Duration::from_secs(
seconds_until_expiry(&stored.expires_at).max(1),
),
};
let token = cached.token.clone();
{
let mut cache = self.cache.write().await;
*cache = Some(cached);
}
return Ok(token);
}
let refresh_token = match &stored.refresh_token {
Some(rt) => rt.clone(),
None => {
return Err(PepError::BadRequest(format!(
"Session expired. Run: trustee mcp auth {}",
self.config.credential_name
)));
}
};
tracing::debug!(
"Refreshing expired token for credential '{}'",
self.config.credential_name
);
let response = self
.oidc_client
.refresh_access_token(
&self.config.issuer_url,
&self.config.client_id,
self.config.client_secret.as_deref(),
&refresh_token,
Some(&self.config.scope),
)
.await;
match response {
Ok(token_response) => {
let new_expires_at = compute_expires_at(token_response.expires_in);
let updated = StoredToken::new(
&token_response.access_token,
token_response.refresh_token.clone().or(Some(refresh_token)),
&token_response.token_type,
&new_expires_at,
token_response.scope.clone().or(stored.scope.clone()),
);
if let Err(e) = self.token_store.save(&self.config.credential_name, &updated) {
tracing::warn!(
"Failed to persist refreshed token: {}. Using in-memory only.",
e
);
}
let cached = CachedToken::from_response(&token_response);
let token = cached.token.clone();
{
let mut cache = self.cache.write().await;
*cache = Some(cached);
}
Ok(token)
}
Err(PepError::TokenRefreshFailed { .. }) => {
Err(PepError::BadRequest(format!(
"Session expired. Run: trustee mcp auth {}",
self.config.credential_name
)))
}
Err(e) => Err(e),
}
}
}
pub(crate) fn seconds_until_expiry(expires_at: &str) -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let expires_epoch = crate::token_store::parse_rfc3339_to_epoch_public(expires_at);
expires_epoch.saturating_sub(now)
}
pub(crate) fn compute_expires_at(expires_in: Option<u64>) -> String {
use std::time::{SystemTime, UNIX_EPOCH};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let expires_epoch = now + expires_in.unwrap_or(900);
epoch_to_rfc3339(expires_epoch)
}
pub(crate) fn compute_expires_at_from_jwt(
access_token: &str,
expires_in_fallback: Option<u64>,
) -> String {
let parts: Vec<&str> = access_token.split('.').collect();
if parts.len() < 2 {
return compute_expires_at(expires_in_fallback);
}
use base64::Engine;
let payload = match base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(parts[1]) {
Ok(bytes) => bytes,
Err(_) => {
match base64::engine::general_purpose::STANDARD_NO_PAD.decode(parts[1]) {
Ok(bytes) => bytes,
Err(_) => return compute_expires_at(expires_in_fallback),
}
}
};
let payload_json: serde_json::Value = match serde_json::from_slice(&payload) {
Ok(v) => v,
Err(_) => return compute_expires_at(expires_in_fallback),
};
if let Some(exp) = payload_json.get("exp").and_then(|v| v.as_u64()) {
tracing::debug!(
jwt_exp = exp,
"Extracted real exp from JWT payload for expiry tracking"
);
epoch_to_rfc3339(exp)
} else {
tracing::debug!("No exp claim in JWT, falling back to expires_in estimate");
compute_expires_at(expires_in_fallback)
}
}
pub(crate) fn epoch_to_rfc3339(epoch: u64) -> String {
let days = epoch / 86400;
let remainder = epoch % 86400;
let hour = remainder / 3600;
let min = (remainder % 3600) / 60;
let sec = remainder % 60;
let (year, month, day) = epoch_to_civil(days);
format!(
"{:04}-{:02}-{:02}T{:02}:{:02}:{:02}Z",
year, month, day, hour, min, sec
)
}
fn epoch_to_civil(days: u64) -> (u32, u32, u32) {
let z = days as i64 + 719468;
let era = if z >= 0 { z } else { z - 146096 } / 146097;
let doe = (z - era * 146097) as u64; let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365; let y = yoe as i64 + era * 400;
let doy = doe - (365 * yoe + yoe / 4 - yoe / 100); let mp = (5 * doy + 2) / 153; let d = doy - (153 * mp + 2) / 5 + 1; let m = if mp < 10 { mp + 3 } else { mp - 9 }; let y = if m <= 2 { y + 1 } else { y };
(y as u32, m as u32, d as u32)
}
#[derive(Clone, Debug)]
pub enum TokenProviderEnum {
Static(StaticTokenProvider),
ServiceAccount(ServiceAccountTokenProvider),
Interactive(InteractiveTokenProvider),
}
impl TokenProvider for TokenProviderEnum {
async fn get_token(&self) -> Result<String> {
match self {
Self::Static(p) => p.get_token().await,
Self::ServiceAccount(p) => p.get_token().await,
Self::Interactive(p) => p.get_token().await,
}
}
}
impl From<String> for TokenProviderEnum {
fn from(token: String) -> Self {
Self::Static(StaticTokenProvider::new(token))
}
}
impl From<Option<String>> for TokenProviderEnum {
fn from(token: Option<String>) -> Self {
Self::Static(StaticTokenProvider::new(token.unwrap_or_default()))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_static_token_provider() {
let rt = tokio::runtime::Runtime::new().unwrap();
let provider = StaticTokenProvider::new("test-token".to_string());
let token = rt.block_on(provider.get_token()).unwrap();
assert_eq!(token, "test-token");
}
#[test]
fn test_cached_token_from_response() {
let response = TokenResponse {
access_token: "abc123".to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(900),
refresh_token: None,
id_token: None,
scope: None,
};
let cached = CachedToken::from_response(&response);
assert_eq!(cached.token, "abc123");
assert!(cached.is_valid());
}
#[test]
fn test_cached_token_expiry() {
let response = TokenResponse {
access_token: "abc123".to_string(),
token_type: "Bearer".to_string(),
expires_in: Some(0), refresh_token: None,
id_token: None,
scope: None,
};
let cached = CachedToken::from_response(&response);
assert!(!cached.is_valid());
}
#[test]
fn test_token_provider_enum_static() {
let rt = tokio::runtime::Runtime::new().unwrap();
let provider: TokenProviderEnum = TokenProviderEnum::Static(
StaticTokenProvider::new("enum-test".to_string()),
);
let token = rt.block_on(provider.get_token()).unwrap();
assert_eq!(token, "enum-test");
}
#[test]
fn test_token_provider_enum_from_string() {
let rt = tokio::runtime::Runtime::new().unwrap();
let provider: TokenProviderEnum = "direct-string".to_string().into();
let token = rt.block_on(provider.get_token()).unwrap();
assert_eq!(token, "direct-string");
}
#[test]
fn test_token_provider_enum_from_option() {
let rt = tokio::runtime::Runtime::new().unwrap();
let provider: TokenProviderEnum = Some("some-token".to_string()).into();
let token = rt.block_on(provider.get_token()).unwrap();
assert_eq!(token, "some-token");
}
#[test]
fn test_service_account_config_builder() {
let config = ServiceAccountConfig {
service_token: "svc-token".to_string(),
issuer_url: "https://idm.example.com/oauth2/openid/pdt-api".to_string(),
client_id: "pdt-api".to_string(),
client_secret: None,
audience: "pdt-api".to_string(),
scope: Some("openid profile email".to_string()),
};
let _provider = ServiceAccountTokenProvider::new(config);
}
#[test]
fn test_compute_expires_at() {
let ts = compute_expires_at(Some(0));
assert!(ts.ends_with("Z"));
assert_eq!(ts.len(), 20);
let now_epoch = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let parsed = crate::token_store::parse_rfc3339_to_epoch_public(&ts);
assert!(parsed <= now_epoch + 1);
}
#[test]
fn test_epoch_to_rfc3339_round_trip() {
assert_eq!(epoch_to_rfc3339(0), "1970-01-01T00:00:00Z");
assert_eq!(epoch_to_rfc3339(86400), "1970-01-02T00:00:00Z");
}
#[test]
fn test_interactive_config_builder() {
let config = InteractiveConfig {
issuer_url: "https://idm.example.com/oauth2/openid/test".to_string(),
client_id: "test-client".to_string(),
client_secret: None,
redirect_uri: "http://localhost:8765/callback".to_string(),
scope: "openid profile".to_string(),
credential_name: "test_interactive".to_string(),
};
assert_eq!(config.credential_name, "test_interactive");
}
}