pub mod crypto;
pub(crate) mod session;
pub mod system_roots;
pub mod x509;
#[cfg(test)]
pub(crate) mod testdata;
mod client_auth;
mod handshake;
mod identity;
mod key_schedule;
mod pem;
pub(crate) mod quic;
mod record;
mod tls12;
pub use client_auth::ClientAuth;
pub use identity::Identity;
use alloc::string::String;
use alloc::vec::Vec;
pub use x509::RootStore;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum TlsVersion {
Tls12,
Tls13,
}
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()),
}
}
}
impl From<TlsError> for crate::courierust_error::Error {
fn from(e: TlsError) -> Self {
use crate::courierust_error::{Error, ErrorKind};
match e {
TlsError::Io(message) => Error::io(message),
TlsError::Protocol(message) => Error::protocol(message),
TlsError::Timeout => Error::new(ErrorKind::Timeout),
TlsError::UnexpectedEof => Error::new(ErrorKind::UnexpectedEof),
other => Error::with_message(ErrorKind::Other, alloc::format!("{other}")),
}
}
}
pub type TlsResult<T> = core::result::Result<T, TlsError>;
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<()> {
self.write_plaintext_record_v([0x03, 0x01], content_type, payload)
}
pub(crate) fn write_plaintext_record_v(
&mut self,
version: [u8; 2],
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,
version[0],
version[1],
(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 write_tls12_record(
&mut self,
suite: tls12::Tls12Suite,
keys: &key_schedule::TrafficKeys,
content_type: u8,
payload: &[u8],
) -> TlsResult<()> {
let seq = self.write_seq.next()?;
let rec = tls12::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_tls12_record_buffered(
&mut self,
suite: tls12::Tls12Suite,
keys: &key_schedule::TrafficKeys,
content_type: u8,
payload: &[u8],
) -> TlsResult<()> {
let seq = self.write_seq.next()?;
let rec = tls12::seal_record(suite, keys, seq, content_type, payload)?;
self.writer.write_all(&rec).map_err(TlsError::from)
}
pub(crate) fn read_tls12_record(
&mut self,
suite: tls12::Tls12Suite,
keys: &key_schedule::TrafficKeys,
) -> TlsResult<(u8, Vec<u8>)> {
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 > record::MAX_RECORD_PAYLOAD + 8 + 16 {
return Err(TlsError::Protocol("record too large".into()));
}
let body = self.reader.read_exact(len).map_err(TlsError::from)?;
if header[0] == record::CONTENT_CHANGE_CIPHER_SPEC {
if len != 1 || body.first() != Some(&1) {
return Err(TlsError::Protocol("malformed ChangeCipherSpec".into()));
}
return Ok((record::CONTENT_CHANGE_CIPHER_SPEC, body));
}
let seq = self.read_seq.next()?;
tls12::open_record(suite, keys, seq, &header, &body)
}
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,
},
}
fn key_use_limit(suite: key_schedule::CipherSuite) -> u64 {
match suite {
key_schedule::CipherSuite::TlsChaCha20Poly1305Sha256 => u64::MAX,
_ => 1 << 24,
}
}
pub struct TlsStream<R, W> {
io: TlsIo<R, W>,
version: TlsVersion,
suite: key_schedule::CipherSuite,
suite12: Option<tls12::Tls12Suite>,
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,
resumption_master: Option<(Vec<u8>, key_schedule::CipherSuite)>,
resumed: bool,
session_store: Option<std::sync::Arc<std::sync::Mutex<Vec<session::ClientSession>>>>,
hostname: String,
now: i64,
write_app_secret: Option<Vec<u8>>,
read_app_secret: Option<Vec<u8>>,
pending_key_update: bool,
write_records: u64,
key_use_limit: u64,
key_read_gen: u64,
key_write_gen: u64,
}
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 fn key_generations(&self) -> (u64, u64) {
(self.key_read_gen, self.key_write_gen)
}
pub fn request_key_update(&mut self) -> TlsResult<()> {
if self.version != TlsVersion::Tls13 {
return Err(TlsError::Protocol("KeyUpdate requires TLS 1.3".into()));
}
if self.closed {
return Err(TlsError::Protocol("connection closed".into()));
}
self.send_key_update(true)?;
self.io.writer.flush().map_err(TlsError::from)
}
fn send_key_update(&mut self, request: bool) -> TlsResult<()> {
let Some(secret) = self.write_app_secret.clone() else {
return Err(TlsError::Protocol(
"KeyUpdate without application keys".into(),
));
};
let msg = [handshake::HS_KEY_UPDATE, 0, 0, 1, u8::from(request)];
self.io.write_encrypted_record_buffered(
self.suite,
&self.write_keys,
record::CONTENT_HANDSHAKE,
&msg,
)?;
let next = key_schedule::update_traffic_secret(self.suite.hash(), &secret);
self.write_keys = key_schedule::TrafficKeys::from_secret(self.suite, &next);
self.write_app_secret = Some(next);
self.io.write_seq = Sequence::default();
self.write_records = 0;
self.key_write_gen += 1;
Ok(())
}
fn apply_key_update(&mut self, body: &[u8]) -> TlsResult<()> {
let request = match body {
[0] => false,
[1] => true,
_ => {
return Err(TlsError::Alert {
level: 2,
description: 47,
})
}
};
let Some(secret) = self.read_app_secret.clone() else {
return Err(TlsError::Protocol(
"KeyUpdate without application keys".into(),
));
};
let next = key_schedule::update_traffic_secret(self.suite.hash(), &secret);
self.read_keys = key_schedule::TrafficKeys::from_secret(self.suite, &next);
self.read_app_secret = Some(next);
self.io.read_seq = Sequence::default();
self.key_read_gen += 1;
if request {
self.pending_key_update = true;
}
Ok(())
}
#[cfg(test)]
pub(crate) fn set_key_use_limit(&mut self, limit: u64) {
self.key_use_limit = limit;
}
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 version(&self) -> TlsVersion {
self.version
}
pub fn cipher_suite(&self) -> u16 {
match self.version {
TlsVersion::Tls13 => self.suite.wire(),
TlsVersion::Tls12 => self.suite12.map(|s| s.wire()).unwrap_or(0),
}
}
pub fn write_all(&mut self, data: &[u8]) -> TlsResult<()> {
if self.closed {
return Err(TlsError::Protocol("connection closed".into()));
}
if self.version == TlsVersion::Tls13 {
if self.pending_key_update {
self.pending_key_update = false;
self.send_key_update(false)?;
} else if self.write_records >= self.key_use_limit {
self.send_key_update(false)?;
}
}
let mut off = 0;
while off < data.len() {
let take = core::cmp::min(data.len() - off, record::MAX_RECORD_PAYLOAD - 2);
match self.version {
TlsVersion::Tls13 => self.io.write_encrypted_record_buffered(
self.suite,
&self.write_keys,
record::CONTENT_APPLICATION_DATA,
&data[off..off + take],
)?,
TlsVersion::Tls12 => {
let suite12 = self
.suite12
.ok_or_else(|| TlsError::Internal("missing TLS 1.2 suite".into()))?;
self.io.write_tls12_record_buffered(
suite12,
&self.write_keys,
record::CONTENT_APPLICATION_DATA,
&data[off..off + take],
)?
}
}
self.write_records = self.write_records.saturating_add(1);
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());
}
if self.pending_pos < self.pending.len() {
let out = self.pending[self.pending_pos..].to_vec();
self.pending_pos = self.pending.len();
return Ok(out);
}
loop {
match self.version {
TlsVersion::Tls13 => match self.read_one_tls13_record()? {
None => return Ok(Vec::new()), Some((ct, payload)) => 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 => {
let mut rest = &payload[..];
while let Some(m) = handshake::peek_complete_hs(rest) {
match m.msg_type {
handshake::HS_NEW_SESSION_TICKET => {
self.capture_ticket(rest)?;
}
handshake::HS_KEY_UPDATE => {
self.apply_key_update(m.body)?;
}
_ => {
return Err(TlsError::Protocol(
"unexpected handshake after handshake".into(),
))
}
}
rest = &rest[4 + m.body.len()..];
}
}
record::CONTENT_CHANGE_CIPHER_SPEC => continue,
_ => {
return Err(TlsError::Protocol("unexpected record type".into()));
}
},
},
TlsVersion::Tls12 => {
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 + 8 + 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 suite12 = self
.suite12
.ok_or_else(|| TlsError::Internal("missing TLS 1.2 suite".into()))?;
let seq = self.io.read_seq.next()?;
let (ct, payload) =
tls12::open_record(suite12, &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 => {
return Err(TlsError::Protocol(
"unexpected handshake after handshake".into(),
));
}
_ => {
return Err(TlsError::Protocol("unexpected record type".into()));
}
}
}
}
}
}
fn read_one_tls13_record(&mut self) -> TlsResult<Option<(u8, Vec<u8>)>> {
loop {
let payload_len = match &self.rec {
RecState::Payload { payload, .. } => payload.len(),
_ => match self.read_record_header()? {
Some((_, len)) => len,
None => return Ok(None),
},
};
if payload_len > record::MAX_RECORD_PAYLOAD + 8 + 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;
}
if header[0] != record::CONTENT_APPLICATION_DATA {
return Err(TlsError::Protocol("unexpected record type".into()));
}
let seq = self.io.read_seq.next()?;
let (ct, payload) =
record::open_record(self.suite, &self.read_keys, seq, &header, &encrypted)?;
return Ok(Some((ct, payload)));
}
}
fn capture_ticket(&mut self, payload: &[u8]) -> TlsResult<()> {
let Some((res_master, suite)) = self.resumption_master.clone() else {
return Ok(()); };
let mut rest = payload;
while !rest.is_empty() {
let Some(m) = handshake::peek_complete_hs(rest) else {
break;
};
if m.msg_type == handshake::HS_NEW_SESSION_TICKET {
if let Ok(ticket) = session::parse_new_session_ticket(m.body) {
let psk = session::derive_psk(suite, &res_master, &ticket.nonce);
let sess = session::ClientSession {
hostname: self.hostname.clone(),
ticket: ticket.ticket,
psk,
suite,
issued_at: self.now,
lifetime: (ticket.lifetime as i64).min(session::SESSION_LIFETIME_SECS),
};
if let Some(store) = &self.session_store {
cache_session(&mut crate::lock(store), sess);
}
}
}
rest = &rest[4 + m.body.len()..];
}
Ok(())
}
pub fn resumed(&self) -> bool {
self.resumed
}
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]; match self.version {
TlsVersion::Tls13 => self.io.write_encrypted_record(
self.suite,
&self.write_keys,
record::CONTENT_ALERT,
&alert,
)?,
TlsVersion::Tls12 => {
let suite12 = self
.suite12
.ok_or_else(|| TlsError::Internal("missing TLS 1.2 suite".into()))?;
self.io.write_tls12_record(
suite12,
&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(),
)
})
}
}
#[derive(Clone)]
pub struct ClientConfig {
pub roots: RootStore,
pub verify: bool,
pub alpn: Vec<Vec<u8>>,
pub now: i64,
pub min_version: TlsVersion,
pub max_version: TlsVersion,
pub identity: Option<Identity>,
}
impl Default for ClientConfig {
fn default() -> Self {
Self {
roots: RootStore::new(),
verify: true,
alpn: Vec::new(),
now: 0,
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls13,
identity: None,
}
}
}
#[derive(Clone)]
pub struct TlsConnector {
config: ClientConfig,
sessions: std::sync::Arc<std::sync::Mutex<Vec<session::ClientSession>>>,
}
impl TlsConnector {
pub fn new(config: ClientConfig) -> Self {
Self {
config,
sessions: std::sync::Arc::new(std::sync::Mutex::new(Vec::new())),
}
}
pub fn config(&self) -> &ClientConfig {
&self.config
}
pub fn clear_sessions(&self) {
crate::lock(&self.sessions).clear();
}
pub fn session_count(&self) -> usize {
crate::lock(&self.sessions).len()
}
fn find_session(&self, hostname: &str) -> Option<session::ClientSession> {
let sessions = crate::lock(&self.sessions);
sessions
.iter()
.find(|s| s.hostname == hostname)
.filter(|s| s.is_fresh(self.config.now))
.cloned()
}
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 allow13 = self.config.max_version >= TlsVersion::Tls13;
let allow12 = self.config.min_version <= TlsVersion::Tls12;
let mut random = [0u8; 32];
handshake::fill_entropy(&mut random)?;
let mut priv13 = [0u8; 32];
handshake::fill_entropy(&mut priv13)?;
let pub13 = crypto::x25519::x25519(&priv13, &crypto::x25519::BASE_POINT);
let resume_session = allow13.then(|| self.find_session(hostname)).flatten();
let ch = match &resume_session {
Some(s) => session::build_client_hello_with_psk(
&random,
&pub13,
&self.config.alpn,
Some(hostname),
s.suite,
&s.psk,
&s.ticket,
)?,
None => handshake::build_client_hello_negotiated(
&random,
Some(&pub13),
&self.config.alpn,
Some(hostname),
None,
allow13,
allow12,
),
};
io.write_plaintext_record_v(tls12::VERSION_12, record::CONTENT_HANDSHAKE, &ch)?;
let (ct, first_payload) = io.read_plaintext_record()?;
if ct != record::CONTENT_HANDSHAKE {
return Err(TlsError::Protocol("expected handshake record".into()));
}
if first_payload.len() < 4 {
return Err(TlsError::Protocol("bad handshake record".into()));
}
let sh_body = &first_payload[4..];
if handshake::is_hello_retry_request(sh_body) {
let hrr = handshake::parse_hello_retry_request(sh_body)?;
if hrr.selected_group == handshake::GROUP_X25519
|| (hrr.selected_group != tls12::GROUP_SECP256R1)
{
let _ = io.write_plaintext_record(record::CONTENT_ALERT, &[2, 47]); return Err(TlsError::Protocol(
"HelloRetryRequest selected an invalid key exchange group".into(),
));
}
let _ = io.write_plaintext_record(record::CONTENT_ALERT, &[2, 40]); return Err(TlsError::Unsupported(format!(
"HelloRetryRequest requested unsupported group 0x{:04x}",
hrr.selected_group
)));
}
let is13 = handshake::server_hello_negotiates_tls13(sh_body);
if is13 {
if !allow13 {
return Err(TlsError::Protocol(
"server negotiated TLS 1.3 but the client did not offer it".into(),
));
}
let hs = handshake::ClientHandshake {
server_name: Some(hostname.to_string()),
verify: self.config.verify,
psk: resume_session.map(|s| (s.psk, s.suite)),
identity: self.config.identity.clone(),
};
let result = hs.run_from_server_hello(
&mut io,
&self.config.roots,
self.config.now,
&ch,
&random,
&priv13,
sh_body,
None,
)?;
io.reset_sequences();
let stream = TlsStream {
io,
version: TlsVersion::Tls13,
suite: result.suite,
suite12: None,
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,
resumption_master: result.resumption_master.map(|m| (m, result.suite)),
resumed: result.resumed,
session_store: Some(self.sessions.clone()),
hostname: hostname.to_string(),
now: self.config.now,
write_app_secret: Some(result.keys.write_secret),
read_app_secret: Some(result.keys.read_secret),
pending_key_update: false,
write_records: 0,
key_use_limit: key_use_limit(result.suite),
key_read_gen: 0,
key_write_gen: 0,
};
Ok(stream)
} else {
if !allow12 {
return Err(TlsError::Protocol(
"server negotiated TLS 1.2 but the client refused it".into(),
));
}
let result = tls12::client_handshake(
&mut io,
&self.config.roots,
self.config.now,
self.config.verify,
Some(hostname),
&ch,
&random,
&first_payload,
allow13,
)?;
Ok(TlsStream {
io,
version: TlsVersion::Tls12,
suite: key_schedule::CipherSuite::TlsAes128GcmSha256,
suite12: Some(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,
resumption_master: None,
resumed: false,
session_store: None,
hostname: hostname.to_string(),
now: self.config.now,
write_app_secret: None,
read_app_secret: None,
pending_key_update: false,
write_records: 0,
key_use_limit: u64::MAX,
key_read_gen: 0,
key_write_gen: 0,
})
}
}
}
fn cache_session(store: &mut Vec<session::ClientSession>, sess: session::ClientSession) {
store.retain(|s| s.hostname != sess.hostname);
store.push(sess);
const MAX_SESSIONS: usize = 8;
if store.len() > MAX_SESSIONS {
let overflow = store.len() - MAX_SESSIONS;
store.drain(..overflow);
}
}
fn unix_now() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
pub struct ServerConfig {
pub identity: Identity,
pub alpn: Vec<Vec<u8>>,
pub min_version: TlsVersion,
pub max_version: TlsVersion,
pub session_ticket_key: Option<[u8; 32]>,
pub client_auth: Option<ClientAuth>,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
identity: Identity::empty(),
alpn: Vec::new(),
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
}
}
}
pub struct TlsAcceptor {
config: ServerConfig,
ticket_key: Option<[u8; 32]>,
}
impl TlsAcceptor {
pub fn new(config: ServerConfig) -> Self {
let ticket_key = match config.session_ticket_key {
Some(k) if k != [0u8; 32] => Some(k),
Some(_) => None,
None => {
let mut k = [0u8; 32];
crypto::rng::fill_random(&mut k).then_some(k)
}
};
Self { config, ticket_key }
}
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 (_, ch_body) = io.read_plaintext_handshake()?;
let allow13 = self.config.max_version >= TlsVersion::Tls13;
let allow12 = self.config.min_version <= TlsVersion::Tls12;
let client_offers_13 = handshake::client_hello_offers_tls13(&ch_body)?;
if client_offers_13 && allow13 {
let hs = handshake::ServerHandshake {
identity: self.config.identity.clone(),
alpn: self.config.alpn.clone(),
ticket_key: self.ticket_key,
now: unix_now(),
client_auth: self.config.client_auth.clone(),
};
let result = hs.run_from_client_hello(&mut io, &ch_body)?;
Ok(TlsStream {
io,
version: TlsVersion::Tls13,
suite: result.suite,
suite12: None,
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,
resumption_master: None,
resumed: result.resumed,
session_store: None,
hostname: String::new(),
now: 0,
write_app_secret: Some(result.keys.write_secret),
read_app_secret: Some(result.keys.read_secret),
pending_key_update: false,
write_records: 0,
key_use_limit: key_use_limit(result.suite),
key_read_gen: 0,
key_write_gen: 0,
})
} else if allow12 {
if self.config.client_auth.is_some() {
return Err(TlsError::Unsupported(
"client authentication requires TLS 1.3".into(),
));
}
let result = match tls12::server_handshake(
&mut io,
&self.config.identity,
&self.config.alpn,
&ch_body,
) {
Ok(r) => r,
Err(e) => {
let _ = io.write_plaintext_record(
record::CONTENT_ALERT,
&[2, 40], );
return Err(e);
}
};
Ok(TlsStream {
io,
version: TlsVersion::Tls12,
suite: key_schedule::CipherSuite::TlsAes128GcmSha256,
suite12: Some(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,
resumption_master: None,
resumed: false,
session_store: None,
hostname: String::new(),
now: 0,
write_app_secret: None,
read_app_secret: None,
pending_key_update: false,
write_records: 0,
key_use_limit: u64::MAX,
key_read_gen: 0,
key_write_gen: 0,
})
} else {
let _ = io.write_plaintext_record(
record::CONTENT_ALERT,
&[2, 70], );
Err(TlsError::Protocol(
"no acceptable TLS version between client and server".into(),
))
}
}
}
pub(crate) fn server_sign(
identity: &Identity,
message: &[u8],
suite: key_schedule::CipherSuite,
) -> TlsResult<Option<(u16, Vec<u8>)>> {
sign::sign_cert_verify(identity, message, suite)
}
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()],
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
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,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
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_key_update_rekeys_both_directions() {
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::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.read_record().unwrap(), b"before");
tls.write_all(b"after").unwrap();
assert_eq!(tls.read_record().unwrap(), b"again");
tls.write_all(b"end").unwrap();
assert_eq!(tls.key_generations(), (1, 1));
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
tls.write_all(b"before").unwrap();
tls.request_key_update().unwrap();
tls.write_all(b"again").unwrap();
assert_eq!(tls.read_record().unwrap(), b"after");
assert_eq!(tls.read_record().unwrap(), b"end");
assert_eq!(tls.key_generations(), (1, 1));
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_key_update_record_budget_forces_a_rekey() {
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::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
for expected in [&b"one"[..], b"two", b"three"] {
assert_eq!(tls.read_record().unwrap(), expected);
}
assert_eq!(tls.key_generations(), (1, 0));
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
tls.set_key_use_limit(2);
tls.write_all(b"one").unwrap();
tls.write_all(b"two").unwrap();
tls.write_all(b"three").unwrap();
assert_eq!(tls.key_generations(), (0, 1));
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(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
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,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
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,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
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();
}
#[test]
fn tls13_mtls_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 mut roots = RootStore::new();
roots.add_der(testdata::SERVER_CERT_DER.to_vec());
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: Some(ClientAuth::required(roots)),
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(
tls.peer_certificate(),
Some(testdata::SERVER_CERT_DER),
"the authenticating client's leaf must reach the server"
);
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::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: Some(testdata::server_identity()),
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
tls.write_all(b"ping").unwrap();
assert_eq!(tls.read_record().unwrap(), b"pong");
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_mtls_required_refuses_an_anonymous_client() {
for required in [true, false] {
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 mut roots = RootStore::new();
roots.add_der(testdata::SERVER_CERT_DER.to_vec());
let auth = if required {
ClientAuth::required(roots)
} else {
ClientAuth::optional(roots)
};
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: Some(auth),
});
acceptor
.accept(&stream, &stream)
.map(|tls| tls.peer_certificate().is_some())
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let outcome = connector.connect("localhost", &stream, &stream);
let server_outcome = server.join().unwrap();
if required {
assert!(
matches!(
server_outcome,
Err(TlsError::Alert {
level: 2,
description: 116
})
),
"a required certificate must be refused with certificate_required"
);
let err = match outcome {
Err(e) => e,
Ok(mut tls) => tls
.read_record()
.expect_err("the server required a client certificate"),
};
assert!(
matches!(
err,
TlsError::Alert {
level: 2,
description: 116
}
),
"got {err:?}"
);
} else {
assert!(
outcome.is_ok() && matches!(server_outcome, Ok(false)),
"an optional request admits an anonymous client"
);
if let Ok(mut tls) = outcome {
let _ = tls.close_notify();
}
}
}
}
#[test]
fn tls13_mtls_refuses_an_untrusted_client_chain() {
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 mut roots = RootStore::new();
roots.add_der(testdata::RSA_SERVER_CERT_DER.to_vec());
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: Some(ClientAuth::required(roots)),
});
acceptor.accept(&stream, &stream).map(|_| ())
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: Some(testdata::server_identity()),
});
let client_outcome = connector.connect("localhost", &stream, &stream);
let server_outcome = server.join().unwrap();
assert!(
server_outcome.is_err(),
"an untrusted client chain must be refused"
);
let err = match client_outcome {
Err(e) => e,
Ok(mut tls) => tls
.read_record()
.expect_err("the server refused this client chain"),
};
assert!(
matches!(
err,
TlsError::Alert {
level: 2,
description: 42
}
),
"got {err:?}"
);
}
#[test]
fn tls13_with_rsa_identity_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::rsa_server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
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::rsa_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
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 tls12_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::rsa_server_identity(),
alpn: vec![b"h2".to_vec()],
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
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::rsa_root_store(),
verify: true,
alpn: vec![b"h2".to_vec()],
now: testdata::NOW,
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
assert!(tls.peer_certificate().is_some());
assert_eq!(tls.alpn(), Some(&b"h2"[..]));
let suite = tls.cipher_suite();
assert!(
matches!(suite, 0xc02f | 0xc030 | 0xcca8),
"unexpected TLS 1.2 suite {suite:#06x}"
);
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_with_p384_identity_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::p384_server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
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::p384_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
assert_eq!(tls.cipher_suite(), 0x1302);
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 tls12_with_p384_identity_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::p384_server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
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::p384_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
assert_eq!(tls.cipher_suite(), 0xc02c);
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 tls_version_autonegotiates_to_13_when_offered() {
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::rsa_server_identity(),
alpn: Vec::new(),
..Default::default()
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::rsa_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
..Default::default()
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls13);
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls12_negotiated_with_sentinel_when_client_offers_13() {
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::rsa_server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::rsa_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
..Default::default()
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_only_client_refuses_tls12_server() {
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::rsa_server_identity(),
alpn: Vec::new(),
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
session_ticket_key: None,
client_auth: None,
});
let _ = acceptor.accept(&stream, &stream);
});
let stream = TcpStream::connect(addr).unwrap();
let connector = TlsConnector::new(ClientConfig {
roots: testdata::rsa_root_store(),
verify: true,
alpn: Vec::new(),
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let err = match connector.connect("localhost", &stream, &stream) {
Ok(_) => panic!("TLS 1.3-only client accepted a TLS 1.2 server"),
Err(e) => e,
};
assert!(
matches!(
err,
TlsError::Protocol(_) | TlsError::Alert { .. } | TlsError::UnexpectedEof
),
"got {err:?}"
);
drop(stream);
let _ = server.join();
}
#[test]
fn tls12_ed25519_identity_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()],
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
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,
min_version: TlsVersion::Tls12,
max_version: TlsVersion::Tls12,
identity: None,
});
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert_eq!(tls.version(), TlsVersion::Tls12);
assert!(tls.peer_certificate().is_some());
let suite = tls.cipher_suite();
assert!(
matches!(suite, 0xc02b | 0xc02c | 0xcca9),
"unexpected TLS 1.2 suite {suite:#06x}"
);
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_session_resumption_roundtrip() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let ticket_key = [0x5au8; 32];
let server = std::thread::spawn(move || {
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: vec![b"h2".to_vec()],
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: Some(ticket_key),
client_auth: None,
});
for _ in 0..2 {
let (stream, _) = listener.accept().unwrap();
let mut tls = acceptor.accept(&stream, &stream).unwrap();
let data = tls.read_record().unwrap();
assert_eq!(data, b"ping");
tls.write_all(b"pong").unwrap();
tls.close_notify().unwrap();
}
});
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: vec![b"h2".to_vec()],
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert!(!tls.resumed(), "first handshake must be full");
tls.write_all(b"ping").unwrap();
match tls.read_record() {
Ok(p) if p == b"pong" => {}
other => panic!("bad first read: {other:?}"),
}
tls.close_notify().unwrap();
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert!(tls.resumed(), "second handshake must resume from the PSK");
tls.write_all(b"ping").unwrap();
assert_eq!(tls.read_record().unwrap(), b"pong");
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_session_resumption_tampered_ticket_falls_back() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let ticket_key = [0x5bu8; 32];
let server = std::thread::spawn(move || {
let acceptor = TlsAcceptor::new(ServerConfig {
identity: testdata::server_identity(),
alpn: vec![b"h2".to_vec()],
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: Some(ticket_key),
client_auth: None,
});
for _ in 0..2 {
let (stream, _) = listener.accept().unwrap();
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.read_record().unwrap(), b"ping");
tls.write_all(b"pong").unwrap();
tls.close_notify().unwrap();
}
});
let connector = TlsConnector::new(ClientConfig {
roots: testdata::root_store(),
verify: true,
alpn: vec![b"h2".to_vec()],
now: testdata::NOW,
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
identity: None,
});
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert!(!tls.resumed());
tls.write_all(b"ping").unwrap();
assert_eq!(tls.read_record().unwrap(), b"pong");
tls.close_notify().unwrap();
{
let mut sessions = connector.sessions.lock().unwrap();
let s = sessions.last_mut().unwrap();
let i = s.ticket.len() - 1;
s.ticket[i] ^= 0x01;
}
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let mut tls = connector.connect("localhost", &stream, &stream).unwrap();
assert!(
!tls.resumed(),
"tampered ticket must fall back to a full handshake"
);
tls.write_all(b"ping").unwrap();
assert_eq!(tls.read_record().unwrap(), b"pong");
tls.close_notify().unwrap();
server.join().unwrap();
}
#[test]
fn tls13_server_sends_hello_retry_request_and_completes() {
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()],
min_version: TlsVersion::Tls13,
max_version: TlsVersion::Tls13,
session_ticket_key: None,
client_auth: None,
});
let mut tls = acceptor.accept(&stream, &stream).unwrap();
assert_eq!(tls.read_record().unwrap(), b"ping");
tls.write_all(b"pong").unwrap();
tls.close_notify().unwrap();
});
let stream = TcpStream::connect(addr).unwrap();
stream
.set_read_timeout(Some(std::time::Duration::from_secs(5)))
.unwrap();
let mut io = TlsIo::new(&stream, &stream);
let mut random = [0u8; 32];
handshake::fill_entropy(&mut random).unwrap();
let ch1 = handshake::build_client_hello_negotiated(
&random,
None,
&[b"h2".to_vec()],
Some("localhost"),
None,
true,
false,
);
io.write_plaintext_record_v(tls12::VERSION_12, record::CONTENT_HANDSHAKE, &ch1)
.unwrap();
let (ct, hrr_payload) = io.read_plaintext_record().unwrap();
assert_eq!(ct, record::CONTENT_HANDSHAKE);
assert!(handshake::is_hello_retry_request(&hrr_payload[4..]));
let hrr = handshake::parse_hello_retry_request(&hrr_payload[4..]).unwrap();
assert_eq!(hrr.selected_group, handshake::GROUP_X25519);
let mut priv2 = [0u8; 32];
handshake::fill_entropy(&mut priv2).unwrap();
let pub2 = crypto::x25519::x25519(&priv2, &crypto::x25519::BASE_POINT);
let ch2 = handshake::build_client_hello_negotiated(
&random,
Some(&pub2),
&[b"h2".to_vec()],
Some("localhost"),
None,
true,
false,
);
io.write_plaintext_record_v(tls12::VERSION_12, record::CONTENT_HANDSHAKE, &ch2)
.unwrap();
let (ct, sh_payload) = io.read_plaintext_record().unwrap();
assert_eq!(ct, record::CONTENT_HANDSHAKE);
let sh_body = &sh_payload[4..];
assert!(handshake::server_hello_negotiates_tls13(sh_body));
let hs = handshake::ClientHandshake {
server_name: Some("localhost".to_string()),
verify: true,
psk: None,
identity: None,
};
let result = hs
.run_from_server_hello(
&mut io,
&testdata::root_store(),
testdata::NOW,
&ch2,
&random,
&priv2,
sh_body,
Some((&ch1, &hrr_payload, &ch2)),
)
.unwrap();
io.reset_sequences();
let mut tls = TlsStream {
io,
version: TlsVersion::Tls13,
suite: result.suite,
suite12: None,
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,
resumption_master: None,
resumed: result.resumed,
session_store: None,
hostname: "localhost".to_string(),
now: testdata::NOW,
write_app_secret: Some(result.keys.write_secret),
read_app_secret: Some(result.keys.read_secret),
pending_key_update: false,
write_records: 0,
key_use_limit: key_use_limit(result.suite),
key_read_gen: 0,
key_write_gen: 0,
};
tls.write_all(b"ping").unwrap();
assert_eq!(tls.read_record().unwrap(), b"pong");
tls.close_notify().unwrap();
server.join().unwrap();
}
}