use sha1::Sha1;
use sha2::{Digest, Sha256};
use rand::rand_core::UnwrapErr;
use rand::rngs::SysRng;
#[cfg(feature = "rsa-auth")]
use rsa::RsaPublicKey;
#[cfg(feature = "rsa-auth")]
use rsa::pkcs1::DecodeRsaPublicKey;
#[cfg(feature = "rsa-auth")]
use rsa::pkcs8::DecodePublicKey;
pub mod plugins {
pub const MYSQL_NATIVE_PASSWORD: &str = "mysql_native_password";
pub const CACHING_SHA2_PASSWORD: &str = "caching_sha2_password";
pub const SHA256_PASSWORD: &str = "sha256_password";
pub const MYSQL_CLEAR_PASSWORD: &str = "mysql_clear_password";
}
pub mod caching_sha2 {
pub const AUTH_MORE_DATA: u8 = 0x01;
pub const REQUEST_PUBLIC_KEY: u8 = 0x02;
pub const FAST_AUTH_SUCCESS: u8 = 0x03;
pub const PERFORM_FULL_AUTH: u8 = 0x04;
}
#[must_use]
pub fn strip_auth_more_data_marker(payload: &[u8]) -> &[u8] {
match payload {
[caching_sha2::AUTH_MORE_DATA, rest @ ..] => rest,
_ => payload,
}
}
pub fn mysql_native_password(password: &str, auth_data: &[u8]) -> Vec<u8> {
if password.is_empty() {
return vec![];
}
let seed = if auth_data.len() > 20 {
&auth_data[..20]
} else {
auth_data
};
let mut hasher = Sha1::new();
hasher.update(password.as_bytes());
let stage1: [u8; 20] = hasher.finalize().into();
let mut hasher = Sha1::new();
hasher.update(stage1);
let stage2: [u8; 20] = hasher.finalize().into();
let mut hasher = Sha1::new();
hasher.update(seed);
hasher.update(stage2);
let stage3: [u8; 20] = hasher.finalize().into();
stage1
.iter()
.zip(stage3.iter())
.map(|(a, b)| a ^ b)
.collect()
}
pub fn caching_sha2_password(password: &str, auth_data: &[u8]) -> Vec<u8> {
if password.is_empty() {
return vec![];
}
let seed = if auth_data.len() == 21 && auth_data.last() == Some(&0) {
&auth_data[..20]
} else {
auth_data
};
let mut hasher = Sha256::new();
hasher.update(password.as_bytes());
let password_hash: [u8; 32] = hasher.finalize().into();
let mut hasher = Sha256::new();
hasher.update(password_hash);
let password_hash_hash: [u8; 32] = hasher.finalize().into();
let mut hasher = Sha256::new();
hasher.update(password_hash_hash);
hasher.update(seed);
let scramble: [u8; 32] = hasher.finalize().into();
password_hash
.iter()
.zip(scramble.iter())
.map(|(a, b)| a ^ b)
.collect()
}
pub fn generate_nonce(length: usize) -> Vec<u8> {
use rand::Rng;
let mut bytes = vec![0u8; length];
UnwrapErr(SysRng).fill_bytes(&mut bytes);
bytes
}
#[cfg(feature = "rsa-auth")]
pub fn sha256_password_rsa(
password: &str,
seed: &[u8],
public_key_pem: &[u8],
use_oaep: bool,
) -> Result<Vec<u8>, String> {
let mut pw = password.as_bytes().to_vec();
pw.push(0);
if seed.is_empty() {
return Err("Seed is empty".to_string());
}
for (i, b) in pw.iter_mut().enumerate() {
*b ^= seed[i % seed.len()];
}
let pem = std::str::from_utf8(public_key_pem)
.map_err(|e| format!("Public key is not valid UTF-8 PEM: {e}"))?;
let pub_key = RsaPublicKey::from_public_key_pem(pem)
.or_else(|_| RsaPublicKey::from_pkcs1_pem(pem))
.map_err(|e| format!("Failed to parse RSA public key PEM: {e}"))?;
let encrypted = if use_oaep {
let padding = rsa::Oaep::<Sha1>::new();
pub_key
.encrypt(&mut UnwrapErr(SysRng), padding, &pw)
.map_err(|e| format!("RSA OAEP encryption failed: {e}"))?
} else {
let padding = rsa::Pkcs1v15Encrypt;
pub_key
.encrypt(&mut UnwrapErr(SysRng), padding, &pw)
.map_err(|e| format!("RSA PKCS1v1.5 encryption failed: {e}"))?
};
Ok(encrypted)
}
#[cfg(not(feature = "rsa-auth"))]
pub fn sha256_password_rsa(
_password: &str,
_seed: &[u8],
_public_key_pem: &[u8],
_use_oaep: bool,
) -> Result<Vec<u8>, String> {
Err(
"MySQL full authentication requires TLS (SslMode::Required or stronger) \
or the `rsa-auth` feature; the server requested RSA password exchange \
on a plaintext connection"
.to_owned(),
)
}
pub fn xor_password_with_seed(password: &str, seed: &[u8]) -> Vec<u8> {
let password_bytes = password.as_bytes();
let mut result = Vec::with_capacity(password_bytes.len() + 1);
for (i, &byte) in password_bytes.iter().enumerate() {
let seed_byte = seed.get(i % seed.len()).copied().unwrap_or(0);
result.push(byte ^ seed_byte);
}
result.push(0);
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn auth_more_data_marker_is_stripped_from_live_server_frames() {
assert_eq!(
strip_auth_more_data_marker(&[
caching_sha2::AUTH_MORE_DATA,
caching_sha2::FAST_AUTH_SUCCESS
]),
&[caching_sha2::FAST_AUTH_SUCCESS]
);
assert_eq!(
strip_auth_more_data_marker(&[
caching_sha2::AUTH_MORE_DATA,
caching_sha2::PERFORM_FULL_AUTH
]),
&[caching_sha2::PERFORM_FULL_AUTH]
);
let mut frame = vec![caching_sha2::AUTH_MORE_DATA];
frame.extend_from_slice(b"-----BEGIN PUBLIC KEY-----\n");
assert_eq!(
strip_auth_more_data_marker(&frame),
b"-----BEGIN PUBLIC KEY-----\n"
);
}
#[test]
fn auth_more_data_marker_stripping_leaves_other_frames_alone() {
assert_eq!(
strip_auth_more_data_marker(&[caching_sha2::FAST_AUTH_SUCCESS]),
&[caching_sha2::FAST_AUTH_SUCCESS]
);
assert_eq!(
strip_auth_more_data_marker(&[0x00, 0x00, 0x00]),
&[0x00, 0x00, 0x00]
);
assert_eq!(strip_auth_more_data_marker(&[0xFE, 0x00]), &[0xFE, 0x00]);
assert!(strip_auth_more_data_marker(&[caching_sha2::AUTH_MORE_DATA]).is_empty());
assert!(strip_auth_more_data_marker(&[]).is_empty());
}
#[test]
fn test_mysql_native_password_empty() {
let result = mysql_native_password("", &[0; 20]);
assert!(result.is_empty());
}
#[test]
fn test_mysql_native_password() {
let seed = [0u8; 20];
let result = mysql_native_password("secret", &seed);
assert_eq!(result.len(), 20);
let result2 = mysql_native_password("secret", &seed);
assert_eq!(result, result2);
}
#[test]
fn test_mysql_native_password_real_seed() {
let seed = [
0x3d, 0x4c, 0x5e, 0x2f, 0x1a, 0x0b, 0x7c, 0x8d, 0x9e, 0xaf, 0x10, 0x21, 0x32, 0x43,
0x54, 0x65, 0x76, 0x87, 0x98, 0xa9,
];
let result = mysql_native_password("mypassword", &seed);
assert_eq!(result.len(), 20);
let result2 = mysql_native_password("otherpassword", &seed);
assert_ne!(result, result2);
}
#[test]
fn test_caching_sha2_password_empty() {
let result = caching_sha2_password("", &[0; 20]);
assert!(result.is_empty());
}
#[test]
fn test_caching_sha2_password() {
let seed = [0u8; 20];
let result = caching_sha2_password("secret", &seed);
assert_eq!(result.len(), 32);
let result2 = caching_sha2_password("secret", &seed);
assert_eq!(result, result2);
}
#[test]
fn test_caching_sha2_password_with_nul() {
let mut seed = vec![0u8; 20];
seed.push(0);
let result = caching_sha2_password("secret", &seed);
assert_eq!(result.len(), 32);
let result2 = caching_sha2_password("secret", &seed[..20]);
assert_eq!(result, result2);
}
#[test]
fn test_generate_nonce() {
let nonce1 = generate_nonce(20);
let nonce2 = generate_nonce(20);
assert_eq!(nonce1.len(), 20);
assert_eq!(nonce2.len(), 20);
assert_ne!(nonce1, nonce2);
}
#[test]
fn test_xor_password_with_seed() {
let password = "test";
let seed = [1, 2, 3, 4, 5, 6, 7, 8];
let result = xor_password_with_seed(password, &seed);
assert_eq!(result.len(), 5);
assert_eq!(result[4], 0);
let recovered: Vec<u8> = result[..4]
.iter()
.enumerate()
.map(|(i, &b)| b ^ seed[i % seed.len()])
.collect();
assert_eq!(recovered, password.as_bytes());
}
#[cfg(feature = "rsa-auth")]
const SPKI_PUBLIC_KEY: &str = "-----BEGIN PUBLIC KEY-----\n\
MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEArJ59U1RtRan3BEJqsItb\n\
tD3nDU7ThwlNaJ42vRSEt/UzFO5/yemxUz3ogTUcsDMgXeLVjQjwwV0+lh+s9IWc\n\
9fU6nu4Q8yW7Pc/SDpDkFBdEtLAOIjfSMv0CzoPQB0A0njVFe7l7SuyrMWQ/19N5\n\
iEtqZQmP2y7h5a23XgPGogXHm0XnKpueZn9KFXhK2lNhZj9IUuQLmhzrH0pov8Ae\n\
FknazXZxL5aoAG+cJIoHKsf9NcGTzj0Hewb36YlgVi+yZ5NRkhgjklQ8E6IL+aaW\n\
yEOsCBtS8kCd/nHP4t6ZeyExdpggPGQ2nJq18jG+sttI+3AnzQhXtR2Adq9LKVdN\n\
RQIDAQAB\n\
-----END PUBLIC KEY-----\n";
#[cfg(feature = "rsa-auth")]
const PKCS1_PUBLIC_KEY: &str = "-----BEGIN RSA PUBLIC KEY-----\n\
MIIBCgKCAQEArJ59U1RtRan3BEJqsItbtD3nDU7ThwlNaJ42vRSEt/UzFO5/yemx\n\
Uz3ogTUcsDMgXeLVjQjwwV0+lh+s9IWc9fU6nu4Q8yW7Pc/SDpDkFBdEtLAOIjfS\n\
Mv0CzoPQB0A0njVFe7l7SuyrMWQ/19N5iEtqZQmP2y7h5a23XgPGogXHm0XnKpue\n\
Zn9KFXhK2lNhZj9IUuQLmhzrH0pov8AeFknazXZxL5aoAG+cJIoHKsf9NcGTzj0H\n\
ewb36YlgVi+yZ5NRkhgjklQ8E6IL+aaWyEOsCBtS8kCd/nHP4t6ZeyExdpggPGQ2\n\
nJq18jG+sttI+3AnzQhXtR2Adq9LKVdNRQIDAQAB\n\
-----END RSA PUBLIC KEY-----\n";
#[cfg(feature = "rsa-auth")]
#[test]
fn test_sha256_password_rsa_accepts_both_pem_encodings() {
let seed = [
0x3d, 0x4c, 0x5e, 0x2f, 0x1a, 0x0b, 0x7c, 0x8d, 0x9e, 0xaf, 0x10, 0x21, 0x32, 0x43,
0x54, 0x65, 0x76, 0x87, 0x98, 0xa9,
];
for pem in [SPKI_PUBLIC_KEY, PKCS1_PUBLIC_KEY] {
for use_oaep in [true, false] {
let out = sha256_password_rsa("hunter2", &seed, pem.as_bytes(), use_oaep)
.expect("RSA encryption with a 2048-bit MySQL-style key");
assert_eq!(out.len(), 256, "pem={pem} oaep={use_oaep}");
let again = sha256_password_rsa("hunter2", &seed, pem.as_bytes(), use_oaep)
.expect("second encryption");
assert_ne!(out, again, "padding must be randomized");
}
}
}
#[cfg(feature = "rsa-auth")]
#[test]
fn test_sha256_password_rsa_rejects_bad_input() {
assert!(sha256_password_rsa("pw", &[], SPKI_PUBLIC_KEY.as_bytes(), true).is_err());
assert!(sha256_password_rsa("pw", &[1, 2, 3], b"not a pem", true).is_err());
}
#[cfg(not(feature = "rsa-auth"))]
#[test]
fn test_sha256_password_rsa_requires_tls_or_feature() {
let error = sha256_password_rsa("pw", &[1, 2, 3], b"unused", true)
.expect_err("no-TLS full auth is unavailable without the rsa-auth feature");
assert!(
error.contains("requires TLS")
&& error.contains("rsa-auth")
&& error.contains("RSA password exchange"),
"error must name the remediation: {error}"
);
}
#[test]
fn test_plugin_names() {
assert_eq!(plugins::MYSQL_NATIVE_PASSWORD, "mysql_native_password");
assert_eq!(plugins::CACHING_SHA2_PASSWORD, "caching_sha2_password");
assert_eq!(plugins::SHA256_PASSWORD, "sha256_password");
}
#[test]
fn rsa_is_used_for_public_key_encryption_only() {
let sources = [
include_str!("auth.rs"),
include_str!("connection.rs"),
include_str!("async_connection.rs"),
include_str!("tls.rs"),
];
let forbidden_tokens = [
concat!("RsaPriv", "ateKey"),
concat!(".decr", "ypt("),
concat!("Signing", "Key"),
concat!(".si", "gn("),
concat!("DecodePriv", "ateKey"),
];
for (name, src) in ["auth.rs", "connection.rs", "async_connection.rs", "tls.rs"]
.iter()
.zip(sources)
{
for forbidden in forbidden_tokens {
let hit = src
.lines()
.filter(|l| !l.trim_start().starts_with("//"))
.any(|l| l.contains(forbidden));
assert!(
!hit,
"{name} references `{forbidden}`: the RUSTSEC-2023-0071 audit ignore assumes \
rsa is used only for public-key encryption; update .cargo/audit.toml"
);
}
}
}
}