use alloc::boxed::Box;
use alloc::vec::Vec;
use core::ops::Deref;
use core::{fmt, mem};
use std::io;
use pki_types::{DnsName, FipsStatus};
use super::config::{ClientHello, ServerConfig};
use crate::common_state::{
CommonState, ConnectionOutputs, EarlyDataEvent, Event, Protocol, Side, maybe_send_fatal_alert,
};
use crate::conn::private::SideOutput;
use crate::conn::split::SplitConnection;
use crate::conn::{
Connection, ConnectionCommon, ConnectionCore, KeyingMaterialExporter, MessageIter, Reader,
SideData, StateMachine, TlsInputBuffer, Writer,
};
#[cfg(doc)]
use crate::crypto;
use crate::crypto::cipher::Payload;
use crate::error::Error;
use crate::log::trace;
use crate::msgs::ServerExtensionsInput;
use crate::server::hs::{ChooseConfig, ExpectClientHello, ReadClientHello, ServerState};
use crate::suites::ExtractedSecrets;
use crate::sync::Arc;
use crate::vecbuf::ChunkVecBuffer;
pub struct ServerConnection {
pub(super) inner: ConnectionCommon<ServerSide>,
}
impl ServerConnection {
pub fn new(config: Arc<ServerConfig>) -> Result<Self, Error> {
Ok(Self {
inner: ConnectionCommon::new(ConnectionCore::for_server(
config,
ServerExtensionsInput::default(),
Protocol::Tcp,
)?),
})
}
pub fn split(self) -> Result<SplitConnection<ServerSide>, Error> {
self.inner.split()
}
pub fn server_name(&self) -> Option<&DnsName<'_>> {
self.inner.core.side.server_name()
}
pub fn received_resumption_data(&self) -> Option<&[u8]> {
self.inner
.core
.side
.received_resumption_data()
}
pub fn set_resumption_data(&mut self, data: &[u8]) -> Result<(), Error> {
assert!(data.len() < 2usize.pow(15));
match &mut self.inner.core.state {
Ok(st) => st.set_resumption_data(data),
Err(e) => Err(e.clone()),
}
}
pub fn early_data(&mut self) -> Option<ReadEarlyData<'_>> {
if self
.inner
.core
.side
.early_data
.was_accepted()
{
Some(ReadEarlyData::new(&mut self.inner))
} else {
None
}
}
}
impl Connection for ServerConnection {
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<crate::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 ServerConnection {
type Target = ConnectionOutputs;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl fmt::Debug for ServerConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ServerConnection")
.finish_non_exhaustive()
}
}
#[non_exhaustive]
#[derive(Debug)]
pub enum ServerHandshake {
NeedsInput(NeedsInput),
Accepted(Accepted),
Complete(SplitConnection<ServerSide>),
}
impl ServerHandshake {
pub fn start() -> NeedsInput {
NeedsInput {
inner: ConnectionCore::for_acceptor(Protocol::Tcp),
}
}
}
impl TryFrom<ConnectionCore<ServerSide>> for ServerHandshake {
type Error = Error;
fn try_from(mut inner: ConnectionCore<ServerSide>) -> Result<Self, Error> {
const MISUSED: Error = Error::Unreachable("forgot to restore state");
Ok(match mem::replace(&mut inner.state, Err(MISUSED))? {
ServerState::ChooseConfig(choose_config) => Self::Accepted(Accepted {
inner,
choose_config,
}),
state if state.is_traffic() => {
inner.state = Ok(state);
Self::Complete(SplitConnection::try_from(inner)?)
}
state => {
inner.state = Ok(state);
Self::NeedsInput(NeedsInput { inner })
}
})
}
}
pub struct NeedsInput {
inner: ConnectionCore<ServerSide>,
}
impl NeedsInput {
pub fn process(
mut self,
input: &mut dyn TlsInputBuffer,
output: &mut Vec<Vec<u8>>,
) -> Result<ServerHandshake, Error> {
let mut iter = MessageIter::new(input, None, &mut self.inner);
let r = loop {
match iter.next() {
Some(Ok(_)) => {}
Some(Err(e)) => break Err(e),
None => break Ok(()),
};
if iter
.state()
.as_ref()
.map(|st| st.is_traffic())
.unwrap_or_default()
{
break Ok(());
}
};
input.discard(
self.inner
.common
.recv
.deframer
.take_discard(),
);
while let Some(chunk) = self
.inner
.common
.send
.sendable_tls
.pop()
{
output.push(chunk);
}
r?;
ServerHandshake::try_from(self.inner)
}
pub fn into_buffered_connection(self) -> ServerConnection {
ServerConnection {
inner: ConnectionCommon::new(self.inner),
}
}
}
impl fmt::Debug for NeedsInput {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("NeedsInput")
.finish_non_exhaustive()
}
}
pub struct Accepted {
inner: ConnectionCore<ServerSide>,
choose_config: Box<ChooseConfig>,
}
impl Accepted {
pub fn client_hello(&self) -> ClientHello<'_> {
let ch = self.choose_config.client_hello();
trace!("Accepted::client_hello(): {ch:#?}");
ch
}
pub fn choose_config(
mut self,
config: Arc<ServerConfig>,
output: &mut Vec<Vec<u8>>,
) -> Result<ServerHandshake, Error> {
let result = self.inner.accepted(
self.choose_config,
ServerExtensionsInput::default(),
None,
config,
);
let send_path = &mut self.inner.common.send;
if let Err(err) = &result {
maybe_send_fatal_alert(send_path, err);
}
while let Some(chunk) = send_path.sendable_tls.pop() {
output.push(chunk);
}
result?;
Ok(ServerHandshake::NeedsInput(NeedsInput {
inner: self.inner,
}))
}
}
impl fmt::Debug for Accepted {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Accepted")
.finish_non_exhaustive()
}
}
pub struct ReadEarlyData<'a> {
common: &'a mut ConnectionCommon<ServerSide>,
}
impl<'a> ReadEarlyData<'a> {
fn new(common: &'a mut ConnectionCommon<ServerSide>) -> Self {
ReadEarlyData { common }
}
pub fn exporter(&mut self) -> Result<KeyingMaterialExporter, Error> {
self.common.core.early_exporter()
}
}
impl io::Read for ReadEarlyData<'_> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.common
.core
.side
.early_data
.read(buf)
}
}
#[derive(Default)]
pub(super) enum EarlyDataState {
#[default]
New,
Accepted {
received: ChunkVecBuffer,
},
}
impl EarlyDataState {
fn accept(&mut self) {
*self = Self::Accepted {
received: ChunkVecBuffer::new(None),
};
}
fn was_accepted(&self) -> bool {
matches!(self, Self::Accepted { .. })
}
#[expect(dead_code)]
fn peek(&self) -> Option<&[u8]> {
match self {
Self::Accepted { received, .. } => received.peek(),
_ => None,
}
}
#[expect(dead_code)]
fn pop(&mut self) -> Option<Vec<u8>> {
match self {
Self::Accepted { received, .. } => received.pop(),
_ => None,
}
}
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Self::Accepted { received, .. } => Ok(received.read(buf)),
_ => Err(io::Error::from(io::ErrorKind::BrokenPipe)),
}
}
fn take_received_plaintext(&mut self, bytes: Payload<'_>) {
let Self::Accepted { received } = self else {
return;
};
received.append(bytes.into_vec());
}
}
impl ConnectionCore<ServerSide> {
pub(crate) fn for_server(
config: Arc<ServerConfig>,
extra_exts: ServerExtensionsInput,
protocol: Protocol,
) -> Result<Self, Error> {
let mut common = CommonState::new(Side::Server, config.fips());
common
.send
.set_max_fragment_size(config.max_fragment_size)?;
Ok(Self::new(
Box::new(ExpectClientHello::new(
config,
extra_exts,
Vec::new(),
protocol,
))
.into(),
ServerConnectionData::default(),
common,
))
}
pub(crate) fn for_acceptor(protocol: Protocol) -> Self {
Self::new(
ReadClientHello::new(protocol).into(),
ServerConnectionData::default(),
CommonState::new(Side::Server, FipsStatus::Unvalidated),
)
}
}
#[derive(Default)]
pub(crate) struct ServerConnectionData {
sni: Option<DnsName<'static>>,
received_resumption_data: Option<Vec<u8>>,
early_data: EarlyDataState,
}
impl ServerConnectionData {
pub(crate) fn received_resumption_data(&self) -> Option<&[u8]> {
self.received_resumption_data.as_deref()
}
pub(crate) fn server_name(&self) -> Option<&DnsName<'static>> {
self.sni.as_ref()
}
}
impl SideOutput for ServerConnectionData {
fn emit(&mut self, ev: Event<'_>) {
match ev {
Event::EarlyApplicationData(data) => self
.early_data
.take_received_plaintext(data),
Event::EarlyData(EarlyDataEvent::Accepted) => self.early_data.accept(),
Event::ReceivedServerName(sni) => self.sni = sni,
Event::ResumptionData(data) => self.received_resumption_data = Some(data),
_ => unreachable!(),
}
}
}
#[expect(clippy::exhaustive_structs)]
#[derive(Debug)]
pub struct ServerSide;
impl SideData for ServerSide {}
impl crate::conn::private::Side for ServerSide {
type Data = ServerConnectionData;
type State = ServerState;
}
#[cfg(test)]
mod tests {
use std::format;
use super::*;
#[test]
fn test_read_in_new_state() {
assert_eq!(
format!("{:?}", EarlyDataState::default().read(&mut [0u8; 5])),
"Err(Kind(BrokenPipe))"
);
}
}