//! Base64 encoding (RFC 4648 4, the standard alphabet with padding), for
//! the ws handshake's `Sec-WebSocket-Key` and `Sec-WebSocket-Accept`.
//!
//! Only encoding: the handshake never needs to decode either value.  The
//! output goes into the caller's buffer, as everything the protocols
//! compose does.

/// The standard alphabet.
const ALPHABET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";

/// The output buffer cannot hold the encoding.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct BufferTooSmall;

impl core::fmt::Display for BufferTooSmall {
    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
        f.write_str("buffer too small for the base64 encoding")
    }
}

impl core::error::Error for BufferTooSmall {}

/// The length of the encoding of `n` bytes, or `None` if it would not fit
/// a `usize`.
#[must_use]
pub const fn encoded_len(n: usize) -> Option<usize> {
    n.div_ceil(3).checked_mul(4)
}

/// Encodes `src` into the start of `dst`, returning how many bytes of `dst`
/// it used.
///
/// ```
/// use npro_core::base64;
///
/// let mut out = [0; 8];
/// let n = base64::encode(b"foob", &mut out)?;
/// assert_eq!(&out[..n], b"Zm9vYg==");
/// # Ok::<(), base64::BufferTooSmall>(())
/// ```
///
/// # Errors
///
/// [`BufferTooSmall`] if `dst` is shorter than [`encoded_len`] of `src`;
/// nothing is written then.
pub fn encode(src: &[u8], dst: &mut [u8]) -> Result<usize, BufferTooSmall> {
    let len = encoded_len(src.len()).ok_or(BufferTooSmall)?;
    let out = dst.get_mut(..len).ok_or(BufferTooSmall)?;

    for (o, i) in out.chunks_exact_mut(4).zip(src.chunks(3)) {
        let (b0, b1, b2, have) = match *i {
            [a, b, c] => (a, b, c, 4),
            [a, b] => (a, b, 0, 3),
            [a] => (a, 0, 0, 2),
            _ => (0, 0, 0, 0),
        };
        let sextets = [b0 >> 2, (b0 << 4 | b1 >> 4), (b1 << 2 | b2 >> 6), b2];
        for (n, (d, s)) in o.iter_mut().zip(sextets).enumerate() {
            *d = if n < have { symbol(s) } else { b'=' };
        }
    }

    Ok(len)
}

/// The alphabet's symbol for the low six bits of `v`.
fn symbol(v: u8) -> u8 {
    // masked to six bits, the index is always inside the 64-symbol alphabet
    ALPHABET.get(usize::from(v & 0x3f)).copied().unwrap_or(b'=')
}

#[cfg(test)]
mod tests {
    use super::*;

    fn enc(src: &[u8]) -> String {
        let mut out = vec![0; encoded_len(src.len()).unwrap()];
        let n = encode(src, &mut out).unwrap();
        assert_eq!(n, out.len());
        String::from_utf8(out).unwrap()
    }

    #[test]
    fn rfc_4648_vectors() {
        for (src, want) in [
            ("", ""),
            ("f", "Zg=="),
            ("fo", "Zm8="),
            ("foo", "Zm9v"),
            ("foob", "Zm9vYg=="),
            ("fooba", "Zm9vYmE="),
            ("foobar", "Zm9vYmFy"),
        ] {
            assert_eq!(enc(src.as_bytes()), want);
        }
    }

    #[test]
    fn every_symbol_and_the_high_bits() {
        // 0x00..=0xff covers every sextet value in every position
        let all: Vec<u8> = (0..=255).collect();
        let s = enc(&all);
        for c in ALPHABET {
            assert!(s.as_bytes().contains(c));
        }
        assert!(s.starts_with("AAECAwQF"));
        assert!(s.ends_with("+/w=="));
    }

    #[test]
    fn a_short_buffer_is_refused_untouched() {
        let mut out = [b'x'; 7];
        assert_eq!(encode(b"foob", &mut out), Err(BufferTooSmall));
        assert_eq!(out, [b'x'; 7]);
    }

    #[test]
    fn the_ws_accept_of_rfc_6455() {
        // RFC 6455 1.3: base64(SHA-1(key + GUID))
        let mut h = crate::sha1::Sha1::new();
        h.update(b"dGhlIHNhbXBsZSBub25jZQ==");
        h.update(b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
        assert_eq!(enc(&h.finish()), "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=");
    }

    #[cfg(feature = "replay")]
    #[test]
    fn the_ws_client_transcripts_key() {
        // C's ws-client transcript: the key its client sends, drawn as the
        // first 16 bytes of lws' random seeded with 1
        use crate::random::{Random, SeededRandom};
        let mut key = [0; 16];
        SeededRandom::new(1).fill(&mut key).unwrap();
        assert_eq!(enc(&key), "OvomtQpKCWUnZW7tMR6Gqw==");
    }
}