use alloc::boxed::Box;
use alloc::vec;
use alloc::vec::Vec;
use core::borrow::Borrow;
use core::fmt;
use core::ops::Deref;
use pki_types::ServerName;
use super::config::{ClientSessionKey, Tls12Resumption};
use super::ech::{EchMode, EchState, EchStatus};
use super::{
ClientHelloDetails, ClientSessionCommon, Retrieved, Tls12Session, Tls13Session, tls12, tls13,
};
use crate::check::inappropriate_handshake_message;
use crate::common_state::{EarlyDataEvent, Event, Output, OutputEvent, Protocol};
use crate::conn::{Input, StateMachine};
use crate::crypto::cipher::Payload;
use crate::crypto::kx::{KeyExchangeAlgorithm, StartedKeyExchange, SupportedKxGroup};
use crate::crypto::{CipherSuite, CryptoProvider, rand};
use crate::enums::{
ApplicationProtocol, CertificateType, ContentType, HandshakeType, ProtocolVersion,
};
use crate::error::{ApiMisuse, Error, PeerIncompatible, PeerMisbehaved};
use crate::hash_hs::HandshakeHashBuffer;
use crate::kernel::KernelState;
use crate::log::{debug, trace};
use crate::msgs::{
CertificateStatusRequest, ClientExtensions, ClientExtensionsInput, ClientHelloPayload,
ClientSessionTicket, ClientTicketRequest, Compression, EncryptedClientHello, ExtensionType,
HandshakeMessagePayload, HandshakePayload, HelloRetryRequest, KeyShareEntry, Message,
MessagePayload, PskKeyExchangeModes, Random, ServerHelloPayload, ServerNamePayload, SessionId,
SupportedEcPointFormats, SupportedProtocolVersions, TransportParameters,
};
use crate::sealed::Sealed;
use crate::suites::{PartiallyExtractedSecrets, Suite, SupportedCipherSuite};
use crate::sync::Arc;
use crate::tls12::Tls12CipherSuite;
use crate::tls13::Tls13CipherSuite;
use crate::tls13::key_schedule::{KeyScheduleEarlyClient, KeyScheduleTrafficSend};
use crate::{ClientConfig, bs_debug};
#[expect(private_interfaces)]
pub(crate) enum ClientState {
ServerHello(Box<ExpectServerHello>),
ServerHelloOrHelloRetryRequest(Box<ExpectServerHelloOrHelloRetryRequest>),
Tls12(tls12::Tls12State),
Tls13(tls13::Tls13State),
}
impl StateMachine for ClientState {
fn handle<'m>(self, input: Input<'m>, output: &mut dyn Output<'m>) -> Result<Self, Error> {
match self {
Self::ServerHello(e) => e.handle(input, output),
Self::ServerHelloOrHelloRetryRequest(e) => e.handle(input, output),
Self::Tls12(sm) => sm.handle(input, output),
Self::Tls13(sm) => sm.handle(input, output),
}
}
fn wants_input(&self) -> bool {
true
}
fn is_traffic(&self) -> bool {
matches!(
self,
Self::Tls12(tls12::Tls12State::Traffic(..))
| Self::Tls13(tls13::Tls13State::Traffic(..) | tls13::Tls13State::QuicTraffic(..))
)
}
fn handle_decrypt_error(&mut self) {
if let Self::Tls12(tls12::Tls12State::Finished(e)) = self {
e.handle_decrypt_error();
}
}
fn into_external_state(
self,
send_keys: &Option<Box<KeyScheduleTrafficSend>>,
) -> Result<(PartiallyExtractedSecrets, Box<dyn KernelState + 'static>), Error> {
match self {
Self::Tls12(tls12::Tls12State::Traffic(e)) => e.into_external_state(send_keys),
Self::Tls13(tls13::Tls13State::Traffic(e)) => e.into_external_state(send_keys),
Self::Tls13(tls13::Tls13State::QuicTraffic(e)) => e.into_external_state(send_keys),
_ => Err(Error::HandshakeNotComplete),
}
}
}
pub(crate) struct ExpectServerHello {
pub(super) input: ClientHelloInput,
pub(super) transcript_buffer: HandshakeHashBuffer,
pub(super) early_data_key_schedule: Option<(KeyScheduleEarlyClient, bool)>,
pub(super) offered_key_share: Option<GroupAndKeyShare>,
pub(super) suite: Option<SupportedCipherSuite>,
pub(super) ech_state: Option<EchState>,
pub(super) ech_status: EchStatus,
pub(super) done_retry: bool,
}
impl ExpectServerHello {
fn with_version<T: Suite + 'static>(
mut self,
server_hello: &ServerHelloPayload,
input: &Input<'_>,
output: &mut dyn Output<'_>,
) -> Result<ClientState, Error>
where
CryptoProvider: Borrow<[&'static T]>,
SupportedCipherSuite: From<&'static T>,
{
if server_hello.compression_method != Compression::Null {
return Err(PeerMisbehaved::SelectedUnofferedCompression.into());
}
let allowed_unsolicited = [ExtensionType::RenegotiationInfo];
if self
.input
.hello
.server_sent_unsolicited_extensions(server_hello, &allowed_unsolicited)
{
return Err(PeerMisbehaved::UnsolicitedServerHelloExtension.into());
}
output.output(OutputEvent::ProtocolVersion(T::VERSION));
if T::VERSION != ProtocolVersion::TLSv1_3 {
process_alpn_protocol(
output,
&self.input.hello.alpn_protocols,
server_hello
.selected_protocol
.as_ref()
.map(|s| s.as_ref()),
self.input.config.check_selected_alpn,
)?;
}
if let Some(point_fmts) = &server_hello.ec_point_formats {
if !point_fmts.uncompressed {
return Err(PeerMisbehaved::ServerHelloMustOfferUncompressedEcPoints.into());
}
}
let suite = <CryptoProvider as Borrow<[&'static T]>>::borrow(self.input.config.provider())
.iter()
.find(|cs| cs.common().suite == server_hello.cipher_suite)
.ok_or(PeerMisbehaved::SelectedUnofferedCipherSuite)?;
match self.suite {
Some(prev_suite) if prev_suite.suite() != suite.common().suite => {
return Err(PeerMisbehaved::SelectedDifferentCipherSuiteAfterRetry.into());
}
_ => {
debug!("Using ciphersuite {suite:?}");
self.suite = Some(SupportedCipherSuite::from(suite));
output.output(OutputEvent::CipherSuite(SupportedCipherSuite::from(suite)));
}
}
suite
.client_handler()
.handle_server_hello(suite, server_hello, input, self, output)
}
}
impl ExpectServerHello {
fn handle(
self: Box<Self>,
input: Input<'_>,
output: &mut dyn Output<'_>,
) -> Result<ClientState, Error> {
let server_hello = require_handshake_msg!(
&input.message,
HandshakeType::ServerHello,
HandshakePayload::ServerHello
)?;
trace!("We got ServerHello {server_hello:#?}");
let config = &self.input.config;
let tls13_supported = config.supports_version(ProtocolVersion::TLSv1_3);
let server_version = if server_hello.legacy_version == ProtocolVersion::TLSv1_2 {
server_hello
.selected_version
.unwrap_or(server_hello.legacy_version)
} else {
server_hello.legacy_version
};
match server_version {
ProtocolVersion::TLSv1_3 if tls13_supported => {
self.with_version::<Tls13CipherSuite>(server_hello, &input, output)
}
ProtocolVersion::TLSv1_2 if config.supports_version(ProtocolVersion::TLSv1_2) => {
if let Some((_, true)) = &self.early_data_key_schedule {
return Err(PeerMisbehaved::OfferedEarlyDataWithOldProtocolVersion.into());
}
if server_hello.selected_version.is_some() {
return Err(PeerMisbehaved::SelectedTls12UsingTls13VersionExtension.into());
}
self.with_version::<Tls12CipherSuite>(server_hello, &input, output)
}
_ => {
let reason = match server_version {
ProtocolVersion::TLSv1_2 | ProtocolVersion::TLSv1_3 => {
PeerIncompatible::ServerTlsVersionIsDisabledByOurConfig
}
_ => PeerIncompatible::ServerDoesNotSupportTls12Or13,
};
Err(reason.into())
}
}
}
}
struct ExpectServerHelloOrHelloRetryRequest {
next: Box<ExpectServerHello>,
extra_exts: ClientExtensionsInput,
}
impl ExpectServerHelloOrHelloRetryRequest {
fn into_expect_server_hello(self) -> ClientState {
ClientState::ServerHello(self.next)
}
fn handle_hello_retry_request(
mut self,
input: Input<'_>,
output: &mut dyn Output<'_>,
) -> Result<ClientState, Error> {
let hrr = require_handshake_msg!(
input.message,
HandshakeType::HelloRetryRequest,
HandshakePayload::HelloRetryRequest
)?;
trace!("Got HRR {hrr:?}");
let proof = input.check_aligned_handshake()?;
let offered_key_share = self.next.offered_key_share.unwrap();
let config = &self.next.input.config;
if let (None, Some(req_group)) = (&hrr.cookie, hrr.key_share) {
let offered_hybrid = offered_key_share
.share
.as_hybrid_checked(&config.provider().kx_groups, ProtocolVersion::TLSv1_3)
.map(|(hybrid, _)| hybrid.component().0);
if req_group == offered_key_share.share.group() || Some(req_group) == offered_hybrid {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithOfferedGroup.into());
}
}
if let Some(cookie) = &hrr.cookie {
if cookie.is_empty() {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithEmptyCookie.into());
}
}
if hrr.cookie.is_none() && hrr.key_share.is_none() {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithNoChanges.into());
}
if hrr.session_id != self.next.input.session_id {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithWrongSessionId.into());
}
let Some(ProtocolVersion::TLSv1_3) = hrr.supported_versions else {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithUnsupportedVersion.into());
};
output.output(OutputEvent::ProtocolVersion(ProtocolVersion::TLSv1_3));
let Some(cs) = config.find_cipher_suite(hrr.cipher_suite) else {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithUnofferedCipherSuite.into());
};
if self.next.ech_status == EchStatus::NotOffered && hrr.encrypted_client_hello.is_some() {
return Err(PeerMisbehaved::IllegalHelloRetryRequestWithInvalidEch.into());
}
output.output(OutputEvent::CipherSuite(cs));
match (self.next.ech_state.as_ref(), cs) {
(Some(ech_state), SupportedCipherSuite::Tls13(tls13_cs))
if !ech_state.confirm_hrr_acceptance(hrr, tls13_cs)? =>
{
self.next.ech_status = EchStatus::Rejected;
output.emit(Event::EchStatus(EchStatus::Rejected));
}
(Some(_), SupportedCipherSuite::Tls12(_)) => {
unreachable!("ECH state should only be set when TLS 1.3 was negotiated")
}
_ => {}
};
let transcript = self
.next
.transcript_buffer
.start_hash(cs.hash_provider());
let mut transcript_buffer = transcript.into_hrr_buffer(&proof);
transcript_buffer.add_message(&input.message);
if let Some(ech_state) = self.next.ech_state.as_mut() {
ech_state.transcript_hrr_update(cs.hash_provider(), &input.message, &proof);
}
let key_share = match hrr.key_share {
Some(group) if group != offered_key_share.share.group() => {
let Some(skxg) = config
.provider()
.find_kx_group(group, ProtocolVersion::TLSv1_3)
else {
return Err(
PeerMisbehaved::IllegalHelloRetryRequestWithUnofferedNamedGroup.into(),
);
};
GroupAndKeyShare::new(skxg)?
}
_ => offered_key_share,
};
emit_client_hello_for_retry(
transcript_buffer,
Some(hrr),
Some(key_share),
self.extra_exts,
Some(cs),
self.next.input,
output,
self.next.ech_state,
self.next.ech_status,
)
}
}
impl ExpectServerHelloOrHelloRetryRequest {
fn handle<'m>(
self,
input: Input<'m>,
output: &mut dyn Output<'m>,
) -> Result<ClientState, Error> {
match input.message.payload {
MessagePayload::Handshake {
parsed: HandshakeMessagePayload(HandshakePayload::ServerHello(..)),
..
} => self
.into_expect_server_hello()
.handle(input, output),
MessagePayload::Handshake {
parsed: HandshakeMessagePayload(HandshakePayload::HelloRetryRequest(..)),
..
} => self.handle_hello_retry_request(input, output),
payload => Err(inappropriate_handshake_message(
&payload,
&[ContentType::Handshake],
&[HandshakeType::ServerHello, HandshakeType::HelloRetryRequest],
)),
}
}
}
pub(crate) struct ClientHelloInput {
pub(super) config: Arc<ClientConfig>,
pub(super) resuming: Option<Retrieved<ClientSessionValue>>,
pub(super) random: Random,
pub(super) sent_tls13_fake_ccs: bool,
pub(super) hello: ClientHelloDetails,
pub(super) protocol: Protocol,
pub(super) session_id: SessionId,
pub(super) session_key: ClientSessionKey<'static>,
pub(super) prev_ech_ext: Option<EncryptedClientHello>,
}
impl ClientHelloInput {
pub(super) fn new(
server_name: ServerName<'static>,
extra_exts: &ClientExtensionsInput,
protocol: Protocol,
output: &mut dyn Output<'_>,
config: Arc<ClientConfig>,
) -> Result<Self, Error> {
let session_key = ClientSessionKey {
config_hash: config.config_hash(),
server_name,
};
let mut resuming = ClientSessionValue::retrieve(&session_key, &config, output);
let session_id = match &mut resuming {
Some(resuming) => {
debug!("Resuming session");
match &mut resuming.value {
ClientSessionValue::Tls12(inner) => {
if !inner.ticket().is_empty() {
inner.session_id = SessionId::random(config.provider().secure_random)?;
}
Some(inner.session_id)
}
_ => None,
}
}
_ => {
debug!("Not resuming any session");
None
}
};
let session_id = match session_id {
Some(session_id) => session_id,
None if output.quic().is_some() => SessionId::empty(),
None if !config.supports_version(ProtocolVersion::TLSv1_3) => SessionId::empty(),
None => SessionId::random(config.provider().secure_random)?,
};
let hello = ClientHelloDetails::new(
extra_exts
.protocols
.clone()
.unwrap_or_default(),
rand::random_u16(config.provider().secure_random)?,
);
let random = Random::new(config.provider().secure_random)?;
Ok(Self {
config,
resuming,
random,
sent_tls13_fake_ccs: false,
hello,
protocol,
session_id,
session_key,
prev_ech_ext: None,
})
}
pub(super) fn start_handshake(
self,
extra_exts: ClientExtensionsInput,
output: &mut dyn Output<'_>,
) -> Result<ClientState, Error> {
let mut transcript_buffer = HandshakeHashBuffer::new();
if !self
.config
.resolver()
.supported_certificate_types()
.is_empty()
{
transcript_buffer.set_client_auth_enabled();
}
let key_share = if self
.config
.supports_version(ProtocolVersion::TLSv1_3)
{
Some(tls13::initial_key_share(&self.config, &self.session_key)?)
} else {
None
};
let ech_state = match self.config.ech_mode.as_ref() {
Some(EchMode::Enable(ech_config)) => Some(ech_config.state(
self.session_key.server_name.clone(),
self.protocol,
&self.config,
)?),
_ => None,
};
emit_client_hello_for_retry(
transcript_buffer,
None,
key_share,
extra_exts,
None,
self,
output,
ech_state,
EchStatus::default(),
)
}
}
fn emit_client_hello_for_retry(
mut transcript_buffer: HandshakeHashBuffer,
retryreq: Option<&HelloRetryRequest>,
key_share: Option<GroupAndKeyShare>,
extra_exts: ClientExtensionsInput,
suite: Option<SupportedCipherSuite>,
mut input: ClientHelloInput,
output: &mut dyn Output<'_>,
mut ech_state: Option<EchState>,
mut ech_status: EchStatus,
) -> Result<ClientState, Error> {
let config = &input.config;
let forbids_tls12 = input.protocol.is_quic() || ech_state.is_some();
let supported_versions = SupportedProtocolVersions {
tls13: config.supports_version(ProtocolVersion::TLSv1_3),
tls12: config.supports_version(ProtocolVersion::TLSv1_2) && !forbids_tls12,
};
assert!(supported_versions.any(|_| true));
let mut exts = Box::new(ClientExtensions {
certificate_status_request: match config
.verifier()
.request_ocsp_response()
{
true => Some(CertificateStatusRequest::build_ocsp()),
false => None,
},
named_groups: Some(
config
.provider()
.kx_groups
.iter()
.filter_map(|skxg| {
let named_group = skxg.name();
supported_versions
.any(|v| named_group.usable_for_version(v))
.then_some(named_group)
})
.collect(),
),
signature_schemes: Some(
config
.verifier()
.supported_verify_schemes(),
),
protocols: extra_exts.protocols.clone(),
extended_master_secret_request: Some(()),
supported_versions: Some(supported_versions),
..Default::default()
});
if let Some(TransportParameters::Quic(v)) = &extra_exts.transport_parameters {
exts.transport_parameters = Some(v.clone());
}
if supported_versions.tls13 {
if let Some(cas_extension) = config.verifier().root_hint_subjects() {
exts.certificate_authority_names = Some(cas_extension.to_vec());
}
}
if config
.provider()
.kx_groups
.iter()
.any(|skxg| skxg.name().key_exchange_algorithm() == KeyExchangeAlgorithm::ECDHE)
{
exts.ec_point_formats = Some(SupportedEcPointFormats::default());
}
exts.server_name = match (ech_state.as_ref(), config.enable_sni) {
(Some(ech_state), _) => Some(ServerNamePayload::from(&ech_state.outer_name)),
(None, true) => match &input.session_key.server_name {
ServerName::DnsName(dns_name) => Some(ServerNamePayload::from(dns_name)),
_ => None,
},
(None, false) => None,
};
if let Some(GroupAndKeyShare { share, .. }) = &key_share {
debug_assert!(supported_versions.tls13);
let mut shares = vec![KeyShareEntry::new(share.group(), share.pub_key())];
if !retryreq
.map(|rr| rr.key_share.is_some())
.unwrap_or_default()
{
if let Some((hybrid, _)) =
share.as_hybrid_checked(&config.provider().kx_groups, ProtocolVersion::TLSv1_3)
{
let (component_group, component_share) = hybrid.component();
shares.push(KeyShareEntry::new(component_group, component_share));
}
}
exts.key_shares = Some(shares);
}
if let Some(cookie) = retryreq.and_then(|hrr| hrr.cookie.as_ref()) {
exts.cookie = Some(cookie.to_vec().into());
}
if supported_versions.tls13 {
exts.preshared_key_modes = Some(PskKeyExchangeModes {
psk_dhe: true,
psk: false,
});
if let Some(ticket_req) = &config.send_ticket_request {
exts.ticket_request = Some(ClientTicketRequest {
new_session_count: ticket_req.new_session_count,
resumption_count: ticket_req.resumption_count,
});
}
}
input.hello.offered_cert_compression =
if supported_versions.tls13 && !config.cert_decompressors.is_empty() {
exts.certificate_compression_algorithms = Some(
config
.cert_decompressors
.iter()
.map(|dec| dec.algorithm())
.collect(),
);
true
} else {
false
};
let client_certificate_types = config
.resolver()
.supported_certificate_types();
match client_certificate_types {
&[] | &[CertificateType::X509] => {}
supported => {
exts.client_certificate_types = Some(supported.to_vec());
}
}
let server_certificate_types = config
.verifier()
.supported_certificate_types();
match server_certificate_types {
[] => return Err(ApiMisuse::NoSupportedCertificateTypes.into()),
[CertificateType::X509] => {}
supported => {
exts.server_certificate_types = Some(supported.to_vec());
}
}
if matches!(ech_status, EchStatus::Rejected | EchStatus::Grease) && retryreq.is_some() {
if let Some(prev_ech_ext) = input.prev_ech_ext.take() {
exts.encrypted_client_hello = Some(prev_ech_ext);
}
}
let tls13_session = prepare_resumption(&input.resuming, &mut exts, suite, output, config);
let (tls13_session, early_data_enabled) = match tls13_session {
Some((tls13_session, early_data_enabled)) => (Some(tls13_session), early_data_enabled),
_ => (None, false),
};
exts.order_seed = input.hello.extension_order_seed;
let mut cipher_suites: Vec<_> = config
.provider()
.iter_cipher_suites()
.filter_map(|cs| match cs.usable_for_protocol(input.protocol) {
true => Some(cs.suite()),
false => None,
})
.collect();
if supported_versions.tls12 {
cipher_suites.push(CipherSuite::TLS_EMPTY_RENEGOTIATION_INFO_SCSV);
}
let mut chp_payload = ClientHelloPayload {
client_version: ProtocolVersion::TLSv1_2,
random: input.random,
session_id: input.session_id,
cipher_suites,
compression_methods: vec![Compression::Null],
extensions: exts,
};
let ech_grease_ext = config
.ech_mode
.as_ref()
.and_then(|mode| match mode {
EchMode::Grease(cfg) => Some(cfg.grease_ext(
config.provider().secure_random,
input.protocol,
input.session_key.server_name.clone(),
&chp_payload,
)),
_ => None,
});
match (ech_status, &mut ech_state) {
(EchStatus::NotOffered | EchStatus::Offered, Some(ech_state)) => {
chp_payload = ech_state.ech_hello(chp_payload, retryreq, tls13_session.as_ref())?;
ech_status = EchStatus::Offered;
output.emit(Event::EchStatus(ech_status));
input.prev_ech_ext = chp_payload
.encrypted_client_hello
.clone();
}
(EchStatus::NotOffered, None) => {
if let Some(grease_ext) = ech_grease_ext {
let grease_ext = grease_ext?;
chp_payload.encrypted_client_hello = Some(grease_ext.clone());
ech_status = EchStatus::Grease;
output.emit(Event::EchStatus(ech_status));
input.prev_ech_ext = Some(grease_ext);
}
}
_ => {}
}
input.hello.sent_extensions = chp_payload.collect_used();
let mut chp = HandshakeMessagePayload(HandshakePayload::ClientHello(chp_payload));
let tls13_early_data_key_schedule = match (ech_state.as_mut(), tls13_session) {
(Some(ech_state), Some(tls13_session)) => ech_state
.early_data_key_schedule
.take()
.map(|schedule| (tls13_session.suite, schedule)),
(_, Some(tls13_session)) => {
let key_schedule = KeyScheduleEarlyClient::new(
input.protocol,
tls13_session.suite,
tls13_session.secret.bytes(),
);
tls13::fill_in_psk_binder(&key_schedule, &transcript_buffer, &mut chp);
Some((tls13_session.suite, key_schedule))
}
_ => None,
};
let ch = Message {
version: match retryreq {
Some(_) => ProtocolVersion::TLSv1_2,
None => ProtocolVersion::TLSv1_0,
},
payload: MessagePayload::handshake(chp),
};
if retryreq.is_some() {
tls13::emit_fake_ccs(&mut input.sent_tls13_fake_ccs, output);
}
trace!("Sending ClientHello {ch:#?}");
transcript_buffer.add_message(&ch);
output.send_msg(ch, false);
let early_data_key_schedule =
tls13_early_data_key_schedule.map(|(resuming_suite, schedule)| {
if !early_data_enabled {
output.emit(Event::EarlyData(EarlyDataEvent::Rejected));
return (schedule, false);
}
let (transcript_buffer, random) = match &ech_state {
Some(ech_state) => (
&ech_state.inner_hello_transcript,
&ech_state.inner_hello_random.0,
),
None => (&transcript_buffer, &input.random.0),
};
tls13::derive_early_traffic_secret(
&*config.key_log,
output,
resuming_suite.common.hash_provider,
&schedule,
&mut input.sent_tls13_fake_ccs,
transcript_buffer,
random,
);
(schedule, true)
});
let mut next = Box::new(ExpectServerHello {
input,
transcript_buffer,
early_data_key_schedule,
offered_key_share: key_share,
suite,
ech_state,
ech_status,
done_retry: false,
});
Ok(if supported_versions.tls13 && retryreq.is_none() {
ClientState::ServerHelloOrHelloRetryRequest(Box::new(
ExpectServerHelloOrHelloRetryRequest { next, extra_exts },
))
} else {
next.done_retry = retryreq.is_some();
ClientState::ServerHello(next)
})
}
pub(super) struct GroupAndKeyShare {
pub(super) group: &'static dyn SupportedKxGroup,
pub(super) share: StartedKeyExchange,
}
impl GroupAndKeyShare {
pub(super) fn new(group: &'static dyn SupportedKxGroup) -> Result<Self, Error> {
Ok(Self {
group,
share: group.start()?,
})
}
}
fn prepare_resumption<'a>(
resuming: &'a Option<Retrieved<ClientSessionValue>>,
exts: &mut ClientExtensions<'_>,
suite: Option<SupportedCipherSuite>,
output: &mut dyn Output<'_>,
config: &ClientConfig,
) -> Option<(Retrieved<&'a Tls13Session>, bool)> {
let resuming = match resuming {
Some(resuming) if !resuming.ticket().is_empty() => resuming,
_ => {
if config.supports_version(ProtocolVersion::TLSv1_2)
&& config.resumption.tls12_resumption == Tls12Resumption::SessionIdOrTickets
{
exts.session_ticket = Some(ClientSessionTicket::Request);
}
return None;
}
};
let Some(tls13) = resuming.map(|csv| csv.tls13()) else {
if config.supports_version(ProtocolVersion::TLSv1_2)
&& config.resumption.tls12_resumption == Tls12Resumption::SessionIdOrTickets
{
exts.session_ticket = Some(ClientSessionTicket::Offer(Payload::new(resuming.ticket())));
}
return None; };
if !config.supports_version(ProtocolVersion::TLSv1_3) {
return None;
}
let suite = match suite {
Some(SupportedCipherSuite::Tls13(suite)) => Some(suite),
Some(SupportedCipherSuite::Tls12(_)) => return None,
None => None,
};
if let Some(suite) = suite {
suite.can_resume_from(tls13.suite)?;
}
let early_data_enabled =
tls13::prepare_resumption(config, output, &tls13, exts, suite.is_some());
Some((tls13, early_data_enabled))
}
pub(super) fn process_alpn_protocol(
output: &mut dyn Output<'_>,
offered_protocols: &[ApplicationProtocol<'_>],
selected: Option<&ApplicationProtocol<'_>>,
check_selected_offered: bool,
) -> Result<(), Error> {
if let Some(alpn_protocol) = selected {
if check_selected_offered && !offered_protocols.contains(alpn_protocol) {
return Err(PeerMisbehaved::SelectedUnofferedApplicationProtocol.into());
}
}
debug!(
"ALPN protocol is {:?}",
selected
.as_ref()
.map(|v| bs_debug::BsDebug(v.as_ref()))
);
if let Some(protocol) = selected {
output.output(OutputEvent::ApplicationProtocol(protocol.to_owned()));
}
Ok(())
}
pub(super) enum ClientSessionValue {
Tls13(Tls13Session),
Tls12(Tls12Session),
}
impl ClientSessionValue {
fn retrieve(
key: &ClientSessionKey<'static>,
config: &ClientConfig,
output: &mut dyn Output<'_>,
) -> Option<Retrieved<Self>> {
let found = config
.resumption
.store
.take_tls13_ticket(key)
.map(ClientSessionValue::Tls13)
.or_else(|| {
config
.resumption
.store
.tls12_session(key)
.map(ClientSessionValue::Tls12)
})
.and_then(|resuming| {
let now = config
.current_time()
.map_err(|_err| debug!("Could not get current time: {_err}"))
.ok()?;
let retrieved = Retrieved::new(resuming, now);
match retrieved.has_expired() {
false => Some(retrieved),
true => None,
}
})
.or_else(|| {
debug!("No cached session for {key:?}");
None
});
if let Some(quic) = output.quic() {
if let Some(quic_params) = found
.as_ref()
.and_then(|r| r.tls13().map(|v| &v.quic_params))
{
quic.transport_parameters(quic_params.bytes().to_vec());
}
}
found
}
fn common(&self) -> &ClientSessionCommon {
match self {
Self::Tls13(inner) => &inner.common,
Self::Tls12(inner) => &inner.common,
}
}
fn tls13(&self) -> Option<&Tls13Session> {
match self {
Self::Tls13(v) => Some(v),
Self::Tls12(_) => None,
}
}
}
impl Deref for ClientSessionValue {
type Target = ClientSessionCommon;
fn deref(&self) -> &Self::Target {
self.common()
}
}
pub(crate) trait ClientHandler<T>: fmt::Debug + Sealed + Send + Sync {
fn handle_server_hello(
&self,
suite: &'static T,
server_hello: &ServerHelloPayload,
input: &Input<'_>,
st: ExpectServerHello,
output: &mut dyn Output<'_>,
) -> Result<ClientState, Error>;
}