pub mod crypto;
pub mod x509;
#[cfg(test)]
mod testdata;
mod handshake;
mod key_schedule;
mod record;
use alloc::string::String;
use alloc::vec::Vec;
pub use x509::RootStore;
pub const TLS_VERSION_1_3: u16 = 0x0304;
pub const TLS_VERSION_1_2: u16 = 0x0303;
const MAX_HANDSHAKE_BUFFER: usize = 16 * 1024 * 1024;
#[derive(Debug, Clone, Default)]
pub struct ClientHelloInfo {
pub server_name: Option<String>,
pub alpn: Vec<Vec<u8>>,
pub negotiated_alpn: Option<Vec<u8>>,
pub cipher_suites: Vec<u16>,
}
#[derive(Debug, Clone)]
pub enum TlsError {
Io(String),
Protocol(String),
Alert {
level: u8,
description: u8,
},
Certificate(String),
Timeout,
UnexpectedEof,
Unsupported(String),
Internal(String),
}
impl core::fmt::Display for TlsError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
TlsError::Io(m) => write!(f, "TLS I/O error: {m}"),
TlsError::Protocol(m) => write!(f, "TLS protocol error: {m}"),
TlsError::Alert { level, description } => {
write!(f, "TLS alert level={level} description={description}")
}
TlsError::Certificate(m) => write!(f, "TLS certificate error: {m}"),
TlsError::Timeout => write!(f, "TLS handshake timeout"),
TlsError::UnexpectedEof => write!(f, "TLS unexpected EOF"),
TlsError::Unsupported(m) => write!(f, "TLS unsupported: {m}"),
TlsError::Internal(m) => write!(f, "TLS internal error: {m}"),
}
}
}
impl From<crate::courierust_error::Error> for TlsError {
fn from(e: crate::courierust_error::Error) -> Self {
use crate::courierust_error::ErrorKind;
match e.kind {
ErrorKind::Timeout => TlsError::Timeout,
ErrorKind::UnexpectedEof => TlsError::UnexpectedEof,
_ => TlsError::Io(e.to_string()),
}
}
}
pub type TlsResult<T> = core::result::Result<T, TlsError>;
#[derive(Debug, Clone)]
pub struct Identity {
pub cert_chain: Vec<Vec<u8>>,
pub private_key: Vec<u8>,
pub is_rsa: bool,
}
use crate::courierust_io::{BufReader, BufWriter};
use record::{open_record, seal_record, Sequence, CONTENT_HANDSHAKE, MAX_RECORD_PAYLOAD};
pub(crate) struct TlsIo<R, W> {
reader: BufReader<R>,
writer: BufWriter<W>,
read_seq: Sequence,
write_seq: Sequence,
}
impl<R: crate::courierust_io::Read, W: crate::courierust_io::Write> TlsIo<R, W> {
pub(crate) fn new(reader: R, writer: W) -> Self {
Self {
reader: BufReader::new(reader, 65536),
writer: BufWriter::new(writer, 65536),
read_seq: Sequence::default(),
write_seq: Sequence::default(),
}
}
pub(crate) fn write_plaintext_record(
&mut self,
content_type: u8,
payload: &[u8],
) -> TlsResult<()> {
if payload.len() > u16::MAX as usize {
return Err(TlsError::Protocol("record too large".into()));
}
let header = [
content_type,
0x03,
0x01,
(payload.len() >> 8) as u8,
payload.len() as u8,
];
self.writer.write_all(&header).map_err(TlsError::from)?;
self.writer.write_all(payload).map_err(TlsError::from)?;
self.writer.flush().map_err(TlsError::from)
}
pub(crate) fn read_plaintext_record(&mut self) -> TlsResult<(u8, Vec<u8>)> {
let mut header = [0u8; 5];
self.reader
.read_exact_into(&mut header)
.map_err(TlsError::from)?;
let content_type = header[0];
let len = ((header[3] as usize) << 8) | header[4] as usize;
if len > MAX_RECORD_PAYLOAD {
return Err(TlsError::Protocol("record too large".into()));
}
let payload = self.reader.read_exact(len).map_err(TlsError::from)?;
Ok((content_type, payload))
}
pub(crate) fn read_plaintext_handshake(&mut self) -> TlsResult<(u8, Vec<u8>)> {
let (ct, payload) = self.read_plaintext_record()?;
if ct != CONTENT_HANDSHAKE {
return Err(TlsError::Protocol("expected handshake record".into()));
}
if payload.len() < 4 {
return Err(TlsError::Protocol("bad handshake record".into()));
}
Ok((payload[0], payload[4..].to_vec()))
}
pub(crate) fn write_encrypted_record(
&mut self,
suite: key_schedule::CipherSuite,
keys: &key_schedule::TrafficKeys,
content_type: u8,
payload: &[u8],
) -> TlsResult<()> {
let seq = self.write_seq.next()?;
let rec = seal_record(suite, keys, seq, content_type, payload)?;
self.writer.write_all(&rec).map_err(TlsError::from)?;
self.writer.flush().map_err(TlsError::from)
}
pub(crate) fn write_encrypted_record_buffered(
&mut self,
suite: key_schedule::CipherSuite,
keys: &key_schedule::TrafficKeys,
content_type: u8,
payload: &[u8],
) -> TlsResult<()> {
let seq = self.write_seq.next()?;
let rec = seal_record(suite, keys, seq, content_type, payload)?;
self.writer.write_all(&rec).map_err(TlsError::from)
}
pub(crate) fn read_encrypted_handshake(
&mut self,
suite: key_schedule::CipherSuite,
keys: &key_schedule::TrafficKeys,
) -> TlsResult<Vec<u8>> {
let mut plain = Vec::new();
loop {
let mut header = [0u8; 5];
self.reader
.read_exact_into(&mut header)
.map_err(TlsError::from)?;
let len = ((header[3] as usize) << 8) | header[4] as usize;
if len > MAX_RECORD_PAYLOAD + 16 {
return Err(TlsError::Protocol("record too large".into()));
}
let encrypted = self.reader.read_exact(len).map_err(TlsError::from)?;
if header[0] == record::CONTENT_CHANGE_CIPHER_SPEC {
if len != 1 || encrypted.first() != Some(&1) {
return Err(TlsError::Protocol("malformed ChangeCipherSpec".into()));
}
continue;
}
let seq = self.read_seq.next()?;
let (ct, payload) = open_record(suite, keys, seq, &header, &encrypted)?;
if ct == CONTENT_HANDSHAKE {
if plain.len() > MAX_HANDSHAKE_BUFFER - payload.len() {
return Err(TlsError::Protocol(
"handshake message exceeds the 16 MiB protocol maximum".into(),
));
}
plain.extend_from_slice(&payload);
if handshake::has_complete_finished(&plain) {
return Ok(plain);
}
}
}
}
pub(crate) fn reset_sequences(&mut self) {
self.read_seq = Sequence::default();
self.write_seq = Sequence::default();
}
}
enum RecState {
Idle,
Header { hdr: [u8; 5], filled: usize },
Payload {
header: [u8; 5],
payload: Vec<u8>,
filled: usize,
},
}
pub struct TlsStream<R, W> {
io: TlsIo<R, W>,
suite: key_schedule::CipherSuite,
write_keys: key_schedule::TrafficKeys,
read_keys: key_schedule::TrafficKeys,
negotiated_alpn: Option<Vec<u8>>,
server_name: Option<String>,
peer_certificate: Option<Vec<u8>>,
closed: bool,
pending: Vec<u8>,
pending_pos: usize,
rec: RecState,
}
impl<R: crate::courierust_io::Read, W: crate::courierust_io::Write> TlsStream<R, W> {
pub fn alpn(&self) -> Option<&[u8]> {
self.negotiated_alpn.as_deref()
}
pub(crate) fn underlying(&self) -> &R {
self.io.reader.get_ref()
}
pub fn server_name(&self) -> Option<&str> {
self.server_name.as_deref()
}
pub fn peer_certificate(&self) -> Option<&[u8]> {
self.peer_certificate.as_deref()
}
pub fn cipher_suite(&self) -> u16 {
self.suite.wire()
}
pub fn write_all(&mut self, data: &[u8]) -> TlsResult<()> {
if self.closed {
return Err(TlsError::Protocol("connection closed".into()));
}
let mut off = 0;
while off < data.len() {
let take = core::cmp::min(data.len() - off, record::MAX_RECORD_PAYLOAD - 2);
self.io.write_encrypted_record_buffered(
self.suite,
&self.write_keys,
record::CONTENT_APPLICATION_DATA,
&data[off..off + take],
)?;
off += take;
}
self.io.writer.flush().map_err(TlsError::from)?;
Ok(())
}
pub fn read_record(&mut self) -> TlsResult<Vec<u8>> {
if self.closed {
return Ok(Vec::new());
}
loop {
let payload_len = match &self.rec {
RecState::Payload { payload, .. } => payload.len(),
_ => match self.read_record_header()? {
Some((_, len)) => len,
None => return Ok(Vec::new()), },
};
if payload_len > record::MAX_RECORD_PAYLOAD + 16 {
return Err(TlsError::Protocol("record too large".into()));
}
let (header, encrypted) = self.read_record_payload()?;
if header[0] == record::CONTENT_CHANGE_CIPHER_SPEC {
if encrypted.len() != 1 || encrypted[0] != 1 {
return Err(TlsError::Protocol("malformed ChangeCipherSpec".into()));
}
continue;
}
let seq = self.io.read_seq.next()?;
let (ct, payload) = open_record(self.suite, &self.read_keys, seq, &header, &encrypted)?;
match ct {
record::CONTENT_APPLICATION_DATA => return Ok(payload),
record::CONTENT_ALERT => {
if payload.first() == Some(&1) && payload.get(1) == Some(&0) {
self.closed = true;
return Ok(Vec::new());
}
return Err(TlsError::Alert {
level: payload.first().copied().unwrap_or(2),
description: payload.get(1).copied().unwrap_or(0),
});
}
record::CONTENT_HANDSHAKE => {
if let Some(m) = handshake::peek_complete_hs(&payload) {
if m.msg_type != handshake::HS_NEW_SESSION_TICKET {
return Err(TlsError::Protocol(
"unexpected handshake after handshake".into(),
));
}
}
}
record::CONTENT_CHANGE_CIPHER_SPEC => {
continue;
}
_ => {
return Err(TlsError::Protocol("unexpected record type".into()));
}
}
}
}
fn read_record_header(&mut self) -> TlsResult<Option<([u8; 5], usize)>> {
let (mut hdr, mut filled) = match &self.rec {
RecState::Header { hdr, filled } => (*hdr, *filled),
_ => ([0u8; 5], 0),
};
loop {
if filled == 5 {
let payload_len = ((hdr[3] as usize) << 8) | hdr[4] as usize;
self.rec = RecState::Payload {
header: hdr,
payload: vec![0u8; payload_len],
filled: 0,
};
return Ok(Some((hdr, payload_len)));
}
match self.io.reader.fill_buf() {
Ok([]) => {
self.closed = true;
return Ok(None);
}
Ok(b) => {
let take = core::cmp::min(5 - filled, b.len());
hdr[filled..filled + take].copy_from_slice(&b[..take]);
self.io.reader.consume(take);
filled += take;
}
Err(e)
if e.kind == crate::courierust_error::ErrorKind::Timeout
|| e.kind == crate::courierust_error::ErrorKind::WouldBlock =>
{
self.rec = RecState::Header { hdr, filled };
return Err(TlsError::Timeout);
}
Err(e) if e.kind == crate::courierust_error::ErrorKind::UnexpectedEof => {
self.closed = true;
return Ok(None);
}
Err(e) => return Err(TlsError::from(e)),
}
}
}
fn read_record_payload(&mut self) -> TlsResult<([u8; 5], Vec<u8>)> {
let (header, mut payload, mut filled) =
match core::mem::replace(&mut self.rec, RecState::Idle) {
RecState::Payload {
header,
payload,
filled,
} => (header, payload, filled),
_ => return Err(TlsError::Internal("payload read without header".into())),
};
let total = payload.len();
loop {
if filled == total {
return Ok((header, payload));
}
match self.io.reader.fill_buf() {
Ok([]) => {
self.closed = true;
return Err(TlsError::UnexpectedEof);
}
Ok(b) => {
let take = core::cmp::min(total - filled, b.len());
payload[filled..filled + take].copy_from_slice(&b[..take]);
self.io.reader.consume(take);
filled += take;
}
Err(e)
if e.kind == crate::courierust_error::ErrorKind::Timeout
|| e.kind == crate::courierust_error::ErrorKind::WouldBlock =>
{
self.rec = RecState::Payload {
header,
payload,
filled,
};
return Err(TlsError::Timeout);
}
Err(e) if e.kind == crate::courierust_error::ErrorKind::UnexpectedEof => {
self.closed = true;
return Err(TlsError::UnexpectedEof);
}
Err(e) => return Err(TlsError::from(e)),
}
}
}
pub fn close_notify(&mut self) -> TlsResult<()> {
if !self.closed {
let alert = [1u8, 0u8]; self.io.write_encrypted_record(
self.suite,
&self.write_keys,
record::CONTENT_ALERT,
&alert,
)?;
}
self.closed = true;
Ok(())
}
}
impl<R: crate::courierust_io::Read, W: crate::courierust_io::Write> crate::courierust_io::Read
for TlsStream<R, W>
{
fn read(&mut self, buf: &mut [u8]) -> crate::courierust_error::Result<usize> {
if self.pending_pos >= self.pending.len() {
self.pending = match self.read_record() {
Ok(p) => p,
Err(TlsError::Timeout) => {
return Err(crate::courierust_error::Error::new(
crate::courierust_error::ErrorKind::Timeout,
))
}
Err(e) => {
return Err(crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
e.to_string(),
))
}
};
self.pending_pos = 0;
if self.pending.is_empty() {
return Ok(0);
}
}
let avail = self.pending.len() - self.pending_pos;
let n = core::cmp::min(buf.len(), avail);
buf[..n].copy_from_slice(&self.pending[self.pending_pos..self.pending_pos + n]);
self.pending_pos += n;
Ok(n)
}
}
impl<R: crate::courierust_io::Read, W: crate::courierust_io::Write> crate::courierust_io::Write
for TlsStream<R, W>
{
fn write(&mut self, buf: &[u8]) -> crate::courierust_error::Result<usize> {
self.write_all(buf).map_err(|e| {
crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
e.to_string(),
)
})?;
Ok(buf.len())
}
fn flush(&mut self) -> crate::courierust_error::Result<()> {
self.io.writer.flush().map_err(|e| {
crate::courierust_error::Error::with_message(
crate::courierust_error::ErrorKind::Other,
e.to_string(),
)
})
}
}
pub struct ClientConfig {
pub roots: RootStore,
pub verify: bool,
pub alpn: Vec<Vec<u8>>,
pub now: i64,
}
impl Default for ClientConfig {
fn default() -> Self {
Self {
roots: RootStore::new(),
verify: true,
alpn: Vec::new(),
now: 0,
}
}
}
pub struct TlsConnector {
config: ClientConfig,
}
impl TlsConnector {
pub fn new(config: ClientConfig) -> Self {
Self { config }
}
pub fn config(&self) -> &ClientConfig {
&self.config
}
pub fn connect<R: crate::courierust_io::Read, W: crate::courierust_io::Write>(
&self,
hostname: &str,
reader: R,
writer: W,
) -> TlsResult<TlsStream<R, W>> {
let mut io = TlsIo::new(reader, writer);
let hs = handshake::ClientHandshake {
alpn: self.config.alpn.clone(),
server_name: Some(hostname.to_string()),
verify: self.config.verify,
};
let result = hs.run(&mut io, &self.config.roots, self.config.now)?;
io.reset_sequences();
Ok(TlsStream {
io,
suite: result.suite,
write_keys: result.keys.write,
read_keys: result.keys.read,
negotiated_alpn: result.alpn,
server_name: result.server_name,
peer_certificate: result.peer_cert,
closed: false,
pending: Vec::new(),
pending_pos: 0,
rec: RecState::Idle,
})
}
}
pub struct ServerConfig {
pub identity: Identity,
pub alpn: Vec<Vec<u8>>,
}
pub struct TlsAcceptor {
config: ServerConfig,
}
impl TlsAcceptor {
pub fn new(config: ServerConfig) -> Self {
Self { config }
}
pub fn config(&self) -> &ServerConfig {
&self.config
}
pub fn accept<R: crate::courierust_io::Read, W: crate::courierust_io::Write>(
&self,
reader: R,
writer: W,
) -> TlsResult<TlsStream<R, W>> {
let mut io = TlsIo::new(reader, writer);
let hs = handshake::ServerHandshake {
identity: self.config.identity.clone(),
alpn: self.config.alpn.clone(),
};
let result = hs.run(&mut io)?;
io.reset_sequences();
Ok(TlsStream {
io,
suite: result.suite,
write_keys: result.keys.write,
read_keys: result.keys.read,
negotiated_alpn: result.alpn,
server_name: result.server_name,
peer_certificate: result.peer_cert,
closed: false,
pending: Vec::new(),
pending_pos: 0,
rec: RecState::Idle,
})
}
}
pub(crate) fn server_sign(
identity: &Identity,
message: &[u8],
) -> TlsResult<Option<(u16, Vec<u8>)>> {
sign::sign_server_cert_verify(identity, message)
}
mod sign;
#[cfg(test)]
mod tests {
use super::*;
use std::net::{TcpListener, TcpStream};
#[test]
fn tls13_handshake_roundtrip() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (stream, _) = listener.accept().unwrap();
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: vec![b"h2".to_vec()],
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.alpn(), Some(&b"h2"[..]));
let data = tls.read_record().unwrap();
assert_eq!(data, b"ping");
tls.write_all(b"pong").unwrap();
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: vec![b"h2".to_vec()],
now: testdata::NOW,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.alpn(), Some(&b"h2"[..]));
assert!(tls.peer_certificate().is_some());
tls.write_all(b"ping").unwrap();
let data = tls.read_record().unwrap();
assert_eq!(data, b"pong");
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_client_rejects_untrusted_and_hostname() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
for _ in 0..2 {
let (stream, _) = listener.accept().unwrap();
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: Vec::new(),
});
let _ = acceptor.accept(&stream, &stream);
}
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: RootStore::new(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
});
let err = match connector.connect("localhost", &stream, &stream) {
Ok(_) => panic!("untrusted root accepted"),
Err(e) => e,
};
assert!(matches!(err, TlsError::Certificate(_)), "got {err:?}");
drop(stream);
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
});
let err = match connector.connect("not-localhost", &stream, &stream) {
Ok(_) => panic!("hostname mismatch accepted"),
Err(e) => e,
};
assert!(matches!(err, TlsError::Certificate(_)), "got {err:?}");
drop(stream);
let _ = server.join();
}
}