use std::sync::Arc;
use bytes::Bytes;
use crate::api::common::Limits;
use crate::models::{ConnectionID, Role, Version};
use crate::protocol::base::{AnyConnection, Transport};
use crate::protocol::common::{self, Buffer, Error};
use crate::protocol::h1::H1Connection;
use crate::protocol::h2::{self, H2Connection};
use crate::protocol::h3::{H3Connection, H3Session};
use crate::tls;
pub type QuicIncoming = tokio_quiche::InitialQuicConnection<tokio::net::UdpSocket, tokio_quiche::metrics::DefaultMetrics>;
#[allow(clippy::large_enum_variant)]
pub enum Incoming {
Stream {
transport: Box<dyn Transport>,
id: ConnectionID,
},
QUIC(QuicIncoming),
}
#[derive(Clone)]
pub struct Negotiation {
pub versions: Vec<Version>,
pub limits: Limits,
pub acceptor: Option<Arc<boring::ssl::SslAcceptor>>,
pub hsts: Option<crate::helpers::hsts::HstsPolicy>,
}
impl Negotiation {
pub async fn accept(&self, incoming: Incoming) -> Result<AnyConnection, Error> {
match incoming {
Incoming::Stream { transport, id } => {
let assembling = std::pin::pin!(self.assemble(transport, id));
common::within(self.limits.read_timeout, assembling).await?
}
Incoming::QUIC(incoming) => {
let id = ConnectionID(Bytes::from(incoming.peer_addr().to_string()));
let session = H3Session::new(Role::Origin, id, self.limits);
let (mut connection, worker) = H3Connection::pair(session, self.hsts);
let quic = incoming.start(worker);
connection.guard = Some(std::sync::Arc::new(quic));
Ok(AnyConnection::H3(connection))
}
}
}
pub async fn assemble(&self, transport: Box<dyn Transport>, id: ConnectionID) -> Result<AnyConnection, Error> {
let Some(acceptor) = &self.acceptor else {
return self.assemble_plain(transport, id).await;
};
let stream = tokio_boring::accept(acceptor, transport).await.map_err(|err| Error::Tls(err.to_string()))?;
let version = tls::negotiated(stream.ssl().selected_alpn_protocol(), &self.versions)?;
let security = tls::security(stream.ssl());
let transport = Box::new(stream) as Box<dyn Transport>;
Ok(match version {
Version::V2_0 => AnyConnection::H2(H2Connection::new(transport, Role::Origin, id, self.limits).with_hsts(self.hsts).with_security(security)),
_ => AnyConnection::H1(H1Connection::new(transport, Role::Origin, id, self.limits).with_hsts(self.hsts).with_security(security)),
})
}
pub async fn assemble_plain(&self, mut transport: Box<dyn Transport>, id: ConnectionID) -> Result<AnyConnection, Error> {
let mut buffer = Buffer::new();
let probe = h2::PREFACE.len().min(4);
while buffer.len() < probe && buffer.fill(&mut transport, self.limits.read_timeout).await? {}
let sniffed = buffer.len().min(probe);
let h2 = self.versions.contains(&Version::V2_0)
&& sniffed > 0
&& buffer.as_slice()[..sniffed] == h2::PREFACE[..sniffed];
if h2 {
return Ok(AnyConnection::H2(H2Connection::resume(transport, Role::Origin, id, self.limits, buffer)));
}
Ok(AnyConnection::H1(H1Connection::resume(transport, Role::Origin, id, self.limits, buffer)))
}
}