use std::{fmt, fs, io::Cursor, sync::Arc};
use base64::{
Engine as _,
engine::general_purpose::{STANDARD, STANDARD_NO_PAD},
};
use hmac::{Hmac, Mac};
use pbkdf2::pbkdf2_hmac;
use sha2::{Digest, Sha256, Sha512};
use tokio_rustls::rustls::{
ClientConfig, RootCertStore,
crypto::aws_lc_rs,
pki_types::{CertificateDer, PrivateKeyDer},
};
use crate::{KafkaConfig, MqError, MqResult};
use super::{KafkaClientError, KafkaClientResult};
const MIN_SCRAM_ITERATIONS: u32 = 4_096;
const MAX_SCRAM_ITERATIONS: u32 = 1_000_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum SecurityProtocol {
Plaintext,
Ssl,
SaslPlaintext,
SaslSsl,
}
impl SecurityProtocol {
pub(crate) fn uses_tls(self) -> bool {
matches!(self, Self::Ssl | Self::SaslSsl)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum SaslMechanism {
Plain,
ScramSha256,
ScramSha512,
}
impl SaslMechanism {
pub(crate) fn name(self) -> &'static str {
match self {
Self::Plain => "PLAIN",
Self::ScramSha256 => "SCRAM-SHA-256",
Self::ScramSha512 => "SCRAM-SHA-512",
}
}
}
#[derive(Clone)]
pub(crate) struct SaslConfig {
pub(crate) mechanism: SaslMechanism,
pub(crate) username: String,
pub(crate) password: String,
}
impl fmt::Debug for SaslConfig {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("SaslConfig")
.field("mechanism", &self.mechanism)
.field("username", &self.username)
.field("password", &"[redacted]")
.finish()
}
}
#[derive(Debug, Clone)]
pub(crate) struct NativeSecurityConfig {
pub(crate) protocol: SecurityProtocol,
pub(crate) tls: Option<Arc<ClientConfig>>,
pub(crate) sasl: Option<SaslConfig>,
}
impl Default for NativeSecurityConfig {
fn default() -> Self {
Self {
protocol: SecurityProtocol::Plaintext,
tls: None,
sasl: None,
}
}
}
pub(crate) fn native_security_config(config: &KafkaConfig) -> MqResult<NativeSecurityConfig> {
let protocol = match config
.get("security.protocol")
.unwrap_or("plaintext")
.trim()
.to_ascii_lowercase()
.as_str()
{
"" | "plaintext" => SecurityProtocol::Plaintext,
"ssl" => SecurityProtocol::Ssl,
"sasl_plaintext" => SecurityProtocol::SaslPlaintext,
"sasl_ssl" => SecurityProtocol::SaslSsl,
other => return Err(unsupported(format!("security.protocol={other}"))),
};
validate_security_properties(config, protocol)?;
let tls = protocol
.uses_tls()
.then(|| build_tls_config(config))
.transpose()?;
let sasl = if matches!(
protocol,
SecurityProtocol::SaslPlaintext | SecurityProtocol::SaslSsl
) {
Some(parse_sasl_config(config)?)
} else {
None
};
Ok(NativeSecurityConfig {
protocol,
tls,
sasl,
})
}
fn validate_security_properties(config: &KafkaConfig, protocol: SecurityProtocol) -> MqResult<()> {
const TLS_KEYS: &[&str] = &[
"ssl.ca.location",
"ssl.ca.pem",
"ssl.certificate.location",
"ssl.certificate.pem",
"ssl.key.location",
"ssl.key.pem",
"ssl.endpoint.identification.algorithm",
];
const SASL_KEYS: &[&str] = &[
"sasl.mechanism",
"sasl.mechanisms",
"sasl.username",
"sasl.password",
];
for (key, value) in config.properties() {
let key = key.trim().to_ascii_lowercase();
if key.starts_with("ssl.") && !TLS_KEYS.contains(&key.as_str()) {
return Err(unsupported(format!("unsupported TLS property {key}")));
}
if key.starts_with("sasl.") && !SASL_KEYS.contains(&key.as_str()) {
return Err(unsupported(format!("unsupported SASL property {key}")));
}
if key == "ssl.endpoint.identification.algorithm"
&& !value.trim().eq_ignore_ascii_case("https")
{
return Err(unsupported(
"disabling or replacing rustls hostname verification is not supported",
));
}
}
if !protocol.uses_tls()
&& config
.properties()
.keys()
.any(|key| key.to_ascii_lowercase().starts_with("ssl."))
{
return Err(MqError::InvalidConfig(
"native Kafka TLS properties require security.protocol=ssl or sasl_ssl".to_owned(),
));
}
if matches!(
protocol,
SecurityProtocol::Plaintext | SecurityProtocol::Ssl
) && config
.properties()
.keys()
.any(|key| key.to_ascii_lowercase().starts_with("sasl."))
{
return Err(MqError::InvalidConfig(
"native Kafka SASL properties require security.protocol=sasl_plaintext or sasl_ssl"
.to_owned(),
));
}
Ok(())
}
fn parse_sasl_config(config: &KafkaConfig) -> MqResult<SaslConfig> {
let mechanism = config
.get("sasl.mechanism")
.or_else(|| config.get("sasl.mechanisms"))
.unwrap_or("PLAIN")
.trim()
.to_ascii_uppercase();
let mechanism = match mechanism.as_str() {
"PLAIN" => SaslMechanism::Plain,
"SCRAM-SHA-256" => SaslMechanism::ScramSha256,
"SCRAM-SHA-512" => SaslMechanism::ScramSha512,
other => {
return Err(unsupported(format!(
"SASL mechanism {other}; native supports PLAIN, SCRAM-SHA-256, and SCRAM-SHA-512"
)));
}
};
let username = required(config, "sasl.username")?;
let password = required(config, "sasl.password")?;
validate_scram_username(&username)
.map_err(|error| MqError::InvalidConfig(error.to_string()))?;
Ok(SaslConfig {
mechanism,
username,
password,
})
}
fn required(config: &KafkaConfig, key: &str) -> MqResult<String> {
config
.get(key)
.filter(|value| !value.is_empty())
.map(str::to_owned)
.ok_or_else(|| MqError::InvalidConfig(format!("native Kafka SASL requires {key}")))
}
fn build_tls_config(config: &KafkaConfig) -> MqResult<Arc<ClientConfig>> {
let native = rustls_native_certs::load_native_certs();
let mut roots = RootCertStore::empty();
roots.add_parsable_certificates(native.certs);
if let Some(pem) = pem_property(config, "ssl.ca.location", "ssl.ca.pem")? {
let certificates = parse_certificates(&pem, "Kafka CA")?;
if certificates.is_empty() {
return Err(MqError::InvalidConfig(
"native Kafka ssl.ca PEM contains no certificates".to_owned(),
));
}
for certificate in certificates {
roots.add(certificate).map_err(|error| {
MqError::InvalidConfig(format!("invalid native Kafka CA certificate: {error}"))
})?;
}
}
if roots.is_empty() {
let detail = native
.errors
.first()
.map(ToString::to_string)
.unwrap_or_else(|| "no system certificates were found".to_owned());
return Err(MqError::InvalidConfig(format!(
"native Kafka TLS root store is empty: {detail}"
)));
}
let provider = Arc::new(aws_lc_rs::default_provider());
let builder = ClientConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|error| MqError::InvalidConfig(format!("invalid rustls protocol set: {error}")))?
.with_root_certificates(roots);
let certificate = pem_property(config, "ssl.certificate.location", "ssl.certificate.pem")?;
let key = pem_property(config, "ssl.key.location", "ssl.key.pem")?;
let tls = match (certificate, key) {
(None, None) => builder.with_no_client_auth(),
(Some(certificate), Some(key)) => {
let certificates = parse_certificates(&certificate, "Kafka client certificate")?;
let key = parse_private_key(&key)?;
builder
.with_client_auth_cert(certificates, key)
.map_err(|error| {
MqError::InvalidConfig(format!(
"invalid native Kafka mTLS certificate/key: {error}"
))
})?
}
_ => {
return Err(MqError::InvalidConfig(
"native Kafka mTLS requires both ssl.certificate and ssl.key PEM".to_owned(),
));
}
};
Ok(Arc::new(tls))
}
fn pem_property(
config: &KafkaConfig,
path_key: &str,
inline_key: &str,
) -> MqResult<Option<Vec<u8>>> {
match (config.get(path_key), config.get(inline_key)) {
(Some(_), Some(_)) => Err(MqError::InvalidConfig(format!(
"native Kafka config must set only one of {path_key} and {inline_key}"
))),
(Some(path), None) => fs::read(path).map(Some).map_err(|error| {
MqError::InvalidConfig(format!(
"failed to read native Kafka {path_key}={path}: {error}"
))
}),
(None, Some(pem)) => Ok(Some(pem.as_bytes().to_vec())),
(None, None) => Ok(None),
}
}
fn parse_certificates(pem: &[u8], what: &str) -> MqResult<Vec<CertificateDer<'static>>> {
rustls_pemfile::certs(&mut Cursor::new(pem))
.collect::<Result<Vec<_>, _>>()
.map_err(|error| MqError::InvalidConfig(format!("invalid {what} PEM: {error}")))
}
fn parse_private_key(pem: &[u8]) -> MqResult<PrivateKeyDer<'static>> {
rustls_pemfile::private_key(&mut Cursor::new(pem))
.map_err(|error| MqError::InvalidConfig(format!("invalid Kafka client key PEM: {error}")))?
.ok_or_else(|| {
MqError::InvalidConfig("Kafka client key PEM contains no private key".to_owned())
})
}
fn unsupported(message: impl fmt::Display) -> MqError {
MqError::InvalidConfig(format!(
"native Kafka security boundary: {message}; {} or {}",
crate::consumer_rdkafka_hint(),
crate::producer_rdkafka_hint()
))
}
pub(crate) struct ScramClient {
mechanism: SaslMechanism,
password: String,
nonce: String,
client_first_bare: String,
}
impl ScramClient {
pub(crate) fn new(config: &SaslConfig) -> KafkaClientResult<Self> {
let mut nonce = [0_u8; 24];
aws_lc_rs::default_provider()
.secure_random
.fill(&mut nonce)
.map_err(|_| KafkaClientError::security("SCRAM nonce generation failed"))?;
Self::with_nonce(config, STANDARD_NO_PAD.encode(nonce))
}
fn with_nonce(config: &SaslConfig, nonce: String) -> KafkaClientResult<Self> {
validate_scram_username(&config.username)?;
if nonce.is_empty() || nonce.contains(',') || !nonce.is_ascii() {
return Err(KafkaClientError::security("invalid SCRAM client nonce"));
}
let username = config.username.replace('=', "=3D").replace(',', "=2C");
let client_first_bare = format!("n={username},r={nonce}");
Ok(Self {
mechanism: config.mechanism,
password: config.password.clone(),
nonce,
client_first_bare,
})
}
pub(crate) fn first_message(&self) -> String {
format!("n,,{}", self.client_first_bare)
}
pub(crate) fn handle_server_first(
&self,
server_first: &str,
) -> KafkaClientResult<(String, ScramServerVerifier)> {
let attributes = parse_attributes(server_first)?;
if attributes.contains_key(&'m') {
return Err(KafkaClientError::security(
"SCRAM server sent an unsupported mandatory extension",
));
}
let nonce = required_attribute(&attributes, 'r')?;
if !nonce.starts_with(&self.nonce) || nonce.len() <= self.nonce.len() {
return Err(KafkaClientError::security(
"SCRAM server nonce does not extend the client nonce",
));
}
let salt = STANDARD
.decode(required_attribute(&attributes, 's')?)
.map_err(|_| KafkaClientError::security("SCRAM server salt is not valid base64"))?;
let iterations = required_attribute(&attributes, 'i')?
.parse::<u32>()
.map_err(|_| {
KafkaClientError::security("SCRAM iteration count is not a positive integer")
})?;
if !(MIN_SCRAM_ITERATIONS..=MAX_SCRAM_ITERATIONS).contains(&iterations) {
return Err(KafkaClientError::security(format!(
"SCRAM iteration count {iterations} is outside the safe range {MIN_SCRAM_ITERATIONS}..={MAX_SCRAM_ITERATIONS}"
)));
}
let final_without_proof = format!("c=biws,r={nonce}");
let auth_message = format!(
"{},{server_first},{final_without_proof}",
self.client_first_bare
);
match self.mechanism {
SaslMechanism::ScramSha256 => {
let mut salted = [0_u8; 32];
pbkdf2_hmac::<Sha256>(self.password.as_bytes(), &salt, iterations, &mut salted);
let client_key = hmac_sha256(&salted, b"Client Key")?;
let stored_key = Sha256::digest(&client_key);
let client_signature = hmac_sha256(&stored_key, auth_message.as_bytes())?;
let proof = xor(&client_key, &client_signature);
let mut verifier =
Hmac::<Sha256>::new_from_slice(&hmac_sha256(&salted, b"Server Key")?).map_err(
|_| KafkaClientError::security("invalid SCRAM-SHA-256 HMAC key"),
)?;
verifier.update(auth_message.as_bytes());
Ok((
format!("{final_without_proof},p={}", STANDARD.encode(proof)),
ScramServerVerifier::Sha256(verifier),
))
}
SaslMechanism::ScramSha512 => {
let mut salted = [0_u8; 64];
pbkdf2_hmac::<Sha512>(self.password.as_bytes(), &salt, iterations, &mut salted);
let client_key = hmac_sha512(&salted, b"Client Key")?;
let stored_key = Sha512::digest(&client_key);
let client_signature = hmac_sha512(&stored_key, auth_message.as_bytes())?;
let proof = xor(&client_key, &client_signature);
let mut verifier =
Hmac::<Sha512>::new_from_slice(&hmac_sha512(&salted, b"Server Key")?).map_err(
|_| KafkaClientError::security("invalid SCRAM-SHA-512 HMAC key"),
)?;
verifier.update(auth_message.as_bytes());
Ok((
format!("{final_without_proof},p={}", STANDARD.encode(proof)),
ScramServerVerifier::Sha512(verifier),
))
}
SaslMechanism::Plain => Err(KafkaClientError::security(
"PLAIN cannot enter the SCRAM exchange",
)),
}
}
}
pub(crate) enum ScramServerVerifier {
Sha256(Hmac<Sha256>),
Sha512(Hmac<Sha512>),
}
impl ScramServerVerifier {
pub(crate) fn verify(self, server_final: &str) -> KafkaClientResult<()> {
let attributes = parse_attributes(server_final)?;
if let Some(error) = attributes.get(&'e') {
return Err(KafkaClientError::security(format!(
"SCRAM server rejected authentication: {error}"
)));
}
let signature = STANDARD
.decode(required_attribute(&attributes, 'v')?)
.map_err(|_| {
KafkaClientError::security("SCRAM server signature is not valid base64")
})?;
let verified = match self {
Self::Sha256(verifier) => verifier.verify_slice(&signature),
Self::Sha512(verifier) => verifier.verify_slice(&signature),
};
verified
.map_err(|_| KafkaClientError::security("SCRAM server signature verification failed"))
}
}
fn hmac_sha256(key: &[u8], data: &[u8]) -> KafkaClientResult<Vec<u8>> {
let mut mac = Hmac::<Sha256>::new_from_slice(key)
.map_err(|_| KafkaClientError::security("invalid SCRAM-SHA-256 HMAC key"))?;
mac.update(data);
Ok(mac.finalize().into_bytes().to_vec())
}
fn hmac_sha512(key: &[u8], data: &[u8]) -> KafkaClientResult<Vec<u8>> {
let mut mac = Hmac::<Sha512>::new_from_slice(key)
.map_err(|_| KafkaClientError::security("invalid SCRAM-SHA-512 HMAC key"))?;
mac.update(data);
Ok(mac.finalize().into_bytes().to_vec())
}
fn xor(left: &[u8], right: &[u8]) -> Vec<u8> {
left.iter()
.zip(right)
.map(|(left, right)| left ^ right)
.collect()
}
fn validate_scram_username(username: &str) -> KafkaClientResult<()> {
if username.is_empty() || username.chars().any(char::is_control) {
return Err(KafkaClientError::security(
"SASL username must be non-empty and contain no control characters",
));
}
Ok(())
}
fn parse_attributes(input: &str) -> KafkaClientResult<std::collections::BTreeMap<char, &str>> {
let mut attributes = std::collections::BTreeMap::new();
for attribute in input.split(',') {
let (name, value) = attribute
.split_once('=')
.ok_or_else(|| KafkaClientError::security("malformed SCRAM attribute"))?;
let mut chars = name.chars();
let name = chars
.next()
.filter(|_| chars.next().is_none())
.ok_or_else(|| KafkaClientError::security("invalid SCRAM attribute name"))?;
if attributes.insert(name, value).is_some() {
return Err(KafkaClientError::security(format!(
"duplicate SCRAM attribute {name}"
)));
}
}
Ok(attributes)
}
fn required_attribute<'a>(
attributes: &'a std::collections::BTreeMap<char, &'a str>,
name: char,
) -> KafkaClientResult<&'a str> {
attributes
.get(&name)
.copied()
.filter(|value| !value.is_empty())
.ok_or_else(|| KafkaClientError::security(format!("SCRAM message is missing {name}=")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rfc_7677_scram_sha_256_vector() {
let config = SaslConfig {
mechanism: SaslMechanism::ScramSha256,
username: "user".to_owned(),
password: "pencil".to_owned(),
};
let client = ScramClient::with_nonce(&config, "rOprNGfwEbeRWgbNEkqO".to_owned())
.expect("RFC client");
assert_eq!(client.first_message(), "n,,n=user,r=rOprNGfwEbeRWgbNEkqO");
let server_first = "r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,s=W22ZaJ0SNY7soEsUEjb6gQ==,i=4096";
let (client_final, verifier) = client
.handle_server_first(server_first)
.expect("RFC server first");
assert_eq!(
client_final,
"c=biws,r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="
);
verifier
.verify("v=6rriTRBi23WpRR/wtup+mMhUZUn/dB5nLTJRsjl95G4=")
.expect("RFC server signature");
}
#[test]
fn scram_rejects_nonce_substitution_and_bad_server_signature() {
let config = SaslConfig {
mechanism: SaslMechanism::ScramSha512,
username: "user".to_owned(),
password: "pencil".to_owned(),
};
let client = ScramClient::with_nonce(&config, "clientnonce".to_owned()).expect("client");
assert!(
client
.handle_server_first("r=othernonce,s=W22ZaJ0SNY7soEsUEjb6gQ==,i=4096")
.is_err()
);
let (_, verifier) = client
.handle_server_first("r=clientnonce-server,s=W22ZaJ0SNY7soEsUEjb6gQ==,i=4096")
.expect("server first");
assert!(verifier.verify("v=AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA==").is_err());
}
#[test]
fn scram_escapes_username_and_rejects_unsafe_iteration_counts() {
let config = SaslConfig {
mechanism: SaslMechanism::ScramSha256,
username: "a,b=c".to_owned(),
password: "secret".to_owned(),
};
let client = ScramClient::with_nonce(&config, "nonce".to_owned()).expect("client");
assert_eq!(client.first_message(), "n,,n=a=2Cb=3Dc,r=nonce");
assert!(
client
.handle_server_first("r=nonce-server,s=c2FsdA==,i=1")
.is_err()
);
assert!(
client
.handle_server_first("r=nonce-server,s=c2FsdA==,i=1000001")
.is_err()
);
assert!(
client
.handle_server_first("r=nonce-server,r=duplicate,s=c2FsdA==,i=4096")
.is_err()
);
}
}