use alloc::vec::Vec;
use core::ops::Deref;
use core::{fmt, mem};
use std::io;
use pki_types::{FipsStatus, ServerName};
use super::config::ClientConfig;
use super::hs::ClientHelloInput;
use crate::TlsInputBuffer;
use crate::client::EchStatus;
use crate::common_state::{CommonState, ConnectionOutputs, EarlyDataEvent, Event, Protocol, Side};
use crate::conn::private::SideOutput;
use crate::conn::split::SplitConnection;
use crate::conn::{
Connection, ConnectionCommon, ConnectionCore, IoState, KeyingMaterialExporter, Reader,
SideCommonOutput, SideData, Writer,
};
#[cfg(doc)]
use crate::crypto;
use crate::enums::ApplicationProtocol;
use crate::error::Error;
use crate::log::trace;
use crate::msgs::ClientExtensionsInput;
use crate::quic::QuicOutput;
use crate::suites::ExtractedSecrets;
use crate::sync::Arc;
pub struct ClientConnection {
inner: ConnectionCommon<ClientSide>,
}
impl fmt::Debug for ClientConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ClientConnection")
.finish_non_exhaustive()
}
}
impl ClientConnection {
pub fn split(self) -> Result<SplitConnection<ClientSide>, Error> {
self.inner.split()
}
pub fn early_data(&mut self) -> Option<WriteEarlyData<'_>> {
if self
.inner
.core
.side
.early_data
.is_enabled()
{
Some(WriteEarlyData::new(self))
} else {
None
}
}
pub fn is_early_data_accepted(&self) -> bool {
self.inner.core.is_early_data_accepted()
}
pub fn ech_status(&self) -> EchStatus {
self.inner.core.side.ech_status
}
fn write_early_data(&mut self, data: &[u8]) -> io::Result<usize> {
self.inner
.core
.side
.early_data
.check_write(data.len())
.map(|sz| {
self.inner
.send
.send_early_plaintext(&data[..sz])
})
}
pub fn tls13_tickets_received(&self) -> u32 {
self.inner
.core
.common
.recv
.tls13_tickets_received
}
}
impl Connection for ClientConnection {
fn write_tls(&mut self, wr: &mut dyn io::Write) -> Result<usize, io::Error> {
self.inner.write_tls(wr)
}
fn wants_read(&self) -> bool {
self.inner.wants_read()
}
fn wants_write(&self) -> bool {
self.inner.wants_write()
}
fn reader(&mut self) -> Reader<'_> {
self.inner.reader()
}
fn writer(&mut self) -> Writer<'_> {
self.inner.writer()
}
fn process_new_packets(&mut self, input: &mut dyn TlsInputBuffer) -> Result<IoState, Error> {
self.inner.process_new_packets(input)
}
fn exporter(&mut self) -> Result<KeyingMaterialExporter, Error> {
self.inner.exporter()
}
fn dangerous_extract_secrets(self) -> Result<ExtractedSecrets, Error> {
self.inner.dangerous_extract_secrets()
}
fn set_buffer_limit(&mut self, limit: Option<usize>) {
self.inner.set_buffer_limit(limit)
}
fn set_plaintext_buffer_limit(&mut self, limit: Option<usize>) {
self.inner
.set_plaintext_buffer_limit(limit)
}
fn refresh_traffic_keys(&mut self) -> Result<(), Error> {
self.inner.refresh_traffic_keys()
}
fn send_close_notify(&mut self) {
self.inner.send_close_notify();
}
fn is_handshaking(&self) -> bool {
self.inner.is_handshaking()
}
fn fips(&self) -> FipsStatus {
self.inner.fips
}
}
impl Deref for ClientConnection {
type Target = ConnectionOutputs;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
pub struct ClientConnectionBuilder {
pub(crate) config: Arc<ClientConfig>,
pub(crate) name: ServerName<'static>,
pub(crate) alpn_protocols: Option<Vec<ApplicationProtocol<'static>>>,
}
impl ClientConnectionBuilder {
pub fn with_alpn(mut self, alpn_protocols: Vec<ApplicationProtocol<'static>>) -> Self {
self.alpn_protocols = Some(alpn_protocols);
self
}
pub fn build(self) -> Result<ClientConnection, Error> {
let Self {
config,
name,
alpn_protocols,
} = self;
let alpn_protocols = alpn_protocols.unwrap_or_else(|| config.alpn_protocols.clone());
Ok(ClientConnection {
inner: ConnectionCommon::new(ConnectionCore::for_client(
config,
name,
ClientExtensionsInput::from_alpn(alpn_protocols),
None,
Protocol::Tcp,
)?),
})
}
}
pub struct WriteEarlyData<'a> {
sess: &'a mut ClientConnection,
}
impl<'a> WriteEarlyData<'a> {
fn new(sess: &'a mut ClientConnection) -> Self {
WriteEarlyData { sess }
}
pub fn bytes_left(&self) -> usize {
self.sess
.inner
.core
.side
.early_data
.bytes_left()
}
pub fn exporter(&mut self) -> Result<KeyingMaterialExporter, Error> {
self.sess.inner.core.early_exporter()
}
}
impl io::Write for WriteEarlyData<'_> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.sess.write_early_data(buf)
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl ConnectionCore<ClientSide> {
pub(crate) fn for_client(
config: Arc<ClientConfig>,
name: ServerName<'static>,
extra_exts: ClientExtensionsInput,
quic: Option<&mut dyn QuicOutput>,
protocol: Protocol,
) -> Result<Self, Error> {
let mut common_state = CommonState::new(Side::Client, config.fips());
common_state
.send
.set_max_fragment_size(config.max_fragment_size)?;
let mut data = ClientConnectionData::new();
let mut output = SideCommonOutput {
side: &mut data,
quic,
common: &mut common_state,
};
let input = ClientHelloInput::new(name, &extra_exts, protocol, &mut output, config)?;
let state = input.start_handshake(extra_exts, &mut output)?;
Ok(Self::new(state, data, common_state))
}
pub(crate) fn is_early_data_accepted(&self) -> bool {
self.side.early_data.is_accepted()
}
}
pub(super) struct EarlyData {
state: EarlyDataState,
left: usize,
}
impl EarlyData {
fn new() -> Self {
Self {
state: EarlyDataState::Disabled,
left: 0,
}
}
fn is_enabled(&self) -> bool {
matches!(
self.state,
EarlyDataState::Ready | EarlyDataState::Sending | EarlyDataState::Accepted
)
}
fn is_accepted(&self) -> bool {
matches!(
self.state,
EarlyDataState::Accepted | EarlyDataState::AcceptedFinished
)
}
fn enable(&mut self, max_data: usize) {
assert_eq!(self.state, EarlyDataState::Disabled);
self.state = EarlyDataState::Ready;
self.left = max_data;
}
fn start(&mut self) {
assert_eq!(self.state, EarlyDataState::Ready);
self.state = EarlyDataState::Sending;
}
fn rejected(&mut self) {
trace!("EarlyData rejected");
self.state = EarlyDataState::Rejected;
}
fn accepted(&mut self) {
trace!("EarlyData accepted");
assert_eq!(self.state, EarlyDataState::Sending);
self.state = EarlyDataState::Accepted;
}
pub(super) fn finished(&mut self) {
trace!("EarlyData finished");
self.state = match self.state {
EarlyDataState::Accepted => EarlyDataState::AcceptedFinished,
_ => panic!("bad EarlyData state"),
}
}
fn check_write(&mut self, sz: usize) -> io::Result<usize> {
self.check_write_opt(sz)
.ok_or_else(|| io::Error::from(io::ErrorKind::InvalidInput))
}
fn check_write_opt(&mut self, sz: usize) -> Option<usize> {
match self.state {
EarlyDataState::Disabled => unreachable!(),
EarlyDataState::Ready | EarlyDataState::Sending | EarlyDataState::Accepted => {
let take = if self.left < sz {
mem::replace(&mut self.left, 0)
} else {
self.left -= sz;
sz
};
Some(take)
}
EarlyDataState::Rejected | EarlyDataState::AcceptedFinished => None,
}
}
fn bytes_left(&self) -> usize {
self.left
}
}
#[derive(Debug, PartialEq)]
enum EarlyDataState {
Disabled,
Ready,
Sending,
Accepted,
AcceptedFinished,
Rejected,
}
pub(crate) struct ClientConnectionData {
early_data: EarlyData,
ech_status: EchStatus,
}
impl ClientConnectionData {
fn new() -> Self {
Self {
early_data: EarlyData::new(),
ech_status: EchStatus::default(),
}
}
}
#[expect(clippy::exhaustive_structs)]
#[derive(Debug)]
pub struct ClientSide;
impl SideData for ClientSide {}
impl crate::conn::private::Side for ClientSide {
type Data = ClientConnectionData;
type State = super::hs::ClientState;
}
impl SideOutput for ClientConnectionData {
fn emit(&mut self, ev: Event<'_>) {
match ev {
Event::EchStatus(ech) => self.ech_status = ech,
Event::EarlyData(EarlyDataEvent::Accepted) => self.early_data.accepted(),
Event::EarlyData(EarlyDataEvent::Enable(sz)) => self.early_data.enable(sz),
Event::EarlyData(EarlyDataEvent::Finished) => self.early_data.finished(),
Event::EarlyData(EarlyDataEvent::Start) => self.early_data.start(),
Event::EarlyData(EarlyDataEvent::Rejected) => self.early_data.rejected(),
_ => unreachable!(),
}
}
}