use std::cell::{Cell, RefCell};
use std::collections::VecDeque;
use std::net::SocketAddr;
use std::rc::Rc;
use std::sync::Arc;
use std::task::Poll;
use std::time::Instant;
use bytes::BytesMut;
use moq_noq_proto::{ConnectionHandle, DatagramEvent, Incoming, Transmit};
use rustc_hash::FxHashMap;
use super::super::{Error, endpoint::Config};
use super::connection;
use crate::quic::Connection;
use crate::shared::Shared;
use crate::udp;
use crate::worker::Owner;
struct Accepting {
queue: VecDeque<Connection>,
}
pub(crate) struct Inner {
owner: Owner,
socket: Rc<udp::Socket>,
local: SocketAddr,
endpoint: RefCell<moq_noq_proto::Endpoint>,
accepting: RefCell<Option<Accepting>>,
conns: RefCell<FxHashMap<ConnectionHandle, connection::Shared>>,
pending: Cell<usize>,
backlog: usize,
handles: Cell<usize>,
closed: RefCell<Option<Error>>,
task_waiters: RefCell<kio::WaiterList>,
accept_waiters: RefCell<kio::WaiterList>,
}
pub struct Endpoint {
inner: Rc<Inner>,
}
impl Endpoint {
pub fn new(socket: udp::Socket, config: Config) -> Result<Self, Error> {
let owner = socket.owner();
let Some(handle) = owner.handle() else {
return Err(Error::Io(Shared::gone_error().to_string()));
};
let local = socket.local_addr().map_err(|err| Error::Io(err.to_string()))?;
let server = match &config.server {
Some(server) => {
let mut server = super::server_config(server)?;
server.max_incoming(config.backlog);
Some(Arc::new(server))
}
None => None,
};
let accepting = server.is_some().then(|| Accepting { queue: VecDeque::new() });
let endpoint = moq_noq_proto::Endpoint::new(super::endpoint_config(socket.shard())?, server, false);
let inner = Rc::new(Inner {
owner,
socket: Rc::new(socket),
local,
endpoint: RefCell::new(endpoint),
accepting: RefCell::new(accepting),
conns: RefCell::new(FxHashMap::default()),
pending: Cell::new(0),
backlog: config.backlog,
handles: Cell::new(1),
closed: RefCell::new(None),
task_waiters: RefCell::new(kio::WaiterList::new()),
accept_waiters: RefCell::new(kio::WaiterList::new()),
});
let task = inner.clone();
handle.spawn(async move { kio::wait(|waiter| task.poll_run(waiter)).await });
Ok(Self { inner })
}
pub fn local_addr(&self) -> SocketAddr {
self.inner.local
}
pub async fn accept(&self) -> Result<Connection, Error> {
kio::wait(|waiter| {
if let Some(err) = self.inner.closed() {
return Poll::Ready(Err(err));
}
let mut accepting = self.inner.accepting.borrow_mut();
let Some(accepting) = accepting.as_mut() else {
return Poll::Ready(Err(Error::NotServer));
};
if let Some(conn) = accepting.queue.pop_front() {
return Poll::Ready(Ok(conn));
}
waiter.register(&mut self.inner.accept_waiters.borrow_mut());
Poll::Pending
})
.await
}
pub async fn connect(&self, config: &crate::quic::client::Config) -> Result<Connection, Error> {
if let Some(err) = self.inner.closed() {
return Err(err);
}
let client = super::client_config(config)?;
let (key, conn) = self
.inner
.endpoint
.borrow_mut()
.connect(Instant::now(), client, config.peer, &config.server_name)
.map_err(|err| Error::Quic(err.to_string()))?;
let shared = self.inner.launch(key, conn);
shared.kick();
connection::establish(shared).await
}
}
impl Clone for Endpoint {
fn clone(&self) -> Self {
self.inner.handles.set(self.inner.handles.get() + 1);
Self {
inner: self.inner.clone(),
}
}
}
impl Drop for Endpoint {
fn drop(&mut self) {
let handles = self.inner.handles.get() - 1;
self.inner.handles.set(handles);
if handles > 0 {
return;
}
if let Some(accepting) = self.inner.accepting.borrow_mut().as_mut() {
for mut conn in accepting.queue.drain(..) {
web_transport_trait::poll::Session::close(&mut conn, 0, "endpoint closed");
}
}
self.inner.task_waiters.borrow_mut().wake();
}
}
impl std::fmt::Debug for Endpoint {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Endpoint")
.field("local", &self.inner.local)
.field("conns", &self.inner.conns.borrow().len())
.finish()
}
}
impl Inner {
fn closed(&self) -> Option<Error> {
if let Some(err) = &*self.closed.borrow() {
return Some(err.clone());
}
self.owner
.handle()
.is_none()
.then(|| Error::Io(Shared::gone_error().to_string()))
}
fn poll_run(self: &Rc<Self>, waiter: &kio::Waiter) -> Poll<()> {
waiter.register(&mut self.task_waiters.borrow_mut());
loop {
if self.handles.get() == 0 && self.conns.borrow().is_empty() {
return Poll::Ready(());
}
match self.socket.poll_recv(waiter) {
Poll::Ready(Ok(mut packet)) => self.demux(&mut packet),
Poll::Ready(Err(err)) => {
self.fail(Error::Io(err.to_string()));
return Poll::Ready(());
}
Poll::Pending => return Poll::Pending,
}
}
}
fn demux(self: &Rc<Self>, packet: &mut udp::Packet) {
let from = packet.from();
let ecn = packet.ecn().map(super::ecn_to_noq);
let mut buf = Vec::new();
let mut fed: Vec<ConnectionHandle> = Vec::new();
for segment in packet.segments() {
buf.clear();
let event = self.endpoint.borrow_mut().handle(
Instant::now(),
from.into(),
ecn,
BytesMut::from(&segment[..]),
&mut buf,
);
match event {
Some(DatagramEvent::ConnectionEvent(key, event)) => {
let conn = self.conns.borrow().get(&key).cloned();
if let Some(conn) = conn {
conn.conn.borrow_mut().handle_event(event);
if !fed.contains(&key) {
fed.push(key);
}
}
}
Some(DatagramEvent::NewConnection(incoming)) => {
if let Some(key) = self.greet(incoming, &mut buf)
&& !fed.contains(&key)
{
fed.push(key);
}
}
Some(DatagramEvent::Response(transmit)) if self.accepting.borrow().is_some() => {
self.respond(&transmit, &buf)
}
Some(DatagramEvent::Response(_)) => {}
None => {}
}
}
for key in fed {
let conn = self.conns.borrow().get(&key).cloned();
if let Some(conn) = conn {
conn.kick();
}
}
}
fn greet(self: &Rc<Self>, incoming: Incoming, buf: &mut Vec<u8>) -> Option<ConnectionHandle> {
if self.handles.get() == 0 {
self.endpoint.borrow_mut().ignore(incoming);
return None;
}
let queued = self
.accepting
.borrow()
.as_ref()
.map(|accepting| accepting.queue.len())
.unwrap_or(0);
if self.pending.get() + queued >= self.backlog {
tracing::debug!(from = %incoming.remote_address(), "dropping a handshake over the backlog");
self.endpoint.borrow_mut().ignore(incoming);
return None;
}
buf.clear();
let accepted = self.endpoint.borrow_mut().accept(incoming, Instant::now(), buf, None);
let (key, conn) = match accepted {
Ok(accepted) => accepted,
Err(err) => {
tracing::debug!(err = %err.cause, "failed to accept a connection");
if let Some(transmit) = err.response {
self.respond(&transmit, buf);
}
return None;
}
};
let shared = self.launch(key, conn);
self.pending.set(self.pending.get() + 1);
let inner = self.clone();
self.owner.spawn(async move {
let outcome = connection::establish(shared).await;
inner.pending.set(inner.pending.get() - 1);
let mut conn = match outcome {
Ok(conn) => conn,
Err(err) => {
tracing::debug!(%err, "incoming handshake failed");
return;
}
};
if inner.handles.get() == 0 {
web_transport_trait::poll::Session::close(&mut conn, 0, "endpoint closed");
return;
}
{
let mut accepting = inner.accepting.borrow_mut();
let accepting = accepting.as_mut().expect("accepted without a server config");
accepting.queue.push_back(conn);
}
inner.accept_waiters.borrow_mut().wake();
});
Some(key)
}
fn respond(&self, transmit: &Transmit, buf: &[u8]) {
let Poll::Ready(Ok(mut tx)) = self.socket.poll_acquire(&kio::Waiter::noop()) else {
return;
};
if transmit.size > tx.len() {
tracing::debug!(size = transmit.size, "dropping an oversized endpoint response");
return;
}
tx[..transmit.size].copy_from_slice(&buf[..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) {
tracing::debug!(%err, "failed to send an endpoint response");
}
}
fn launch(self: &Rc<Self>, key: ConnectionHandle, conn: moq_noq_proto::Connection) -> connection::Shared {
let (shared, driver) = connection::launch(&self.owner, self.socket.clone(), Rc::downgrade(self), key, conn);
self.conns.borrow_mut().insert(key, shared.clone());
let inner = self.clone();
self.owner.spawn(async move {
driver.await;
inner.release(key);
});
shared
}
pub(crate) fn on_connection_event(
&self,
key: ConnectionHandle,
event: moq_noq_proto::EndpointEvent,
) -> Option<moq_noq_proto::ConnectionEvent> {
self.endpoint.borrow_mut().handle_event(key, event)
}
fn release(&self, key: ConnectionHandle) {
self.conns.borrow_mut().remove(&key);
self.task_waiters.borrow_mut().wake();
}
fn fail(&self, err: Error) {
*self.closed.borrow_mut() = Some(err.clone());
let conns: Vec<connection::Shared> = self.conns.borrow().values().cloned().collect();
for conn in conns {
conn.fail(err.clone());
}
self.accept_waiters.borrow_mut().wake();
}
}