use std::future::Future;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use tracing::{debug, warn};
use zeroize::Zeroizing;
use crate::error::{KrafkaError, Result};
use crate::http::{HttpClient, base64_encode};
use super::oauthbearer::{OAuthBearerToken, OAuthBearerTokenProvider};
const JWT_BEARER_ASSERTION_TYPE: &str = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer";
const MAX_TOKEN_RESPONSE_BYTES: usize = 1024 * 1024;
const MAX_ASSERTION_FILE_BYTES: u64 = 64 * 1024;
#[derive(Clone)]
pub enum AssertionSource {
File(PathBuf),
Static(Zeroizing<String>),
#[allow(clippy::type_complexity)]
Callback(Arc<dyn Fn() -> Pin<Box<dyn Future<Output = Result<String>> + Send>> + Send + Sync>),
}
impl std::fmt::Debug for AssertionSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::File(path) => f.debug_tuple("File").field(path).finish(),
Self::Static(_) => f.debug_tuple("Static").field(&"[REDACTED]").finish(),
Self::Callback(_) => f.write_str("Callback(<fn>)"),
}
}
}
impl AssertionSource {
async fn resolve(&self) -> Result<Zeroizing<String>> {
match self {
Self::Static(jwt) => Ok(jwt.clone()),
Self::Callback(f) => Ok(Zeroizing::new(f().await?)),
Self::File(path) => {
let metadata = tokio::fs::metadata(path).await.map_err(|e| {
KrafkaError::auth(format!(
"cannot stat client-assertion file {}: {e}",
path.display()
))
})?;
if metadata.len() > MAX_ASSERTION_FILE_BYTES {
return Err(KrafkaError::auth(format!(
"client-assertion file {} is {} bytes, above the {MAX_ASSERTION_FILE_BYTES} \
byte limit; a JWT is never this large, so this path is probably wrong",
path.display(),
metadata.len()
)));
}
let raw = tokio::fs::read_to_string(path).await.map_err(|e| {
KrafkaError::auth(format!(
"cannot read client-assertion file {}: {e}",
path.display()
))
})?;
let trimmed = raw.trim();
if trimmed.is_empty() {
return Err(KrafkaError::auth(format!(
"client-assertion file {} is empty; if a sidecar rotates it, \
this is the truncate-then-write window and the next attempt \
should succeed",
path.display()
)));
}
Ok(Zeroizing::new(trimmed.to_string()))
}
}
}
}
#[derive(Clone)]
pub enum ClientCredentials {
Secret {
client_id: String,
client_secret: Zeroizing<String>,
},
Assertion {
source: AssertionSource,
},
}
impl std::fmt::Debug for ClientCredentials {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Secret { client_id, .. } => f
.debug_struct("Secret")
.field("client_id", client_id)
.field("client_secret", &"[REDACTED]")
.finish(),
Self::Assertion { source } => {
f.debug_struct("Assertion").field("source", source).finish()
}
}
}
}
impl ClientCredentials {
pub fn secret(client_id: impl Into<String>, client_secret: impl Into<String>) -> Self {
Self::Secret {
client_id: client_id.into(),
client_secret: Zeroizing::new(client_secret.into()),
}
}
pub fn assertion(source: AssertionSource) -> Self {
Self::Assertion { source }
}
}
pub struct OidcTokenProvider {
token_endpoint: String,
credentials: ClientCredentials,
client_id: Option<String>,
scope: Option<String>,
form_parameters: Vec<(String, String)>,
sasl_extensions: Vec<(String, String)>,
http: HttpClient,
}
impl std::fmt::Debug for OidcTokenProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OidcTokenProvider")
.field("token_endpoint", &self.token_endpoint)
.field("credentials", &self.credentials)
.field("client_id", &self.client_id)
.field("scope", &self.scope)
.field("form_parameters", &self.form_parameters.len())
.field("sasl_extensions", &self.sasl_extensions.len())
.finish()
}
}
impl OidcTokenProvider {
pub fn builder(token_endpoint: impl Into<String>) -> OidcTokenProviderBuilder {
OidcTokenProviderBuilder {
token_endpoint: token_endpoint.into(),
credentials: None,
client_id: None,
scope: None,
form_parameters: Vec::new(),
sasl_extensions: Vec::new(),
request_timeout: None,
}
}
async fn build_request(&self) -> Result<(Zeroizing<String>, Option<Zeroizing<String>>)> {
let mut form = String::from("grant_type=client_credentials");
let auth_header = match &self.credentials {
ClientCredentials::Secret {
client_id,
client_secret,
} => {
let raw = format!(
"{}:{}",
form_urlencode(client_id),
form_urlencode(client_secret)
);
Some(Zeroizing::new(format!(
"Basic {}",
base64_encode(raw.as_bytes())
)))
}
ClientCredentials::Assertion { source } => {
let jwt = source.resolve().await?;
form.push_str("&client_assertion_type=");
form.push_str(&form_urlencode(JWT_BEARER_ASSERTION_TYPE));
form.push_str("&client_assertion=");
form.push_str(&form_urlencode(&jwt));
None
}
};
if let Some(client_id) = &self.client_id {
form.push_str("&client_id=");
form.push_str(&form_urlencode(client_id));
}
if let Some(scope) = &self.scope {
form.push_str("&scope=");
form.push_str(&form_urlencode(scope));
}
for (key, value) in &self.form_parameters {
form.push('&');
form.push_str(&form_urlencode(key));
form.push('=');
form.push_str(&form_urlencode(value));
}
Ok((Zeroizing::new(form), auth_header))
}
async fn fetch_token(&self) -> Result<OAuthBearerToken> {
let (form, auth_header) = self.build_request().await?;
let response = self
.http
.request(
"POST",
&self.token_endpoint,
&[
("Content-Type", "application/x-www-form-urlencoded"),
("Accept", "application/json"),
],
Some(form.as_bytes()),
auth_header.as_ref().map(|h| h.as_str()),
)
.await?;
if response.body.len() > MAX_TOKEN_RESPONSE_BYTES {
return Err(KrafkaError::auth(format!(
"token endpoint returned {} bytes, above the {MAX_TOKEN_RESPONSE_BYTES} byte limit",
response.body.len()
)));
}
if !(200..300).contains(&response.status) {
let detail = describe_oauth_error(&response.body);
return Err(KrafkaError::auth(format!(
"token endpoint {} returned HTTP {}{detail}",
self.token_endpoint, response.status
)));
}
let parsed: TokenResponse = serde_json::from_slice(&response.body).map_err(|e| {
KrafkaError::auth(format!(
"token endpoint {} returned a body that is not a valid OAuth token \
response: {e}",
self.token_endpoint
))
})?;
if parsed.access_token.is_empty() {
return Err(KrafkaError::auth(format!(
"token endpoint {} returned an empty access_token",
self.token_endpoint
)));
}
let mut token = OAuthBearerToken::new(parsed.access_token);
match parsed.expires_in {
Some(seconds) if seconds > 0 => {
let now_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_err(|_| KrafkaError::auth("system clock predates Unix epoch"))?
.as_millis();
let expiry_ms = i64::try_from(now_ms)
.ok()
.and_then(|now| seconds.checked_mul(1000).and_then(|d| now.checked_add(d)));
match expiry_ms {
Some(ms) => token = token.with_lifetime_ms(ms),
None => warn!(
expires_in = seconds,
"token endpoint reported an expires_in that overflows i64 \
milliseconds; treating the token as having no known expiry"
),
}
}
Some(seconds) => warn!(
expires_in = seconds,
"token endpoint reported a non-positive expires_in; treating the token \
as having no known expiry"
),
None => debug!(
"token endpoint returned no expires_in; the token will be re-fetched on \
the provider store's unknown-expiry schedule"
),
}
for (key, value) in &self.sasl_extensions {
token = token.with_extension(key, value);
}
token.validate()?;
Ok(token)
}
}
impl OAuthBearerTokenProvider for OidcTokenProvider {
fn provide_token(&self) -> Pin<Box<dyn Future<Output = Result<OAuthBearerToken>> + Send + '_>> {
Box::pin(self.fetch_token())
}
}
#[must_use = "builders do nothing until .build() is called"]
#[derive(Debug)]
pub struct OidcTokenProviderBuilder {
token_endpoint: String,
credentials: Option<ClientCredentials>,
client_id: Option<String>,
scope: Option<String>,
form_parameters: Vec<(String, String)>,
sasl_extensions: Vec<(String, String)>,
request_timeout: Option<Duration>,
}
impl OidcTokenProviderBuilder {
pub fn credentials(mut self, credentials: ClientCredentials) -> Self {
self.credentials = Some(credentials);
self
}
pub fn client_id(mut self, client_id: impl Into<String>) -> Self {
self.client_id = Some(client_id.into());
self
}
pub fn scope(mut self, scope: impl Into<String>) -> Self {
self.scope = Some(scope.into());
self
}
pub fn form_parameter(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.form_parameters.push((key.into(), value.into()));
self
}
pub fn sasl_extension(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
self.sasl_extensions.push((key.into(), value.into()));
self
}
pub fn request_timeout(mut self, timeout: Duration) -> Self {
self.request_timeout = Some(timeout);
self
}
pub fn build(self) -> Result<OidcTokenProvider> {
if self.token_endpoint.is_empty() {
return Err(KrafkaError::config("OIDC token_endpoint must not be empty"));
}
if self.token_endpoint.starts_with("http://") {
return Err(KrafkaError::config(format!(
"OIDC token endpoint {} uses plain HTTP; the request carries a client \
credential and the response carries an access token, so https is \
required",
self.token_endpoint
)));
}
if !self.token_endpoint.starts_with("https://") {
return Err(KrafkaError::config(format!(
"OIDC token endpoint {} is not an absolute https URL",
self.token_endpoint
)));
}
let credentials = self.credentials.ok_or_else(|| {
KrafkaError::config(
"OIDC token provider needs credentials: ClientCredentials::secret(..) \
for KIP-768, or ClientCredentials::assertion(..) for KIP-1258",
)
})?;
Ok(OidcTokenProvider {
token_endpoint: self.token_endpoint,
credentials,
client_id: self.client_id,
scope: self.scope,
form_parameters: self.form_parameters,
sasl_extensions: self.sasl_extensions,
http: HttpClient::with_webpki_roots(self.request_timeout)?,
})
}
}
#[derive(serde::Deserialize)]
struct TokenResponse {
access_token: String,
#[serde(default)]
expires_in: Option<i64>,
}
fn describe_oauth_error(body: &[u8]) -> String {
#[derive(serde::Deserialize)]
struct OAuthError {
error: Option<String>,
error_description: Option<String>,
}
let Ok(parsed) = serde_json::from_slice::<OAuthError>(body) else {
return String::new();
};
match (parsed.error, parsed.error_description) {
(Some(code), Some(description)) => format!(" — {code}: {description}"),
(Some(code), None) => format!(" — {code}"),
(None, Some(description)) => format!(" — {description}"),
(None, None) => String::new(),
}
}
fn form_urlencode(value: &str) -> String {
let mut out = String::with_capacity(value.len());
for byte in value.as_bytes() {
match byte {
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' => {
out.push(*byte as char);
}
b' ' => out.push('+'),
other => {
out.push('%');
out.push(
char::from_digit((other >> 4) as u32, 16)
.unwrap_or('0')
.to_ascii_uppercase(),
);
out.push(
char::from_digit((other & 0x0F) as u32, 16)
.unwrap_or('0')
.to_ascii_uppercase(),
);
}
}
}
out
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn form_urlencode_escapes_separators() {
assert_eq!(form_urlencode("a&b=c"), "a%26b%3Dc");
assert_eq!(form_urlencode("plain-value_1.0~x"), "plain-value_1.0~x");
assert_eq!(form_urlencode("a b"), "a+b");
assert_eq!(form_urlencode("sl/ash"), "sl%2Fash");
assert_eq!(form_urlencode("ü"), "%C3%BC");
assert_eq!(form_urlencode("nl\n"), "nl%0A");
}
#[test]
fn form_urlencode_leaves_a_jwt_intact() {
let jwt = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJhYmMtMTIzIn0.c2lnbmF0dXJl-_x";
assert_eq!(form_urlencode(jwt), jwt);
}
fn provider(credentials: ClientCredentials) -> OidcTokenProvider {
OidcTokenProvider::builder("https://idp.example.com/token")
.credentials(credentials)
.build()
.expect("valid provider")
}
#[tokio::test]
async fn secret_credentials_use_http_basic() {
let p = provider(ClientCredentials::secret("id", "secret"));
let (form, auth) = p.build_request().await.unwrap();
assert_eq!(&*form, "grant_type=client_credentials");
let auth = auth.expect("Basic header present");
assert!(auth.starts_with("Basic "), "got: {}", *auth);
assert_eq!(&*auth, "Basic aWQ6c2VjcmV0");
}
#[tokio::test]
async fn secret_with_colon_is_encoded_before_joining() {
let p = provider(ClientCredentials::secret("id", "pa:ss"));
let (_, auth) = p.build_request().await.unwrap();
let auth = auth.unwrap();
let encoded = auth.strip_prefix("Basic ").unwrap();
let decoded = base64_decode_for_test(encoded);
assert_eq!(decoded, "id:pa%3Ass");
}
#[tokio::test]
async fn assertion_credentials_use_the_jwt_bearer_type() {
let jwt = "header.payload.signature";
let p = provider(ClientCredentials::assertion(AssertionSource::Static(
Zeroizing::new(jwt.to_string()),
)));
let (form, auth) = p.build_request().await.unwrap();
assert!(auth.is_none(), "assertion flow sends no Basic header");
assert!(form.starts_with("grant_type=client_credentials"));
assert!(
form.contains(
"client_assertion_type=\
urn%3Aietf%3Aparams%3Aoauth%3Aclient-assertion-type%3Ajwt-bearer"
),
"got: {}",
*form
);
assert!(
form.contains(&format!("client_assertion={jwt}")),
"got: {}",
*form
);
}
#[tokio::test]
async fn scope_and_client_id_are_appended() {
let p = OidcTokenProvider::builder("https://idp.example.com/token")
.credentials(ClientCredentials::assertion(AssertionSource::Static(
Zeroizing::new("a.b.c".into()),
)))
.client_id("my client")
.scope("kafka:write kafka:read")
.form_parameter("audience", "kafka")
.build()
.unwrap();
let (form, _) = p.build_request().await.unwrap();
assert!(form.contains("&client_id=my+client"), "got: {}", *form);
assert!(
form.contains("&scope=kafka%3Awrite+kafka%3Aread"),
"got: {}",
*form
);
assert!(form.contains("&audience=kafka"), "got: {}", *form);
}
#[tokio::test]
async fn file_assertion_is_trimmed() {
let dir = std::env::temp_dir().join(format!("krafka-oidc-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("assertion.jwt");
std::fs::write(&path, " a.b.c\n").unwrap();
let source = AssertionSource::File(path.clone());
assert_eq!(&*source.resolve().await.unwrap(), "a.b.c");
std::fs::remove_file(&path).ok();
}
#[tokio::test]
async fn empty_assertion_file_names_the_rotation_window() {
let dir = std::env::temp_dir().join(format!("krafka-oidc-empty-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("assertion.jwt");
std::fs::write(&path, "\n \n").unwrap();
let err = AssertionSource::File(path.clone())
.resolve()
.await
.expect_err("empty file must error");
assert!(err.to_string().contains("empty"), "got: {err}");
std::fs::remove_file(&path).ok();
}
#[tokio::test]
async fn missing_assertion_file_names_the_path() {
let path = PathBuf::from("/nonexistent/krafka/assertion.jwt");
let err = AssertionSource::File(path)
.resolve()
.await
.expect_err("missing file must error");
assert!(err.to_string().contains("assertion.jwt"), "got: {err}");
}
#[tokio::test]
async fn callback_assertion_is_invoked_each_time() {
use std::sync::atomic::{AtomicUsize, Ordering};
let calls = Arc::new(AtomicUsize::new(0));
let counter = calls.clone();
let source = AssertionSource::Callback(Arc::new(move || {
let counter = counter.clone();
Box::pin(async move {
let n = counter.fetch_add(1, Ordering::SeqCst);
Ok(format!("jwt-{n}"))
})
}));
assert_eq!(&*source.resolve().await.unwrap(), "jwt-0");
assert_eq!(&*source.resolve().await.unwrap(), "jwt-1");
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[test]
fn plain_http_endpoint_is_rejected() {
let err = OidcTokenProvider::builder("http://idp.example.com/token")
.credentials(ClientCredentials::secret("id", "secret"))
.build()
.expect_err("http must be rejected")
.to_string();
assert!(err.contains("https is required"), "got: {err}");
}
#[test]
fn relative_endpoint_is_rejected() {
assert!(
OidcTokenProvider::builder("/token")
.credentials(ClientCredentials::secret("id", "secret"))
.build()
.is_err()
);
assert!(
OidcTokenProvider::builder("")
.credentials(ClientCredentials::secret("id", "secret"))
.build()
.is_err()
);
}
#[test]
fn missing_credentials_are_rejected() {
let err = OidcTokenProvider::builder("https://idp.example.com/token")
.build()
.expect_err("credentials are required")
.to_string();
assert!(err.contains("ClientCredentials"), "got: {err}");
}
#[test]
fn debug_redacts_the_credential() {
let secret = format!(
"{:?}",
provider(ClientCredentials::secret("id", "super-secret"))
);
assert!(!secret.contains("super-secret"), "got: {secret}");
let assertion = format!(
"{:?}",
AssertionSource::Static(Zeroizing::new("a.b.c".into()))
);
assert!(!assertion.contains("a.b.c"), "got: {assertion}");
assert!(assertion.contains("REDACTED"), "got: {assertion}");
}
#[tokio::test]
async fn sasl_extensions_do_not_leak_into_the_token_request() {
let p = OidcTokenProvider::builder("https://idp.example.com/token")
.credentials(ClientCredentials::secret("id", "secret"))
.sasl_extension("logicalCluster", "lkc-123")
.form_parameter("audience", "kafka")
.build()
.unwrap();
let (form, _) = p.build_request().await.unwrap();
assert!(form.contains("&audience=kafka"), "got: {}", *form);
assert!(
!form.contains("logicalCluster"),
"SASL extensions must not be sent to the identity provider: {}",
*form
);
assert_eq!(p.sasl_extensions.len(), 1);
}
#[test]
fn oauth_error_body_is_surfaced() {
let body = br#"{"error":"invalid_client","error_description":"unknown client id"}"#;
assert_eq!(
describe_oauth_error(body),
" — invalid_client: unknown client id"
);
let code_only = br#"{"error":"invalid_scope"}"#;
assert_eq!(describe_oauth_error(code_only), " — invalid_scope");
}
#[test]
fn non_json_error_body_degrades_quietly() {
assert_eq!(describe_oauth_error(b"<html>502 Bad Gateway</html>"), "");
assert_eq!(describe_oauth_error(b""), "");
}
#[test]
fn token_response_expires_in_is_optional() {
let with: TokenResponse =
serde_json::from_slice(br#"{"access_token":"t","expires_in":3600}"#).unwrap();
assert_eq!(with.access_token, "t");
assert_eq!(with.expires_in, Some(3600));
let without: TokenResponse =
serde_json::from_slice(br#"{"access_token":"t","token_type":"Bearer"}"#).unwrap();
assert_eq!(without.expires_in, None);
}
fn base64_decode_for_test(input: &str) -> String {
use base64::Engine as _;
let bytes = base64::engine::general_purpose::STANDARD
.decode(input)
.expect("valid base64");
String::from_utf8(bytes).expect("valid utf8")
}
}