diff --git a/base32ct/src/encoding.rs b/base32ct/src/encoding.rs index 6b50cd3bb..96c60c54b 100644 --- a/base32ct/src/encoding.rs +++ b/base32ct/src/encoding.rs @@ -123,6 +123,18 @@ impl Encoding for T { err |= ((c[0] | c[1] | c[2] | c[3] | c[4] | c[5] | c[6]) >> 8) as u8; + // RFC 4648 3.5: the unused low bits of the last symbol of a partial block must be zero + // (otherwise several encodings decode to the same bytes). A remainder of 2/4/5/7 + // symbols leaves 2/4/1/3 unused bits in its last symbol. + let (last, unused_mask) = match src_rem.len() { + 2 => (c[1], 0b11), + 4 => (c[3], 0b1111), + 5 => (c[4], 0b1), + 7 => (c[6], 0b111), + _ => (0, 0), + }; + err |= (last & unused_mask != 0) as u8; + if err == 0 { Ok(dst) } else { diff --git a/base32ct/tests/proptests.rs b/base32ct/tests/proptests.rs index 89f1ebb7f..186bdffee 100644 --- a/base32ct/tests/proptests.rs +++ b/base32ct/tests/proptests.rs @@ -65,4 +65,28 @@ proptest! { assert_eq!(a, b); } } + + /// Every accepted input is the canonical encoding of its output (RFC 4648 3.5: the unused + /// bits of the last symbol are zero), so decode-then-encode gives the input back. + #[test] + fn decode_is_canonical(string in string_regex("[a-z2-7]{0,32}").unwrap()) { + if let Ok(bytes) = Base32UnpaddedCt::decode_vec(&string) { + prop_assert_eq!(Base32UnpaddedCt::encode_string(&bytes), string.clone()); + } + let padded = format!("{string}{}", "=".repeat((8 - string.len() % 8) % 8)); + if let Ok(bytes) = Base32Ct::decode_vec(&padded) { + prop_assert_eq!(Base32Ct::encode_string(&bytes), padded); + } + } +} + +/// "me" is the canonical encoding of "a" (0x61 = 01100 001|00); "mf" and "mh" set the unused +/// low bits of the last symbol and must be rejected. +#[test] +fn reject_non_zero_trailing_bits() { + assert_eq!(Base32UnpaddedCt::decode_vec("me").unwrap(), b"a"); + assert!(Base32UnpaddedCt::decode_vec("mf").is_err()); + assert!(Base32UnpaddedCt::decode_vec("mh").is_err()); + assert_eq!(Base32Ct::decode_vec("me======").unwrap(), b"a"); + assert!(Base32Ct::decode_vec("mf======").is_err()); }