#[cfg(test)]
use md5::{Digest, Md5};
const MD5_A0: u32 = 0x6745_2301;
const MD5_B0: u32 = 0xEFCD_AB89;
const MD5_C0: u32 = 0x98BA_DCFE;
const MD5_D0: u32 = 0x1032_5476;
const S: [u32; 64] = [
7, 12, 17, 22, 7, 12, 17, 22, 7, 12, 17, 22, 7, 12, 17, 22, 5, 9, 14, 20, 5, 9, 14, 20, 5, 9, 14, 20, 5, 9, 14, 20, 4, 11, 16, 23, 4, 11, 16, 23, 4, 11, 16, 23, 4, 11, 16, 23, 6, 10, 15, 21, 6, 10, 15, 21, 6, 10, 15, 21, 6, 10, 15, 21, ];
const K: [u32; 64] = [
0xd76aa478, 0xe8c7b756, 0x242070db, 0xc1bdceee, 0xf57c0faf, 0x4787c62a, 0xa8304613, 0xfd469501,
0x698098d8, 0x8b44f7af, 0xffff5bb1, 0x895cd7be, 0x6b901122, 0xfd987193, 0xa679438e, 0x49b40821,
0xf61e2562, 0xc040b340, 0x265e5a51, 0xe9b6c7aa, 0xd62f105d, 0x02441453, 0xd8a1e681, 0xe7d3fbc8,
0x21e1cde6, 0xc33707d6, 0xf4d50d87, 0x455a14ed, 0xa9e3e905, 0xfcefa3f8, 0x676f02d9, 0x8d2a4c8a,
0xfffa3942, 0x8771f681, 0x6d9d6122, 0xfde5380c, 0xa4beea44, 0x4bdecfa9, 0xf6bb4b60, 0xbebfbc70,
0x289b7ec6, 0xeaa127fa, 0xd4ef3085, 0x04881d05, 0xd9d4d039, 0xe6db99e5, 0x1fa27cf8, 0xc4ac5665,
0xf4292244, 0x432aff97, 0xab9423a7, 0xfc93a039, 0x655b59c3, 0x8f0ccc92, 0xffeff47d, 0x85845dd1,
0x6fa87e4f, 0xfe2ce6e0, 0xa3014314, 0x4e0811a1, 0xf7537e82, 0xbd3af235, 0x2ad7d2bb, 0xeb86d391,
];
const G: [usize; 64] = [
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 1, 6, 11, 0, 5, 10, 15, 4, 9, 14, 3, 8, 13, 2, 7, 12, 5, 8, 11, 14, 1, 4, 7, 10, 13, 0, 3, 6, 9, 12, 15, 2, 0, 7, 14, 5, 12, 3, 10, 1, 8, 15, 6, 13, 4, 11, 2, 9, ];
pub fn md5_multi(inputs: &[&[u8]], pad_to: Option<u64>) -> Vec<[u8; 16]> {
assert!(!inputs.is_empty() && inputs.len() <= 4);
#[cfg(target_arch = "aarch64")]
{
return unsafe { md5_multi_neon(inputs, pad_to) };
}
#[cfg(target_arch = "x86_64")]
{
return unsafe { md5_multi_x86(inputs, pad_to) };
}
#[cfg(target_arch = "x86")]
{
if std::is_x86_feature_detected!("sse2") {
return unsafe { md5_multi_x86(inputs, pad_to) };
}
}
#[allow(unreachable_code)]
md5_multi_scalar(inputs, pad_to)
}
fn md5_multi_scalar(inputs: &[&[u8]], pad_to: Option<u64>) -> Vec<[u8; 16]> {
inputs
.iter()
.map(|inp| {
let effective_len = match pad_to {
Some(p) if p > inp.len() as u64 => p,
_ => inp.len() as u64,
};
md5_single_scalar(inp, effective_len)
})
.collect()
}
fn md5_single_scalar(data: &[u8], effective_len: u64) -> [u8; 16] {
let padded = md5_pad(data, effective_len);
let num_blocks = padded.len() / 64;
let mut a = MD5_A0;
let mut b = MD5_B0;
let mut c = MD5_C0;
let mut d = MD5_D0;
for block_idx in 0..num_blocks {
let block = &padded[block_idx * 64..(block_idx + 1) * 64];
let mut m = [0u32; 16];
for w in 0..16 {
m[w] = u32::from_le_bytes(block[w * 4..w * 4 + 4].try_into().unwrap());
}
let (oa, ob, oc, od) = (a, b, c, d);
for r in 0..64 {
let f = match r {
0..16 => (b & c) | (!b & d),
16..32 => (d & b) | (!d & c),
32..48 => b ^ c ^ d,
_ => c ^ (b | !d),
};
let tmp = f.wrapping_add(a).wrapping_add(K[r]).wrapping_add(m[G[r]]);
a = d;
d = c;
c = b;
b = b.wrapping_add(tmp.rotate_left(S[r]));
}
a = a.wrapping_add(oa);
b = b.wrapping_add(ob);
c = c.wrapping_add(oc);
d = d.wrapping_add(od);
}
let mut digest = [0u8; 16];
digest[0..4].copy_from_slice(&a.to_le_bytes());
digest[4..8].copy_from_slice(&b.to_le_bytes());
digest[8..12].copy_from_slice(&c.to_le_bytes());
digest[12..16].copy_from_slice(&d.to_le_bytes());
digest
}
fn md5_pad(data: &[u8], effective_len: u64) -> Vec<u8> {
let bit_len = effective_len * 8;
let mut padded = Vec::with_capacity((effective_len as usize + 72) & !63);
padded.extend_from_slice(data);
if (data.len() as u64) < effective_len {
padded.resize(effective_len as usize, 0);
}
padded.push(0x80);
while padded.len() % 64 != 56 {
padded.push(0);
}
padded.extend_from_slice(&bit_len.to_le_bytes());
padded
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "sse2")]
unsafe fn md5_multi_x86(inputs: &[&[u8]], pad_to: Option<u64>) -> Vec<[u8; 16]> {
#[cfg(target_arch = "x86")]
use std::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
let n = inputs.len();
let effective_lens: Vec<u64> = inputs
.iter()
.map(|inp| {
let raw = inp.len() as u64;
match pad_to {
Some(p) if p > raw => p,
_ => raw,
}
})
.collect();
let padded: Vec<Vec<u8>> = (0..n)
.map(|i| md5_pad(inputs[i], effective_lens[i]))
.collect();
let block_counts: Vec<usize> = padded.iter().map(|p| p.len() / 64).collect();
let uniform_blocks = block_counts.iter().all(|&c| c == block_counts[0]);
if !uniform_blocks {
return md5_multi_scalar(inputs, pad_to);
}
let num_blocks = block_counts[0];
let lane_ptrs: [&[u8]; 4] = [
&padded[0],
if n > 1 { &padded[1] } else { &padded[0] },
if n > 2 { &padded[2] } else { &padded[0] },
if n > 3 { &padded[3] } else { &padded[0] },
];
#[inline(always)]
unsafe fn sse2_rotl(v: __m128i, amount: u32) -> __m128i {
#[cfg(target_arch = "x86")]
use std::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
let left = _mm_sll_epi32(v, _mm_cvtsi32_si128(amount as i32));
let right = _mm_srl_epi32(v, _mm_cvtsi32_si128((32 - amount) as i32));
_mm_or_si128(left, right)
}
}
unsafe {
let mut a = _mm_set1_epi32(MD5_A0 as i32);
let mut b = _mm_set1_epi32(MD5_B0 as i32);
let mut c = _mm_set1_epi32(MD5_C0 as i32);
let mut d = _mm_set1_epi32(MD5_D0 as i32);
let all_ones = _mm_set1_epi32(-1);
for block_idx in 0..num_blocks {
let mut m = [_mm_setzero_si128(); 16];
let off = block_idx * 64;
for w in 0..4 {
let idx = w * 4;
let w_off = idx * 4;
let in0 = _mm_loadu_si128(lane_ptrs[0].as_ptr().add(off + w_off) as *const __m128i);
let in1 = _mm_loadu_si128(lane_ptrs[1].as_ptr().add(off + w_off) as *const __m128i);
let in2 = _mm_loadu_si128(lane_ptrs[2].as_ptr().add(off + w_off) as *const __m128i);
let in3 = _mm_loadu_si128(lane_ptrs[3].as_ptr().add(off + w_off) as *const __m128i);
let z01_lo = _mm_unpacklo_epi32(in0, in1);
let z01_hi = _mm_unpackhi_epi32(in0, in1);
let z23_lo = _mm_unpacklo_epi32(in2, in3);
let z23_hi = _mm_unpackhi_epi32(in2, in3);
m[idx] = _mm_unpacklo_epi64(z01_lo, z23_lo);
m[idx + 1] = _mm_unpackhi_epi64(z01_lo, z23_lo);
m[idx + 2] = _mm_unpacklo_epi64(z01_hi, z23_hi);
m[idx + 3] = _mm_unpackhi_epi64(z01_hi, z23_hi);
}
let oa = a;
let ob = b;
let oc = c;
let od = d;
for r in 0..16 {
let f = _mm_or_si128(_mm_and_si128(b, c), _mm_andnot_si128(b, d));
let tmp = _mm_add_epi32(
_mm_add_epi32(a, f),
_mm_add_epi32(_mm_set1_epi32(K[r] as i32), m[G[r]]),
);
a = d;
d = c;
c = b;
b = _mm_add_epi32(b, sse2_rotl(tmp, S[r]));
}
for r in 16..32 {
let f = _mm_or_si128(_mm_and_si128(d, b), _mm_andnot_si128(d, c));
let tmp = _mm_add_epi32(
_mm_add_epi32(a, f),
_mm_add_epi32(_mm_set1_epi32(K[r] as i32), m[G[r]]),
);
a = d;
d = c;
c = b;
b = _mm_add_epi32(b, sse2_rotl(tmp, S[r]));
}
for r in 32..48 {
let f = _mm_xor_si128(_mm_xor_si128(b, c), d);
let tmp = _mm_add_epi32(
_mm_add_epi32(a, f),
_mm_add_epi32(_mm_set1_epi32(K[r] as i32), m[G[r]]),
);
a = d;
d = c;
c = b;
b = _mm_add_epi32(b, sse2_rotl(tmp, S[r]));
}
for r in 48..64 {
let f = _mm_xor_si128(c, _mm_or_si128(b, _mm_andnot_si128(d, all_ones)));
let tmp = _mm_add_epi32(
_mm_add_epi32(a, f),
_mm_add_epi32(_mm_set1_epi32(K[r] as i32), m[G[r]]),
);
a = d;
d = c;
c = b;
b = _mm_add_epi32(b, sse2_rotl(tmp, S[r]));
}
a = _mm_add_epi32(a, oa);
b = _mm_add_epi32(b, ob);
c = _mm_add_epi32(c, oc);
d = _mm_add_epi32(d, od);
}
let mut a_words = [0u32; 4];
let mut b_words = [0u32; 4];
let mut c_words = [0u32; 4];
let mut d_words = [0u32; 4];
_mm_storeu_si128(a_words.as_mut_ptr() as *mut __m128i, a);
_mm_storeu_si128(b_words.as_mut_ptr() as *mut __m128i, b);
_mm_storeu_si128(c_words.as_mut_ptr() as *mut __m128i, c);
_mm_storeu_si128(d_words.as_mut_ptr() as *mut __m128i, d);
let mut results = Vec::with_capacity(n);
for lane in 0..n {
let mut digest = [0u8; 16];
digest[0..4].copy_from_slice(&a_words[lane].to_le_bytes());
digest[4..8].copy_from_slice(&b_words[lane].to_le_bytes());
digest[8..12].copy_from_slice(&c_words[lane].to_le_bytes());
digest[12..16].copy_from_slice(&d_words[lane].to_le_bytes());
results.push(digest);
}
results
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn md5_multi_neon(inputs: &[&[u8]], pad_to: Option<u64>) -> Vec<[u8; 16]> {
use std::arch::aarch64::*;
let n = inputs.len();
let effective_lens: Vec<u64> = inputs
.iter()
.map(|inp| {
let raw = inp.len() as u64;
match pad_to {
Some(p) if p > raw => p,
_ => raw,
}
})
.collect();
let padded: Vec<Vec<u8>> = (0..n)
.map(|i| md5_pad(inputs[i], effective_lens[i]))
.collect();
let block_counts: Vec<usize> = padded.iter().map(|p| p.len() / 64).collect();
let uniform_blocks = block_counts.iter().all(|&c| c == block_counts[0]);
if !uniform_blocks {
return md5_multi_scalar(inputs, pad_to);
}
let num_blocks = block_counts[0];
let lane_ptrs: [&[u8]; 4] = [
&padded[0],
if n > 1 { &padded[1] } else { &padded[0] },
if n > 2 { &padded[2] } else { &padded[0] },
if n > 3 { &padded[3] } else { &padded[0] },
];
unsafe {
let mut a = vdupq_n_u32(MD5_A0);
let mut b = vdupq_n_u32(MD5_B0);
let mut c = vdupq_n_u32(MD5_C0);
let mut d = vdupq_n_u32(MD5_D0);
for block_idx in 0..num_blocks {
let mut m = [vdupq_n_u32(0); 16];
let off = block_idx * 64;
for w in 0..4 {
let idx = w * 4;
let w_off = idx * 4;
let in0 = vreinterpretq_u32_u8(vld1q_u8(lane_ptrs[0].as_ptr().add(off + w_off)));
let in1 = vreinterpretq_u32_u8(vld1q_u8(lane_ptrs[1].as_ptr().add(off + w_off)));
let in2 = vreinterpretq_u32_u8(vld1q_u8(lane_ptrs[2].as_ptr().add(off + w_off)));
let in3 = vreinterpretq_u32_u8(vld1q_u8(lane_ptrs[3].as_ptr().add(off + w_off)));
let z01 = vzipq_u32(in0, in1);
let z23 = vzipq_u32(in2, in3);
m[idx] = vcombine_u32(vget_low_u32(z01.0), vget_low_u32(z23.0));
m[idx + 1] = vcombine_u32(vget_high_u32(z01.0), vget_high_u32(z23.0));
m[idx + 2] = vcombine_u32(vget_low_u32(z01.1), vget_low_u32(z23.1));
m[idx + 3] = vcombine_u32(vget_high_u32(z01.1), vget_high_u32(z23.1));
}
let oa = a;
let ob = b;
let oc = c;
let od = d;
#[inline(always)]
fn neon_rotl(v: uint32x4_t, amount: u32) -> uint32x4_t {
unsafe {
let left = vshlq_u32(v, vdupq_n_s32(amount as i32));
let right = vshlq_u32(v, vdupq_n_s32(-(32i32 - amount as i32)));
vorrq_u32(left, right)
}
}
for r in 0..16 {
let f = vbslq_u32(b, c, d);
let tmp = vaddq_u32(vaddq_u32(a, f), vaddq_u32(vdupq_n_u32(K[r]), m[G[r]]));
a = d;
d = c;
c = b;
b = vaddq_u32(b, neon_rotl(tmp, S[r]));
}
for r in 16..32 {
let f = vbslq_u32(d, b, c);
let tmp = vaddq_u32(vaddq_u32(a, f), vaddq_u32(vdupq_n_u32(K[r]), m[G[r]]));
a = d;
d = c;
c = b;
b = vaddq_u32(b, neon_rotl(tmp, S[r]));
}
for r in 32..48 {
let f = veorq_u32(veorq_u32(b, c), d);
let tmp = vaddq_u32(vaddq_u32(a, f), vaddq_u32(vdupq_n_u32(K[r]), m[G[r]]));
a = d;
d = c;
c = b;
b = vaddq_u32(b, neon_rotl(tmp, S[r]));
}
for r in 48..64 {
let f = veorq_u32(c, vornq_u32(b, d));
let tmp = vaddq_u32(vaddq_u32(a, f), vaddq_u32(vdupq_n_u32(K[r]), m[G[r]]));
a = d;
d = c;
c = b;
b = vaddq_u32(b, neon_rotl(tmp, S[r]));
}
a = vaddq_u32(a, oa);
b = vaddq_u32(b, ob);
c = vaddq_u32(c, oc);
d = vaddq_u32(d, od);
}
let mut results = Vec::with_capacity(n);
for lane in 0..n {
let mut digest = [0u8; 16];
digest[0..4].copy_from_slice(&extract_lane(a, lane).to_le_bytes());
digest[4..8].copy_from_slice(&extract_lane(b, lane).to_le_bytes());
digest[8..12].copy_from_slice(&extract_lane(c, lane).to_le_bytes());
digest[12..16].copy_from_slice(&extract_lane(d, lane).to_le_bytes());
results.push(digest);
}
results
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
fn extract_lane(v: std::arch::aarch64::uint32x4_t, lane: usize) -> u32 {
unsafe {
use std::arch::aarch64::*;
match lane {
0 => vgetq_lane_u32::<0>(v),
1 => vgetq_lane_u32::<1>(v),
2 => vgetq_lane_u32::<2>(v),
3 => vgetq_lane_u32::<3>(v),
_ => unreachable!(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn reference_md5(data: &[u8]) -> [u8; 16] {
Md5::digest(data).into()
}
fn reference_md5_padded(data: &[u8], pad_to: u64) -> [u8; 16] {
let mut padded = data.to_vec();
if (padded.len() as u64) < pad_to {
padded.resize(pad_to as usize, 0);
}
Md5::digest(&padded).into()
}
#[test]
fn single_input_matches_reference() {
let data = b"hello world";
let result = md5_multi(&[data], None);
assert_eq!(result[0], reference_md5(data));
}
#[test]
fn four_inputs_match_reference() {
let inputs: Vec<Vec<u8>> = (0..4)
.map(|i| {
(0..256u32)
.map(|j| ((j * 7 + i * 13) % 256) as u8)
.collect()
})
.collect();
let refs: Vec<&[u8]> = inputs.iter().map(|v| v.as_slice()).collect();
let results = md5_multi(&refs, None);
for (i, input) in inputs.iter().enumerate() {
assert_eq!(results[i], reference_md5(input), "mismatch for input {i}");
}
}
#[test]
fn variable_input_counts() {
let data: Vec<Vec<u8>> = (0..4)
.map(|i| vec![(i * 37 % 256) as u8; 100 + i * 50])
.collect();
for count in 1..=4 {
let refs: Vec<&[u8]> = data[..count].iter().map(|v| v.as_slice()).collect();
let results = md5_multi(&refs, None);
assert_eq!(results.len(), count);
for (i, input) in data[..count].iter().enumerate() {
assert_eq!(
results[i],
reference_md5(input),
"mismatch for count={count}, input {i}"
);
}
}
}
#[test]
fn pad_to_semantics() {
let data = b"short data";
let pad_to = 128u64;
let result = md5_multi(&[data], Some(pad_to));
assert_eq!(result[0], reference_md5_padded(data, pad_to));
}
#[test]
fn pad_to_multiple_inputs() {
let inputs: Vec<Vec<u8>> = vec![
vec![0xAA; 50],
vec![0xBB; 64],
vec![0xCC; 100],
vec![0xDD; 128],
];
let pad_to = 128u64;
let refs: Vec<&[u8]> = inputs.iter().map(|v| v.as_slice()).collect();
let results = md5_multi(&refs, Some(pad_to));
for (i, input) in inputs.iter().enumerate() {
assert_eq!(
results[i],
reference_md5_padded(input, pad_to),
"pad_to mismatch for input {i} (len={})",
input.len()
);
}
}
#[test]
fn empty_input() {
let data: &[u8] = b"";
let result = md5_multi(&[data], None);
assert_eq!(result[0], reference_md5(data));
}
#[test]
fn exact_block_sizes() {
for size in [64, 128, 192, 256] {
let data: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let result = md5_multi(&[&data], None);
assert_eq!(result[0], reference_md5(&data), "mismatch for size={size}");
}
}
#[test]
fn non_block_aligned_sizes() {
for size in [1, 7, 55, 56, 63, 65, 100, 119, 120, 127, 129] {
let data: Vec<u8> = (0..size).map(|i| (i % 256) as u8).collect();
let result = md5_multi(&[&data], None);
assert_eq!(result[0], reference_md5(&data), "mismatch for size={size}");
}
}
#[test]
fn large_input() {
let data: Vec<u8> = (0..65536u32).map(|i| (i % 256) as u8).collect();
let result = md5_multi(&[&data], None);
assert_eq!(result[0], reference_md5(&data));
}
#[test]
fn four_large_inputs() {
let inputs: Vec<Vec<u8>> = (0..4)
.map(|i| {
(0..8192u32)
.map(|j| ((j * 7 + i * 1337) % 256) as u8)
.collect()
})
.collect();
let refs: Vec<&[u8]> = inputs.iter().map(|v| v.as_slice()).collect();
let results = md5_multi(&refs, None);
for (i, input) in inputs.iter().enumerate() {
assert_eq!(results[i], reference_md5(input), "mismatch for input {i}");
}
}
#[test]
fn pad_to_with_exact_length_is_noop() {
let data = vec![0xAB; 256];
let result_no_pad = md5_multi(&[&data], None);
let result_pad = md5_multi(&[&data], Some(256));
assert_eq!(result_no_pad[0], result_pad[0]);
}
#[test]
fn scalar_matches_neon() {
let inputs: Vec<Vec<u8>> = (0..4)
.map(|i| {
(0..500u32)
.map(|j| ((j * 11 + i * 97) % 256) as u8)
.collect()
})
.collect();
let refs: Vec<&[u8]> = inputs.iter().map(|v| v.as_slice()).collect();
let scalar = md5_multi_scalar(&refs, None);
let dispatched = md5_multi(&refs, None);
for i in 0..4 {
assert_eq!(
scalar[i], dispatched[i],
"scalar vs dispatched mismatch for input {i}"
);
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[test]
fn scalar_matches_x86_backend() {
let inputs: Vec<Vec<u8>> = (0..4)
.map(|i| {
(0..500u32)
.map(|j| ((j * 17 + i * 101) % 256) as u8)
.collect()
})
.collect();
let refs: Vec<&[u8]> = inputs.iter().map(|v| v.as_slice()).collect();
let scalar = md5_multi_scalar(&refs, None);
let x86 = unsafe { md5_multi_x86(&refs, None) };
for i in 0..4 {
assert_eq!(scalar[i], x86[i], "scalar vs x86 mismatch for input {i}");
}
}
}