use std::{
collections::VecDeque,
fmt,
future::Future,
io,
io::IoSliceMut,
mem,
net::{SocketAddr, SocketAddrV6},
pin::Pin,
str,
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
task::{Context, Poll, Waker},
};
#[cfg(not(wasm_browser))]
use super::runtime::default_runtime;
use super::{
runtime::{AsyncUdpSocket, Runtime},
udp_transmit,
};
use crate::{
ClientConfig, ConnectError, ConnectionError, ConnectionHandle, DatagramEvent, EndpointEvent,
ServerConfig,
};
use crate::{Duration, Instant};
use bytes::{Bytes, BytesMut};
use pin_project_lite::pin_project;
use quinn_udp::{BATCH_SIZE, RecvMeta};
use rustc_hash::FxHashMap;
#[cfg(all(not(wasm_browser), feature = "network-discovery"))]
use socket2::{Domain, Protocol, Socket, Type};
use tokio::sync::mpsc::error::TrySendError;
use tokio::sync::{Notify, futures::Notified, mpsc};
use tracing::{Instrument, Span, debug, error};
use super::{
ConnectionEvent, IO_LOOP_BOUND, RECV_TIME_BOUND, connection::Connecting,
work_limiter::WorkLimiter,
};
use crate::{EndpointConfig, VarInt};
pub(crate) fn is_transient_socket_error(error: &io::Error) -> bool {
matches!(
error.kind(),
io::ErrorKind::AddrNotAvailable
| io::ErrorKind::ConnectionRefused
| io::ErrorKind::ConnectionReset
| io::ErrorKind::HostUnreachable
| io::ErrorKind::NetworkUnreachable
| io::ErrorKind::NotConnected
| io::ErrorKind::TimedOut
) || matches!(error.raw_os_error(), Some(49 | 51 | 65 | 101 | 113))
}
const DRIVER_INFRA_REFS: usize = 1;
const DRIVER_RESPAWN_INITIAL_BACKOFF: Duration = Duration::from_millis(100);
const DRIVER_RESPAWN_MAX_BACKOFF: Duration = Duration::from_secs(5);
const DRIVER_RESPAWN_HEALTHY_RUN: Duration = Duration::from_secs(30);
#[derive(Debug, Clone)]
pub struct Endpoint {
pub(crate) inner: EndpointRef,
pub(crate) default_client_config: Option<ClientConfig>,
runtime: Arc<dyn Runtime>,
}
impl Endpoint {
#[cfg(all(not(wasm_browser), feature = "network-discovery"))]
pub fn client(addr: SocketAddr) -> io::Result<Self> {
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))?;
if addr.is_ipv6() {
if let Err(e) = socket.set_only_v6(false) {
tracing::debug!(%e, "unable to make socket dual-stack");
}
}
use crate::config::buffer_defaults;
let buffer_size = buffer_defaults::PLATFORM_DEFAULT;
if let Err(e) = socket.set_send_buffer_size(buffer_size) {
tracing::debug!(%e, "unable to set send buffer size to {}", buffer_size);
}
if let Err(e) = socket.set_recv_buffer_size(buffer_size) {
tracing::debug!(%e, "unable to set recv buffer size to {}", buffer_size);
}
socket.bind(&addr.into())?;
let runtime =
default_runtime().ok_or_else(|| io::Error::other("no async runtime found"))?;
Self::new_with_abstract_socket(
EndpointConfig::default(),
None,
runtime.wrap_udp_socket(socket.into())?,
runtime,
)
}
#[cfg(all(not(wasm_browser), not(feature = "network-discovery")))]
pub fn client(addr: SocketAddr) -> io::Result<Self> {
let socket = std::net::UdpSocket::bind(addr)?;
let runtime =
default_runtime().ok_or_else(|| io::Error::other("no async runtime found"))?;
Self::new_with_abstract_socket(
EndpointConfig::default(),
None,
runtime.wrap_udp_socket(socket)?,
runtime,
)
}
pub fn stats(&self) -> EndpointStats {
self.inner
.state
.lock()
.map(|state| state.stats)
.unwrap_or_else(|_| {
error!("Endpoint state mutex poisoned");
EndpointStats::default()
})
}
#[cfg(all(not(wasm_browser), feature = "network-discovery"))]
pub fn server(config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))?;
if addr.is_ipv6() {
if let Err(e) = socket.set_only_v6(false) {
tracing::debug!(%e, "unable to make server socket dual-stack");
}
}
socket.set_nonblocking(true)?;
use crate::config::buffer_defaults;
let buffer_size = buffer_defaults::PLATFORM_DEFAULT;
if let Err(e) = socket.set_send_buffer_size(buffer_size) {
tracing::debug!(%e, "unable to set send buffer size to {}", buffer_size);
}
if let Err(e) = socket.set_recv_buffer_size(buffer_size) {
tracing::debug!(%e, "unable to set recv buffer size to {}", buffer_size);
}
socket.bind(&addr.into())?;
let runtime =
default_runtime().ok_or_else(|| io::Error::other("no async runtime found"))?;
Self::new_with_abstract_socket(
EndpointConfig::default(),
Some(config),
runtime.wrap_udp_socket(socket.into())?,
runtime,
)
}
#[cfg(all(not(wasm_browser), not(feature = "network-discovery")))]
pub fn server(config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
let socket = std::net::UdpSocket::bind(addr)?;
let runtime =
default_runtime().ok_or_else(|| io::Error::other("no async runtime found"))?;
Self::new_with_abstract_socket(
EndpointConfig::default(),
Some(config),
runtime.wrap_udp_socket(socket)?,
runtime,
)
}
#[cfg(not(wasm_browser))]
pub fn new(
config: EndpointConfig,
server_config: Option<ServerConfig>,
socket: std::net::UdpSocket,
runtime: Arc<dyn Runtime>,
) -> io::Result<Self> {
let socket = runtime.wrap_udp_socket(socket)?;
Self::new_with_abstract_socket(config, server_config, socket, runtime)
}
pub fn new_with_abstract_socket(
config: EndpointConfig,
server_config: Option<ServerConfig>,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
) -> io::Result<Self> {
let addr = socket.local_addr()?;
let allow_mtud = !socket.may_fragment();
let rc = EndpointRef::new(
socket,
crate::endpoint::Endpoint::new(
Arc::new(config),
server_config.map(Arc::new),
allow_mtud,
None,
),
addr.is_ipv6(),
runtime.clone(),
);
spawn_supervised_driver(rc.clone(), runtime.clone());
Ok(Self {
inner: rc,
default_client_config: None,
runtime,
})
}
pub fn accept(&self) -> Accept<'_> {
Accept {
endpoint: self,
notify: self.inner.shared.incoming.notified(),
}
}
pub fn set_default_client_config(&mut self, config: ClientConfig) {
self.default_client_config = Some(config);
}
pub fn connect(&self, addr: SocketAddr, server_name: &str) -> Result<Connecting, ConnectError> {
let config = match &self.default_client_config {
Some(config) => config.clone(),
None => return Err(ConnectError::NoDefaultClientConfig),
};
self.connect_with(config, addr, server_name)
}
pub fn connect_with(
&self,
config: ClientConfig,
addr: SocketAddr,
server_name: &str,
) -> Result<Connecting, ConnectError> {
let mut endpoint = self
.inner
.state
.lock()
.map_err(|_| ConnectError::EndpointStopping)?;
if endpoint.driver_lost || endpoint.recv_state.connections.close.is_some() {
return Err(ConnectError::EndpointStopping);
}
if addr.is_ipv6() && !endpoint.ipv6 {
return Err(ConnectError::InvalidRemoteAddress(addr));
}
let addr = if endpoint.ipv6 {
SocketAddr::V6(ensure_ipv6(addr))
} else {
addr
};
let (ch, conn) = endpoint
.inner
.connect(self.runtime.now(), config, addr, server_name)?;
let socket = endpoint.socket.clone();
endpoint.stats.outgoing_handshakes += 1;
let registry = endpoint.buffered_registry.clone();
Ok(endpoint
.recv_state
.connections
.insert(ch, conn, socket, self.runtime.clone(), registry))
}
#[cfg(not(wasm_browser))]
pub fn rebind(&self, socket: std::net::UdpSocket) -> io::Result<()> {
self.rebind_abstract(self.runtime.wrap_udp_socket(socket)?)
}
pub fn rebind_abstract(&self, socket: Arc<dyn AsyncUdpSocket>) -> io::Result<()> {
let addr = socket.local_addr()?;
let mut inner = self
.inner
.state
.lock()
.map_err(|_| io::Error::other("Endpoint state mutex poisoned"))?;
inner.prev_socket = Some(mem::replace(&mut inner.socket, socket));
inner.ipv6 = addr.is_ipv6();
let socket = inner.socket.clone();
inner
.recv_state
.connections
.broadcast_control(move || ConnectionEvent::Rebind(socket.clone()));
Ok(())
}
pub fn set_server_config(&self, server_config: Option<ServerConfig>) {
if let Ok(mut state) = self.inner.state.lock() {
state.inner.set_server_config(server_config.map(Arc::new));
} else {
error!("Failed to set server config: endpoint state mutex poisoned");
}
}
pub fn register_connection_peer_id(
&self,
addr: SocketAddr,
peer_id: crate::nat_traversal_api::PeerId,
) {
if let Ok(mut state) = self.inner.state.lock() {
let handle = state.inner.connection_handle_for_addr(&addr);
if let Some(ch) = handle {
state.inner.set_connection_peer_id(ch, peer_id);
tracing::info!(
"Registered peer ID {} for connection {} at low-level endpoint",
hex::encode(&peer_id.0[..8]),
addr
);
} else {
tracing::debug!(
"No connection handle found for {} — peer ID not registered",
addr
);
}
}
}
pub fn set_peer_address_update_tx(&self, tx: mpsc::UnboundedSender<(SocketAddr, SocketAddr)>) {
if let Ok(mut state) = self.inner.state.lock() {
state.peer_address_update_tx = Some(tx);
}
}
pub fn peer_connection_addr_by_id(&self, peer_id: &[u8; 32]) -> Option<SocketAddr> {
let state = self.inner.state.lock().ok()?;
let pid = crate::nat_traversal_api::PeerId(*peer_id);
state.inner.peer_connection_addr(&pid)
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.inner
.state
.lock()
.map_err(|_| io::Error::other("Endpoint state mutex poisoned"))?
.socket
.local_addr()
}
#[cfg(not(wasm_browser))]
pub(crate) fn release_socket_for_shutdown(&self) -> io::Result<Vec<Arc<dyn AsyncUdpSocket>>> {
let (old_addr, runtime) = {
let state = self
.inner
.state
.lock()
.map_err(|_| io::Error::other("Endpoint state mutex poisoned"))?;
if state.socket_released_for_shutdown {
return Ok(Vec::new());
}
(state.socket.local_addr()?, state.runtime.clone())
};
let replacement_addr = if old_addr.is_ipv6() {
SocketAddr::from((std::net::Ipv6Addr::LOCALHOST, 0))
} else {
SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, 0))
};
let replacement = std::net::UdpSocket::bind(replacement_addr).or_else(|first_error| {
if old_addr.is_ipv6() {
let fallback_addr = SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, 0));
std::net::UdpSocket::bind(fallback_addr)
} else {
Err(first_error)
}
})?;
replacement.set_nonblocking(true)?;
let replacement = runtime.wrap_udp_socket(replacement)?;
let replacement_addr = replacement.local_addr()?;
let mut state = self
.inner
.state
.lock()
.map_err(|_| io::Error::other("Endpoint state mutex poisoned"))?;
if state.socket_released_for_shutdown {
return Ok(Vec::new());
}
let mut released = Vec::with_capacity(2);
released.push(mem::replace(&mut state.socket, replacement));
if let Some(prev_socket) = state.prev_socket.take() {
released.push(prev_socket);
}
state.ipv6 = replacement_addr.is_ipv6();
state.socket_released_for_shutdown = true;
Ok(released)
}
pub fn open_connections(&self) -> usize {
self.inner
.state
.lock()
.map(|state| state.inner.open_connections())
.unwrap_or(0)
}
#[doc(hidden)]
pub fn buffered_bytes_totals(&self) -> (u64, u64, usize, usize) {
self.inner
.state
.lock()
.map(|state| state.buffered_registry.totals())
.unwrap_or((0, 0, 0, 0))
}
pub fn set_max_connections(&self, max: usize) {
if let Ok(mut state) = self.inner.state.lock() {
state.recv_state.connections.max_connections = max.max(1);
}
}
pub fn close(&self, error_code: VarInt, reason: &[u8]) {
let reason = Bytes::copy_from_slice(reason);
let mut endpoint = match self.inner.state.lock() {
Ok(endpoint) => endpoint,
Err(_) => {
error!("Failed to close endpoint: state mutex poisoned");
return;
}
};
endpoint.recv_state.connections.close = Some((error_code, reason.clone()));
endpoint
.recv_state
.connections
.broadcast_control(move || ConnectionEvent::Close {
error_code,
reason: reason.clone(),
});
self.inner.shared.incoming.notify_waiters();
}
pub async fn wait_idle(&self) {
loop {
{
let endpoint = match self.inner.state.lock() {
Ok(endpoint) => endpoint,
Err(_) => {
error!("Failed to wait for idle: state mutex poisoned");
break;
}
};
if endpoint.recv_state.connections.is_empty() {
break;
}
self.inner.shared.idle.notified()
}
.await;
}
}
}
#[non_exhaustive]
#[derive(Debug, Default, Copy, Clone)]
pub struct EndpointStats {
pub accepted_handshakes: u64,
pub outgoing_handshakes: u64,
pub refused_handshakes: u64,
pub ignored_handshakes: u64,
}
#[must_use = "endpoint drivers must be spawned for I/O to occur"]
#[derive(Debug)]
pub(crate) struct EndpointDriver(pub(crate) EndpointRef);
impl Future for EndpointDriver {
type Output = Result<(), io::Error>;
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let mut endpoint = match self.0.state.lock() {
Ok(endpoint) => endpoint,
Err(_) => {
return Poll::Ready(Err(io::Error::other("Endpoint state mutex poisoned")));
}
};
if endpoint.driver.is_none() {
endpoint.driver = Some(cx.waker().clone());
}
let now = endpoint.runtime.now();
let mut keep_going = false;
keep_going |= endpoint.drive_recv(cx, now)?;
keep_going |= endpoint.handle_events(cx, &self.0.shared);
if !endpoint.recv_state.incoming.is_empty() {
self.0.shared.incoming.notify_waiters();
}
if endpoint.ref_count == DRIVER_INFRA_REFS && endpoint.recv_state.connections.is_empty() {
Poll::Ready(Ok(()))
} else {
drop(endpoint);
if keep_going {
cx.waker().wake_by_ref();
}
Poll::Pending
}
}
}
impl Drop for EndpointDriver {
fn drop(&mut self) {
if let Ok(mut endpoint) = self.0.state.lock() {
endpoint.driver = None;
if endpoint.recv_state.connections.close.is_some() {
endpoint.driver_lost = true;
endpoint.recv_state.connections.senders.clear();
} else {
debug!("endpoint driver exited without close; supervisor will respawn it");
}
self.0.shared.incoming.notify_waiters();
} else {
error!("Failed to lock endpoint state in drop - mutex poisoned");
}
}
}
fn spawn_supervised_driver(rc: EndpointRef, runtime: Arc<dyn Runtime>) {
let task_runtime = runtime.clone();
runtime.spawn(Box::pin(
async move {
let mut backoff = DRIVER_RESPAWN_INITIAL_BACKOFF;
loop {
let started = task_runtime.now();
let result = EndpointDriver(rc.clone()).await;
let Err(error) = result else {
return;
};
let closed = match rc.state.lock() {
Ok(state) => {
state.recv_state.connections.close.is_some()
|| state.socket_released_for_shutdown
}
Err(_) => {
error!(
"endpoint driver terminated ({error}) and the state mutex is poisoned; \
not respawning"
);
return;
}
};
if closed {
debug!("endpoint driver terminated after intentional close: {error}");
return;
}
error!("endpoint driver terminated: {error} — respawning in {backoff:?}");
if task_runtime.now() >= started + DRIVER_RESPAWN_HEALTHY_RUN {
backoff = DRIVER_RESPAWN_INITIAL_BACKOFF;
}
sleep_for(&*task_runtime, backoff).await;
backoff = (backoff * 2).min(DRIVER_RESPAWN_MAX_BACKOFF);
let closed_during_backoff = match rc.state.lock() {
Ok(state) => {
state.recv_state.connections.close.is_some()
|| state.socket_released_for_shutdown
}
Err(_) => true,
};
if closed_during_backoff {
debug!("endpoint driver not respawned: endpoint closed during backoff");
return;
}
#[cfg(not(wasm_browser))]
rebind_for_respawn(&rc, &*task_runtime);
}
}
.instrument(Span::current()),
));
}
async fn sleep_for(runtime: &dyn Runtime, duration: Duration) {
let mut timer = runtime.new_timer(runtime.now() + duration);
std::future::poll_fn(|cx| timer.as_mut().poll(cx)).await;
}
#[cfg(not(wasm_browser))]
fn rebind_for_respawn(rc: &EndpointRef, runtime: &dyn Runtime) {
let old_addr = match rc.state.lock() {
Ok(state) => match state.socket.local_addr() {
Ok(addr) => addr,
Err(e) => {
debug!("driver respawn: old socket address unreadable ({e}); keeping old socket");
return;
}
},
Err(_) => {
error!("endpoint state mutex poisoned; cannot rebind for driver respawn");
return;
}
};
let socket = match bind_replacement_socket(old_addr)
.and_then(|socket| runtime.wrap_udp_socket(socket))
{
Ok(socket) => socket,
Err(e) => {
debug!("driver respawn: rebind to {old_addr} failed ({e}); keeping old socket");
return;
}
};
let mut state = match rc.state.lock() {
Ok(state) => state,
Err(_) => {
error!("endpoint state mutex poisoned; cannot rebind for driver respawn");
return;
}
};
state.prev_socket = Some(mem::replace(&mut state.socket, socket));
state.ipv6 = old_addr.is_ipv6();
let socket = state.socket.clone();
state
.recv_state
.connections
.broadcast_control(move || ConnectionEvent::Rebind(socket.clone()));
tracing::info!("endpoint driver respawn rebound socket on {old_addr}");
}
#[cfg(all(not(wasm_browser), feature = "network-discovery"))]
fn bind_replacement_socket(addr: SocketAddr) -> io::Result<std::net::UdpSocket> {
let socket = Socket::new(Domain::for_address(addr), Type::DGRAM, Some(Protocol::UDP))?;
if addr.is_ipv6() {
if let Err(e) = socket.set_only_v6(false) {
debug!(%e, "unable to make replacement socket dual-stack");
}
}
socket.set_nonblocking(true)?;
use crate::config::buffer_defaults;
let buffer_size = buffer_defaults::PLATFORM_DEFAULT;
if let Err(e) = socket.set_send_buffer_size(buffer_size) {
debug!(%e, "unable to set send buffer size on replacement socket");
}
if let Err(e) = socket.set_recv_buffer_size(buffer_size) {
debug!(%e, "unable to set recv buffer size on replacement socket");
}
socket.bind(&addr.into())?;
Ok(socket.into())
}
#[cfg(all(not(wasm_browser), not(feature = "network-discovery")))]
fn bind_replacement_socket(addr: SocketAddr) -> io::Result<std::net::UdpSocket> {
let socket = std::net::UdpSocket::bind(addr)?;
socket.set_nonblocking(true)?;
Ok(socket)
}
#[derive(Debug)]
pub(crate) struct EndpointInner {
pub(crate) state: Mutex<State>,
pub(crate) shared: Shared,
}
impl EndpointInner {
pub(crate) fn accept(
&self,
incoming: crate::Incoming,
server_config: Option<Arc<ServerConfig>>,
) -> Result<Connecting, ConnectionError> {
let mut state = self.state.lock().map_err(|_| {
ConnectionError::TransportError(crate::transport_error::Error::INTERNAL_ERROR(
"Endpoint state mutex poisoned".to_string(),
))
})?;
let mut response_buffer = Vec::new();
let now = state.runtime.now();
match state
.inner
.accept(incoming, now, &mut response_buffer, server_config)
{
Ok((handle, conn)) => {
state.stats.accepted_handshakes += 1;
let socket = state.socket.clone();
let runtime = state.runtime.clone();
let registry = state.buffered_registry.clone();
Ok(state
.recv_state
.connections
.insert(handle, conn, socket, runtime, registry))
}
Err(error) => {
if let Some(transmit) = error.response {
respond(transmit, &response_buffer, &*state.socket);
}
Err(error.cause)
}
}
}
pub(crate) fn refuse(&self, incoming: crate::Incoming) {
let mut state = match self.state.lock() {
Ok(state) => state,
Err(_) => {
error!("Failed to refuse connection: endpoint state mutex poisoned");
return;
}
};
state.stats.refused_handshakes += 1;
let mut response_buffer = Vec::new();
let transmit = state.inner.refuse(incoming, &mut response_buffer);
respond(transmit, &response_buffer, &*state.socket);
}
pub(crate) fn retry(
&self,
incoming: crate::Incoming,
) -> Result<(), crate::endpoint::RetryError> {
let mut state = match self.state.lock() {
Ok(state) => state,
Err(_) => {
error!("Failed to retry connection: endpoint state mutex poisoned");
return Err(crate::endpoint::RetryError::incoming(incoming));
}
};
let mut response_buffer = Vec::new();
let transmit = state.inner.retry(incoming, &mut response_buffer)?;
respond(transmit, &response_buffer, &*state.socket);
Ok(())
}
pub(crate) fn ignore(&self, incoming: crate::Incoming) {
if let Ok(mut state) = self.state.lock() {
state.stats.ignored_handshakes += 1;
state.inner.ignore(incoming);
} else {
error!("Failed to ignore incoming connection: endpoint state mutex poisoned");
}
}
}
#[derive(Debug)]
pub(crate) struct State {
socket: Arc<dyn AsyncUdpSocket>,
prev_socket: Option<Arc<dyn AsyncUdpSocket>>,
inner: crate::endpoint::Endpoint,
recv_state: RecvState,
driver: Option<Waker>,
ipv6: bool,
events: mpsc::UnboundedReceiver<(ConnectionHandle, EndpointEvent)>,
ref_count: usize,
driver_lost: bool,
runtime: Arc<dyn Runtime>,
stats: EndpointStats,
peer_address_update_tx: Option<mpsc::UnboundedSender<(SocketAddr, SocketAddr)>>,
socket_released_for_shutdown: bool,
buffered_registry: Arc<BufferedBytesRegistry>,
}
#[derive(Debug)]
pub(crate) struct Shared {
incoming: Notify,
idle: Notify,
}
impl State {
fn drive_recv(&mut self, cx: &mut Context, now: Instant) -> Result<bool, io::Error> {
let get_time = || self.runtime.now();
self.recv_state.recv_limiter.start_cycle(get_time);
if let Some(socket) = &self.prev_socket {
let poll_res =
self.recv_state
.poll_socket(cx, &mut self.inner, &**socket, &*self.runtime, now);
if poll_res.is_err() {
self.prev_socket = None;
}
};
let poll_res =
self.recv_state
.poll_socket(cx, &mut self.inner, &*self.socket, &*self.runtime, now);
self.recv_state.recv_limiter.finish_cycle(get_time);
let poll_res = poll_res?;
if poll_res.received_connection_packet {
self.prev_socket = None;
}
Ok(poll_res.keep_going)
}
fn handle_events(&mut self, cx: &mut Context, shared: &Shared) -> bool {
let mut did_work = false;
for _ in 0..IO_LOOP_BOUND {
let (ch, event) = match self.events.poll_recv(cx) {
Poll::Ready(Some(x)) => x,
Poll::Ready(None) => unreachable!("EndpointInner owns one sender"),
Poll::Pending => {
break;
}
};
did_work = true;
if event.is_drained() {
self.recv_state.connections.senders.remove(&ch);
if self.recv_state.connections.is_empty() {
shared.idle.notify_waiters();
}
}
let Some(event) = self.inner.handle_event(ch, event) else {
continue;
};
self.recv_state.connections.send_proto(ch, event);
}
for (ch, event) in self.inner.drain_relay_events() {
did_work = true;
if self.recv_state.connections.senders.contains_key(&ch) {
tracing::debug!("Sending relay event to connection {:?}", ch);
self.recv_state.connections.send_proto(ch, event);
} else {
tracing::warn!(
"Cannot send relay event: connection {:?} not found in senders",
ch
);
}
}
let address_updates: Vec<(SocketAddr, SocketAddr)> =
self.inner.drain_peer_address_updates().collect();
for (peer_addr, advertised_addr) in address_updates {
did_work = true;
if let Some(ref tx) = self.peer_address_update_tx {
let _ = tx.send((peer_addr, advertised_addr));
}
}
did_work
}
}
impl Drop for State {
fn drop(&mut self) {
for incoming in self.recv_state.incoming.drain(..) {
self.inner.ignore(incoming);
}
}
}
fn respond(transmit: crate::Transmit, response_buffer: &[u8], socket: &dyn AsyncUdpSocket) {
let mut sender = socket.create_sender();
let waker = futures_util::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let _ = sender.as_mut().poll_send(
&udp_transmit(&transmit, &response_buffer[..transmit.size]),
&mut cx,
);
}
#[inline]
fn proto_ecn(ecn: quinn_udp::EcnCodepoint) -> crate::EcnCodepoint {
match ecn {
quinn_udp::EcnCodepoint::Ect0 => crate::EcnCodepoint::Ect0,
quinn_udp::EcnCodepoint::Ect1 => crate::EcnCodepoint::Ect1,
quinn_udp::EcnCodepoint::Ce => crate::EcnCodepoint::Ce,
}
}
const RECV_EVENT_BOUND: usize = 256;
const RECV_OVERFLOW_KILL_THRESHOLD: u64 = 1024;
const DEFAULT_MAX_CONNECTIONS: usize = 4096;
#[derive(Debug)]
struct ConnectionChannels {
sender: mpsc::Sender<ConnectionEvent>,
recv_overflows: AtomicU64,
}
#[derive(Debug, Default)]
pub(crate) struct BufferedBytesRegistry {
entries: std::sync::Mutex<FxHashMap<usize, ConnBufferedSnapshot>>,
}
#[derive(Debug, Clone, Copy, Default)]
pub(crate) struct ConnBufferedSnapshot {
pub(crate) send_unacked: u64,
pub(crate) recv_buffered: u64,
pub(crate) recv_streams_with_unread: usize,
}
impl BufferedBytesRegistry {
pub(crate) fn update(&self, handle: usize, snap: ConnBufferedSnapshot) {
if let Ok(mut m) = self.entries.lock() {
m.insert(handle, snap);
}
}
pub(crate) fn remove(&self, handle: usize) {
if let Ok(mut m) = self.entries.lock() {
m.remove(&handle);
}
}
pub(crate) fn totals(&self) -> (u64, u64, usize, usize) {
match self.entries.lock() {
Ok(m) => {
let mut send_unacked = 0u64;
let mut recv_buffered = 0u64;
let mut streams = 0usize;
for s in m.values() {
send_unacked += s.send_unacked;
recv_buffered += s.recv_buffered;
streams += s.recv_streams_with_unread;
}
(send_unacked, recv_buffered, streams, m.len())
}
Err(poisoned) => {
let m = poisoned.into_inner();
let n = m.len();
(0, 0, 0, n)
}
}
}
}
struct ConnectionSet {
senders: FxHashMap<ConnectionHandle, ConnectionChannels>,
sender: mpsc::UnboundedSender<(ConnectionHandle, EndpointEvent)>,
close: Option<(VarInt, Bytes)>,
max_connections: usize,
}
impl ConnectionSet {
fn insert(
&mut self,
handle: ConnectionHandle,
conn: crate::Connection,
socket: Arc<dyn AsyncUdpSocket>,
runtime: Arc<dyn Runtime>,
buffered_registry: Arc<BufferedBytesRegistry>,
) -> Connecting {
let (send, recv) = mpsc::channel(RECV_EVENT_BOUND);
if let Some((error_code, ref reason)) = self.close {
let _ = send.try_send(ConnectionEvent::Close {
error_code,
reason: reason.clone(),
});
}
self.senders.insert(
handle,
ConnectionChannels {
sender: send,
recv_overflows: AtomicU64::new(0),
},
);
Connecting::new(
handle,
conn,
self.sender.clone(),
recv,
socket,
runtime,
buffered_registry,
)
}
fn is_empty(&self) -> bool {
self.senders.is_empty()
}
fn send_proto(&mut self, ch: ConnectionHandle, event: crate::shared::ConnectionEvent) {
let kill = match self.senders.get(&ch) {
None => return,
Some(channels) => match channels.sender.try_send(ConnectionEvent::Proto(event)) {
Ok(()) => {
channels.recv_overflows.store(0, Ordering::Relaxed);
false
}
Err(TrySendError::Closed(_)) => return,
Err(TrySendError::Full(_)) => {
channels.recv_overflows.fetch_add(1, Ordering::Relaxed) + 1
>= RECV_OVERFLOW_KILL_THRESHOLD
}
},
};
if kill {
self.senders.remove(&ch);
}
}
fn broadcast_control<F>(&mut self, mut make_event: F)
where
F: FnMut() -> ConnectionEvent,
{
let mut kill = Vec::new();
for (ch, channels) in self.senders.iter() {
match channels.sender.try_send(make_event()) {
Ok(()) => {}
Err(TrySendError::Closed(_)) => {}
Err(TrySendError::Full(_)) => kill.push(*ch),
}
}
for ch in kill {
self.senders.remove(&ch);
}
}
}
fn ensure_ipv6(x: SocketAddr) -> SocketAddrV6 {
match x {
SocketAddr::V6(x) => x,
SocketAddr::V4(x) => SocketAddrV6::new(x.ip().to_ipv6_mapped(), x.port(), 0, 0),
}
}
pin_project! {
pub struct Accept<'a> {
endpoint: &'a Endpoint,
#[pin]
notify: Notified<'a>,
}
}
impl Future for Accept<'_> {
type Output = Option<super::incoming::Incoming>;
fn poll(self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Self::Output> {
let mut this = self.project();
let mut endpoint = match this.endpoint.inner.state.lock() {
Ok(endpoint) => endpoint,
Err(_) => return Poll::Ready(None),
};
if endpoint.driver_lost {
return Poll::Ready(None);
}
if let Some(incoming) = endpoint.recv_state.incoming.pop_front() {
drop(endpoint);
let incoming = super::incoming::Incoming::new(incoming, this.endpoint.inner.clone());
return Poll::Ready(Some(incoming));
}
if endpoint.recv_state.connections.close.is_some() {
return Poll::Ready(None);
}
loop {
match this.notify.as_mut().poll(ctx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(()) => this
.notify
.set(this.endpoint.inner.shared.incoming.notified()),
}
}
}
}
#[derive(Debug)]
pub(crate) struct EndpointRef(Arc<EndpointInner>);
impl EndpointRef {
pub(crate) fn new(
socket: Arc<dyn AsyncUdpSocket>,
inner: crate::endpoint::Endpoint,
ipv6: bool,
runtime: Arc<dyn Runtime>,
) -> Self {
let (sender, events) = mpsc::unbounded_channel();
let recv_state = RecvState::new(sender, socket.max_receive_segments(), &inner);
Self(Arc::new(EndpointInner {
shared: Shared {
incoming: Notify::new(),
idle: Notify::new(),
},
state: Mutex::new(State {
socket,
prev_socket: None,
inner,
ipv6,
events,
driver: None,
ref_count: 0,
driver_lost: false,
recv_state,
runtime,
stats: EndpointStats::default(),
peer_address_update_tx: None,
socket_released_for_shutdown: false,
buffered_registry: Arc::new(BufferedBytesRegistry::default()),
}),
}))
}
}
impl Clone for EndpointRef {
fn clone(&self) -> Self {
if let Ok(mut state) = self.0.state.lock() {
state.ref_count += 1;
}
Self(self.0.clone())
}
}
impl Drop for EndpointRef {
fn drop(&mut self) {
if let Ok(mut endpoint) = self.0.state.lock() {
if let Some(x) = endpoint.ref_count.checked_sub(1) {
endpoint.ref_count = x;
if x == DRIVER_INFRA_REFS {
if let Some(task) = endpoint.driver.take() {
task.wake();
}
}
}
} else {
error!("Failed to drop EndpointRef: state mutex poisoned");
}
}
}
impl std::ops::Deref for EndpointRef {
type Target = EndpointInner;
fn deref(&self) -> &Self::Target {
&self.0
}
}
struct RecvState {
incoming: VecDeque<crate::Incoming>,
connections: ConnectionSet,
recv_buf: Box<[u8]>,
recv_limiter: WorkLimiter,
}
impl RecvState {
fn new(
sender: mpsc::UnboundedSender<(ConnectionHandle, EndpointEvent)>,
max_receive_segments: usize,
endpoint: &crate::endpoint::Endpoint,
) -> Self {
const PQC_MIN_RECV_SIZE: u64 = 4096;
let configured_size = endpoint.config().get_max_udp_payload_size();
let effective_size = configured_size.max(PQC_MIN_RECV_SIZE).min(64 * 1024) as usize;
let recv_buf = vec![0; effective_size * max_receive_segments * BATCH_SIZE];
Self {
connections: ConnectionSet {
senders: FxHashMap::default(),
sender,
close: None,
max_connections: DEFAULT_MAX_CONNECTIONS,
},
incoming: VecDeque::new(),
recv_buf: recv_buf.into(),
recv_limiter: WorkLimiter::new(RECV_TIME_BOUND),
}
}
fn poll_socket(
&mut self,
cx: &mut Context,
endpoint: &mut crate::endpoint::Endpoint,
socket: &dyn AsyncUdpSocket,
runtime: &dyn Runtime,
now: Instant,
) -> Result<PollProgress, io::Error> {
let mut received_connection_packet = false;
let mut metas = [RecvMeta::default(); BATCH_SIZE];
let mut iovs: [IoSliceMut; BATCH_SIZE] = {
let mut bufs = self
.recv_buf
.chunks_mut(self.recv_buf.len() / BATCH_SIZE)
.map(IoSliceMut::new);
std::array::from_fn(|_| {
bufs.next().unwrap_or_else(|| {
error!("Insufficient buffers for BATCH_SIZE");
IoSliceMut::new(&mut [])
})
})
};
loop {
match socket.poll_recv(cx, &mut iovs, &mut metas) {
Poll::Ready(Ok(msgs)) => {
self.recv_limiter.record_work(msgs);
for (meta, buf) in metas.iter().zip(iovs.iter()).take(msgs) {
let mut data: BytesMut = buf[0..meta.len].into();
while !data.is_empty() {
let buf = data.split_to(meta.stride.min(data.len()));
let mut response_buffer = Vec::new();
match endpoint.handle(
now,
meta.addr,
meta.dst_ip,
meta.ecn.map(proto_ecn),
buf,
&mut response_buffer,
) {
Some(DatagramEvent::NewConnection(incoming)) => {
if self.connections.close.is_none()
&& self.connections.senders.len()
< self.connections.max_connections
{
self.incoming.push_back(incoming);
} else {
let transmit =
endpoint.refuse(incoming, &mut response_buffer);
respond(transmit, &response_buffer, socket);
}
}
Some(DatagramEvent::ConnectionEvent(handle, event)) => {
received_connection_packet = true;
self.connections.send_proto(handle, event);
}
Some(DatagramEvent::Response(transmit)) => {
respond(transmit, &response_buffer, socket);
}
None => {}
}
}
}
}
Poll::Pending => {
return Ok(PollProgress {
received_connection_packet,
keep_going: false,
});
}
Poll::Ready(Err(ref e)) if is_transient_socket_error(e) => {
debug!("ignoring transient socket recv error: {}", e);
continue;
}
Poll::Ready(Err(e)) => {
return Err(e);
}
}
if !self.recv_limiter.allow_work(|| runtime.now()) {
return Ok(PollProgress {
received_connection_packet,
keep_going: true,
});
}
}
}
}
impl fmt::Debug for RecvState {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.debug_struct("RecvState")
.field("incoming", &self.incoming)
.field("connections", &self.connections.senders.len())
.field("recv_limiter", &self.recv_limiter)
.finish_non_exhaustive()
}
}
#[derive(Default)]
struct PollProgress {
received_connection_packet: bool,
keep_going: bool,
}
#[cfg(test)]
mod driver_error_tests {
use super::is_transient_socket_error;
use std::io;
#[test]
fn icmp_derived_errors_are_transient() {
for kind in [
io::ErrorKind::HostUnreachable,
io::ErrorKind::NetworkUnreachable,
io::ErrorKind::ConnectionRefused,
io::ErrorKind::ConnectionReset,
io::ErrorKind::AddrNotAvailable,
io::ErrorKind::NotConnected,
io::ErrorKind::TimedOut,
] {
assert!(
is_transient_socket_error(&io::Error::new(kind, "test")),
"{kind:?} must not kill the endpoint driver"
);
}
for errno in [49, 51, 65, 101, 113] {
assert!(
is_transient_socket_error(&io::Error::from_raw_os_error(errno)),
"errno {errno} must not kill the endpoint driver"
);
}
}
#[test]
fn fatal_errors_still_terminate() {
for kind in [
io::ErrorKind::BrokenPipe,
io::ErrorKind::PermissionDenied,
io::ErrorKind::NotFound,
io::ErrorKind::InvalidInput,
] {
assert!(
!is_transient_socket_error(&io::Error::new(kind, "test")),
"{kind:?} must remain driver-fatal"
);
}
}
}
#[cfg(test)]
mod recv_event_backpressure_tests {
use super::*;
use crate::shared::ConnectionEventInner;
fn proto_event() -> crate::shared::ConnectionEvent {
crate::shared::ConnectionEvent(ConnectionEventInner::NewIdentifiers(
Vec::new(),
Instant::now(),
))
}
fn empty_set() -> ConnectionSet {
let (sender, _events) = mpsc::unbounded_channel();
ConnectionSet {
senders: FxHashMap::default(),
sender,
close: None,
max_connections: DEFAULT_MAX_CONNECTIONS,
}
}
fn insert_stuck_connection(
cs: &mut ConnectionSet,
) -> (ConnectionHandle, mpsc::Receiver<ConnectionEvent>) {
let (tx, rx) = mpsc::channel(RECV_EVENT_BOUND);
let handle = ConnectionHandle(0);
cs.senders.insert(
handle,
ConnectionChannels {
sender: tx,
recv_overflows: AtomicU64::new(0),
},
);
(handle, rx)
}
#[test]
fn full_channel_drops_events_and_counts_overflows() {
let mut cs = empty_set();
let (handle, _rx) = insert_stuck_connection(&mut cs);
for _ in 0..RECV_EVENT_BOUND {
cs.send_proto(handle, proto_event());
}
assert_eq!(
cs.senders
.get(&handle)
.unwrap()
.recv_overflows
.load(Ordering::Relaxed),
0,
"successful sends must reset the overflow counter"
);
for _ in 0..10 {
cs.send_proto(handle, proto_event());
}
assert_eq!(
cs.senders
.get(&handle)
.unwrap()
.recv_overflows
.load(Ordering::Relaxed),
10,
"each dropped event must increment the overflow counter"
);
assert!(
cs.senders.contains_key(&handle),
"below the kill threshold the connection must stay alive"
);
}
#[test]
fn overflow_past_threshold_force_closes_connection() {
let mut cs = empty_set();
let (tx, _rx) = mpsc::channel(RECV_EVENT_BOUND);
let handle = ConnectionHandle(7);
cs.senders.insert(
handle,
ConnectionChannels {
sender: tx,
recv_overflows: AtomicU64::new(0),
},
);
for _ in 0..RECV_EVENT_BOUND {
cs.send_proto(handle, proto_event());
}
for _ in 0..RECV_OVERFLOW_KILL_THRESHOLD {
cs.send_proto(handle, proto_event());
}
assert!(
!cs.senders.contains_key(&handle),
"past the kill threshold the sender must be dropped (force-close)"
);
assert!(cs.senders.is_empty());
}
#[test]
fn draining_consumer_never_killed() {
let mut cs = empty_set();
let (tx, mut rx) = mpsc::channel(RECV_EVENT_BOUND);
let handle = ConnectionHandle(1);
cs.senders.insert(
handle,
ConnectionChannels {
sender: tx,
recv_overflows: AtomicU64::new(0),
},
);
for _ in 0..(RECV_EVENT_BOUND + 50) {
cs.send_proto(handle, proto_event());
}
assert!(cs.senders.contains_key(&handle));
while rx.try_recv().is_ok() {}
for _ in 0..5 {
cs.send_proto(handle, proto_event());
}
assert_eq!(
cs.senders
.get(&handle)
.unwrap()
.recv_overflows
.load(Ordering::Relaxed),
0,
"a successful send after draining must reset the overflow counter"
);
assert!(cs.senders.contains_key(&handle));
}
}
#[cfg(test)]
mod driver_supervisor_tests {
use super::*;
use crate::high_level::runtime::UdpSender;
use std::sync::atomic::AtomicUsize;
#[derive(Debug)]
struct ControllableSocket {
addr: SocketAddr,
fail_with: Mutex<Option<io::ErrorKind>>,
recv_calls: AtomicUsize,
}
#[derive(Debug)]
struct NoopSender;
impl UdpSender for NoopSender {
fn poll_send(
self: Pin<&mut Self>,
_transmit: &quinn_udp::Transmit,
_cx: &mut Context<'_>,
) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
impl AsyncUdpSocket for ControllableSocket {
fn create_sender(&self) -> Pin<Box<dyn UdpSender>> {
Box::pin(NoopSender)
}
fn poll_recv(
&self,
_cx: &mut Context,
_bufs: &mut [IoSliceMut<'_>],
_meta: &mut [RecvMeta],
) -> Poll<io::Result<usize>> {
self.recv_calls.fetch_add(1, Ordering::Relaxed);
match *self.fail_with.lock().unwrap() {
Some(kind) => Poll::Ready(Err(io::Error::new(kind, "injected fatal recv error"))),
None => Poll::Pending,
}
}
fn local_addr(&self) -> io::Result<SocketAddr> {
Ok(self.addr)
}
}
fn test_endpoint_ref(
socket: Arc<ControllableSocket>,
) -> (EndpointRef, Arc<EndpointInner>, Arc<dyn Runtime>) {
let runtime = default_runtime().expect("tests run inside a tokio runtime");
let rc = EndpointRef::new(
socket,
crate::endpoint::Endpoint::new(Arc::new(EndpointConfig::default()), None, false, None),
false,
runtime.clone(),
);
let inner = rc.0.clone();
(rc, inner, runtime)
}
fn unbindable_addr() -> SocketAddr {
SocketAddr::from((std::net::Ipv4Addr::new(192, 0, 2, 1), 12345))
}
async fn wait_until(mut condition: impl FnMut() -> bool, timeout: Duration) -> bool {
let deadline = std::time::Instant::now() + timeout;
while !condition() {
if std::time::Instant::now() >= deadline {
return false;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
true
}
fn recv_calls(socket: &ControllableSocket) -> usize {
socket.recv_calls.load(Ordering::Relaxed)
}
fn insert_fake_connection(
inner: &EndpointInner,
) -> (ConnectionHandle, mpsc::Receiver<ConnectionEvent>) {
let (tx, rx) = mpsc::channel(RECV_EVENT_BOUND);
let handle = ConnectionHandle(0);
inner
.state
.lock()
.unwrap()
.recv_state
.connections
.senders
.insert(
handle,
ConnectionChannels {
sender: tx,
recv_overflows: AtomicU64::new(0),
},
);
(handle, rx)
}
#[tokio::test]
async fn fatal_driver_exit_is_respawned_and_shutdown_stays_clean() {
let socket = Arc::new(ControllableSocket {
addr: unbindable_addr(),
fail_with: Mutex::new(Some(io::ErrorKind::PermissionDenied)),
recv_calls: AtomicUsize::new(0),
});
let (rc, inner, runtime) = test_endpoint_ref(socket.clone());
let user_handle = rc.clone();
spawn_supervised_driver(rc, runtime);
assert!(
wait_until(|| recv_calls(&socket) >= 1, Duration::from_secs(5)).await,
"driver must poll the socket"
);
assert!(
wait_until(|| recv_calls(&socket) >= 2, Duration::from_secs(5)).await,
"supervisor must respawn the driver after a fatal exit"
);
assert!(
!inner.state.lock().unwrap().driver_lost,
"supervised driver death must not set driver_lost"
);
*socket.fail_with.lock().unwrap() = None;
let calls = recv_calls(&socket);
assert!(
wait_until(|| recv_calls(&socket) > calls, Duration::from_secs(5)).await,
"respawned driver must poll the recovered socket"
);
drop(user_handle);
assert!(
wait_until(|| Arc::strong_count(&inner) == 1, Duration::from_secs(5)).await,
"driver and supervisor must release the endpoint after the last handle drops"
);
}
#[tokio::test]
async fn respawn_preserves_connection_channels() {
let socket = Arc::new(ControllableSocket {
addr: unbindable_addr(),
fail_with: Mutex::new(Some(io::ErrorKind::PermissionDenied)),
recv_calls: AtomicUsize::new(0),
});
let (rc, inner, runtime) = test_endpoint_ref(socket.clone());
let _user_handle = rc.clone();
let (handle, _rx) = insert_fake_connection(&inner);
spawn_supervised_driver(rc, runtime);
assert!(
wait_until(|| recv_calls(&socket) >= 2, Duration::from_secs(5)).await,
"supervisor must respawn the driver after a fatal exit"
);
let state = inner.state.lock().unwrap();
assert!(
state.recv_state.connections.senders.contains_key(&handle),
"respawn must preserve live connection channels"
);
assert!(!state.driver_lost);
*socket.fail_with.lock().unwrap() = None;
}
#[tokio::test]
async fn closed_endpoint_driver_is_not_respawned() {
let socket = Arc::new(ControllableSocket {
addr: unbindable_addr(),
fail_with: Mutex::new(Some(io::ErrorKind::PermissionDenied)),
recv_calls: AtomicUsize::new(0),
});
let (rc, inner, runtime) = test_endpoint_ref(socket.clone());
let _user_handle = rc.clone();
let (_handle, _rx) = insert_fake_connection(&inner);
inner.state.lock().unwrap().recv_state.connections.close =
Some((VarInt::from_u32(0), Bytes::new()));
spawn_supervised_driver(rc, runtime);
assert!(
wait_until(|| recv_calls(&socket) >= 1, Duration::from_secs(5)).await,
"driver must poll the socket once"
);
tokio::time::sleep(DRIVER_RESPAWN_INITIAL_BACKOFF * 5).await;
assert_eq!(
recv_calls(&socket),
1,
"driver on a closed endpoint must not be respawned"
);
let state = inner.state.lock().unwrap();
assert!(
state.driver_lost,
"terminal teardown on a closed endpoint must set driver_lost"
);
assert!(
state.recv_state.connections.senders.is_empty(),
"terminal teardown must drop connection channels"
);
}
#[tokio::test]
async fn shutdown_during_backoff_prevents_respawn() {
let socket = Arc::new(ControllableSocket {
addr: unbindable_addr(),
fail_with: Mutex::new(Some(io::ErrorKind::PermissionDenied)),
recv_calls: AtomicUsize::new(0),
});
let (rc, inner, runtime) = test_endpoint_ref(socket.clone());
let _user_handle = rc.clone();
spawn_supervised_driver(rc, runtime);
assert!(
wait_until(|| recv_calls(&socket) >= 1, Duration::from_secs(5)).await,
"driver must poll the socket once"
);
inner.state.lock().unwrap().socket_released_for_shutdown = true;
tokio::time::sleep(DRIVER_RESPAWN_INITIAL_BACKOFF * 5).await;
assert_eq!(
recv_calls(&socket),
1,
"supervisor must not respawn a driver after shutdown released the socket"
);
}
#[tokio::test]
async fn close_marks_endpoint_closed() {
let socket = Arc::new(ControllableSocket {
addr: unbindable_addr(),
fail_with: Mutex::new(None),
recv_calls: AtomicUsize::new(0),
});
let runtime = default_runtime().expect("tests run inside a tokio runtime");
let endpoint =
Endpoint::new_with_abstract_socket(EndpointConfig::default(), None, socket, runtime)
.expect("endpoint construction");
endpoint.close(VarInt::from_u32(0), b"done");
assert!(
endpoint
.inner
.state
.lock()
.unwrap()
.recv_state
.connections
.close
.is_some(),
"close must record the close reason"
);
assert!(
endpoint.accept().await.is_none(),
"accept must drain once the endpoint is closed"
);
}
#[cfg(not(wasm_browser))]
#[test]
fn replacement_socket_rebinds_after_old_socket_gone() {
let addr = SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, 0));
let first = bind_replacement_socket(addr).expect("initial bind");
let bound = first.local_addr().expect("local addr");
assert!(
bind_replacement_socket(bound).is_err(),
"rebind must fail while the old socket owns the address"
);
drop(first);
let second = bind_replacement_socket(bound).expect("rebind after old socket is gone");
assert_eq!(
second.local_addr().expect("local addr"),
bound,
"replacement socket must take over the old address"
);
}
}