use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock, Mutex as StdMutex};
use tokio::io::{AsyncReadExt, AsyncWriteExt, DuplexStream};
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore, mpsc};
use weida_core::{Error, LossCause, MAX_BUS_BYTES};
use weida_protocol::codes;
use weida_runtime::NameRegistry;
const NO_CODE: u64 = u64::MAX;
static BUSES: LazyLock<NameRegistry<LocalConn>> =
LazyLock::new(|| NameRegistry::new(MAX_BUS_BYTES));
pub fn validate_bus(bus: &str) -> Result<(), Error> {
BUSES.validate(bus)?;
if bus.as_bytes().contains(&b'/') {
return Err(Error::InvalidAddress(format!(
"invalid byte in inproc bus name: {bus:?}"
)));
}
Ok(())
}
pub(crate) fn bind(bus: &str) -> Result<mpsc::UnboundedReceiver<LocalConn>, Error> {
validate_bus(bus)?;
BUSES.bind(bus)
}
pub(crate) fn unbind(bus: &str) {
BUSES.unbind(bus);
}
pub(crate) async fn wait_bound(bus: &str) {
BUSES.wait_bound(bus).await;
}
pub(crate) fn dial(bus: &str, max_streams: usize, buffer: usize) -> Result<LocalConn, Error> {
validate_bus(bus)?;
let Some(tx) = BUSES.lookup(bus) else {
return Err(Error::ConnectionLost(LossCause::PeerClosed));
};
let (dialled, accepted) = LocalConn::pair(max_streams, buffer);
tx.send(accepted)
.map_err(|_| Error::ConnectionLost(LossCause::PeerClosed))?;
Ok(dialled)
}
pub(crate) struct LocalConn {
id: usize,
uni_to_peer: mpsc::UnboundedSender<LocalRecv>,
bi_to_peer: mpsc::UnboundedSender<(LocalSend, LocalRecv)>,
uni_from_peer: tokio::sync::Mutex<mpsc::UnboundedReceiver<LocalRecv>>,
bi_from_peer: tokio::sync::Mutex<mpsc::UnboundedReceiver<(LocalSend, LocalRecv)>>,
state: Arc<LinkState>,
slots: Arc<Semaphore>,
buffer: usize,
}
struct LinkState {
code: AtomicU64,
reason: StdMutex<String>,
closed_by: AtomicUsize,
closed: Notify,
}
impl LinkState {
fn close(&self, code: u64, reason: &str, by: usize) {
if self
.code
.compare_exchange(NO_CODE, code, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.closed_by.store(by, Ordering::Release);
*self.reason.lock().expect("close reason poisoned") = reason.to_owned();
self.closed.notify_waiters();
}
}
fn close_reason(&self, me: usize) -> Option<Error> {
let code = self.code.load(Ordering::Acquire);
(code != NO_CODE).then(|| {
let local = self.closed_by.load(Ordering::Acquire) == me;
closed_error(code, local)
})
}
}
fn closed_error(code: u64, local: bool) -> Error {
match code {
codes::NEGOTIATION_FAILED => {
Error::Negotiation("peer closed the connection: negotiation failed".into())
}
codes::LIMIT_EXCEEDED => Error::LimitExceeded,
codes::PROTOCOL_VIOLATION => Error::Protocol("peer reported a protocol violation".into()),
_ if local => Error::ConnectionLost(LossCause::LocallyClosed),
_ => Error::ConnectionLost(LossCause::PeerClosed),
}
}
static NEXT_ID: AtomicUsize = AtomicUsize::new(1);
impl LocalConn {
fn pair(max_streams: usize, buffer: usize) -> (LocalConn, LocalConn) {
let (a_uni_tx, a_uni_rx) = mpsc::unbounded_channel();
let (b_uni_tx, b_uni_rx) = mpsc::unbounded_channel();
let (a_bi_tx, a_bi_rx) = mpsc::unbounded_channel();
let (b_bi_tx, b_bi_rx) = mpsc::unbounded_channel();
let state = Arc::new(LinkState {
code: AtomicU64::new(NO_CODE),
reason: StdMutex::new(String::new()),
closed_by: AtomicUsize::new(usize::MAX),
closed: Notify::new(),
});
let slots = Arc::new(Semaphore::new(max_streams));
let id = NEXT_ID.fetch_add(2, Ordering::Relaxed);
(
LocalConn {
id,
uni_to_peer: b_uni_tx,
bi_to_peer: b_bi_tx,
uni_from_peer: tokio::sync::Mutex::new(a_uni_rx),
bi_from_peer: tokio::sync::Mutex::new(a_bi_rx),
state: Arc::clone(&state),
slots: Arc::clone(&slots),
buffer,
},
LocalConn {
id: id + 1,
uni_to_peer: a_uni_tx,
bi_to_peer: a_bi_tx,
uni_from_peer: tokio::sync::Mutex::new(b_uni_rx),
bi_from_peer: tokio::sync::Mutex::new(b_bi_rx),
state,
slots,
buffer,
},
)
}
pub(crate) fn stable_id(&self) -> usize {
self.id
}
pub(crate) fn close_reason(&self) -> Option<Error> {
self.state.close_reason(self.id)
}
pub(crate) fn close(&self, code: u64, reason: &str) {
self.state.close(code, reason, self.id);
}
pub(crate) async fn closed(&self) -> Error {
loop {
if let Some(reason) = self.state.close_reason(self.id) {
return reason;
}
self.state.closed.notified().await;
}
}
pub(crate) fn slots_exhausted(&self) -> bool {
self.slots.available_permits() == 0
}
async fn slot(&self) -> Result<Arc<StreamSlot>, Error> {
let waiting = Arc::clone(&self.slots).acquire_owned();
tokio::select! {
permit = waiting => match permit {
Ok(permit) => Ok(Arc::new(StreamSlot { _permit: permit })),
Err(_) => Err(Error::ConnectionLost(LossCause::PeerClosed)),
},
reason = self.closed() => Err(reason),
}
}
pub(crate) async fn open_uni(&self) -> Result<LocalSend, Error> {
if let Some(closed) = self.close_reason() {
return Err(closed);
}
let slot = self.slot().await?;
let (send, recv) = stream_pair(self.buffer, slot);
self.uni_to_peer
.send(recv)
.map_err(|_| Error::ConnectionLost(LossCause::PeerClosed))?;
Ok(send)
}
pub(crate) async fn open_bi(&self) -> Result<(LocalSend, LocalRecv), Error> {
if let Some(closed) = self.close_reason() {
return Err(closed);
}
let slot = self.slot().await?;
let (send, peer_recv) = stream_pair(self.buffer, Arc::clone(&slot));
let (peer_send, recv) = stream_pair(self.buffer, slot);
self.bi_to_peer
.send((peer_send, peer_recv))
.map_err(|_| Error::ConnectionLost(LossCause::PeerClosed))?;
Ok((send, recv))
}
pub(crate) async fn accept_uni(&self) -> Result<LocalRecv, Error> {
let mut queue = self.uni_from_peer.lock().await;
tokio::select! {
opened = queue.recv() => opened.ok_or(Error::ConnectionLost(LossCause::PeerClosed)),
reason = self.closed() => Err(reason),
}
}
pub(crate) async fn accept_bi(&self) -> Result<(LocalSend, LocalRecv), Error> {
let mut queue = self.bi_from_peer.lock().await;
tokio::select! {
opened = queue.recv() => opened.ok_or(Error::ConnectionLost(LossCause::PeerClosed)),
reason = self.closed() => Err(reason),
}
}
}
struct StreamSlot {
_permit: OwnedSemaphorePermit,
}
struct Signal {
stop: AtomicU64,
reset: AtomicU64,
finished: AtomicBool,
changed: Notify,
_slot: Arc<StreamSlot>,
}
fn stream_pair(buffer: usize, slot: Arc<StreamSlot>) -> (LocalSend, LocalRecv) {
let (writer, reader) = tokio::io::duplex(buffer);
let signal = Arc::new(Signal {
stop: AtomicU64::new(NO_CODE),
reset: AtomicU64::new(NO_CODE),
finished: AtomicBool::new(false),
changed: Notify::new(),
_slot: slot,
});
(
LocalSend {
io: Some(writer),
signal: Arc::clone(&signal),
},
LocalRecv {
io: Some(reader),
signal,
},
)
}
pub(crate) struct LocalSend {
io: Option<DuplexStream>,
signal: Arc<Signal>,
}
impl LocalSend {
fn stop_code(&self) -> Option<u64> {
let code = self.signal.stop.load(Ordering::Acquire);
(code != NO_CODE).then_some(code)
}
pub(crate) async fn write_all(&mut self, buf: &[u8]) -> Result<(), Error> {
if let Some(code) = self.stop_code() {
return Err(codes::stop_reason(code).into());
}
let Some(io) = self.io.as_mut() else {
return Err(Error::Transport("stream already closed".into()));
};
match io.write_all(buf).await {
Ok(()) => Ok(()),
Err(e) => Err(match self.stop_code() {
Some(code) => codes::stop_reason(code).into(),
None => Error::Transport(format!("local stream write failed: {e}")),
}),
}
}
pub(crate) fn finish(&mut self) -> Result<(), Error> {
if self.io.take().is_none() {
return Err(Error::Transport("stream already closed".into()));
}
self.signal.finished.store(true, Ordering::Release);
self.signal.changed.notify_waiters();
Ok(())
}
pub(crate) fn reset(&mut self, code: u64) {
self.signal.reset.store(code, Ordering::Release);
self.signal.changed.notify_waiters();
self.io = None;
}
pub(crate) fn stopped(
&self,
) -> impl Future<Output = Result<Option<u64>, Error>> + Send + Sync + use<> {
let signal = Arc::clone(&self.signal);
async move {
loop {
let waiting = signal.changed.notified();
let stop = signal.stop.load(Ordering::Acquire);
if stop != NO_CODE {
return Ok(Some(stop));
}
if signal.finished.load(Ordering::Acquire) {
return Ok(None);
}
waiting.await;
}
}
}
}
pub(crate) struct LocalRecv {
io: Option<DuplexStream>,
signal: Arc<Signal>,
}
impl LocalRecv {
pub(crate) async fn read(&mut self, buf: &mut [u8]) -> Result<Option<usize>, Error> {
let Some(io) = self.io.as_mut() else {
return Ok(None);
};
let read = io
.read(buf)
.await
.map_err(|e| Error::Transport(format!("local stream read failed: {e}")))?;
if read == 0 {
let reset = self.signal.reset.load(Ordering::Acquire);
if reset != NO_CODE {
return Err(codes::stop_reason(reset).into());
}
return Ok(None);
}
Ok(Some(read))
}
pub(crate) async fn read_exact(&mut self, buf: &mut [u8]) -> Result<(), Error> {
let mut filled = 0;
while filled < buf.len() {
match self.read(&mut buf[filled..]).await? {
Some(n) => filled += n,
None => return Err(Error::Protocol("stream ended mid-header".into())),
}
}
Ok(())
}
pub(crate) fn stop(&mut self, code: u64) {
self.signal.stop.store(code, Ordering::Release);
self.signal.changed.notify_waiters();
self.io = None;
}
pub(crate) fn io_mut(&mut self) -> Option<&mut DuplexStream> {
self.io.as_mut()
}
pub(crate) fn reset_code(&self) -> Option<u64> {
let code = self.signal.reset.load(Ordering::Acquire);
(code != NO_CODE).then_some(code)
}
}
impl LocalSend {
pub(crate) fn io_mut(&mut self) -> Option<&mut DuplexStream> {
self.io.as_mut()
}
}
impl Drop for LocalRecv {
fn drop(&mut self) {
if self.io.is_some()
&& !self.signal.finished.load(Ordering::Acquire)
&& self.signal.stop.load(Ordering::Acquire) == NO_CODE
{
self.signal.stop.store(codes::CANCELED, Ordering::Release);
self.signal.changed.notify_waiters();
}
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
#[test]
fn a_bus_name_is_bounded_and_has_no_separator() {
assert!(validate_bus("orders").is_ok());
assert!(validate_bus(&"x".repeat(MAX_BUS_BYTES)).is_ok());
assert!(validate_bus(&"x".repeat(MAX_BUS_BYTES + 1)).is_err());
assert!(validate_bus("").is_err());
assert!(validate_bus("has/slash").is_err());
assert!(validate_bus("has\ncontrol").is_err());
}
#[tokio::test]
async fn a_stream_carries_bytes_and_its_fin() {
let (a, _b) = LocalConn::pair(8, 64 * 1024);
let mut send = a.open_uni().await.expect("open");
send.write_all(b"payload").await.expect("write");
send.finish().expect("finish");
let mut recv = _b.accept_uni().await.expect("accept");
let mut buf = [0u8; 16];
let n = recv.read(&mut buf).await.expect("read").expect("bytes");
assert_eq!(&buf[..n], b"payload");
assert_eq!(recv.read(&mut buf).await.expect("read"), None, "the FIN");
}
#[tokio::test]
async fn the_stream_budget_makes_an_open_wait_for_a_slot() {
let (a, b) = LocalConn::pair(2, 1024);
let _one = a.open_uni().await.expect("first");
let _two = a.open_uni().await.expect("second");
let third = tokio::spawn(async move {
a.open_uni().await.expect("the third waits for a slot");
});
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(
!third.is_finished(),
"max_local_streams must bound live transfers"
);
drop(_one);
let accepted = b.accept_uni().await.expect("accept");
drop(accepted);
tokio::time::timeout(Duration::from_secs(5), third)
.await
.expect("a slot came free")
.expect("the waiting open completed");
}
#[tokio::test]
async fn a_refused_stream_reports_its_code_to_the_writer() {
let (a, b) = LocalConn::pair(8, 1024);
let send = a.open_uni().await.expect("open");
let mut recv = b.accept_uni().await.expect("accept");
recv.stop(codes::REJECTED);
let stopped = send.stopped().await.expect("stopped");
assert_eq!(stopped, Some(codes::REJECTED));
}
#[tokio::test]
async fn closing_either_side_is_seen_by_both() {
let (a, b) = LocalConn::pair(8, 1024);
assert!(a.close_reason().is_none());
b.close(codes::SHUTDOWN, "going away");
assert!(matches!(
a.close_reason(),
Some(Error::ConnectionLost(LossCause::PeerClosed))
));
assert!(
a.open_uni().await.is_err(),
"a closed connection opens nothing"
);
}
}