use alloc::boxed::Box;
use alloc::vec::Vec;
use core::fmt;
use core::ops::{DerefMut, Range};
use std::sync::MutexGuard;
use super::receive::{Discard, JoinOutput};
use crate::client::ClientSide;
use crate::common_state::UnborrowedPayload;
use crate::conn::{
ConnectionCore, MessageIter, ReceivePath, SendOutput, SendPath, TlsInputBuffer, WrittenInto,
};
use crate::crypto::cipher::{MessageEncrypter, OutboundPlain};
use crate::enums::ProtocolVersion;
use crate::error::{AlertDescription, ErrorWithAlert};
use crate::lock::Mutex;
use crate::msgs::{AlertLevel, Delocator, Message};
use crate::sync::Arc;
use crate::tls13::key_schedule::KeyScheduleTrafficSend;
use crate::{ConnectionOutputs, Error, SideData};
#[expect(clippy::exhaustive_structs)]
#[derive(Debug)]
pub struct SplitConnection<Side: SideData> {
pub send: SendTraffic,
pub receive: ReceiveTraffic<Side>,
pub outputs: ConnectionOutputs,
}
impl<Side: SideData> TryFrom<ConnectionCore<Side>> for SplitConnection<Side> {
type Error = Error;
fn try_from(conn: ConnectionCore<Side>) -> Result<Self, Error> {
let send = Arc::new(Mutex::new(conn.common.send));
let state = conn.state?;
Ok(Self {
send: SendTraffic(send.clone()),
receive: ReceiveTraffic {
state,
recv: conn.common.recv,
send,
pending_flush_sender: false,
},
outputs: conn.common.outputs,
})
}
}
pub struct SendTraffic(pub(crate) Arc<Mutex<SendPath>>);
impl SendTraffic {
pub fn write(&mut self, application_data: OutboundPlain<'_>) -> Vec<Vec<u8>> {
let mut inner = self.0.lock().unwrap();
inner.maybe_refresh_traffic_keys();
inner.send_appdata_encrypt(application_data);
inner.sendable_tls.take()
}
pub fn write_tls_into(
&mut self,
application_data: OutboundPlain<'_>,
out: &mut [u8],
) -> Result<WrittenInto, Error> {
let mut inner = self.0.lock().unwrap();
inner.maybe_refresh_traffic_keys();
inner.write_appdata_into(application_data, out)
}
pub fn close(mut self) -> Vec<Vec<u8>> {
let mut inner = self.0.lock().unwrap();
inner.send_close_notify();
drop(inner);
self.take_data()
}
pub fn take_data(&mut self) -> Vec<Vec<u8>> {
let mut inner = self.0.lock().unwrap();
inner.maybe_refresh_traffic_keys();
inner.sendable_tls.take()
}
pub fn refresh_traffic_keys(&mut self) -> Result<(), Error> {
self.0
.lock()
.unwrap()
.refresh_traffic_keys()
}
}
impl fmt::Debug for SendTraffic {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("SendTraffic")
.finish_non_exhaustive()
}
}
pub struct ReceiveTraffic<Side: SideData> {
pub(crate) state: Side::State,
pub(crate) recv: ReceivePath,
pub(crate) send: Arc<Mutex<SendPath>>,
pub(crate) pending_flush_sender: bool,
}
impl<Side: SideData> ReceiveTraffic<Side> {
pub fn read<'a>(
self,
input: &'a mut impl TlsInputBuffer,
) -> Result<ReceiveTrafficState<'a, Side>, ErrorWithAlert> {
let Self {
state,
mut recv,
send,
mut pending_flush_sender,
} = self;
let mut send_adapter = SendAdapter::Unlocked(&send);
let mut state = Ok(state);
let output = JoinOutput {
outputs: &mut Discard,
quic: None,
send: &mut send_adapter,
side: &mut Discard,
};
let mut iter = MessageIter::<Side>::receive(input, &mut state, &mut recv, output);
let received_plain = match iter.next() {
Some(Ok(payload)) => Some(payload),
Some(Err(error)) => {
return Err(ErrorWithAlert::new(
error,
send_adapter
.as_locked(false)
.deref_mut(),
));
}
None => None,
};
let state = state.unwrap();
if let Some(unborrowed) = received_plain {
let pending_discard = recv.deframer.take_discard();
let UnborrowedPayload::Unborrowed(range) = unborrowed else {
return Err(Error::Unreachable("decrypted data should be borrowed").into());
};
if let SendAdapter::Locked { send_required, .. } = send_adapter {
pending_flush_sender |= send_required;
}
drop(send_adapter);
return Ok(ReceiveTrafficState::Available(ReceivedApplicationData {
range,
input,
pending_discard,
rt: Self {
state,
recv,
send,
pending_flush_sender,
},
}));
}
input.discard(recv.deframer.take_discard());
if let SendAdapter::Locked { send_required, .. } = send_adapter {
pending_flush_sender |= send_required;
}
drop(send_adapter);
let mut rt = Self {
state,
recv,
send,
pending_flush_sender,
};
if core::mem::take(&mut rt.pending_flush_sender) {
return Ok(ReceiveTrafficState::FlushSender(FlushSender { rt }));
}
Ok(match rt.recv.has_received_close_notify {
true => ReceiveTrafficState::CloseNotify,
false => ReceiveTrafficState::ReadMore(rt),
})
}
}
impl ReceiveTraffic<ClientSide> {
pub fn tls13_tickets_received(&self) -> u32 {
self.recv.tls13_tickets_received
}
}
impl<Side: SideData> fmt::Debug for ReceiveTraffic<Side> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ReceiveTraffic")
.finish_non_exhaustive()
}
}
#[expect(clippy::exhaustive_enums)]
pub enum ReceiveTrafficState<'a, Side: SideData> {
ReadMore(ReceiveTraffic<Side>),
FlushSender(FlushSender<Side>),
Available(ReceivedApplicationData<'a, Side>),
CloseNotify,
}
impl<Side: SideData> fmt::Debug for ReceiveTrafficState<'_, Side> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ReadMore(_) => f
.debug_tuple("ReadMore")
.finish_non_exhaustive(),
Self::FlushSender(_) => f
.debug_tuple("FlushSender")
.finish_non_exhaustive(),
Self::Available(_) => f
.debug_tuple("Available")
.finish_non_exhaustive(),
Self::CloseNotify => write!(f, "CloseNotify"),
}
}
}
pub struct ReceivedApplicationData<'a, Side: SideData> {
input: &'a mut dyn TlsInputBuffer,
range: Range<usize>,
pending_discard: usize,
rt: ReceiveTraffic<Side>,
}
impl<Side: SideData> ReceivedApplicationData<'_, Side> {
pub fn data(&mut self) -> &[u8] {
Delocator::new(self.input.slice_mut()).slice_from_range(&self.range)
}
pub fn into_next(mut self) -> ReceiveTrafficState<'static, Side> {
self.input.discard(self.pending_discard);
if core::mem::take(&mut self.rt.pending_flush_sender) {
return ReceiveTrafficState::FlushSender(FlushSender { rt: self.rt });
}
match self.rt.recv.has_received_close_notify {
true => ReceiveTrafficState::CloseNotify,
false => ReceiveTrafficState::ReadMore(self.rt),
}
}
}
pub struct FlushSender<Side: SideData> {
rt: ReceiveTraffic<Side>,
}
impl<Side: SideData> FlushSender<Side> {
pub fn into_next(self) -> ReceiveTrafficState<'static, Side> {
match self.rt.recv.has_received_close_notify {
true => ReceiveTrafficState::CloseNotify,
false => ReceiveTrafficState::ReadMore(self.rt),
}
}
}
enum SendAdapter<'a> {
Unlocked(&'a Mutex<SendPath>),
Locked {
guard: MutexGuard<'a, SendPath>,
send_required: bool,
},
}
impl<'a> SendAdapter<'a> {
fn as_locked<'b>(&'b mut self, may_send: bool) -> &'b mut MutexGuard<'a, SendPath> {
if let Self::Unlocked(m) = self {
*self = Self::Locked {
guard: m.lock().unwrap(),
send_required: false,
};
}
let Self::Locked {
guard,
send_required,
} = self
else {
unreachable!();
};
*send_required |= may_send;
guard
}
}
impl SendOutput for SendAdapter<'_> {
fn negotiated_version(&mut self, version: ProtocolVersion) {
self.as_locked(false)
.negotiated_version(version);
}
fn ensure_key_update_queued(&mut self) {
self.as_locked(true)
.ensure_key_update_queued();
}
fn set_encrypter(&mut self, cipher: Box<dyn MessageEncrypter>, max_messages: u64) {
self.as_locked(false)
.set_encrypter(cipher, max_messages);
}
fn update_key_schedule(&mut self, schedule: Box<KeyScheduleTrafficSend>) {
self.as_locked(false)
.update_key_schedule(schedule);
}
fn send_alert(&mut self, level: AlertLevel, desc: AlertDescription) {
self.as_locked(true)
.send_alert(level, desc)
}
fn start_traffic(&mut self) {
self.as_locked(false).start_traffic();
}
fn send_msg(&mut self, m: Message<'_>, must_encrypt: bool) {
self.as_locked(true)
.send_msg(m, must_encrypt)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::crypto::test_provider::Tls13Cipher;
#[test]
fn send_adapter_flag() {
assert!(!send_flag_for(
|adapter| adapter.negotiated_version(ProtocolVersion::TLSv1_3)
));
assert!(send_flag_for(|adapter| adapter.ensure_key_update_queued()));
assert!(!send_flag_for(
|adapter| adapter.set_encrypter(Box::new(Tls13Cipher), 1234)
));
assert!(send_flag_for(|adapter| adapter.send_alert(
AlertLevel::Fatal,
AlertDescription::CertificateUnknown
)));
assert!(!send_flag_for(|adapter| adapter.start_traffic()));
assert!(send_flag_for(
|adapter| adapter.send_msg(Message::build_key_update_notify(), false)
));
}
fn send_flag_for(f: impl FnOnce(&mut SendAdapter<'_>)) -> bool {
let mut send = SendPath::default();
send.set_encrypter(Box::new(Tls13Cipher), 1234);
let send = Mutex::new(send);
let mut adapter = SendAdapter::Unlocked(&send);
f(&mut adapter);
let SendAdapter::Locked { send_required, .. } = adapter else {
panic!("expected to find SendAdapter::Locked");
};
send_required
}
}