use std::collections::{HashMap, VecDeque};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex as StdMutex};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore, mpsc};
use weida_core::{Error, LossCause, PeerIdentity};
use weida_protocol::codes;
const KIND_CONTROL: u8 = 0x01;
const KIND_TRANSFER: u8 = 0x02;
const KIND_REVERSE: u8 = 0x03;
const TOKEN_LEN: usize = 16;
const NO_CODE: u64 = u64::MAX;
pub(crate) trait Stream: AsyncRead + AsyncWrite + Unpin + Send + Sized + 'static {
type Endpoint: Send + Sync + 'static;
type Principal: Clone + Send + Sync + 'static;
type Writer: AsyncWrite + Unpin + Send + 'static;
type Reader: AsyncRead + Unpin + Send + 'static;
fn connect(endpoint: &Self::Endpoint) -> impl Future<Output = Result<Self, Error>> + Send;
fn principal(&self) -> Result<Self::Principal, Error>;
fn split(self) -> (Self::Reader, Self::Writer);
fn same_peer(group: &Self::Principal, asking: &Self::Principal) -> bool;
fn identity(principal: &Self::Principal) -> PeerIdentity;
fn finish(writer: Self::Writer);
fn reset(writer: Self::Writer, code: u64);
fn stop(reader: Self::Reader, code: u64);
fn read_error(error: std::io::Error) -> Error;
}
pub(crate) type Halves<S> = (LocalSend<S>, LocalRecv<S>);
pub(crate) struct Groups<S: Stream> {
entries: StdMutex<HashMap<[u8; TOKEN_LEN], Group<S>>>,
}
impl<S: Stream> Default for Groups<S> {
fn default() -> Groups<S> {
Groups {
entries: StdMutex::new(HashMap::new()),
}
}
}
struct Group<S: Stream> {
principal: S::Principal,
transfers: mpsc::UnboundedSender<Halves<S>>,
reverse: Arc<ReversePool<S>>,
}
pub(crate) struct ReversePool<S: Stream> {
parked: StdMutex<VecDeque<LocalSend<S>>>,
max: usize,
slots: Arc<Semaphore>,
}
impl<S: Stream> ReversePool<S> {
fn new(slots: Arc<Semaphore>, max: usize) -> ReversePool<S> {
ReversePool {
parked: StdMutex::new(VecDeque::new()),
max,
slots,
}
}
fn park(&self, send: S::Writer) -> bool {
let mut parked = self.parked.lock().expect("reverse pool poisoned");
if parked.len() >= self.max {
return false;
}
let Some(slot) = StreamSlot::try_acquire(&self.slots) else {
return false;
};
parked.push_back(LocalSend::new(send, Some(Arc::new(slot))));
true
}
fn take(&self) -> Option<LocalSend<S>> {
self.parked
.lock()
.expect("reverse pool poisoned")
.pop_front()
}
}
impl<S: Stream> Groups<S> {
fn insert(
&self,
token: [u8; TOKEN_LEN],
principal: S::Principal,
transfers: mpsc::UnboundedSender<Halves<S>>,
reverse: Arc<ReversePool<S>>,
) {
self.entries
.lock()
.expect("group registry poisoned")
.insert(
token,
Group {
principal,
transfers,
reverse,
},
);
}
fn remove(&self, token: &[u8; TOKEN_LEN]) {
self.entries
.lock()
.expect("group registry poisoned")
.remove(token);
}
fn with_group<T>(
&self,
token: &[u8; TOKEN_LEN],
principal: &S::Principal,
f: impl FnOnce(&Group<S>) -> T,
) -> Option<T> {
let entries = self.entries.lock().expect("group registry poisoned");
let group = entries.get(token)?;
if !S::same_peer(&group.principal, principal) {
return None;
}
Some(f(group))
}
fn admit(&self, token: &[u8; TOKEN_LEN], principal: &S::Principal, halves: Halves<S>) -> bool {
self.with_group(token, principal, |group| {
group.transfers.send(halves).is_ok()
})
.unwrap_or(false)
}
fn park(&self, token: &[u8; TOKEN_LEN], principal: &S::Principal, send: S::Writer) -> bool {
self.with_group(token, principal, |group| group.reverse.park(send))
.unwrap_or(false)
}
}
pub(crate) enum Accepted<S: Stream> {
Control(S, S::Principal),
Transfer([u8; TOKEN_LEN], S, S::Principal),
Reverse([u8; TOKEN_LEN], S, S::Principal),
}
pub(crate) async fn read_accepted<S: Stream>(mut stream: S) -> Result<Accepted<S>, Error> {
let mut kind = [0u8; 1];
stream.read_exact(&mut kind).await.map_err(Error::Io)?;
let principal = stream.principal()?;
match kind[0] {
KIND_CONTROL => Ok(Accepted::Control(stream, principal)),
KIND_TRANSFER | KIND_REVERSE => {
let mut token = [0u8; TOKEN_LEN];
stream.read_exact(&mut token).await.map_err(Error::Io)?;
Ok(if kind[0] == KIND_TRANSFER {
Accepted::Transfer(token, stream, principal)
} else {
Accepted::Reverse(token, stream, principal)
})
}
other => Err(Error::Protocol(format!(
"unknown local connection kind {other:#04x}"
))),
}
}
pub(crate) async fn accept_control<S: Stream>(
mut stream: S,
principal: S::Principal,
groups: Arc<Groups<S>>,
max_streams: usize,
max_parked: usize,
) -> Result<Grouped<S>, Error> {
let token = random_token();
stream.write_all(&token).await.map_err(Error::Io)?;
let (transfers_tx, transfers_rx) = mpsc::unbounded_channel();
let mut link = Grouped::new(
Side::Accept {
groups: Arc::clone(&groups),
token,
},
stream,
Some(principal.clone()),
max_streams,
max_parked,
);
let reverse = Arc::new(ReversePool::new(Arc::clone(&link.slots), max_parked));
link.reverse = Some(Arc::clone(&reverse));
link.transfers = tokio::sync::Mutex::new(Some(transfers_rx));
groups.insert(token, principal, transfers_tx, reverse);
Ok(link)
}
pub(crate) fn admit_reverse<S: Stream>(
groups: &Groups<S>,
token: &[u8; TOKEN_LEN],
principal: &S::Principal,
stream: S,
) -> bool {
let (_recv, send) = stream.split();
groups.park(token, principal, send)
}
pub(crate) fn admit_transfer<S: Stream>(
groups: &Groups<S>,
token: &[u8; TOKEN_LEN],
principal: &S::Principal,
stream: S,
) -> bool {
let (recv, send) = stream.split();
groups.admit(
token,
principal,
(LocalSend::new(send, None), LocalRecv::new(recv, None)),
)
}
pub(crate) async fn dial<S: Stream>(
endpoint: S::Endpoint,
max_streams: usize,
max_parked: usize,
) -> Result<Grouped<S>, Error> {
let mut stream = S::connect(&endpoint).await?;
stream.write_all(&[KIND_CONTROL]).await.map_err(Error::Io)?;
let mut token = [0u8; TOKEN_LEN];
stream.read_exact(&mut token).await.map_err(Error::Io)?;
let principal = stream.principal()?;
Ok(Grouped::new(
Side::Dial { endpoint, token },
stream,
Some(principal),
max_streams,
max_parked,
))
}
fn random_token() -> [u8; TOKEN_LEN] {
use rand::RngCore;
let mut token = [0u8; TOKEN_LEN];
rand::rng().fill_bytes(&mut token);
token
}
enum Side<S: Stream> {
Dial {
endpoint: S::Endpoint,
token: [u8; TOKEN_LEN],
},
Accept {
groups: Arc<Groups<S>>,
token: [u8; TOKEN_LEN],
},
}
pub(crate) struct Grouped<S: Stream> {
side: Side<S>,
control_send: StdMutex<Option<S::Writer>>,
control_recv: StdMutex<Option<S::Reader>>,
transfers: tokio::sync::Mutex<Option<mpsc::UnboundedReceiver<Halves<S>>>>,
reverse: Option<Arc<ReversePool<S>>>,
parked_tx: mpsc::UnboundedSender<LocalRecv<S>>,
parked_rx: tokio::sync::Mutex<mpsc::UnboundedReceiver<LocalRecv<S>>>,
deficit: Arc<Deficit>,
peer: Option<S::Principal>,
slots: Arc<Semaphore>,
max_parked: usize,
closed: AtomicU64,
closed_notify: Notify,
id: usize,
}
#[derive(Default)]
pub(crate) struct Deficit {
count: AtomicUsize,
notify: Notify,
}
impl Deficit {
fn record(&self) {
self.count.fetch_add(1, Ordering::Relaxed);
self.notify.notify_one();
}
async fn take(&self) -> usize {
loop {
let owed = self.count.swap(0, Ordering::Relaxed);
if owed > 0 {
return owed;
}
self.notify.notified().await;
}
}
}
static NEXT_ID: AtomicUsize = AtomicUsize::new(1);
impl<S: Stream> Grouped<S> {
fn new(
side: Side<S>,
control: S,
peer: Option<S::Principal>,
max_streams: usize,
max_parked: usize,
) -> Grouped<S> {
let (recv, send) = control.split();
let (parked_tx, parked_rx) = mpsc::unbounded_channel();
Grouped {
side,
control_send: StdMutex::new(Some(send)),
control_recv: StdMutex::new(Some(recv)),
transfers: tokio::sync::Mutex::new(None),
reverse: None,
parked_tx,
parked_rx: tokio::sync::Mutex::new(parked_rx),
deficit: Arc::new(Deficit::default()),
peer,
slots: Arc::new(Semaphore::new(max_streams)),
max_parked,
closed: AtomicU64::new(NO_CODE),
closed_notify: Notify::new(),
id: NEXT_ID.fetch_add(1, Ordering::Relaxed),
}
}
pub(crate) fn stable_id(&self) -> usize {
self.id
}
pub(crate) fn peer(&self) -> Option<PeerIdentity> {
self.peer.as_ref().map(S::identity)
}
pub(crate) fn close_reason(&self) -> Option<Error> {
let code = self.closed.load(Ordering::Acquire);
(code != NO_CODE).then(|| match code {
codes::NEGOTIATION_FAILED => {
Error::Negotiation("peer closed the connection: negotiation failed".into())
}
_ => Error::ConnectionLost(LossCause::PeerClosed),
})
}
pub(crate) fn close(&self, code: u64, _reason: &str) {
if self
.closed
.compare_exchange(NO_CODE, code, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.control_send.lock().expect("poisoned").take();
self.control_recv.lock().expect("poisoned").take();
if let Side::Accept { groups, token, .. } = &self.side {
groups.remove(token);
}
self.closed_notify.notify_waiters();
}
}
pub(crate) async fn closed(&self) -> Error {
loop {
if let Some(reason) = self.close_reason() {
return reason;
}
self.closed_notify.notified().await;
}
}
pub(crate) fn slots_exhausted(&self) -> bool {
self.slots.available_permits() == 0
}
async fn slot(&self) -> Result<StreamSlot, Error> {
let waiting = Arc::clone(&self.slots).acquire_owned();
tokio::select! {
permit = waiting => match permit {
Ok(permit) => Ok(StreamSlot { _permit: permit }),
Err(_) => Err(Error::ConnectionLost(LossCause::LocallyClosed)),
},
reason = self.closed() => Err(reason),
}
}
pub(crate) async fn open_uni(&self) -> Result<LocalSend<S>, Error> {
if let Side::Accept { .. } = &self.side {
if let Some(closed) = self.close_reason() {
return Err(closed);
}
return self
.reverse
.as_ref()
.and_then(|pool| pool.take())
.ok_or(Error::NoParkedConnection);
}
let (send, _recv) = self.open_transfer().await?;
Ok(send)
}
pub(crate) fn open_control(&self) -> Result<LocalSend<S>, Error> {
match self.take_control_send() {
Some(send) => Ok(LocalSend::new(send, None)),
None => Err(Error::Transport("control stream already used".into())),
}
}
pub(crate) async fn open_bi(&self) -> Result<Halves<S>, Error> {
self.open_transfer().await
}
fn take_control_send(&self) -> Option<S::Writer> {
self.control_send.lock().expect("poisoned").take()
}
async fn open_transfer(&self) -> Result<Halves<S>, Error> {
if let Some(closed) = self.close_reason() {
return Err(closed);
}
let (endpoint, token) = match &self.side {
Side::Dial { endpoint, token } => (endpoint, token),
Side::Accept { .. } => return Err(Error::Unsupported),
};
let slot = self.slot().await?;
let mut stream = match S::connect(endpoint).await {
Ok(stream) => stream,
Err(e @ Error::ConnectionLost(_)) => {
self.close(codes::SHUTDOWN, "nobody answers on the address");
return Err(e);
}
Err(e) => return Err(e),
};
let mut preamble = [0u8; 1 + TOKEN_LEN];
preamble[0] = KIND_TRANSFER;
preamble[1..].copy_from_slice(token);
stream.write_all(&preamble).await.map_err(Error::Io)?;
let (recv, send) = stream.split();
let slot = Arc::new(slot);
Ok((
LocalSend::new(send, Some(Arc::clone(&slot))),
LocalRecv::new(recv, Some(slot)),
))
}
pub(crate) async fn accept_uni(&self) -> Result<LocalRecv<S>, Error> {
if let Some(recv) = self.control_recv.lock().expect("poisoned").take() {
return Ok(LocalRecv::new(recv, None));
}
let mut parked = self.parked_rx.lock().await;
tokio::select! {
arrived = parked.recv() => arrived.ok_or(Error::ConnectionLost(LossCause::PeerClosed)),
reason = self.closed() => Err(reason),
}
}
pub(crate) async fn park_reverse(&self) -> Result<usize, Error> {
let mut parked = 0;
for _ in 0..self.max_parked {
match self.park_one().await {
Ok(()) => parked += 1,
Err(e) if parked == 0 => return Err(e),
Err(e) => {
tracing::debug!(error = %e, parked, "reverse pool filled short");
break;
}
}
}
Ok(parked)
}
pub(crate) async fn maintain_reverse(&self) {
loop {
let owed = tokio::select! {
owed = self.deficit.take() => owed,
_ = self.closed() => return,
};
for _ in 0..owed {
if let Err(e) = self.park_one().await {
tracing::debug!(error = %e, "reverse pool not replenished");
return;
}
}
}
}
async fn park_one(&self) -> Result<(), Error> {
if let Some(closed) = self.close_reason() {
return Err(closed);
}
let Side::Dial { endpoint, token } = &self.side else {
return Err(Error::Unsupported);
};
let slot = Arc::new(StreamSlot::try_acquire(&self.slots).ok_or(Error::LimitExceeded)?);
let mut stream = S::connect(endpoint).await?;
let mut preamble = [0u8; 1 + TOKEN_LEN];
preamble[0] = KIND_REVERSE;
preamble[1..].copy_from_slice(token);
stream.write_all(&preamble).await.map_err(Error::Io)?;
let (recv, _send) = stream.split();
let mut recv = LocalRecv::new(recv, Some(slot));
recv.deficit = Some(Arc::clone(&self.deficit));
self.parked_tx
.send(recv)
.map_err(|_| Error::ConnectionLost(LossCause::LocallyClosed))
}
pub(crate) async fn accept_bi(&self) -> Result<Halves<S>, Error> {
let mut queue = self.transfers.lock().await;
let Some(queue) = queue.as_mut() else {
return Err(self.closed().await);
};
tokio::select! {
accepted = queue.recv() => accepted.ok_or(Error::ConnectionLost(LossCause::PeerClosed)),
reason = self.closed() => Err(reason),
}
}
}
struct StreamSlot {
_permit: OwnedSemaphorePermit,
}
impl StreamSlot {
fn try_acquire(slots: &Arc<Semaphore>) -> Option<StreamSlot> {
Arc::clone(slots)
.try_acquire_owned()
.ok()
.map(|permit| StreamSlot { _permit: permit })
}
}
pub(crate) struct LocalSend<S: Stream> {
io: Option<S::Writer>,
_slot: Option<Arc<StreamSlot>>,
}
impl<S: Stream> LocalSend<S> {
fn new(io: S::Writer, slot: Option<Arc<StreamSlot>>) -> LocalSend<S> {
LocalSend {
io: Some(io),
_slot: slot,
}
}
pub(crate) fn io_mut(&mut self) -> Option<&mut S::Writer> {
self.io.as_mut()
}
pub(crate) async fn write_all(&mut self, buf: &[u8]) -> Result<(), Error> {
let Some(io) = self.io.as_mut() else {
return Err(Error::Transport("stream already closed".into()));
};
io.write_all(buf).await.map_err(|e| match e.kind() {
std::io::ErrorKind::BrokenPipe | std::io::ErrorKind::ConnectionReset => Error::Canceled,
_ => Error::Transport(format!("local stream write failed: {e}")),
})
}
pub(crate) fn finish(&mut self) -> Result<(), Error> {
match self.io.take() {
Some(io) => {
S::finish(io);
Ok(())
}
None => Err(Error::Transport("stream already closed".into())),
}
}
pub(crate) fn reset(&mut self, code: u64) {
if let Some(io) = self.io.take() {
S::reset(io, code);
}
}
pub(crate) fn stopped(
&self,
) -> impl Future<Output = Result<Option<u64>, Error>> + Send + Sync + use<S> {
async move { Ok(None) }
}
}
impl<S: Stream> Drop for LocalSend<S> {
fn drop(&mut self) {
if let Some(io) = self.io.take() {
S::reset(io, codes::CANCELED);
}
}
}
pub(crate) struct LocalRecv<S: Stream> {
io: Option<S::Reader>,
_slot: Option<Arc<StreamSlot>>,
deficit: Option<Arc<Deficit>>,
spent: bool,
}
impl<S: Stream> LocalRecv<S> {
fn new(io: S::Reader, slot: Option<Arc<StreamSlot>>) -> LocalRecv<S> {
LocalRecv {
io: Some(io),
_slot: slot,
deficit: None,
spent: false,
}
}
pub(crate) fn io_mut(&mut self) -> Option<&mut S::Reader> {
self.io.as_mut()
}
fn spend(&mut self) {
if self.spent {
return;
}
self.spent = true;
if let Some(deficit) = &self.deficit {
deficit.record();
}
}
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(S::read_error)?;
if read > 0 {
self.spend();
}
Ok((read > 0).then_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) {
if let Some(io) = self.io.take() {
S::stop(io, code);
}
}
}
impl<S: Stream> Drop for LocalRecv<S> {
fn drop(&mut self) {
self.spend();
if let Some(io) = self.io.take() {
S::stop(io, codes::CANCELED);
}
}
}