taler-rust

GNU Taler code in Rust. Largely core banking integrations.
Log | Files | Refs | Submodules | README | LICENSE

commit 0b262a884fa9afaf4b718cf7b841f0fa0a9dc62f
parent 0a8cd6865484b261e7f86d4d0bc680fa12330b14
Author: Antoine A <>
Date:   Tue, 28 Jul 2026 14:14:41 +0200

common: bug fixes and improvements

Diffstat:
Mcommon/taler-common/src/encoding/base64.rs | 68++++++++++++++++++++++++++++++++++++++++++--------------------------
Mcommon/taler-common/src/encoding/hex.rs | 8++++----
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 {