rustpython-host_env 0.6.0

Host OS API abstractions for RustPython
Documentation
//! Server acceptor state used before rustls has a `ServerConfig`.

use core::fmt;
use rustls::server::Acceptor;

/// Server configuration is selected only after receiving ClientHello and
/// invoking SNI. A failed connection stays terminal while its alert drains.
/// ShuttingDown means our close_notify is queued and must not be sent twice.
pub enum TlsState<E> {
    WaitingForClientHello(Box<Acceptor>),
    InProgress,
    Handshaking,
    Connected,
    ShuttingDown,
    ShutDown,
    SendingAlert { error: E },
}

impl<E> TlsState<E> {
    pub fn new(server_side: bool) -> Self {
        if server_side {
            Self::WaitingForClientHello(Box::default())
        } else {
            Self::Handshaking
        }
    }
}

impl<E> fmt::Debug for TlsState<E> {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        f.write_str(match self {
            Self::WaitingForClientHello(_) => "WaitingForClientHello",
            Self::InProgress => "InProgress",
            Self::Handshaking => "Handshaking",
            Self::Connected => "Connected",
            Self::ShuttingDown => "ShuttingDown",
            Self::ShutDown => "ShutDown",
            Self::SendingAlert { .. } => "SendingAlert",
        })
    }
}

/// SNI rejection happens before ServerHello, so the fatal alert is plaintext
/// even when the ClientHello offers TLS 1.3 (RFC 8446 section 5.1).
#[must_use]
pub fn sni_alert(description: u8) -> Vec<u8> {
    vec![21, 3, 3, 0, 2, 2, description]
}

pub fn feed_acceptor(acceptor: &mut Acceptor, bytes: &[u8]) -> std::io::Result<()> {
    let mut reader = std::io::Cursor::new(bytes);
    while reader.position() < bytes.len() as u64 {
        if acceptor.read_tls(&mut reader)? == 0 {
            return Err(std::io::ErrorKind::UnexpectedEof.into());
        }
    }
    Ok(())
}

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

    fn large_client_hello() -> Vec<u8> {
        let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider());
        let mut config = rustls::ClientConfig::builder_with_provider(provider)
            .with_safe_default_protocol_versions()
            .unwrap()
            .with_root_certificates(rustls::RootCertStore::empty())
            .with_no_client_auth();
        config.alpn_protocols = (0..70).map(|i| format!("{i:0100}").into_bytes()).collect();
        let mut client =
            rustls::ClientConnection::new(Arc::new(config), "localhost".try_into().unwrap())
                .unwrap();
        let mut bytes = Vec::new();
        client.write_tls(&mut bytes).unwrap();
        assert!(bytes.len() > 4096);
        bytes
    }

    #[test]
    fn consumes_large_client_hello_completely() {
        let mut acceptor = Acceptor::default();
        feed_acceptor(&mut acceptor, &large_client_hello()).unwrap();
        let accepted = acceptor.accept().unwrap().unwrap();
        assert_eq!(accepted.client_hello().server_name(), Some("localhost"));
        assert_eq!(accepted.client_hello().alpn().unwrap().count(), 70);
    }

    #[test]
    fn waits_for_every_fragment_of_client_hello() {
        let hello = large_client_hello();
        let mut acceptor = Acceptor::default();
        for byte in &hello[..hello.len() - 1] {
            feed_acceptor(&mut acceptor, &[*byte]).unwrap();
            assert!(acceptor.accept().unwrap().is_none());
        }
        feed_acceptor(&mut acceptor, &hello[hello.len() - 1..]).unwrap();
        assert!(acceptor.accept().unwrap().is_some());
    }

    #[test]
    fn fatal_sni_alert_is_one_plaintext_record() {
        assert_eq!(sni_alert(49), [21, 3, 3, 0, 2, 2, 49]);
        assert_eq!(sni_alert(40), [21, 3, 3, 0, 2, 2, 40]);
        assert_eq!(sni_alert(80), [21, 3, 3, 0, 2, 2, 80]);
    }
}