1use crate::{ensure, Result};
32use core::hint::black_box;
33
34#[inline(always)]
37fn neg_mask(x: i16) -> u8 {
38 (x >> 8) as u8
39}
40
41#[inline(always)]
44fn hex_char(n: u8) -> u8 {
45 let n16 = n as i16;
46 n + b'0' + (neg_mask(9 - n16) & 39)
47}
48
49#[inline(always)]
53fn base64_char(v: u8) -> u8 {
54 let v16 = v as i16;
55 let mut d = b'A' as i16;
56 d += (neg_mask(25 - v16) & 6) as i16;
57 d -= (neg_mask(51 - v16) & 75) as i16;
58 d -= (neg_mask(61 - v16) & 15) as i16;
59 d += (neg_mask(62 - v16) & 3) as i16;
60 (v16 + d) as u8
61}
62
63pub fn hex_encode(input: &[u8], out: &mut [u8]) -> Result<()> {
67 ensure!(
68 out.len() == input.len() * 2,
69 InvalidLength,
70 "hex output buffer"
71 );
72 for (i, &b) in input.iter().enumerate() {
73 out[i * 2] = hex_char(b >> 4);
74 out[i * 2 + 1] = hex_char(b & 0x0f);
75 }
76 Ok(())
77}
78
79pub fn hex_decode(input: &[u8], out: &mut [u8]) -> Result<()> {
83 ensure!(
84 input.len().is_multiple_of(2),
85 MalformedEncoding,
86 "hex length must be even"
87 );
88 ensure!(
89 out.len() == input.len() / 2,
90 InvalidLength,
91 "hex output buffer"
92 );
93 let mut valid = 0xFFu8;
95 for i in 0..out.len() {
96 let (hi, ok_hi) = hex_nibble(input[i * 2]);
97 let (lo, ok_lo) = hex_nibble(input[i * 2 + 1]);
98 valid = black_box(valid & ok_hi & ok_lo);
99 out[i] = (hi << 4) | lo;
100 }
101 if valid != 0xFF {
102 out.fill(0);
104 }
105 ensure!(valid == 0xFF, MalformedEncoding, "hex digit");
106 Ok(())
107}
108
109#[inline(always)]
113fn hex_nibble(c: u8) -> (u8, u8) {
114 let digit = c.wrapping_sub(b'0');
115 let lower = c.wrapping_sub(b'a');
116 let upper = c.wrapping_sub(b'A');
117 let m_digit = neg_mask(digit as i16 - 10);
118 let m_lower = neg_mask(lower as i16 - 6);
119 let m_upper = neg_mask(upper as i16 - 6);
120 let value =
121 (digit & m_digit) | (lower.wrapping_add(10) & m_lower) | (upper.wrapping_add(10) & m_upper);
122 (value, m_digit | m_lower | m_upper)
123}
124
125pub const fn base64_encoded_len(n: usize) -> usize {
127 n.div_ceil(3) * 4
128}
129
130pub fn base64_encode(input: &[u8], out: &mut [u8]) -> Result<()> {
132 ensure!(
133 out.len() == base64_encoded_len(input.len()),
134 InvalidLength,
135 "base64 output buffer"
136 );
137 let mut oi = 0;
138 let mut chunks = input.chunks_exact(3);
139 for c in &mut chunks {
140 let n = ((c[0] as u32) << 16) | ((c[1] as u32) << 8) | c[2] as u32;
141 out[oi] = base64_char((n >> 18) as u8 & 63);
142 out[oi + 1] = base64_char((n >> 12) as u8 & 63);
143 out[oi + 2] = base64_char((n >> 6) as u8 & 63);
144 out[oi + 3] = base64_char(n as u8 & 63);
145 oi += 4;
146 }
147 let rem = chunks.remainder();
149 match rem.len() {
150 0 => {}
151 1 => {
152 let n = (rem[0] as u32) << 16;
153 out[oi] = base64_char((n >> 18) as u8 & 63);
154 out[oi + 1] = base64_char((n >> 12) as u8 & 63);
155 out[oi + 2] = b'=';
156 out[oi + 3] = b'=';
157 }
158 _ => {
159 let n = ((rem[0] as u32) << 16) | ((rem[1] as u32) << 8);
160 out[oi] = base64_char((n >> 18) as u8 & 63);
161 out[oi + 1] = base64_char((n >> 12) as u8 & 63);
162 out[oi + 2] = base64_char((n >> 6) as u8 & 63);
163 out[oi + 3] = b'=';
164 }
165 }
166 Ok(())
167}
168
169pub fn base64_decode(input: &[u8], out: &mut [u8]) -> Result<usize> {
174 ensure!(
175 input.len().is_multiple_of(4),
176 MalformedEncoding,
177 "base64 length"
178 );
179 if input.is_empty() {
180 return Ok(0);
181 }
182 let pad = usize::from(input[input.len() - 1] == b'=')
185 + usize::from(input.len() >= 2 && input[input.len() - 2] == b'=');
186 let decoded = input.len() / 4 * 3 - pad;
187 ensure!(out.len() >= decoded, InvalidLength, "base64 output buffer");
188
189 let mut valid = 0xFFu8;
191 let data_chars = input.len() - pad;
192 let mut oi = 0;
193 for (bi, block) in input.chunks_exact(4).enumerate() {
194 let mut n: u32 = 0;
195 for (j, &c) in block.iter().enumerate() {
196 let (v, ok) = base64_value(c);
197 let is_pad_pos = bi * 4 + j >= data_chars;
201 let ok = if is_pad_pos {
202 neg_mask(-i16::from(c == b'='))
203 } else {
204 ok
205 };
206 valid = black_box(valid & ok);
207 n |= (v as u32) << (18 - 6 * j);
208 }
209 let bytes = [(n >> 16) as u8, (n >> 8) as u8, n as u8];
210 for &b in &bytes {
211 if oi < decoded {
212 out[oi] = b;
213 oi += 1;
214 }
215 }
216 }
217 if valid != 0xFF {
218 out[..decoded].fill(0);
219 }
220 ensure!(valid == 0xFF, MalformedEncoding, "base64 character");
221 Ok(decoded)
222}
223
224#[inline(always)]
227fn base64_value(c: u8) -> (u8, u8) {
228 let upper = c.wrapping_sub(b'A');
229 let lower = c.wrapping_sub(b'a');
230 let digit = c.wrapping_sub(b'0');
231 let m_upper = neg_mask(upper as i16 - 26);
232 let m_lower = neg_mask(lower as i16 - 26);
233 let m_digit = neg_mask(digit as i16 - 10);
234 let m_plus = neg_mask((c ^ b'+') as i16 - 1);
237 let m_slash = neg_mask((c ^ b'/') as i16 - 1);
238 let value = (upper & m_upper)
239 | (lower.wrapping_add(26) & m_lower)
240 | (digit.wrapping_add(52) & m_digit)
241 | (62 & m_plus)
242 | (63 & m_slash);
243 (value, m_upper | m_lower | m_digit | m_plus | m_slash)
244}
245
246#[cfg(feature = "std")]
247mod alloc_helpers {
248 use super::*;
249
250 pub fn hex(input: &[u8]) -> String {
252 let mut buf = vec![0u8; input.len() * 2];
253 hex_encode(input, &mut buf).expect("buffer sized exactly");
254 String::from_utf8(buf).expect("hex alphabet is ASCII")
255 }
256
257 pub fn unhex(input: &str) -> Result<Vec<u8>> {
259 let mut buf = vec![0u8; input.len() / 2];
260 hex_decode(input.as_bytes(), &mut buf)?;
261 Ok(buf)
262 }
263
264 pub fn b64(input: &[u8]) -> String {
266 let mut buf = vec![0u8; base64_encoded_len(input.len())];
267 base64_encode(input, &mut buf).expect("buffer sized exactly");
268 String::from_utf8(buf).expect("base64 alphabet is ASCII")
269 }
270
271 pub fn unb64(input: &str) -> Result<Vec<u8>> {
273 let mut buf = vec![0u8; input.len() / 4 * 3];
274 let n = base64_decode(input.as_bytes(), &mut buf)?;
275 buf.truncate(n);
276 Ok(buf)
277 }
278}
279
280#[cfg(feature = "std")]
281pub use alloc_helpers::{b64, hex, unb64, unhex};
282
283#[cfg(test)]
284mod tests {
285 use super::*;
286
287 #[test]
288 fn hex_roundtrip() {
289 let data = [0x00u8, 0x0f, 0xf0, 0xff, 0x42];
290 let mut enc = [0u8; 10];
291 hex_encode(&data, &mut enc).unwrap();
292 assert_eq!(&enc, b"000ff0ff42");
293 let mut dec = [0u8; 5];
294 hex_decode(&enc, &mut dec).unwrap();
295 assert_eq!(dec, data);
296 }
297
298 #[test]
299 fn hex_accepts_uppercase_and_rejects_junk() {
300 let mut dec = [0u8; 2];
301 hex_decode(b"AbCd", &mut dec).unwrap();
302 assert_eq!(dec, [0xab, 0xcd]);
303 assert!(hex_decode(b"zz", &mut dec[..1]).is_err());
304 assert!(hex_decode(b"abc", &mut dec).is_err());
305 }
306
307 #[test]
308 fn base64_matches_rfc4648_vectors() {
309 for (plain, encoded) in [
310 (&b""[..], ""),
311 (&b"f"[..], "Zg=="),
312 (&b"fo"[..], "Zm8="),
313 (&b"foo"[..], "Zm9v"),
314 (&b"foob"[..], "Zm9vYg=="),
315 (&b"fooba"[..], "Zm9vYmE="),
316 (&b"foobar"[..], "Zm9vYmFy"),
317 ] {
318 let mut enc = vec![0u8; base64_encoded_len(plain.len())];
319 base64_encode(plain, &mut enc).unwrap();
320 assert_eq!(
321 core::str::from_utf8(&enc).unwrap(),
322 encoded,
323 "encoding {plain:?}"
324 );
325
326 let mut dec = vec![0u8; plain.len() + 3];
327 let n = base64_decode(encoded.as_bytes(), &mut dec).unwrap();
328 assert_eq!(&dec[..n], plain, "decoding {encoded}");
329 }
330 }
331
332 #[test]
335 fn classification_agrees_with_the_alphabets_on_every_byte() {
336 const HEX: &[u8; 16] = b"0123456789abcdef";
337 const B64: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
338 for v in 0..16u8 {
339 assert_eq!(hex_char(v), HEX[v as usize], "hex {v}");
340 }
341 for v in 0..64u8 {
342 assert_eq!(base64_char(v), B64[v as usize], "base64 {v}");
343 }
344 for c in 0..=255u8 {
345 let want_hex = HEX
346 .iter()
347 .position(|&h| h == c.to_ascii_lowercase())
348 .map(|i| i as u8);
349 let (v, ok) = hex_nibble(c);
350 match want_hex {
351 Some(w) => assert_eq!((v, ok), (w, 0xFF), "hex {c:#04x}"),
352 None => assert_eq!(ok, 0, "hex {c:#04x} accepted"),
353 }
354 let want_b64 = B64.iter().position(|&b| b == c).map(|i| i as u8);
355 let (v, ok) = base64_value(c);
356 match want_b64 {
357 Some(w) => assert_eq!((v, ok), (w, 0xFF), "base64 {c:#04x}"),
358 None => assert_eq!((v, ok), (0, 0), "base64 {c:#04x} accepted"),
359 }
360 }
361 }
362
363 #[test]
364 fn padding_is_accepted_only_where_it_belongs() {
365 let mut out = [0u8; 6];
366 assert!(base64_decode(b"Zm9vYg==", &mut out).is_ok());
367 assert!(base64_decode(b"Zm9vYmE=", &mut out).is_ok());
368 for bad in [
369 &b"Zm=vYmFy"[..],
370 b"=m9vYmFy",
371 b"Zm9v=mFy",
372 b"Zm9vY=E=",
373 b"Zm9vYmF!",
374 ] {
375 let mut out = [0xAAu8; 6];
376 assert!(base64_decode(bad, &mut out).is_err(), "{bad:?} accepted");
377 assert!(
378 out.iter().all(|&b| b == 0 || b == 0xAA),
379 "{bad:?} left decoded bytes behind"
380 );
381 }
382 }
383
384 #[test]
385 fn a_bad_digit_anywhere_is_rejected_and_nothing_is_left() {
386 let mut out = [0u8; 4];
387 for i in 0..8 {
388 let mut text = *b"00112233";
389 text[i] = b'g';
390 assert!(hex_decode(&text, &mut out).is_err(), "position {i}");
391 assert_eq!(out, [0; 4], "position {i} left bytes behind");
392 }
393 }
394
395 #[test]
396 fn string_helpers_roundtrip() {
397 assert_eq!(hex(b"\xde\xad\xbe\xef"), "deadbeef");
398 assert_eq!(unhex("deadbeef").unwrap(), b"\xde\xad\xbe\xef");
399 assert_eq!(b64(b"foobar"), "Zm9vYmFy");
400 assert_eq!(unb64("Zm9vYmE=").unwrap(), b"fooba");
401 }
402}