use std::collections::hash_map::Entry;
use std::collections::{HashMap, HashSet};
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, RwLock, Weak};
use quinn::VarInt;
use tokio::sync::{mpsc, watch};
use weida_core::{Error, Fingerprint, Limits, validate_endpoint_path};
use weida_protocol::codes;
use weida_protocol::header::GuaranteeSet;
use crate::config::{ClientTls, ServerTls};
use crate::conn::ConnCtx;
use crate::endpoint::{
BusMember, BusState, Endpoint, PairState, Paired, PubState, Publisher, PullState, Puller,
RepState, Replier, RespondState, Respondent,
};
use crate::inproc;
use crate::pubsub::SubRegistry;
use crate::runtime::{Exec, RuntimeInner, Shared};
use crate::stream::{Acceptor, Incoming};
use crate::tls;
use crate::transfer::{IncomingRequest, IncomingTransfer};
use crate::transport::Link;
pub(crate) enum Route {
Request(mpsc::Sender<IncomingRequest>),
Transfer(mpsc::Sender<IncomingTransfer>),
Raw(mpsc::Sender<Incoming>),
Pub,
Pair {
queue: mpsc::Sender<IncomingTransfer>,
owner: Arc<PairOwner>,
},
}
pub(crate) struct PairOwner {
id: AtomicU64,
conn: watch::Sender<Option<Weak<ConnCtx>>>,
}
impl PairOwner {
pub(crate) fn new() -> PairOwner {
PairOwner {
id: AtomicU64::new(0),
conn: watch::channel(None).0,
}
}
pub(crate) fn claim(&self, conn: &Arc<ConnCtx>) -> bool {
let id = conn.conn.stable_id() as u64 + 1;
loop {
match self
.id
.compare_exchange(0, id, Ordering::AcqRel, Ordering::Acquire)
{
Ok(_) => {
self.conn.send_replace(Some(Arc::downgrade(conn)));
return true;
}
Err(held) if held == id => return true,
Err(held) => {
if self.holder_is_live() {
return false;
}
if self
.id
.compare_exchange(held, 0, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.conn.send_replace(None);
}
}
}
}
}
fn holder_is_live(&self) -> bool {
let held = self.conn.borrow().clone();
held.and_then(|weak| weak.upgrade())
.is_some_and(|ctx| ctx.conn.close_reason().is_none())
}
pub(crate) async fn peer(&self) -> Result<Arc<ConnCtx>, Error> {
let mut rx = self.conn.subscribe();
loop {
if let Some(weak) = rx.borrow_and_update().clone() {
return weak.upgrade().ok_or(Error::NotConnected);
}
rx.changed().await.map_err(|_| Error::NotConnected)?;
}
}
}
pub(crate) struct Namespace {
routes: RwLock<HashMap<Arc<str>, Route>>,
consumers: RwLock<HashMap<usize, HashSet<Arc<str>>>>,
}
impl Namespace {
pub(crate) fn new() -> Namespace {
Namespace {
routes: RwLock::new(HashMap::new()),
consumers: RwLock::new(HashMap::new()),
}
}
pub(crate) fn lookup(&self, path: &str) -> Option<Route> {
self.routes
.read()
.expect("namespace lock poisoned")
.get(path)
.map(Route::clone_sender)
}
pub(crate) fn register(&self, path: &str, route: Route) -> Result<(), Error> {
let mut routes = self.routes.write().expect("namespace lock poisoned");
match routes.entry(Arc::from(path)) {
Entry::Occupied(_) => Err(Error::AlreadyRegistered),
Entry::Vacant(slot) => {
slot.insert(route);
Ok(())
}
}
}
pub(crate) fn unregister(&self, path: &str) {
self.routes
.write()
.expect("namespace lock poisoned")
.remove(path);
}
pub(crate) fn note_consumer(&self, conn_id: usize, path: &str) {
self.consumers
.write()
.expect("namespace lock poisoned")
.entry(conn_id)
.or_default()
.insert(Arc::from(path));
}
pub(crate) fn forget_consumer(&self, conn_id: usize, path: &str) {
let mut consumers = self.consumers.write().expect("namespace lock poisoned");
if let Some(paths) = consumers.get_mut(&conn_id) {
paths.remove(path);
if paths.is_empty() {
consumers.remove(&conn_id);
}
}
}
pub(crate) fn take_consumer_routes(&self, conn_id: usize) -> Vec<(Arc<str>, Route)> {
let taken = self
.consumers
.write()
.expect("namespace lock poisoned")
.remove(&conn_id);
let Some(paths) = taken else {
return Vec::new();
};
let routes = self.routes.read().expect("namespace lock poisoned");
paths
.into_iter()
.filter_map(|path| {
let route = routes.get(&path)?.clone_sender();
Some((path, route))
})
.collect()
}
}
impl Route {
fn clone_sender(&self) -> Route {
match self {
Route::Request(tx) => Route::Request(tx.clone()),
Route::Transfer(tx) => Route::Transfer(tx.clone()),
Route::Raw(tx) => Route::Raw(tx.clone()),
Route::Pair { queue, owner } => Route::Pair {
queue: queue.clone(),
owner: Arc::clone(owner),
},
Route::Pub => Route::Pub,
}
}
}
pub(crate) struct ListenerInner {
pub(crate) runtime: Arc<RuntimeInner>,
pub(crate) namespace: Arc<Namespace>,
pub(crate) subs: Arc<SubRegistry>,
}
#[derive(Clone)]
pub struct Listener {
inner: Arc<ListenerInner>,
}
impl Listener {
pub(crate) fn new(runtime: Arc<RuntimeInner>) -> Listener {
let limits = runtime.config.limits;
let ordering = runtime.config.guarantees.ordering;
Listener {
inner: Arc::new(ListenerInner {
runtime,
namespace: Arc::new(Namespace::new()),
subs: Arc::new(SubRegistry::new(limits, ordering)),
}),
}
}
pub async fn bind_quic(
&self,
addr: SocketAddr,
tls: impl Into<ServerTls>,
) -> Result<Binding, Error> {
let tls = tls.into();
let limits = self.inner.runtime.config.limits;
let server_config = tls::server_config(&tls, &limits)?;
let exec = self.inner.runtime.exec.clone();
let endpoint = {
let _guard = exec.enter();
quinn::Endpoint::server(server_config, addr).map_err(Error::Io)?
};
let local_addr = endpoint.local_addr().map_err(Error::Io)?;
self.inner.runtime.track_endpoint(endpoint.clone());
exec.spawn(accept_connections(
endpoint.clone(),
Arc::clone(&self.inner),
));
tracing::info!(%local_addr, "quic binding listening");
Ok(Binding {
endpoint,
local_addr,
})
}
pub fn bind_inproc(&self, bus: &str) -> Result<LocalBinding, Error> {
let incoming = inproc::bind(bus)?;
let exec = self.inner.runtime.exec.clone();
exec.spawn(accept_local(
incoming,
Arc::clone(&self.inner.namespace),
Arc::clone(&self.inner.subs),
self.inner.runtime.config.limits,
exec.clone(),
self.inner.runtime.config.guarantees,
self.inner.runtime.shared(),
));
tracing::info!(bus, "inproc binding listening");
Ok(LocalBinding {
bus: bus.to_owned(),
})
}
#[cfg(unix)]
pub fn bind_unix(&self, path: impl AsRef<std::path::Path>) -> Result<UnixBinding, Error> {
let path = path.as_ref();
let exec = self.inner.runtime.exec.clone();
let (binding, listener) = {
let _guard = exec.enter();
weida_runtime::BoundUnixSocket::bind(path)?
};
exec.spawn(accept_unix(
listener,
self.local_accept(|link| Link::Unix(Box::new(link))),
));
tracing::info!(path = %path.display(), "unix binding listening");
Ok(UnixBinding { inner: binding })
}
#[cfg(windows)]
pub fn bind_pipe(&self, name: &str) -> Result<PipeBinding, Error> {
let addr = weida_core::PipeAddr::parse(&format!("{}://{name}/", weida_core::SCHEME_PIPE))?;
let exec = self.inner.runtime.exec.clone();
let (binding, first) = {
let _guard = exec.enter();
weida_runtime::BoundPipe::bind(addr.os_path())?
};
let name = addr.name.clone();
let (stop_tx, stop_rx) = tokio::sync::oneshot::channel();
exec.spawn(accept_pipe(
binding,
first,
stop_rx,
self.local_accept(|link| Link::Pipe(Box::new(link))),
));
tracing::info!(pipe = %name, "pipe binding listening");
Ok(PipeBinding {
name,
_stop: stop_tx,
})
}
#[cfg(any(unix, windows))]
fn local_accept<S: crate::grouped::Stream>(
&self,
link: fn(crate::grouped::Grouped<S>) -> Link,
) -> LocalAccept<S> {
LocalAccept {
groups: Arc::new(crate::grouped::Groups::default()),
namespace: Arc::clone(&self.inner.namespace),
subs: Arc::clone(&self.inner.subs),
limits: self.inner.runtime.config.limits,
exec: self.inner.runtime.exec.clone(),
guarantees: self.inner.runtime.config.guarantees,
shared: self.inner.runtime.shared(),
link,
}
}
pub fn replier(&self, path: &str) -> Result<Replier, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
self.inner.namespace.register(path, Route::Request(tx))?;
Ok(Endpoint::from_state(RepState::new(path, rx)))
}
pub fn puller(&self, path: &str) -> Result<Puller, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
self.inner.namespace.register(path, Route::Transfer(tx))?;
Ok(Endpoint::from_state(PullState::new(path, rx)))
}
pub fn publisher(&self, path: &str) -> Result<Publisher, Error> {
validate_endpoint_path(path)?;
self.inner.namespace.register(path, Route::Pub)?;
Ok(Endpoint::from_state(PubState::new(
path,
Arc::clone(&self.inner.subs),
self.inner.runtime.config.limits.subscriber_buffer_bytes,
)))
}
pub fn acceptor(&self, path: &str) -> Result<Acceptor, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
self.inner.namespace.register(path, Route::Raw(tx))?;
Ok(Acceptor::new(path, rx))
}
pub fn pair(&self, path: &str) -> Result<Paired, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
let owner = Arc::new(PairOwner::new());
self.inner.namespace.register(
path,
Route::Pair {
queue: tx,
owner: Arc::clone(&owner),
},
)?;
Ok(Endpoint::from_state(PairState::bound(path, owner, rx)))
}
pub fn respondent(&self, path: &str) -> Result<Respondent, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
self.inner.namespace.register(path, Route::Request(tx))?;
Ok(Endpoint::from_state(RespondState::new(path, rx)))
}
pub fn bus(&self, path: &str, tls: impl Into<ClientTls>) -> Result<BusMember, Error> {
validate_endpoint_path(path)?;
let (tx, rx) = mpsc::channel(self.inner.runtime.config.endpoint_queue);
self.inner.namespace.register(path, Route::Transfer(tx))?;
Ok(Endpoint::from_state(BusState::new(
path,
Arc::clone(&self.inner.runtime),
Arc::new(tls.into()),
rx,
self.inner.runtime.config.endpoint_queue,
self.inner.runtime.config.limits.subscriber_buffer_bytes,
)))
}
}
impl std::fmt::Debug for Listener {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Listener").finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct Binding {
endpoint: quinn::Endpoint,
local_addr: SocketAddr,
}
impl Binding {
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub async fn close(&self) {
self.endpoint.close(shutdown_code(), b"binding closed");
self.endpoint.wait_idle().await;
}
}
#[derive(Debug)]
pub struct LocalBinding {
bus: String,
}
impl LocalBinding {
pub fn bus(&self) -> &str {
&self.bus
}
}
impl Drop for LocalBinding {
fn drop(&mut self) {
inproc::unbind(&self.bus);
}
}
#[cfg(unix)]
#[derive(Debug)]
pub struct UnixBinding {
inner: weida_runtime::BoundUnixSocket,
}
#[cfg(unix)]
impl UnixBinding {
pub fn path(&self) -> &std::path::Path {
self.inner.path()
}
}
#[cfg(windows)]
#[derive(Debug)]
pub struct PipeBinding {
name: String,
_stop: tokio::sync::oneshot::Sender<()>,
}
#[cfg(windows)]
impl PipeBinding {
pub fn name(&self) -> &str {
&self.name
}
}
#[cfg(any(unix, windows))]
struct LocalAccept<S: crate::grouped::Stream> {
groups: Arc<crate::grouped::Groups<S>>,
namespace: Arc<Namespace>,
subs: Arc<SubRegistry>,
limits: Limits,
exec: Exec,
guarantees: GuaranteeSet,
shared: Arc<Shared>,
link: fn(crate::grouped::Grouped<S>) -> Link,
}
#[cfg(any(unix, windows))]
impl<S: crate::grouped::Stream> LocalAccept<S> {
fn clone_for(&self) -> LocalAccept<S> {
LocalAccept {
groups: Arc::clone(&self.groups),
namespace: Arc::clone(&self.namespace),
subs: Arc::clone(&self.subs),
limits: self.limits,
exec: self.exec.clone(),
guarantees: self.guarantees,
shared: Arc::clone(&self.shared),
link: self.link,
}
}
async fn serve(self, stream: S) {
use crate::grouped::{
Accepted, accept_control, admit_reverse, admit_transfer, read_accepted,
};
let accepted = match read_accepted(stream).await {
Ok(accepted) => accepted,
Err(e) => {
tracing::debug!(error = %e, "local connection preamble rejected");
return;
}
};
match accepted {
Accepted::Control(stream, principal) => {
let link = match accept_control(
stream,
principal,
Arc::clone(&self.groups),
self.limits.max_local_streams,
self.limits.max_parked_reverse,
)
.await
{
Ok(link) => link,
Err(e) => {
tracing::debug!(error = %e, "local control connection failed");
return;
}
};
let ctx = ConnCtx::spawn(
(self.link)(link),
self.limits,
self.namespace,
Some(Arc::clone(&self.subs)),
self.exec,
self.guarantees,
self.shared,
);
let conn_id = ctx.conn.stable_id();
let reason = ctx.conn.closed().await;
self.subs.remove_connection(conn_id);
crate::conn::drop_consumers(&ctx).await;
tracing::debug!(%reason, "local connection closed");
}
Accepted::Transfer(token, stream, principal) => {
if !admit_transfer(&self.groups, &token, &principal, stream) {
tracing::warn!(
"local transfer connection refused: unknown or mismatched group"
);
}
}
Accepted::Reverse(token, stream, principal) => {
if !admit_reverse(&self.groups, &token, &principal, stream) {
tracing::debug!(
"local reverse connection not parked: unknown group or pool full"
);
}
}
}
}
}
#[cfg(unix)]
async fn accept_unix(
listener: tokio::net::UnixListener,
serve: LocalAccept<crate::unix::UnixLocal>,
) {
loop {
let Ok((stream, _)) = listener.accept().await else {
tracing::debug!("unix binding closed; accept loop ending");
return;
};
if serve.shared.drain.is_draining() {
continue;
}
let stream = crate::unix::UnixLocal::accepted(stream, serve.exec.clone());
serve.exec.spawn(serve.clone_for().serve(stream));
}
}
#[cfg(windows)]
async fn accept_pipe(
binding: weida_runtime::BoundPipe,
first: tokio::net::windows::named_pipe::NamedPipeServer,
mut stop: tokio::sync::oneshot::Receiver<()>,
serve: LocalAccept<crate::pipe::PipeStream>,
) {
let mut listening = first;
loop {
let connected = tokio::select! {
connected = listening.connect() => connected,
_ = &mut stop => {
tracing::debug!("pipe binding dropped; accept loop ending");
return;
}
};
if let Err(e) = connected {
tracing::debug!(error = %e, "pipe accept failed; accept loop ending");
return;
}
let accepted = std::mem::replace(
&mut listening,
match binding.next_instance() {
Ok(next) => next,
Err(e) => {
tracing::warn!(error = %e, "pipe instance not created; accept loop ending");
return;
}
},
);
if serve.shared.drain.is_draining() {
continue;
}
let stream = crate::pipe::PipeStream::accepted(accepted, serve.exec.clone());
serve.exec.spawn(serve.clone_for().serve(stream));
}
}
async fn accept_local(
mut incoming: mpsc::UnboundedReceiver<inproc::LocalConn>,
namespace: Arc<Namespace>,
subs: Arc<SubRegistry>,
limits: Limits,
exec: Exec,
guarantees: GuaranteeSet,
shared: Arc<Shared>,
) {
while let Some(conn) = incoming.recv().await {
if shared.drain.is_draining() {
conn.close(codes::SHUTDOWN, "runtime draining");
continue;
}
let namespace = Arc::clone(&namespace);
let subs = Arc::clone(&subs);
let exec_for_conn = exec.clone();
let shared = Arc::clone(&shared);
exec.spawn(async move {
let ctx = ConnCtx::spawn(
Link::Local(conn),
limits,
namespace,
Some(Arc::clone(&subs)),
exec_for_conn,
guarantees,
shared,
);
let conn_id = ctx.conn.stable_id();
let reason = ctx.conn.closed().await;
subs.remove_connection(conn_id);
crate::conn::drop_consumers(&ctx).await;
tracing::debug!(%reason, "local connection closed");
});
}
}
fn shutdown_code() -> VarInt {
VarInt::from_u32(codes::SHUTDOWN as u32)
}
#[derive(Default)]
struct PeerCounts(std::sync::Mutex<HashMap<Fingerprint, usize>>);
impl PeerCounts {
fn admit(&self, peer: Option<Fingerprint>, max: usize) -> bool {
let Some(peer) = peer else {
return true;
};
let mut counts = self.0.lock().expect("peer count poisoned");
let count = counts.entry(peer).or_insert(0);
if *count >= max {
return false;
}
*count += 1;
true
}
fn release(&self, peer: Option<Fingerprint>) {
let Some(peer) = peer else {
return;
};
let mut counts = self.0.lock().expect("peer count poisoned");
if let Some(count) = counts.get_mut(&peer) {
*count -= 1;
if *count == 0 {
counts.remove(&peer);
}
}
}
}
async fn accept_connections(endpoint: quinn::Endpoint, listener: Arc<ListenerInner>) {
let config = &listener.runtime.config;
let limits = config.limits;
let max_connections = config.max_connections;
let max_connections_per_peer = config.max_connections_per_peer;
let guarantees = config.guarantees;
let exec = listener.runtime.exec.clone();
let shared = listener.runtime.shared();
let namespace = Arc::clone(&listener.namespace);
let subs = Arc::clone(&listener.subs);
let live = Arc::new(AtomicUsize::new(0));
let peers = Arc::new(PeerCounts::default());
while let Some(incoming) = endpoint.accept().await {
if shared.drain.is_draining() {
incoming.refuse();
continue;
}
if live.load(Ordering::Relaxed) >= max_connections {
tracing::warn!(max = max_connections, "connection limit reached; refusing");
exec.spawn(async move {
if let Ok(conn) = incoming.await {
conn.close(
VarInt::from_u32(codes::LIMIT_EXCEEDED as u32),
b"connection limit reached",
);
}
});
continue;
}
let namespace = Arc::clone(&namespace);
let subs = Arc::clone(&subs);
let live = Arc::clone(&live);
live.fetch_add(1, Ordering::Relaxed);
let exec_for_conn = exec.clone();
let shared = Arc::clone(&shared);
let peers = Arc::clone(&peers);
exec.spawn(async move {
match incoming.await {
Ok(conn) => {
let remote = conn.remote_address();
let conn_id = conn.stable_id();
let peer = crate::tls::peer_fingerprint(&conn);
if !peers.admit(peer, max_connections_per_peer) {
tracing::warn!(
%remote,
max = max_connections_per_peer,
"per-peer connection limit reached; refusing"
);
conn.close(
VarInt::from_u32(codes::LIMIT_EXCEEDED as u32),
b"per-peer connection limit reached",
);
} else {
tracing::debug!(%remote, "connection accepted");
let ctx = ConnCtx::spawn(
Link::Quic(conn.clone()),
limits,
namespace,
Some(Arc::clone(&subs)),
exec_for_conn,
guarantees,
shared,
);
let reason = conn.closed().await;
subs.remove_connection(conn_id);
crate::conn::drop_consumers(&ctx).await;
peers.release(peer);
tracing::debug!(%remote, %reason, "connection closed");
}
}
Err(e) => tracing::debug!(error = %e, "handshake failed"),
}
live.fetch_sub(1, Ordering::Relaxed);
});
}
tracing::debug!("binding stopped accepting");
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn namespace_rejects_duplicate_paths() {
let ns = Namespace::new();
let (tx, _rx) = mpsc::channel(1);
ns.register("/a", Route::Request(tx.clone())).unwrap();
assert!(matches!(
ns.register("/a", Route::Request(tx.clone())),
Err(Error::AlreadyRegistered)
));
assert!(matches!(
ns.register("/a", Route::Pub),
Err(Error::AlreadyRegistered)
));
ns.register("/b", Route::Request(tx)).unwrap();
}
#[tokio::test]
async fn namespace_lookup_is_exact() {
let ns = Namespace::new();
let (tx, _rx) = mpsc::channel(1);
ns.register("/jobs/a", Route::Request(tx)).unwrap();
assert!(ns.lookup("/jobs/a").is_some());
assert!(ns.lookup("/jobs").is_none());
assert!(ns.lookup("/jobs/a/").is_none());
assert!(ns.lookup("/jobs/*").is_none());
assert!(ns.lookup("/JOBS/A").is_none());
}
#[tokio::test]
async fn unregistering_releases_the_path() {
let ns = Namespace::new();
let (tx, _rx) = mpsc::channel(1);
ns.register("/s", Route::Transfer(tx)).unwrap();
ns.unregister("/s");
assert!(ns.lookup("/s").is_none());
ns.unregister("/s");
ns.register("/s", Route::Pub).unwrap();
}
}