use std::collections::BTreeMap;
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, OnceLock};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::{Mutex, RwLock};
use tokio::task::JoinHandle;
use tracing::{debug, warn};
use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::error::{KrafkaError, Result};
use crate::metrics::ConnectionMetrics;
const OAUTHBEARER_EXPIRY_SKEW_MARGIN_MS: i64 = 30_000;
pub(crate) const OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE: Duration = Duration::from_secs(300);
fn current_epoch_ms() -> Result<i64> {
let d = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_err(|_| KrafkaError::auth("system clock predates Unix epoch"))?;
i64::try_from(d.as_millis()).map_err(|_| KrafkaError::auth("current epoch_ms overflows i64"))
}
pub trait OAuthBearerTokenProvider: Send + Sync {
fn provide_token(&self) -> Pin<Box<dyn Future<Output = Result<OAuthBearerToken>> + Send + '_>>;
}
impl<F, Fut> OAuthBearerTokenProvider for F
where
F: Fn() -> Fut + Send + Sync,
Fut: Future<Output = Result<OAuthBearerToken>> + Send + 'static,
{
fn provide_token(&self) -> Pin<Box<dyn Future<Output = Result<OAuthBearerToken>> + Send + '_>> {
Box::pin(self())
}
}
struct OAuthTokenStoreInner {
provider: Arc<dyn OAuthBearerTokenProvider>,
cached: RwLock<Option<CachedToken>>,
refreshing: Mutex<()>,
metrics: OnceLock<Arc<ConnectionMetrics>>,
}
impl OAuthTokenStoreInner {
async fn fetch(&self) -> Result<OAuthBearerToken> {
let started = Instant::now();
match self.provider.provide_token().await {
Ok(token) => {
if let Some(metrics) = self.metrics.get() {
metrics.record_oauth_token_fetch(started.elapsed(), token.lifetime_ms());
}
debug!(
lifetime_ms = token.lifetime_ms(),
elapsed_ms = started.elapsed().as_millis() as u64,
"OAUTHBEARER token fetched"
);
Ok(token)
}
Err(e) => {
if let Some(metrics) = self.metrics.get() {
metrics.record_oauth_token_fetch_failure();
}
warn!(
error = %e,
elapsed_ms = started.elapsed().as_millis() as u64,
"OAUTHBEARER token fetch failed; broker connections needing a fresh \
token will fail until the provider recovers"
);
Err(e)
}
}
}
}
#[derive(Clone)]
struct CachedToken {
token: OAuthBearerToken,
fetched_at: Instant,
}
impl CachedToken {
fn new(token: OAuthBearerToken) -> Self {
Self {
token,
fetched_at: Instant::now(),
}
}
fn is_stale(&self) -> bool {
if self.token.lifetime_ms().is_none() {
return self.fetched_at.elapsed() >= OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE;
}
self.token.needs_refresh()
}
}
#[derive(Clone)]
pub struct OAuthBearerTokenProviderHandle(Arc<OAuthTokenStoreInner>);
impl OAuthBearerTokenProviderHandle {
pub fn new(provider: impl OAuthBearerTokenProvider + 'static) -> Self {
Self(Arc::new(OAuthTokenStoreInner {
provider: Arc::new(provider),
cached: RwLock::new(None),
refreshing: Mutex::new(()),
metrics: OnceLock::new(),
}))
}
pub(crate) fn bind_metrics(&self, metrics: Arc<ConnectionMetrics>) {
let _ = self.0.metrics.set(metrics);
}
pub async fn provide_token(&self) -> Result<OAuthBearerToken> {
{
let guard = self.0.cached.read().await;
if let Some(entry) = guard.as_ref()
&& !entry.is_stale()
{
return Ok(entry.token.clone());
}
}
let _coalesce = self.0.refreshing.lock().await;
{
let guard = self.0.cached.read().await;
if let Some(entry) = guard.as_ref()
&& !entry.is_stale()
{
return Ok(entry.token.clone());
}
}
let token = self.0.fetch().await?;
*self.0.cached.write().await = Some(CachedToken::new(token.clone()));
Ok(token)
}
#[cfg(test)]
pub(crate) async fn test_seed_cache(&self, token: OAuthBearerToken, age: Duration) {
let fetched_at = Instant::now().checked_sub(age).unwrap_or_else(Instant::now);
*self.0.cached.write().await = Some(CachedToken { token, fetched_at });
}
pub fn start_refresh_task(&self) -> JoinHandle<()> {
let inner = self.0.clone();
tokio::spawn(async move {
loop {
let sleep_duration = {
let guard = inner.cached.read().await;
match guard.as_ref().and_then(|e| e.token.lifetime_ms()) {
Some(lifetime_ms) => {
let wake_at_ms =
lifetime_ms.saturating_sub(OAUTHBEARER_EXPIRY_SKEW_MARGIN_MS);
let now_ms = current_epoch_ms().unwrap_or(i64::MAX);
let remaining_ms = wake_at_ms.saturating_sub(now_ms).max(0);
Duration::from_millis(remaining_ms as u64)
}
None => {
OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE
}
}
};
if !sleep_duration.is_zero() {
tokio::time::sleep(sleep_duration).await;
}
let _coalesce = inner.refreshing.lock().await;
match inner.fetch().await {
Ok(token) => {
*inner.cached.write().await = Some(CachedToken::new(token));
}
Err(_) => {
drop(_coalesce);
tokio::time::sleep(Duration::from_secs(5)).await;
}
}
}
})
}
}
impl fmt::Debug for OAuthBearerTokenProviderHandle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("[OAuthBearerTokenProvider]")
}
}
#[derive(Clone, Zeroize, ZeroizeOnDrop)]
pub struct OAuthBearerToken {
token_value: String,
#[zeroize(skip)]
extensions: BTreeMap<String, String>,
#[zeroize(skip)]
lifetime_ms: Option<i64>,
}
impl OAuthBearerToken {
pub fn new(token_value: impl Into<String>) -> Self {
Self {
token_value: token_value.into(),
extensions: BTreeMap::new(),
lifetime_ms: None,
}
}
pub fn with_extension(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.extensions.insert(key.into(), value.into());
self
}
pub fn with_lifetime_ms(mut self, lifetime_ms: i64) -> Self {
self.lifetime_ms = Some(lifetime_ms);
self
}
pub fn lifetime_ms(&self) -> Option<i64> {
self.lifetime_ms
}
pub fn is_expired(&self) -> bool {
self.lifetime_ms
.is_some_and(|lifetime_ms| current_epoch_ms().map_or(true, |now| now >= lifetime_ms))
}
pub fn needs_refresh(&self) -> bool {
self.lifetime_ms.is_some_and(|lifetime_ms| {
current_epoch_ms().map_or(true, |now| {
now >= lifetime_ms.saturating_sub(OAUTHBEARER_EXPIRY_SKEW_MARGIN_MS)
})
})
}
pub fn validate(&self) -> crate::error::Result<()> {
if self.token_value.is_empty() {
return Err(crate::error::KrafkaError::auth(
"OAuthBearer token value must not be empty",
));
}
if let Some(bad) = self
.token_value
.bytes()
.find(|&b| !(0x20..=0x7E).contains(&b))
{
return Err(crate::error::KrafkaError::auth(format!(
"OAuthBearer token value contains an invalid byte 0x{bad:02X}; \
token must consist of printable ASCII characters (0x20–0x7E) only"
)));
}
for (key, value) in &self.extensions {
for s in [key.as_str(), value.as_str()] {
if let Some(bad) = s.bytes().find(|&b| !(0x20..=0x7E).contains(&b)) {
return Err(crate::error::KrafkaError::auth(format!(
"OAuthBearer extension key/value contains an invalid byte 0x{bad:02X}; \
extension strings must consist of printable ASCII characters only"
)));
}
}
}
Ok(())
}
pub(crate) fn to_gs2_initial_response(&self) -> Vec<u8> {
let mut capacity = 3 + 1 + 12 + self.token_value.len() + 2; for (k, v) in &self.extensions {
capacity += 1 + k.len() + 1 + v.len(); }
let mut response = Vec::with_capacity(capacity);
response.extend_from_slice(b"n,,");
response.push(0x01);
response.extend_from_slice(b"auth=Bearer ");
response.extend_from_slice(self.token_value.as_bytes());
for (key, value) in &self.extensions {
response.push(0x01);
response.extend_from_slice(key.as_bytes());
response.push(b'=');
response.extend_from_slice(value.as_bytes());
}
response.push(0x01);
response.push(0x01);
response
}
pub(crate) fn process_server_response(&self, challenge: &[u8]) -> Result<()> {
if challenge.is_empty() {
return Ok(());
}
let error_msg = String::from_utf8_lossy(challenge);
Err(KrafkaError::auth(format!(
"OAUTHBEARER authentication failed: {error_msg}"
)))
}
}
impl fmt::Debug for OAuthBearerToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OAuthBearerToken")
.field("token_value", &"[REDACTED]")
.field("extensions", &self.extensions)
.field("lifetime_ms", &self.lifetime_ms)
.finish()
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn test_oauthbearer_token_basic() {
let token = OAuthBearerToken::new("my-jwt-token");
let response = token.to_gs2_initial_response();
let expected = b"n,,\x01auth=Bearer my-jwt-token\x01\x01";
assert_eq!(response, expected);
}
#[test]
fn test_oauthbearer_token_with_single_extension() {
let token = OAuthBearerToken::new("my-token").with_extension("logicalCluster", "lkc-123");
let response = token.to_gs2_initial_response();
let response_str = String::from_utf8_lossy(&response);
assert!(response_str.starts_with("n,,\x01auth=Bearer my-token"));
assert!(response_str.contains("\x01logicalCluster=lkc-123"));
assert!(response_str.ends_with("\x01\x01"));
}
#[test]
fn test_oauthbearer_token_with_multiple_extensions() {
let token = OAuthBearerToken::new("tok")
.with_extension("ext1", "val1")
.with_extension("ext2", "val2");
let response = token.to_gs2_initial_response();
let response_str = String::from_utf8_lossy(&response);
assert!(response_str.starts_with("n,,\x01auth=Bearer tok"));
assert!(response_str.contains("ext1=val1"));
assert!(response_str.contains("ext2=val2"));
assert!(response_str.ends_with("\x01\x01"));
}
#[test]
fn test_oauthbearer_debug_redacts_token() {
let token = OAuthBearerToken::new("secret-token-value");
let debug = format!("{token:?}");
assert!(!debug.contains("secret-token-value"));
assert!(debug.contains("[REDACTED]"));
}
#[test]
fn test_oauthbearer_server_response_success_empty() {
let token = OAuthBearerToken::new("tok");
assert!(token.process_server_response(b"").is_ok());
}
#[test]
fn test_oauthbearer_server_response_error_json() {
let token = OAuthBearerToken::new("tok");
let error_json = br#"{"status":"invalid_token","scope":"openid"}"#;
let result = token.process_server_response(error_json);
assert!(result.is_err());
let err = result.unwrap_err().to_string();
assert!(err.contains("invalid_token"));
}
#[test]
fn test_oauthbearer_gs2_format_compliance() {
let token = OAuthBearerToken::new("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJhbGljZSJ9.sig");
let response = token.to_gs2_initial_response();
assert_eq!(&response[..3], b"n,,");
assert_eq!(response[3], 0x01);
assert_eq!(&response[4..16], b"auth=Bearer ");
let len = response.len();
assert_eq!(response[len - 2], 0x01);
assert_eq!(response[len - 1], 0x01);
}
#[test]
fn test_oauthbearer_empty_token_produces_valid_gs2() {
let token = OAuthBearerToken::new("");
let response = token.to_gs2_initial_response();
assert_eq!(response, b"n,,\x01auth=Bearer \x01\x01");
}
#[test]
fn test_oauthbearer_validate_valid_token() {
let token = OAuthBearerToken::new("eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJhbGljZSJ9.sig");
assert!(token.validate().is_ok());
}
#[test]
fn test_oauthbearer_validate_empty_token_is_err() {
let token = OAuthBearerToken::new("");
assert!(token.validate().is_err());
}
#[test]
fn test_oauthbearer_validate_null_byte_is_err() {
let token = OAuthBearerToken::new("tok\x00en");
assert!(token.validate().is_err(), "null byte must be rejected");
}
#[test]
fn test_oauthbearer_validate_gs2_separator_is_err() {
let token = OAuthBearerToken::new("tok\x01en");
assert!(
token.validate().is_err(),
"0x01 (GS2 separator) must be rejected in token value"
);
}
#[test]
fn test_oauthbearer_validate_non_ascii_is_err() {
let token = OAuthBearerToken::new("\u{0080}");
assert!(
token.validate().is_err(),
"non-ASCII (multi-byte UTF-8) byte must be rejected"
);
}
#[test]
fn test_oauthbearer_validate_extension_null_byte_is_err() {
let token = OAuthBearerToken::new("valid-token").with_extension("key\x00", "value");
assert!(
token.validate().is_err(),
"null byte in extension key must be rejected"
);
}
#[test]
fn test_oauthbearer_validate_extension_gs2_separator_in_value_is_err() {
let token = OAuthBearerToken::new("valid-token").with_extension("key", "val\x01ue");
assert!(
token.validate().is_err(),
"0x01 in extension value must be rejected"
);
}
#[test]
fn test_oauthbearer_validate_valid_extension() {
let token = OAuthBearerToken::new("valid-token")
.with_extension("logicalCluster", "lkc-abc123")
.with_extension("identityPoolId", "pool-xyz789");
assert!(token.validate().is_ok());
}
#[test]
fn test_oauthbearer_token_clone() {
let token = OAuthBearerToken::new("tok").with_extension("k", "v");
let cloned = token.clone();
assert_eq!(
cloned.to_gs2_initial_response(),
token.to_gs2_initial_response()
);
}
#[tokio::test]
async fn test_token_provider_closure_impl() {
let provider = || async { Ok(OAuthBearerToken::new("from-closure")) };
let token = provider.provide_token().await.unwrap();
assert_eq!(
token.to_gs2_initial_response(),
OAuthBearerToken::new("from-closure").to_gs2_initial_response()
);
}
#[tokio::test]
async fn test_token_provider_handle() {
let handle = OAuthBearerTokenProviderHandle::new(|| async {
Ok(OAuthBearerToken::new("handle-token"))
});
let token = handle.provide_token().await.unwrap();
assert_eq!(
token.to_gs2_initial_response(),
OAuthBearerToken::new("handle-token").to_gs2_initial_response()
);
}
#[test]
fn test_token_provider_handle_clone() {
let handle =
OAuthBearerTokenProviderHandle::new(|| async { Ok(OAuthBearerToken::new("tok")) });
let cloned = handle.clone();
assert!(Arc::ptr_eq(&handle.0, &cloned.0));
}
#[test]
fn test_token_provider_handle_debug_no_secrets() {
let handle = OAuthBearerTokenProviderHandle::new(|| async {
Ok(OAuthBearerToken::new("super-secret"))
});
let debug = format!("{handle:?}");
assert_eq!(debug, "[OAuthBearerTokenProvider]");
assert!(!debug.contains("super-secret"));
}
#[tokio::test]
async fn test_token_provider_error_propagation() {
let handle = OAuthBearerTokenProviderHandle::new(|| async {
Err(KrafkaError::auth("token expired"))
});
let result = handle.provide_token().await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("token expired"));
}
#[tokio::test]
async fn test_token_provider_struct_impl() {
struct StaticProvider {
token: String,
}
impl OAuthBearerTokenProvider for StaticProvider {
fn provide_token(
&self,
) -> Pin<Box<dyn Future<Output = Result<OAuthBearerToken>> + Send + '_>> {
let token = self.token.clone();
Box::pin(async move { Ok(OAuthBearerToken::new(token)) })
}
}
let provider = StaticProvider {
token: "struct-token".to_string(),
};
let handle = OAuthBearerTokenProviderHandle::new(provider);
let token = handle.provide_token().await.unwrap();
assert_eq!(
token.to_gs2_initial_response(),
OAuthBearerToken::new("struct-token").to_gs2_initial_response()
);
}
#[test]
fn test_oauthbearer_token_not_expired_without_lifetime() {
let token = OAuthBearerToken::new("tok");
assert!(!token.is_expired());
assert!(token.lifetime_ms().is_none());
}
#[test]
fn test_oauthbearer_token_not_expired_future_lifetime() {
let future_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as i64
+ 3_600_000;
let token = OAuthBearerToken::new("tok").with_lifetime_ms(future_ms);
assert!(!token.is_expired());
assert_eq!(token.lifetime_ms(), Some(future_ms));
}
#[test]
fn test_oauthbearer_token_expired_past_lifetime() {
let past_ms = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as i64
- 3_600_000;
let token = OAuthBearerToken::new("tok").with_lifetime_ms(past_ms);
assert!(token.is_expired());
}
#[test]
fn test_oauthbearer_token_needs_refresh_near_expiry() {
let near_future_ms = current_epoch_ms().unwrap() + 10_000;
let token = OAuthBearerToken::new("tok").with_lifetime_ms(near_future_ms);
assert!(!token.is_expired());
assert!(token.needs_refresh());
}
#[test]
fn test_oauthbearer_token_does_not_need_refresh_with_safe_margin() {
let future_ms = current_epoch_ms().unwrap() + 60_000;
let token = OAuthBearerToken::new("tok").with_lifetime_ms(future_ms);
assert!(!token.needs_refresh());
}
#[test]
fn test_cached_token_without_lifetime_is_fresh_when_new() {
let entry = CachedToken::new(OAuthBearerToken::new("jwt"));
assert!(
!entry.is_stale(),
"a just-fetched token must be served from cache"
);
}
#[test]
fn test_cached_token_without_lifetime_goes_stale_with_age() {
let token = OAuthBearerToken::new("jwt");
assert!(
!token.needs_refresh(),
"precondition: an unknown-expiry token never reports needs_refresh"
);
let entry = CachedToken {
token,
fetched_at: Instant::now()
.checked_sub(OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE + Duration::from_secs(1))
.expect("test backdate within system uptime"),
};
assert!(
entry.is_stale(),
"an unknown-expiry token older than the max age must be refetched"
);
}
#[test]
fn test_cached_token_with_lifetime_still_uses_skew_margin() {
let near = current_epoch_ms().unwrap() + 10_000;
let entry = CachedToken::new(OAuthBearerToken::new("jwt").with_lifetime_ms(near));
assert!(entry.is_stale(), "inside the 30s skew window");
let far = current_epoch_ms().unwrap() + 3_600_000;
let entry = CachedToken::new(OAuthBearerToken::new("jwt").with_lifetime_ms(far));
assert!(!entry.is_stale(), "well outside the skew window");
}
#[tokio::test]
async fn a_successful_fetch_is_counted_with_its_expiry() {
let metrics = Arc::new(ConnectionMetrics::new());
let expiry = current_epoch_ms().unwrap() + 3_600_000;
let handle = OAuthBearerTokenProviderHandle::new(move || async move {
Ok(OAuthBearerToken::new("jwt").with_lifetime_ms(expiry))
});
handle.bind_metrics(metrics.clone());
handle.provide_token().await.expect("provider succeeds");
assert_eq!(metrics.oauth_token_fetches.get(), 1);
assert_eq!(metrics.oauth_token_fetch_failures.get(), 0);
assert_eq!(
metrics.oauth_token_expiry_epoch_ms.get(),
expiry as u64,
"the dashboard's remaining-lifetime panel reads this gauge"
);
handle.provide_token().await.expect("cached");
assert_eq!(
metrics.oauth_token_fetches.get(),
1,
"serving from cache must not be counted as a fetch"
);
}
#[tokio::test]
async fn a_failed_fetch_is_counted_separately() {
let metrics = Arc::new(ConnectionMetrics::new());
let handle = OAuthBearerTokenProviderHandle::new(|| async {
Err(KrafkaError::auth("token endpoint returned HTTP 401"))
});
handle.bind_metrics(metrics.clone());
let err = handle
.provide_token()
.await
.expect_err("the provider fails");
assert!(err.to_string().contains("401"), "the error must survive");
assert_eq!(metrics.oauth_token_fetches.get(), 1);
assert_eq!(metrics.oauth_token_fetch_failures.get(), 1);
assert_eq!(
metrics.oauth_token_expiry_epoch_ms.get(),
0,
"a failure must not claim a token expiry"
);
}
#[tokio::test]
async fn a_failed_refresh_leaves_the_previous_expiry_alone() {
use std::sync::atomic::{AtomicUsize, Ordering};
let metrics = Arc::new(ConnectionMetrics::new());
let expiry = current_epoch_ms().unwrap() + 3_600_000;
let calls = Arc::new(AtomicUsize::new(0));
let c = calls.clone();
let handle = OAuthBearerTokenProviderHandle::new(move || {
let c = c.clone();
async move {
if c.fetch_add(1, Ordering::SeqCst) == 0 {
Ok(OAuthBearerToken::new("jwt").with_lifetime_ms(expiry))
} else {
Err(KrafkaError::auth("identity provider unavailable"))
}
}
});
handle.bind_metrics(metrics.clone());
handle.provide_token().await.expect("first fetch succeeds");
handle
.test_seed_cache(
OAuthBearerToken::new("jwt"),
OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE + Duration::from_secs(1),
)
.await;
handle
.provide_token()
.await
.expect_err("second fetch fails");
assert_eq!(metrics.oauth_token_fetches.get(), 2);
assert_eq!(metrics.oauth_token_fetch_failures.get(), 1);
assert_eq!(
metrics.oauth_token_expiry_epoch_ms.get(),
expiry as u64,
"a failed refresh must leave the known expiry in place"
);
}
#[tokio::test]
async fn an_unbound_handle_still_fetches() {
let handle =
OAuthBearerTokenProviderHandle::new(|| async { Ok(OAuthBearerToken::new("jwt")) });
assert!(handle.provide_token().await.is_ok());
}
#[tokio::test]
async fn test_provide_token_refetches_after_max_age() {
use std::sync::atomic::{AtomicUsize, Ordering};
let calls = Arc::new(AtomicUsize::new(0));
let c = calls.clone();
let handle = OAuthBearerTokenProviderHandle::new(move || {
let c = c.clone();
async move {
let n = c.fetch_add(1, Ordering::SeqCst);
Ok(OAuthBearerToken::new(format!("jwt-{n}")))
}
});
let first = handle.provide_token().await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
let _ = handle.provide_token().await.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"fresh entry must be cached"
);
handle
.test_seed_cache(
first,
OAUTHBEARER_UNKNOWN_EXPIRY_MAX_AGE + Duration::from_secs(1),
)
.await;
let refreshed = handle.provide_token().await.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"a token past the unknown-expiry max age must be refetched"
);
assert_eq!(
refreshed.to_gs2_initial_response(),
OAuthBearerToken::new("jwt-1").to_gs2_initial_response()
);
}
}