use crate::{
codec::{Decode as _, Encode as _},
collections::{ArrayVectorCopy, SingleTypeStorage},
misc::Lease,
net::{RoleTy, Stream},
rng::CryptoRng,
tls::{
AlertDescription, Alpn, CHANGE_CIPHER_SPEC, CipherSuite, DLFT_MAX_FRAGMENT_LENGTH,
HandshakePath, MaxFragmentLength, NamedGroup, ProtocolVersion, SignatureScheme, TlsBuffer,
TlsConfig, TlsCtx, TlsCtxSk, TlsError, TlsStream,
key_schedule::KeySchedule,
misc::{
fetch_rec_from_stream, manage_err_handshake, post_handshake_dec_error,
pre_handshake_dec_error, server_sig_msg, write_payloads,
},
protocol::{
certificate::{Certificate, CertificateEntry},
certificate_verify::CertificateVerify,
client_hello::ClientHello,
encrypted_extensions::EncryptedExtensions,
finished::Finished,
handshake::Handshake,
handshake_ty::HandshakeTy,
key_share_entry::KeyShareEntry,
record::Record,
record_content_ty::RecordContentTy,
server_hello::ServerHello,
},
read_record_info::ReadRecordInfo,
tls_decode_wrapper::TlsDecodeWrapper,
tls_encode_wrapper::TlsEncodeWrapper,
tls_hash::TlsHash,
},
};
#[derive(Debug)]
pub struct TlsAcceptor<RNG, S, TCG> {
buffer: TlsBuffer,
config: TCG,
handshake_path: HandshakePath,
key_schedule: KeySchedule,
max_fragment_length: u16,
max_fragment_length_send: u16,
named_group: NamedGroup,
rng: RNG,
signature_algorithms: ArrayVectorCopy<(SignatureScheme, u8), { SignatureScheme::len() }>,
stream: S,
transcript_hash: TlsHash,
}
impl<RNG, S, TCG, TCX> TlsAcceptor<RNG, S, TCG>
where
TCG: Lease<TlsConfig<TCX>> + SingleTypeStorage<Item = TCX>,
{
#[inline]
pub fn new(config: TCG, rng: RNG, stream: S) -> Self {
let cfg_ref = config.lease();
let key_schedule = KeySchedule::default();
let transcript_hash = key_schedule.cipher_suite().hash_new();
let max_fragment_length =
cfg_ref.max_fragment_length().map_or(DLFT_MAX_FRAGMENT_LENGTH, |el| el.num());
let max_fragment_length_send =
cfg_ref.max_fragment_length_send().map_or(DLFT_MAX_FRAGMENT_LENGTH, |el| el.num());
let named_group = cfg_ref
.inner
.supported_groups
.named_group_list
.first()
.copied()
.unwrap_or(NamedGroup::default());
let signature_algorithms = filter_signature_algorithms(cfg_ref);
Self {
buffer: TlsBuffer::new(),
config,
handshake_path: HandshakePath::Full,
key_schedule,
max_fragment_length,
max_fragment_length_send,
named_group,
rng,
signature_algorithms,
stream,
transcript_hash,
}
}
#[inline]
pub const fn handshake_path(&self) -> HandshakePath {
self.handshake_path
}
#[inline]
pub const fn named_group(&self) -> NamedGroup {
self.named_group
}
#[inline]
pub const fn rng(&self) -> &RNG {
&self.rng
}
#[inline]
pub const fn rng_mut(&mut self) -> &mut RNG {
&mut self.rng
}
}
impl<RNG, S, TCG, TCX> TlsAcceptor<RNG, S, TCG>
where
RNG: CryptoRng,
S: Stream,
TCG: Lease<TlsConfig<TCX>> + SingleTypeStorage<Item = TCX>,
TCX: TlsCtxSk,
{
#[inline]
pub async fn accept(mut self) -> crate::Result<TlsAcceptOutput<RNG, S, TCX>> {
if TCX::TY.is_plain_text() {
return Ok(TlsAcceptOutput {
handshake_path: self.handshake_path,
named_group: self.named_group,
rng: self.rng,
tls_stream: TlsStream::new(
self.buffer,
self.key_schedule,
self.max_fragment_length,
self.max_fragment_length_send,
self.stream,
)?,
});
}
_trace!(target: crate::_WTX_TLS_HS, "Start");
let fut = async {
let first_rri = self.fetch_rec_from_stream::<false, true>(false).await?;
_trace!(target: crate::_WTX_TLS_HS, "Read ClientHello: {:?}", &first_rri);
let indices = self.manage_initial_client_record(&first_rri)?;
let buffer = self.buffer.reader_buffer.buffer_mut();
let payloads = match indices.as_slice() {
[idx0, idx1] => {
&[buffer.get(*idx0..*idx1).unwrap_or_default(), buffer.get(*idx1..).unwrap_or_default()][..]
}
[idx0, idx1, idx2, idx3] => &[
buffer.get(*idx0..*idx1).unwrap_or_default(),
buffer.get(*idx1..*idx2).unwrap_or_default(),
buffer.get(*idx2..*idx3).unwrap_or_default(),
buffer.get(*idx3..).unwrap_or_default(),
],
_ => &[],
};
_trace!(target: crate::_WTX_TLS_HS, "Write Records");
write_payloads(
RecordContentTy::Handshake,
self.key_schedule.write_mut(),
self.max_fragment_length_send,
payloads,
&mut self.stream,
&mut self.buffer.writer_buffer,
)
.await?;
buffer.truncate(indices.first().copied().unwrap_or_default());
let mut last_rri = self.fetch_rec_from_stream::<true, false>(true).await?;
if last_rri.outer_ty == RecordContentTy::ChangeCipherSpec {
last_rri = self.fetch_rec_from_stream::<false, false>(true).await?;
}
_trace!(target: crate::_WTX_TLS_HS, "Read Finished: {:?}", &last_rri);
self.manage_final_client_record(&last_rri)?;
Ok(())
};
let rslt = fut.await;
let kss = self.key_schedule.write_mut().state_mut();
manage_err_handshake(true, kss, rslt, &mut self.stream).await?;
_trace!(target: crate::_WTX_TLS_HS, "Successful handshake");
Ok(TlsAcceptOutput {
handshake_path: self.handshake_path,
named_group: self.named_group,
rng: self.rng,
tls_stream: TlsStream::new(
self.buffer,
self.key_schedule,
self.max_fragment_length,
self.max_fragment_length_send,
self.stream,
)?,
})
}
#[inline]
pub fn manage_initial_client_record(
&mut self,
rri: &ReadRecordInfo,
) -> crate::Result<ArrayVectorCopy<usize, 4>> {
let RecordContentTy::Handshake = rri.outer_ty else {
return Err(TlsError::InvalidHandshakeTy.into());
};
let output = self.negotiate(rri)?;
self.buffer.reader_buffer.clear_if_exhausted();
let reader_buffer = self.buffer.reader_buffer.buffer_mut();
let mut curr_idx = reader_buffer.len();
let mut indices = ArrayVectorCopy::new();
let encrypted_extensions = Handshake::new(
HandshakeTy::EncryptedExtensions,
EncryptedExtensions::new(output.alpn, output.max_fragment_length, None, None),
);
encrypted_extensions.encode(&mut TlsEncodeWrapper::from_buffer(reader_buffer))?;
self.transcript_hash.update(reader_buffer.get(curr_idx..).unwrap_or_default());
drop(indices.push(curr_idx));
curr_idx = reader_buffer.len();
drop(indices.push(curr_idx));
let mut cert_list = ArrayVectorCopy::new();
let Some(public_key) = self.config.lease().public_keys().get(output.signature_scheme.1) else {
return Err(TlsError::UnsupportedSignAlgorithm.into());
};
for cert in public_key.certs() {
cert_list.push(CertificateEntry::new(cert))?;
}
{
let ty = HandshakeTy::Certificate;
let certificate = Handshake::new(ty, Certificate::new(cert_list, &[]));
certificate.encode(&mut TlsEncodeWrapper::from_buffer(reader_buffer))?;
self.transcript_hash.update(reader_buffer.get(curr_idx..).unwrap_or_default());
curr_idx = reader_buffer.len();
let _rslt = indices.push(curr_idx);
}
let signature;
let signature_ref = if TCX::TY.is_unverified() {
&[][..]
} else {
signature = self.config.lease().inner.ctx.sign(
&mut self.buffer.writer_buffer,
&server_sig_msg(self.transcript_hash.clone().finalize().lease())?,
&mut self.rng,
output.signature_scheme.0,
)?;
signature.as_ref()
};
{
let certificate_verify = Handshake::new(
HandshakeTy::CertificateVerify,
CertificateVerify::new(output.signature_scheme.0, signature_ref),
);
certificate_verify.encode(&mut TlsEncodeWrapper::from_buffer(reader_buffer))?;
self.transcript_hash.update(reader_buffer.get(curr_idx..).unwrap_or_default());
curr_idx = reader_buffer.len();
let _rslt = indices.push(curr_idx);
}
let verify_data = self
.key_schedule
.write_mut()
.state_mut()
.create_finished_verify_data(self.transcript_hash.clone().finalize().lease())?;
let finished = Handshake::new(HandshakeTy::Finished, Finished::new(verify_data.as_slice()));
finished.encode(&mut TlsEncodeWrapper::from_buffer(reader_buffer))?;
self.transcript_hash.update(reader_buffer.get(curr_idx..).unwrap_or_default());
Ok(indices)
}
#[inline]
pub fn manage_final_client_record(&mut self, rri: &ReadRecordInfo) -> crate::Result<()> {
let rslt = self.do_manage_final_client_record(rri);
self.key_schedule.master_secret::<false>(&self.transcript_hash.clone().finalize())?;
rslt
}
#[inline]
fn do_manage_final_client_record(&mut self, rri: &ReadRecordInfo) -> crate::Result<()> {
if rri.outer_ty != RecordContentTy::ApplicationData
|| rri.inner_ty != RecordContentTy::Handshake
{
return Err(TlsError::InvalidHandshakeTy.into());
}
let current = self.buffer.reader_buffer.current();
let plaintext = current.get(..rri.plaintext_len).unwrap_or_default();
let mut dw = TlsDecodeWrapper::from_bytes(plaintext);
let hs = Handshake::<&[u8]>::decode(&mut dw)?;
if hs.msg_type != HandshakeTy::Finished {
return Err(crate::Error::TlsErrorReply(
TlsError::ClientExpectedFinished,
AlertDescription::UnexpectedMessage,
));
}
let trailing_len = dw.bytes().len();
*dw.bytes_mut() = hs.data;
*dw.cipher_suite_mut() = self.key_schedule.cipher_suite();
let finished = Finished::decode(&mut dw)?;
post_handshake_dec_error(dw.bytes(), HandshakeTy::Finished)?;
if self
.key_schedule
.read_mut()
.state_mut()
.verify_finished_record(
self.transcript_hash.clone().finalize().lease(),
finished.verify_data(),
)
.is_err()
{
return Err(TlsError::DigestCheckFailed.into());
}
if trailing_len > 0 {
return Err(crate::Error::TlsErrorReply(
TlsError::ExcessHandshakeData(RoleTy::Server),
AlertDescription::UnexpectedMessage,
));
}
Ok(())
}
#[inline]
async fn fetch_rec_from_stream<const CHECK_CCS: bool, const IS_CH: bool>(
&mut self,
decrypt: bool,
) -> crate::Result<ReadRecordInfo> {
Ok(
fetch_rec_from_stream::<_, CHECK_CCS, IS_CH>(
decrypt.then(|| self.key_schedule.read_mut().state_mut()),
self.max_fragment_length,
&mut self.buffer.reader_buffer,
&mut self.stream,
)
.await?
.ok_or(TlsError::AbruptDisconnect)?,
)
}
#[inline]
fn negotiate(&mut self, rri: &ReadRecordInfo) -> crate::Result<NegotiateOutput>
where
TCX: TlsCtx,
{
let current = self.buffer.reader_buffer.current();
let plaintext = current.get(..rri.plaintext_len).unwrap_or_default();
let mut dw = TlsDecodeWrapper::from_bytes(plaintext);
let handshake = Handshake::<&[u8]>::decode(&mut dw)?;
*dw.bytes_mut() = handshake.data;
pre_handshake_dec_error(handshake.rec_len() != rri.plaintext_len)?;
let client_hello = ClientHello::decode(&mut dw)?;
post_handshake_dec_error(dw.bytes(), HandshakeTy::ClientHello)?;
if !client_hello
.supported_versions()
.versions
.iter()
.copied()
.any(|el| el == ProtocolVersion::Tls13)
{
return Err(
TlsError::UnsupportedTlsVersion(client_hello.supported_versions().versions.last().copied())
.into(),
);
}
let cipher_suite = seek_cipher_suite(
&client_hello.tls_config().cipher_suites,
&self.config.lease().inner.cipher_suites,
)?;
self.key_schedule.set_cipher_suite(cipher_suite);
self.key_schedule.early_secret()?;
self.transcript_hash = cipher_suite.hash_new();
self.transcript_hash.update(plaintext);
let key_share = seek_key_share(
&client_hello.generic().client_shares,
&self.config.lease().inner.supported_groups.named_group_list,
)?;
let alpn = seek_alpn(&client_hello.tls_config().alpn, &self.config.lease().inner.alpn)?;
self.named_group = key_share.group();
let max_fragment_length = client_hello.tls_config().max_fragment_length;
if let Some(client_mfg) = max_fragment_length {
let client_num = client_mfg.num();
if let Some(server_mfl) = self.config.lease().max_fragment_length()
&& client_num > server_mfl.num()
{
return Err(TlsError::InvalidNegotiatedMaxFragmentLength.into());
}
self.max_fragment_length = client_num;
self.max_fragment_length_send = self.max_fragment_length_send.min(client_num);
}
let Some(signature_scheme) = seek_signature_scheme(
&client_hello.tls_config().signature_algorithms.signature_schemes,
&self.signature_algorithms,
) else {
return Err(crate::Error::TlsErrorReply(
TlsError::ServerHasNoCompatibleSignatureScheme,
AlertDescription::HandshakeFailure,
));
};
let legacy_session_id = *client_hello.legacy_session_id();
let agreement = key_share.group().agreement(&mut self.rng)?;
let ephemeral_pk = agreement.public_key()?;
let secret = agreement.diffie_hellman::<false>(key_share.opaque())?;
let writer_buffer = &mut self.buffer.writer_buffer;
let server_hello_rec = Record::new(
RecordContentTy::Handshake,
ProtocolVersion::Tls12,
Handshake::new(
HandshakeTy::ServerHello,
ServerHello::new(
cipher_suite,
false,
KeyShareEntry::new(key_share.group(), ephemeral_pk.as_ref()),
legacy_session_id,
&mut self.rng,
),
),
);
writer_buffer.clear();
server_hello_rec.encode(&mut TlsEncodeWrapper::from_buffer(writer_buffer))?;
self.transcript_hash.update(writer_buffer.get(5..).unwrap_or_default());
writer_buffer.extend_from_copyable_slice(&CHANGE_CIPHER_SPEC)?;
self
.key_schedule
.handshake_secret::<false>(secret.as_ref(), &self.transcript_hash.clone().finalize())?;
Ok(NegotiateOutput { alpn, max_fragment_length, signature_scheme })
}
}
#[derive(Debug)]
pub struct TlsAcceptOutput<RNG, S, TCX> {
pub handshake_path: HandshakePath,
pub named_group: NamedGroup,
pub rng: RNG,
pub tls_stream: TlsStream<S, TCX, false>,
}
#[derive(Debug)]
struct NegotiateOutput {
alpn: Option<Alpn>,
max_fragment_length: Option<MaxFragmentLength>,
signature_scheme: (SignatureScheme, u8),
}
#[inline]
fn filter_signature_algorithms<TCX>(
cfg_ref: &TlsConfig<TCX>,
) -> ArrayVectorCopy<(SignatureScheme, u8), { SignatureScheme::len() }> {
let mut rslt = ArrayVectorCopy::new();
let mut local_signature_schemes = cfg_ref.signature_algorithms().signature_schemes;
let mut key_tys_idx: u8 = 0;
for cert_kt in cfg_ref.public_keys().key_tys() {
let mut idx = 0;
while idx < local_signature_schemes.len() {
let Some(local_signature_scheme) = local_signature_schemes.get(usize::from(idx)) else {
break;
};
if local_signature_scheme.cert_kt() == cert_kt {
drop(rslt.push((*local_signature_scheme, key_tys_idx)));
let _ = local_signature_schemes.swap_remove(idx);
if cfg_ref.unique_signature_algorithms() {
break;
}
} else {
idx = idx.wrapping_add(1);
}
}
key_tys_idx = key_tys_idx.wrapping_add(1);
}
rslt
}
#[inline]
fn seek_alpn(client_opt: &Option<Alpn>, server_opt: &Option<Alpn>) -> crate::Result<Option<Alpn>> {
let (Some(client), Some(server)) = (client_opt, server_opt) else {
return Ok(None);
};
if server.protocol_name_list.is_empty() {
return Err(crate::Error::TlsErrorReply(
TlsError::EmptyNegotiatedAlpnServer,
AlertDescription::InternalError,
));
}
let mut rslt = None;
for client_el in &client.protocol_name_list {
if client_el.is_empty() {
return Err(crate::Error::TlsErrorReply(
TlsError::EmptyNegotiatedAlpnClient,
AlertDescription::DecodeError,
));
}
if rslt.is_none() && server.protocol_name_list.contains(client_el) {
let mut alpn = Alpn::default();
let _rslt = alpn.protocol_name_list.push(*client_el);
rslt = Some(alpn);
}
}
if let Some(elem) = rslt {
Ok(Some(elem))
} else {
Err(crate::Error::TlsErrorReply(
TlsError::MismatchedNegotiatedAlpnServer,
AlertDescription::NoApplicationProtocol,
))
}
}
#[inline]
fn seek_cipher_suite(client: &[CipherSuite], server: &[CipherSuite]) -> crate::Result<CipherSuite> {
for elem in server {
if client.contains(elem) {
return Ok(*elem);
}
}
Err(TlsError::ServerHasNoCompatibleCypherSuite.into())
}
#[inline]
fn seek_key_share<'client, 'rslt, 'server>(
client: &'client [KeyShareEntry<&'client [u8]>],
server: &'server [NamedGroup],
) -> crate::Result<KeyShareEntry<&'rslt [u8]>>
where
'client: 'rslt,
'server: 'rslt,
{
for server_el in server {
let Some(client_el) = client.iter().find(|client_el| client_el.group() == *server_el) else {
continue;
};
return Ok(*client_el);
}
Err(crate::Error::TlsErrorReply(
TlsError::ServerHasNoCompatibleKeyShare,
AlertDescription::UnexpectedMessage,
))
}
#[inline]
fn seek_signature_scheme(
client_signature_schemes: &[SignatureScheme],
server_signature_schemes: &[(SignatureScheme, u8)],
) -> Option<(SignatureScheme, u8)> {
for (server_signature_scheme, idx) in server_signature_schemes {
if !client_signature_schemes.contains(server_signature_scheme) {
continue;
}
return Some((*server_signature_scheme, *idx));
}
None
}