use std::path::PathBuf;
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::TlsConfig;
use super::oauthbearer::OAuthBearerToken;
use super::{CredentialProvider, ErasedCredentialProvider};
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 struct AssertionSource(Source);
#[derive(Clone)]
enum Source {
File(PathBuf),
Static(Zeroizing<String>),
Provider(Arc<dyn ErasedCredentialProvider<String>>),
}
impl std::fmt::Debug for AssertionSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &self.0 {
Source::File(path) => f.debug_tuple("File").field(path).finish(),
Source::Static(_) => f.debug_tuple("Static").field(&"[REDACTED]").finish(),
Source::Provider(_) => f.write_str("Provider(<fn>)"),
}
}
}
impl AssertionSource {
pub fn file(path: impl Into<PathBuf>) -> Self {
Self(Source::File(path.into()))
}
pub fn fixed(jwt: impl Into<String>) -> Self {
Self(Source::Static(Zeroizing::new(jwt.into())))
}
pub fn provider(provider: impl CredentialProvider<String> + 'static) -> Self {
Self(Source::Provider(Arc::new(provider)))
}
async fn resolve(&self) -> Result<Zeroizing<String>> {
match &self.0 {
Source::Static(jwt) => Ok(jwt.clone()),
Source::Provider(provider) => Ok(Zeroizing::new(provider.credentials_erased().await?)),
Source::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,
trust: TlsConfig::new(),
}
}
async fn build_request(&self) -> Result<(Zeroizing<String>, Option<Zeroizing<String>>)> {
let mut pairs: Vec<(&str, &str)> = vec![("grant_type", "client_credentials")];
let jwt;
let auth_header = match &self.credentials {
ClientCredentials::Secret {
client_id,
client_secret,
} => {
let mut raw = Zeroizing::new(String::with_capacity(
form_urlencoded_len(client_id) + 1 + form_urlencoded_len(client_secret),
));
form_urlencode_into(&mut raw, client_id);
raw.push(':');
form_urlencode_into(&mut raw, client_secret);
let encoded = Zeroizing::new(base64_encode(raw.as_bytes()));
let mut header = Zeroizing::new(String::with_capacity(6 + encoded.len()));
header.push_str("Basic ");
header.push_str(&encoded);
Some(header)
}
ClientCredentials::Assertion { source } => {
jwt = source.resolve().await?;
pairs.push(("client_assertion_type", JWT_BEARER_ASSERTION_TYPE));
pairs.push(("client_assertion", &jwt));
None
}
};
if let Some(client_id) = &self.client_id {
pairs.push(("client_id", client_id));
}
if let Some(scope) = &self.scope {
pairs.push(("scope", scope));
}
for (key, value) in &self.form_parameters {
pairs.push((key, value));
}
Ok((encode_form(&pairs), auth_header))
}
async fn fetch_token(&self) -> Result<OAuthBearerToken> {
let (form, auth_header) = self.build_request().await?;
let response = self
.http
.post_form(
&self.token_endpoint,
form.as_bytes(),
auth_header.as_ref().map(|h| h.as_str()),
)
.await
.map_err(|e| self.endpoint_error(e))?;
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 mut 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(std::mem::take(&mut *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 OidcTokenProvider {
fn endpoint_error(&self, error: KrafkaError) -> KrafkaError {
match error {
KrafkaError::Auth { message, source } => KrafkaError::Auth {
message: format!("token endpoint {}: {message}", self.token_endpoint),
source,
},
other => other,
}
}
}
impl CredentialProvider<OAuthBearerToken> for OidcTokenProvider {
async fn credentials(&self) -> Result<OAuthBearerToken> {
self.fetch_token().await
}
}
#[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>,
trust: TlsConfig,
}
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 ca_cert(mut self, path: impl Into<String>) -> Self {
self.trust.ca_cert_path = Some(path.into());
self
}
#[cfg(feature = "native-tls-roots")]
#[cfg_attr(docsrs, doc(cfg(feature = "native-tls-roots")))]
pub fn native_roots(mut self) -> Self {
self.trust.use_native_roots = true;
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::new(
Arc::new(super::tls::build_tls_config_sync(&self.trust)?),
self.request_timeout,
MAX_TOKEN_RESPONSE_BYTES,
),
})
}
}
#[derive(serde::Deserialize)]
struct TokenResponse {
#[serde(deserialize_with = "zeroizing_string")]
access_token: Zeroizing<String>,
#[serde(default)]
expires_in: Option<i64>,
}
fn zeroizing_string<'de, D: serde::Deserializer<'de>>(
deserializer: D,
) -> std::result::Result<Zeroizing<String>, D::Error> {
<String as serde::Deserialize>::deserialize(deserializer).map(Zeroizing::new)
}
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 encode_form(pairs: &[(&str, &str)]) -> Zeroizing<String> {
let len = pairs
.iter()
.map(|(k, v)| form_urlencoded_len(k) + 1 + form_urlencoded_len(v))
.sum::<usize>()
+ pairs.len().saturating_sub(1);
let mut form = Zeroizing::new(String::with_capacity(len));
for (i, (key, value)) in pairs.iter().enumerate() {
if i > 0 {
form.push('&');
}
form_urlencode_into(&mut form, key);
form.push('=');
form_urlencode_into(&mut form, value);
}
form
}
fn is_unreserved(byte: u8) -> bool {
matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~')
}
fn form_urlencoded_len(value: &str) -> usize {
value
.bytes()
.map(|b| if is_unreserved(b) || b == b' ' { 1 } else { 3 })
.sum()
}
fn form_urlencode_into(out: &mut String, value: &str) {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
for byte in value.bytes() {
if is_unreserved(byte) {
out.push(char::from(byte));
} else if byte == b' ' {
out.push('+');
} else {
out.push('%');
out.push(char::from(HEX[usize::from(byte >> 4)]));
out.push(char::from(HEX[usize::from(byte & 0x0F)]));
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
fn form_urlencode(value: &str) -> String {
let mut out = String::new();
form_urlencode_into(&mut out, value);
assert_eq!(
out.len(),
form_urlencoded_len(value),
"pre-sizing for {value:?}"
);
out
}
#[test]
fn the_form_is_allocated_at_its_final_length() {
let form = encode_form(&[("a b", "x&y"), ("client_assertion", "h.p.s"), ("k", "ü")]);
assert_eq!(&*form, "a+b=x%26y&client_assertion=h.p.s&k=%C3%BC");
assert_eq!(form.capacity(), form.len());
}
#[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::fixed(jwt)));
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::fixed(
"a.b.c",
)))
.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::provider(move || {
let counter = counter.clone();
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::fixed("a.b.c"));
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")
}
fn testdata(name: &str) -> String {
format!("{}/src/auth/testdata/{name}", env!("CARGO_MANIFEST_DIR"))
}
async fn token_endpoint(response: Vec<u8>) -> String {
use rustls::pki_types::pem::PemObject;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let certs = CertificateDer::pem_file_iter(testdata("server.pem"))
.unwrap()
.collect::<std::result::Result<Vec<_>, _>>()
.unwrap();
let key = PrivateKeyDer::from_pem_file(testdata("server.key")).unwrap();
let config = rustls::ServerConfig::builder_with_provider(
crate::auth::tls::resolve_crypto_provider(),
)
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_single_cert(certs, key)
.unwrap();
let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(config));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
let (tcp, _) = listener.accept().await.unwrap();
let Ok(mut tls) = acceptor.accept(tcp).await else {
return;
};
let mut request = Vec::new();
let mut buf = [0u8; 4096];
while !request.windows(4).any(|w| w == b"\r\n\r\n") {
match tls.read(&mut buf).await {
Ok(0) | Err(_) => return,
Ok(n) => request.extend_from_slice(&buf[..n]),
}
}
let _ = tls.write_all(&response).await;
let _ = tls.shutdown().await;
});
format!("https://127.0.0.1:{port}/token")
}
fn ok_response(body: &str) -> Vec<u8> {
format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{body}",
body.len()
)
.into_bytes()
}
#[tokio::test]
async fn an_endpoint_behind_a_private_ca_is_reachable_with_that_ca() {
let endpoint =
token_endpoint(ok_response(r#"{"access_token":"tok-1","expires_in":60}"#)).await;
let token = OidcTokenProvider::builder(&endpoint)
.credentials(ClientCredentials::secret("id", "secret"))
.ca_cert(testdata("ca.pem"))
.build()
.unwrap()
.fetch_token()
.await
.expect("the pinned CA must verify the endpoint");
let initial = token.to_gs2_initial_response();
assert!(
initial.windows(17).any(|w| w == b"auth=Bearer tok-1"),
"the fetched token must be the endpoint's"
);
}
#[tokio::test]
async fn without_its_ca_the_endpoint_fails_verification() {
let endpoint = token_endpoint(ok_response(r#"{"access_token":"tok-1"}"#)).await;
let err = OidcTokenProvider::builder(&endpoint)
.credentials(ClientCredentials::secret("id", "secret"))
.build()
.unwrap()
.fetch_token()
.await
.expect_err("an unknown issuer must be rejected")
.to_string();
assert!(err.contains(&endpoint), "got: {err}");
assert!(err.contains("UnknownIssuer"), "got: {err}");
}
#[test]
fn an_unreadable_ca_bundle_fails_the_build() {
let err = OidcTokenProvider::builder("https://idp.example.com/token")
.credentials(ClientCredentials::secret("id", "secret"))
.ca_cert("/nonexistent/krafka/ca.pem")
.build()
.expect_err("a missing CA bundle must not fall back to other roots")
.to_string();
assert!(err.contains("/nonexistent/krafka/ca.pem"), "got: {err}");
}
#[tokio::test]
async fn an_oversized_token_response_stops_at_the_token_cap() {
let mut response = b"HTTP/1.1 200 OK\r\nConnection: close\r\n\r\n".to_vec();
response.extend(std::iter::repeat_n(b' ', 2 * MAX_TOKEN_RESPONSE_BYTES));
let endpoint = token_endpoint(response).await;
let err = OidcTokenProvider::builder(&endpoint)
.credentials(ClientCredentials::secret("id", "secret"))
.ca_cert(testdata("ca.pem"))
.build()
.unwrap()
.fetch_token()
.await
.expect_err("a 2 MiB token response must be refused")
.to_string();
assert!(
err.contains(&format!("exceeds {MAX_TOKEN_RESPONSE_BYTES}-byte limit")),
"got: {err}"
);
}
}