age 0.11.4

[BETA] A simple, secure, and modern encryption library.
Documentation
use std::io;

use bech32::{FromBase32, Variant};

#[cfg(all(any(feature = "armor", feature = "cli-common"), windows))]
pub(crate) const LINE_ENDING: &str = "\r\n";
#[cfg(all(any(feature = "armor", feature = "cli-common"), not(windows)))]
pub(crate) const LINE_ENDING: &str = "\n";

pub(crate) fn parse_bech32(s: &str) -> Option<(String, Vec<u8>)> {
    bech32::decode(s).ok().and_then(|(hrp, data, variant)| {
        if let Variant::Bech32 = variant {
            Vec::from_base32(&data).ok().map(|d| (hrp, d))
        } else {
            None
        }
    })
}

pub(crate) struct LimitedReader<R> {
    inner: R,
    n: usize,
    limit_exceeded: bool,
}
impl<R> LimitedReader<R> {
    pub(crate) fn new(reader: R, n: usize) -> Self {
        Self {
            inner: reader,
            n,
            limit_exceeded: false,
        }
    }

    fn limit_exceeded() -> io::Error {
        io::Error::new(io::ErrorKind::InvalidData, "reader exceeded size limit")
    }
}

impl<R: io::Read> io::Read for LimitedReader<R> {
    fn read(&mut self, mut buf: &mut [u8]) -> io::Result<usize> {
        if buf.is_empty() {
            return Ok(0);
        }

        if self.limit_exceeded {
            return Err(Self::limit_exceeded());
        }

        if self.n == 0 {
            let mut probe = [0];
            if self.inner.read(&mut probe)? == 0 {
                Ok(0)
            } else {
                self.limit_exceeded = true;
                Err(Self::limit_exceeded())
            }
        } else {
            if buf.len() > self.n {
                buf = &mut buf[..self.n];
            }
            let read = self.inner.read(buf)?;
            self.n -= read;
            Ok(read)
        }
    }
}

impl<R: io::BufRead> io::BufRead for LimitedReader<R> {
    fn fill_buf(&mut self) -> io::Result<&[u8]> {
        if self.limit_exceeded {
            return Err(Self::limit_exceeded());
        }

        if self.n == 0 {
            if self.inner.fill_buf()?.is_empty() {
                Ok(&[])
            } else {
                self.limit_exceeded = true;
                Err(Self::limit_exceeded())
            }
        } else {
            let buf = self.inner.fill_buf()?;
            Ok(&buf[..buf.len().min(self.n)])
        }
    }

    fn consume(&mut self, amount: usize) {
        self.n -= amount;
        self.inner.consume(amount);
    }
}

#[cfg(test)]
mod tests {
    use super::LimitedReader;
    use std::io::{self, BufRead, Cursor, Read};

    #[test]
    fn limited_reader_read() {
        for (input, limit) in [(&b"abc"[..], 4), (&b"abc"[..], 3)] {
            let mut output = vec![];
            LimitedReader::new(input, limit)
                .read_to_end(&mut output)
                .unwrap();
            assert_eq!(output, input);
        }

        let mut output = vec![];
        let mut reader = LimitedReader::new(&b"abcd"[..], 3);
        let err = reader.read_to_end(&mut output).unwrap_err();
        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
        assert_eq!(output, b"abc");
        assert_eq!(
            reader.read(&mut [0]).unwrap_err().kind(),
            io::ErrorKind::InvalidData
        );
    }

    #[test]
    fn limited_reader_bufread() {
        for (input, limit) in [(&b"abc"[..], 4), (&b"abc"[..], 3)] {
            let mut output = vec![];
            LimitedReader::new(Cursor::new(input), limit)
                .read_until(b'\n', &mut output)
                .unwrap();
            assert_eq!(output, input);
        }

        let mut output = vec![];
        let mut reader = LimitedReader::new(Cursor::new(&b"abcd"[..]), 3);
        let err = reader.read_until(b'\n', &mut output).unwrap_err();
        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
        assert_eq!(output, b"abc");
        assert_eq!(
            reader.fill_buf().unwrap_err().kind(),
            io::ErrorKind::InvalidData
        );
    }
}

pub(crate) mod read {
    use std::str::FromStr;

    use base64::{prelude::BASE64_STANDARD_NO_PAD, Engine};
    use nom::{character::complete::digit1, combinator::verify, ParseTo};

    #[cfg(feature = "ssh")]
    use nom::{
        combinator::map_res,
        error::{make_error, ErrorKind},
        multi::separated_list1,
        IResult,
    };

    #[cfg(feature = "ssh")]
    #[cfg_attr(docsrs, doc(cfg(feature = "ssh")))]
    pub(crate) fn encoded_str(
        count: usize,
        engine: impl base64::Engine,
    ) -> impl Fn(&str) -> IResult<&str, Vec<u8>> {
        use nom::bytes::streaming::take;

        // Unpadded encoded length
        let encoded_count = ((4 * count) + 2) / 3;

        move |input: &str| {
            let (i, data) = take(encoded_count)(input)?;
            match engine.decode(data) {
                Ok(decoded) => Ok((i, decoded)),
                Err(_) => Err(nom::Err::Failure(make_error(input, ErrorKind::Eof))),
            }
        }
    }

    #[cfg(feature = "ssh")]
    #[cfg_attr(docsrs, doc(cfg(feature = "ssh")))]
    pub(crate) fn str_while_encoded(
        engine: impl base64::Engine,
    ) -> impl Fn(&str) -> IResult<&str, Vec<u8>> {
        use nom::bytes::complete::take_while1;

        move |input: &str| {
            map_res(
                take_while1(|c| {
                    let c = c as u8;
                    // Substitute the character in twice after AA, so that padding
                    // characters will also be detected as a valid if allowed.
                    engine.decode_slice([65, 65, c, c], &mut [0, 0, 0]).is_ok()
                }),
                |data| engine.decode(data),
            )(input)
        }
    }

    #[cfg(feature = "ssh")]
    #[cfg_attr(docsrs, doc(cfg(feature = "ssh")))]
    pub(crate) fn wrapped_str_while_encoded(
        engine: impl Engine,
    ) -> impl Fn(&str) -> IResult<&str, Vec<u8>> {
        use nom::{bytes::streaming::take_while1, character::streaming::line_ending};

        move |input: &str| {
            map_res(
                separated_list1(
                    line_ending,
                    take_while1(|c| {
                        let c = c as u8;
                        // Substitute the character in twice after AA, so that padding
                        // characters will also be detected as a valid if allowed.
                        engine.decode_slice([65, 65, c, c], &mut [0, 0, 0]).is_ok()
                    }),
                ),
                |chunks| {
                    let data = chunks.join("");
                    engine.decode(&data)
                },
            )(input)
        }
    }

    pub(crate) fn base64_arg<A: AsRef<[u8]>, const N: usize, const B: usize>(
        arg: &A,
    ) -> Option<[u8; N]> {
        if N > B {
            return None;
        }

        let mut buf = [0; B];
        match BASE64_STANDARD_NO_PAD.decode_slice(arg, buf.as_mut()) {
            Ok(n) if n == N => Some(buf[..N].try_into().unwrap()),
            _ => None,
        }
    }

    /// Parses a decimal number composed only of digits with no leading zeros.
    pub(crate) fn decimal_digit_arg<T: FromStr>(arg: &str) -> Option<T> {
        verify::<_, _, _, (), _, _>(digit1, |n: &str| !n.starts_with('0'))(arg)
            .ok()
            .and_then(|(_, n)| n.parse_to())
    }
}

pub(crate) mod write {
    use base64::{prelude::BASE64_STANDARD_NO_PAD, Engine};
    use cookie_factory::{combinator::string, SerializeFn};
    use std::io::Write;

    pub(crate) fn encoded_data<W: Write>(data: &[u8]) -> impl SerializeFn<W> {
        let encoded = BASE64_STANDARD_NO_PAD.encode(data);
        string(encoded)
    }
}