use std::cell::RefCell;
use std::net::SocketAddr;
use std::rc::{Rc, Weak};
use std::task::{Context, Poll};
use std::time::Instant;
use bytes::Bytes;
use moq_noq_proto::{ConnectionHandle, Dir, StreamId, VarInt};
use rustc_hash::FxHashMap;
use super::super::{Error, SEGMENT};
use super::endpoint;
use crate::udp;
use crate::worker::Owner;
pub(crate) type Shared = Rc<Inner>;
const TRAIN_SEGMENTS: usize = 63;
pub(crate) struct Inner {
pub(crate) conn: RefCell<moq_noq_proto::Connection>,
pub(crate) state: RefCell<State>,
pub(crate) owner: Owner,
}
pub(crate) struct State {
driver: kio::WaiterList,
established: bool,
establish_waiters: kio::WaiterList,
accept_bi_waiters: kio::WaiterList,
accept_uni_waiters: kio::WaiterList,
open_waiters: kio::WaiterList,
readable: FxHashMap<StreamId, kio::WaiterList>,
writable: FxHashMap<StreamId, kio::WaiterList>,
finishing: FxHashMap<StreamId, kio::WaiterList>,
sends: FxHashMap<StreamId, Option<End>>,
datagram_recv_waiters: kio::WaiterList,
datagram_send_waiters: kio::WaiterList,
closed: Option<Error>,
closed_waiters: kio::WaiterList,
dead: bool,
local_close: Option<(u64, String)>,
}
impl State {
fn new() -> Self {
Self {
driver: kio::WaiterList::new(),
established: false,
establish_waiters: kio::WaiterList::new(),
accept_bi_waiters: kio::WaiterList::new(),
accept_uni_waiters: kio::WaiterList::new(),
open_waiters: kio::WaiterList::new(),
readable: FxHashMap::default(),
writable: FxHashMap::default(),
finishing: FxHashMap::default(),
sends: FxHashMap::default(),
datagram_recv_waiters: kio::WaiterList::new(),
datagram_send_waiters: kio::WaiterList::new(),
closed: None,
closed_waiters: kio::WaiterList::new(),
dead: false,
local_close: None,
}
}
fn forget_send(&mut self, id: StreamId) {
self.sends.remove(&id);
self.finishing.remove(&id);
self.writable.remove(&id);
}
fn forget_recv(&mut self, id: StreamId) {
self.readable.remove(&id);
}
fn fail(&mut self, err: Error) {
if self.closed.is_none() {
self.closed = Some(err);
}
self.driver.wake();
self.closed_waiters.wake();
self.establish_waiters.wake();
self.accept_bi_waiters.wake();
self.accept_uni_waiters.wake();
self.open_waiters.wake();
self.datagram_recv_waiters.wake();
self.datagram_send_waiters.wake();
for waiters in self.readable.values_mut() {
waiters.wake();
}
for waiters in self.writable.values_mut() {
waiters.wake();
}
for waiters in self.finishing.values_mut() {
waiters.wake();
}
}
}
impl Inner {
pub(crate) fn closed(&self) -> Option<Error> {
self.state.borrow().closed.clone()
}
pub(crate) fn kick(&self) {
self.state.borrow_mut().driver.wake();
}
pub(crate) fn fail(&self, err: Error) {
let mut state = self.state.borrow_mut();
state.dead = true;
state.fail(err);
}
pub(crate) fn park_writable(&self, id: StreamId, waiter: &kio::Waiter) {
let mut state = self.state.borrow_mut();
waiter.register(state.writable.entry(id).or_default());
}
pub(crate) fn park_readable(&self, id: StreamId, waiter: &kio::Waiter) {
let mut state = self.state.borrow_mut();
waiter.register(state.readable.entry(id).or_default());
}
pub(crate) fn park_finishing(&self, id: StreamId, waiter: &kio::Waiter) {
let mut state = self.state.borrow_mut();
waiter.register(state.finishing.entry(id).or_default());
}
pub(crate) fn track(&self, id: StreamId) {
let stopped = self.conn.borrow_mut().send_stream(id).stopped().ok().flatten();
self.state
.borrow_mut()
.sends
.insert(id, stopped.map(|code| End::Stopped(code.into_inner())));
}
pub(crate) fn ended(&self, id: StreamId) -> Option<End> {
self.state.borrow().sends.get(&id).copied().flatten()
}
pub(crate) fn forget_send(&self, id: StreamId) {
self.state.borrow_mut().forget_send(id);
}
pub(crate) fn forget_recv(&self, id: StreamId) {
self.state.borrow_mut().forget_recv(id);
}
pub(crate) fn wake_readable(&self, id: StreamId) {
let mut state = self.state.borrow_mut();
if let Some(mut waiters) = state.readable.remove(&id) {
waiters.wake();
}
}
pub(crate) fn close_code(&self, code: u64, reason: &str) {
self.conn.borrow_mut().close(
Instant::now(),
VarInt::from_u64(code).unwrap_or(VarInt::MAX),
Bytes::copy_from_slice(reason.as_bytes()),
);
let mut state = self.state.borrow_mut();
state.local_close.get_or_insert((code, reason.to_string()));
state.driver.wake();
}
}
pub struct Connection {
shared: Shared,
park: kio::Park,
alpn: Option<String>,
server_name: Option<String>,
remote: SocketAddr,
}
impl Connection {
pub(crate) fn close_code(&self, code: u64, reason: &str) {
self.shared.close_code(code, reason);
}
pub(crate) fn owner(&self) -> &Owner {
&self.shared.owner
}
pub fn peer_chain(&self) -> Option<Vec<Vec<u8>>> {
let conn = self.shared.conn.borrow();
let identity = conn.crypto_session().peer_identity()?;
let chain = identity
.downcast::<Vec<rustls::pki_types::CertificateDer<'static>>>()
.ok()?;
Some(chain.iter().map(|cert| cert.to_vec()).collect())
}
pub fn remote_addr(&self) -> SocketAddr {
self.remote
}
pub fn server_name(&self) -> Option<&str> {
self.server_name.as_deref()
}
}
impl Clone for Connection {
fn clone(&self) -> Self {
Self {
shared: self.shared.clone(),
park: kio::Park::default(),
alpn: self.alpn.clone(),
server_name: self.server_name.clone(),
remote: self.remote,
}
}
}
pub(crate) fn launch(
owner: &Owner,
socket: Rc<udp::Socket>,
endpoint: Weak<endpoint::Inner>,
key: ConnectionHandle,
conn: moq_noq_proto::Connection,
) -> (Shared, impl Future<Output = ()> + use<>) {
let shared = Rc::new(Inner {
conn: RefCell::new(conn),
state: RefCell::new(State::new()),
owner: owner.clone(),
});
let mut driver = Driver {
shared: shared.clone(),
socket,
endpoint,
key,
deadline: owner.timer(),
scratch: Vec::with_capacity(TRAIN_SEGMENTS * SEGMENT),
blocked: false,
};
let future = async move { kio::wait(|waiter| driver.poll(waiter)).await };
(shared, future)
}
struct EstablishGuard {
shared: Option<Shared>,
}
impl Drop for EstablishGuard {
fn drop(&mut self) {
if let Some(shared) = self.shared.take() {
shared.close_code(0, "QUIC handshake abandoned");
}
}
}
pub(crate) async fn establish(shared: Shared) -> Result<Connection, Error> {
let mut guard = EstablishGuard {
shared: Some(shared.clone()),
};
let result = kio::wait(|waiter| {
let mut state = shared.state.borrow_mut();
if state.established {
return Poll::Ready(Ok(()));
}
if let Some(err) = &state.closed {
return Poll::Ready(Err(err.clone()));
}
waiter.register(&mut state.establish_waiters);
Poll::Pending
})
.await;
if result.is_err() {
guard.shared = None;
}
result?;
let (alpn, server_name, remote) = {
let conn = shared.conn.borrow();
let handshake = conn
.crypto_session()
.handshake_data()
.and_then(|data| data.downcast::<moq_noq_proto::crypto::rustls::HandshakeData>().ok());
let (alpn, server_name) = match handshake {
Some(data) => (
data.protocol.map(|proto| String::from_utf8_lossy(&proto).into_owned()),
data.server_name,
),
None => (None, None),
};
let remote = conn
.network_path(moq_noq_proto::PathId::ZERO)
.expect("an established connection has its first path")
.remote();
(alpn, server_name, remote)
};
let conn = Connection {
shared,
park: kio::Park::default(),
alpn,
server_name,
remote,
};
guard.shared = None;
Ok(conn)
}
impl web_transport_trait::poll::Session for Connection {
type SendStream = super::SendStream;
type RecvStream = super::RecvStream;
type Error = Error;
fn poll_accept_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::RecvStream, Self::Error>> {
let waiter = self.park.hold(cx);
if let Some(id) = self.shared.conn.borrow_mut().streams().accept(Dir::Uni) {
return Poll::Ready(Ok(super::RecvStream::new(self.shared.clone(), id)));
}
let mut state = self.shared.state.borrow_mut();
if let Some(err) = &state.closed {
return Poll::Ready(Err(err.clone()));
}
waiter.register(&mut state.accept_uni_waiters);
Poll::Pending
}
fn poll_accept_bi(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let waiter = self.park.hold(cx);
let accepted = self.shared.conn.borrow_mut().streams().accept(Dir::Bi);
if let Some(id) = accepted {
return Poll::Ready(Ok((
super::SendStream::new(self.shared.clone(), id),
super::RecvStream::new(self.shared.clone(), id),
)));
}
let mut state = self.shared.state.borrow_mut();
if let Some(err) = &state.closed {
return Poll::Ready(Err(err.clone()));
}
waiter.register(&mut state.accept_bi_waiters);
Poll::Pending
}
fn poll_open_uni(&mut self, cx: &mut Context<'_>) -> Poll<Result<Self::SendStream, Self::Error>> {
let waiter = self.park.hold(cx);
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
let opened = self.shared.conn.borrow_mut().streams().open(Dir::Uni);
match opened {
Some(id) => Poll::Ready(Ok(super::SendStream::new(self.shared.clone(), id))),
None => {
let mut state = self.shared.state.borrow_mut();
waiter.register(&mut state.open_waiters);
Poll::Pending
}
}
}
fn poll_open_bi(
&mut self,
cx: &mut Context<'_>,
) -> Poll<Result<(Self::SendStream, Self::RecvStream), Self::Error>> {
let waiter = self.park.hold(cx);
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
let opened = self.shared.conn.borrow_mut().streams().open(Dir::Bi);
match opened {
Some(id) => Poll::Ready(Ok((
super::SendStream::new(self.shared.clone(), id),
super::RecvStream::new(self.shared.clone(), id),
))),
None => {
let mut state = self.shared.state.borrow_mut();
waiter.register(&mut state.open_waiters);
Poll::Pending
}
}
}
fn poll_send_datagram(&mut self, cx: &mut Context<'_>, payload: &[u8]) -> Poll<Result<(), Self::Error>> {
let waiter = self.park.hold(cx);
if let Some(err) = self.shared.closed() {
return Poll::Ready(Err(err));
}
let payload = Bytes::copy_from_slice(payload);
match self.shared.conn.borrow_mut().datagrams().send(payload, false) {
Ok(()) => {
self.shared.kick();
Poll::Ready(Ok(()))
}
Err(moq_noq_proto::SendDatagramError::Blocked(_)) => {
let mut state = self.shared.state.borrow_mut();
waiter.register(&mut state.datagram_send_waiters);
Poll::Pending
}
Err(err) => Poll::Ready(Err(Error::Quic(err.to_string()))),
}
}
fn poll_recv_datagram(&mut self, cx: &mut Context<'_>) -> Poll<Result<Bytes, Self::Error>> {
let waiter = self.park.hold(cx);
if let Some(datagram) = self.shared.conn.borrow_mut().datagrams().recv() {
return Poll::Ready(Ok(datagram));
}
let mut state = self.shared.state.borrow_mut();
if let Some(err) = &state.closed {
return Poll::Ready(Err(err.clone()));
}
waiter.register(&mut state.datagram_recv_waiters);
Poll::Pending
}
fn max_datagram_size(&self) -> usize {
self.shared.conn.borrow_mut().datagrams().max_size().unwrap_or(0)
}
fn protocol(&self) -> Option<&str> {
self.alpn.as_deref()
}
fn close(&mut self, code: u32, reason: &str) {
self.shared.close_code(u64::from(code), reason);
}
fn poll_closed(&mut self, cx: &mut Context<'_>) -> Poll<Self::Error> {
let waiter = self.park.hold(cx);
let mut state = self.shared.state.borrow_mut();
if let Some(err) = &state.closed {
return Poll::Ready(err.clone());
}
waiter.register(&mut state.closed_waiters);
Poll::Pending
}
fn stats(&self) -> impl web_transport_trait::Stats {
let (stats, path) = {
let mut conn = self.shared.conn.borrow_mut();
let stats = conn.stats();
let path = conn.path_stats(moq_noq_proto::PathId::ZERO).unwrap_or_default();
(stats, path)
};
Stats {
bytes_sent: stats.udp_tx.bytes,
bytes_received: stats.udp_rx.bytes,
bytes_lost: stats.lost_bytes,
packets_sent: stats.udp_tx.datagrams,
packets_received: stats.udp_rx.datagrams,
packets_lost: stats.lost_packets,
rtt: path.rtt,
}
}
}
impl std::fmt::Debug for Connection {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Connection").field("alpn", &self.alpn).finish()
}
}
struct Stats {
bytes_sent: u64,
bytes_received: u64,
bytes_lost: u64,
packets_sent: u64,
packets_received: u64,
packets_lost: u64,
rtt: std::time::Duration,
}
impl web_transport_trait::Stats for Stats {
fn bytes_sent(&self) -> Option<u64> {
Some(self.bytes_sent)
}
fn bytes_received(&self) -> Option<u64> {
Some(self.bytes_received)
}
fn bytes_lost(&self) -> Option<u64> {
Some(self.bytes_lost)
}
fn packets_sent(&self) -> Option<u64> {
Some(self.packets_sent)
}
fn packets_received(&self) -> Option<u64> {
Some(self.packets_received)
}
fn packets_lost(&self) -> Option<u64> {
Some(self.packets_lost)
}
fn rtt(&self) -> Option<std::time::Duration> {
Some(self.rtt)
}
fn estimated_send_rate(&self) -> Option<u64> {
None
}
}
struct Driver {
shared: Shared,
socket: Rc<udp::Socket>,
endpoint: Weak<endpoint::Inner>,
key: ConnectionHandle,
deadline: crate::Timer,
scratch: Vec<u8>,
blocked: bool,
}
impl Driver {
fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
{
let mut state = self.shared.state.borrow_mut();
if state.dead {
return Poll::Ready(());
}
waiter.register(&mut state.driver);
}
loop {
self.endpoint_events();
self.sweep();
if self.shared.state.borrow().dead {
return Poll::Ready(());
}
if self.shared.conn.borrow().is_drained() {
return Poll::Ready(());
}
if let Poll::Ready(err) = self.flush(waiter) {
self.shared.state.borrow_mut().fail(err);
return Poll::Ready(());
}
self.publish_close();
self.deadline.set(self.shared.conn.borrow_mut().poll_timeout());
if self.deadline.poll(waiter).is_pending() {
return Poll::Pending;
}
self.shared.conn.borrow_mut().handle_timeout(Instant::now());
}
}
fn endpoint_events(&mut self) {
let Some(endpoint) = self.endpoint.upgrade() else {
return;
};
loop {
let event = self.shared.conn.borrow_mut().poll_endpoint_events();
let Some(event) = event else {
return;
};
if let Some(event) = endpoint.on_connection_event(self.key, event) {
self.shared.conn.borrow_mut().handle_event(event);
}
}
}
fn sweep(&mut self) {
loop {
let event = self.shared.conn.borrow_mut().poll();
let Some(event) = event else {
return;
};
let mut state = self.shared.state.borrow_mut();
match event {
moq_noq_proto::Event::Connected => {
state.established = true;
state.establish_waiters.wake();
}
moq_noq_proto::Event::ConnectionLost { reason } => state.fail(reason.into()),
moq_noq_proto::Event::DatagramReceived => state.datagram_recv_waiters.wake(),
moq_noq_proto::Event::DatagramsUnblocked => state.datagram_send_waiters.wake(),
moq_noq_proto::Event::HandshakeDataReady => {}
moq_noq_proto::Event::Stream(event) => sweep_stream(&mut state, event),
moq_noq_proto::Event::HandshakeConfirmed
| moq_noq_proto::Event::Path(_)
| moq_noq_proto::Event::NatTraversal(_) => {}
}
}
}
fn publish_close(&mut self) {
if self.blocked || !self.shared.conn.borrow().is_closed() {
return;
}
let mut state = self.shared.state.borrow_mut();
let Some((code, reason)) = state.local_close.take() else {
return;
};
state.fail(Error::App { code, reason });
}
fn flush(&mut self, waiter: &kio::Waiter) -> Poll<Error> {
match self.flush_one(waiter) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(err)) => return Poll::Ready(err),
Poll::Pending => return Poll::Pending,
}
waiter.waker().wake_by_ref();
Poll::Pending
}
fn flush_one(&mut self, waiter: &kio::Waiter) -> Poll<Result<(), Error>> {
let mut tx = match self.socket.poll_acquire(waiter) {
Poll::Ready(Ok(tx)) => tx,
Poll::Ready(Err(err)) => return Poll::Ready(Err(Error::Io(err.to_string()))),
Poll::Pending => {
self.blocked = true;
return Poll::Pending;
}
};
self.blocked = false;
let segments = (tx.len() / SEGMENT).min(TRAIN_SEGMENTS);
if segments == 0 {
return Poll::Ready(Err(Error::Io(format!(
"transmit buffer of {} bytes holds no {SEGMENT} byte segment",
tx.len()
))));
}
let segments = std::num::NonZeroUsize::new(segments).expect("segments was checked above");
self.scratch.clear();
let transmit = match self
.shared
.conn
.borrow_mut()
.poll_transmit(Instant::now(), segments, &mut self.scratch)
{
Some(transmit) => transmit,
None => return Poll::Pending,
};
tx[..transmit.size].copy_from_slice(&self.scratch[..transmit.size]);
let transmit = udp::Transmit {
to: transmit.destination,
len: transmit.size,
segment: transmit.segment_size.unwrap_or(transmit.size),
ecn: transmit.ecn.map(super::ecn_from_noq),
};
if let Err(err) = tx.send(transmit) {
return Poll::Ready(Err(Error::Io(err.to_string())));
}
self.shared.state.borrow_mut().datagram_send_waiters.wake();
Poll::Ready(Ok(()))
}
}
#[derive(Clone, Copy, Debug)]
pub(crate) enum End {
Delivered,
Stopped(u64),
}
fn end(state: &mut State, id: StreamId, end: End) {
if let Some(slot) = state.sends.get_mut(&id) {
*slot = Some(end);
}
if let Some(mut waiters) = state.finishing.remove(&id) {
waiters.wake();
}
}
fn sweep_stream(state: &mut State, event: moq_noq_proto::StreamEvent) {
match event {
moq_noq_proto::StreamEvent::Opened { dir: Dir::Bi } => state.accept_bi_waiters.wake(),
moq_noq_proto::StreamEvent::Opened { dir: Dir::Uni } => state.accept_uni_waiters.wake(),
moq_noq_proto::StreamEvent::Available { .. } => state.open_waiters.wake(),
moq_noq_proto::StreamEvent::Readable { id } => {
if let Some(mut waiters) = state.readable.remove(&id) {
waiters.wake();
}
}
moq_noq_proto::StreamEvent::Writable { id } => {
if let Some(mut waiters) = state.writable.remove(&id) {
waiters.wake();
}
}
moq_noq_proto::StreamEvent::Finished { id } => end(state, id, End::Delivered),
moq_noq_proto::StreamEvent::Stopped { id, error_code } => {
end(state, id, End::Stopped(error_code.into_inner()));
if let Some(mut waiters) = state.writable.remove(&id) {
waiters.wake();
}
}
}
}
#[cfg(test)]
mod tests {
use moq_noq_proto::Side;
use super::*;
#[test]
fn forgetting_a_handle_clears_its_parking() {
let mut state = State::new();
let id = StreamId::new(Side::Client, Dir::Bi, 0);
state.writable.entry(id).or_default();
state.readable.entry(id).or_default();
state.finishing.entry(id).or_default();
state.sends.insert(id, None);
state.forget_send(id);
assert!(state.writable.is_empty(), "the write half's parking");
assert!(state.finishing.is_empty(), "the finish watch");
assert!(state.sends.is_empty(), "the end bookkeeping");
assert!(!state.readable.is_empty(), "the read half's parking survives");
state.forget_recv(id);
assert!(state.readable.is_empty(), "the read half's parking");
}
}