commit 0b262a884fa9afaf4b718cf7b841f0fa0a9dc62f
parent 0a8cd6865484b261e7f86d4d0bc680fa12330b14
Author: Antoine A <>
Date: Tue, 28 Jul 2026 14:14:41 +0200
common: bug fixes and improvements
Diffstat:
2 files changed, 46 insertions(+), 30 deletions(-)
diff --git a/common/taler-common/src/encoding/base64.rs b/common/taler-common/src/encoding/base64.rs
@@ -90,21 +90,19 @@ pub enum Base64Error {
Length,
}
-const BASE64_INV: [u8; 256] = {
- let mut table = [255u8; 256];
+const fn build_shift_table(shift: u32) -> [u32; 256] {
+ let mut t = [0xFFFF_FFFFu32; 256]; // sentinel: invalid marker in all bits
let mut i = 0;
while i < 64 {
- table[BASE64_ALPHABET[i] as usize] = i as u8;
+ t[BASE64_ALPHABET[i] as usize] = (i as u32) << shift;
i += 1;
}
- table
-};
-
-/** Unpadded decoded length from a padded base64 string */
-fn decoded_len(encoded: &[u8]) -> usize {
- let padding = encoded.iter().rev().take_while(|&&b| b == b'=').count();
- encoded.len() * 3 / 4 - padding
+ t
}
+const T0: [u32; 256] = build_shift_table(26);
+const T1: [u32; 256] = build_shift_table(20);
+const T2: [u32; 256] = build_shift_table(14);
+const T3: [u32; 256] = build_shift_table(8);
/** Decode a standard base64 string (with `=` padding) */
pub fn decode(encoded: impl AsRef<[u8]>) -> Result<Vec<u8>, Base64Error> {
@@ -113,28 +111,41 @@ pub fn decode(encoded: impl AsRef<[u8]>) -> Result<Vec<u8>, Base64Error> {
return Err(Base64Error::Length);
}
- let out_len = decoded_len(encoded);
+ // Padding, if present at all, can only be this trailing run.
+ let pad = encoded.iter().rev().take_while(|&&b| b == b'=').count();
+ if pad > 2 {
+ return Err(Base64Error::Format);
+ }
+
+ let core = &encoded[..encoded.len() - pad];
+
+ let out_len = core.len() / 4 * 3
+ + match core.len() % 4 {
+ 2 => 1,
+ 3 => 2,
+ _ => 0,
+ };
let mut decoded = Vec::with_capacity(out_len);
let mut invalid = false;
- for chunk in encoded.chunks(4) {
- let mut buf = [0u8; 4];
- // Lookup chunk
- for (i, &b) in chunk.iter().enumerate() {
- buf[i] = if b == b'=' { 0 } else { BASE64_INV[b as usize] };
- }
+ let (chunks, tail) = core.as_chunks::<4>();
+ for chunk in chunks {
+ let word = T0[chunk[0] as usize]
+ | T1[chunk[1] as usize]
+ | T2[chunk[2] as usize]
+ | T3[chunk[3] as usize];
+ invalid |= word & 0xFF != 0;
+ decoded.extend_from_slice(&word.to_be_bytes()[..3]);
+ }
- // Check chunk validity
- invalid |= buf.contains(&255);
+ if !tail.is_empty() {
+ let mut word = T0[tail[0] as usize] | T1[tail[1] as usize];
- // Decode chunk
- decoded.push((buf[0] << 2) | (buf[1] >> 4));
- if chunk[2] != b'=' {
- decoded.push((buf[1] << 4) | (buf[2] >> 2));
- }
- if chunk[3] != b'=' {
- decoded.push((buf[2] << 6) | buf[3]);
+ if tail.len() == 3 {
+ word |= T2[tail[2] as usize];
}
+ invalid |= word & 0xFF != 0;
+ decoded.extend_from_slice(&word.to_be_bytes()[..tail.len() - 1]);
}
if invalid {
@@ -172,5 +183,10 @@ mod test {
// Invalid characters
assert_eq!(decode(b"Zg=!"), Err(Base64Error::Format));
assert_eq!(decode(b"Z\x00=="), Err(Base64Error::Format));
+
+ // Invalid padding
+ assert_eq!(decode(b"===="), Err(Base64Error::Format));
+ assert_eq!(decode(b"A=BC"), Err(Base64Error::Format));
+ assert_eq!(decode(b"AA==QUJD"), Err(Base64Error::Format));
}
}
diff --git a/common/taler-common/src/encoding/hex.rs b/common/taler-common/src/encoding/hex.rs
@@ -20,7 +20,7 @@ pub const HEX_ALPHABET: &[u8] = b"0123456789abcdef";
/** Encode a single byte to two hex characters */
#[inline(always)]
-fn encode_byte(byte: u8, encoded: &mut [u8; 2]) {
+const fn encode_byte(byte: u8, encoded: &mut [u8; 2]) {
encoded[0] = HEX_ALPHABET[(byte >> 4) as usize];
encoded[1] = HEX_ALPHABET[(byte & 0x0F) as usize];
}
@@ -82,16 +82,16 @@ pub fn decode(encoded: impl AsRef<[u8]>) -> Result<Vec<u8>, HexError> {
return Err(HexError::Length);
}
- let mut decoded = Vec::with_capacity(encoded.len() / 2);
+ let mut decoded = vec![0u8; encoded.len() / 2];
let mut invalid = false;
- for [hi, lo] in encoded.as_chunks::<2>().0 {
+ for ([hi, lo], decoded) in encoded.as_chunks::<2>().0.iter().zip(decoded.iter_mut()) {
let hi = HEX_INV[*hi as usize];
let lo = HEX_INV[*lo as usize];
invalid |= hi == 255 || lo == 255;
- decoded.push((hi << 4) | lo);
+ *decoded = (hi << 4) | lo;
}
if invalid {