use std::collections::HashMap;
use std::collections::VecDeque;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::MutexGuard;
use bytes::Bytes;
use js_sys::Reflect;
use js_sys::Uint8Array;
use rings_core::dht::Did;
use wasm_bindgen::JsCast;
use wasm_bindgen::JsValue;
use wasm_bindgen_futures::JsFuture;
use web_sys::ReadableStream;
use web_sys::ReadableStreamDefaultReader;
use web_sys::WebTransport;
use web_sys::WritableStream;
use web_sys::WritableStreamDefaultWriter;
use crate::error::Error;
use crate::error::Result;
use crate::extension::ext::Scope;
use crate::extension::protocols::relay::RelayCommand;
use crate::extension::transport::allocate_non_reusing;
use crate::extension::transport::platform::spawn_detached;
use crate::extension::transport::EffectEnqueue;
use crate::extension::transport::Frame;
use crate::extension::transport::Initiator;
use crate::extension::transport::OutboundDrainState;
use crate::extension::transport::OutboundQueueBudget;
use crate::extension::transport::SessionKey;
use crate::extension::transport::TransportKind;
enum Outbound {
Data(Bytes),
Shutdown,
}
enum OutboundDrainStep {
Operation(WritableStreamDefaultWriter, Outbound),
Complete,
Superseded,
InvariantViolation,
}
impl Outbound {
fn data_bytes(&self) -> usize {
match self {
Self::Data(bytes) => bytes.len(),
Self::Shutdown => 0,
}
}
}
enum SessionHandle {
Opening {
queue: VecDeque<Outbound>,
budget: OutboundQueueBudget,
generation: u64,
},
Ready {
writer: WritableStreamDefaultWriter,
transport: WebTransport,
queue: VecDeque<Outbound>,
budget: OutboundQueueBudget,
drain: OutboundDrainState,
generation: u64,
},
}
impl SessionHandle {
fn generation(&self) -> u64 {
match self {
SessionHandle::Opening { generation, .. } => *generation,
SessionHandle::Ready { generation, .. } => *generation,
}
}
}
#[derive(Default)]
pub(crate) struct WtSessions {
map: Mutex<HashMap<SessionKey, SessionHandle>>,
generations: AtomicU64,
}
impl WtSessions {
pub fn new() -> Self {
Self::default()
}
pub fn connect(
self: Arc<Self>,
scope: Scope,
key: SessionKey,
url: String,
kind: TransportKind,
) -> EffectEnqueue {
debug_assert_eq!(
scope.namespace(),
key.namespace.as_str(),
"relay engine acted with a scope outside the session's namespace"
);
let Some(generation) = self.open_slot(key.clone()) else {
return EffectEnqueue::Failed;
};
spawn_detached(async move {
self.finish_connect(scope, key, url, kind, generation).await;
});
EffectEnqueue::Enqueued
}
async fn finish_connect(
self: Arc<Self>,
scope: Scope,
key: SessionKey,
url: String,
kind: TransportKind,
generation: u64,
) {
match open(url.as_str(), kind).await {
Ok((transport, readable, writer)) => {
if let Some(start_drain) = self.promote(&key, generation, writer, transport.clone())
{
if start_drain {
self.spawn_writer_loop(scope.clone(), key.clone(), generation);
}
self.spawn_read_loop(scope, key, readable, generation);
} else {
transport.close();
}
}
Err(e) => {
tracing::error!("WebTransport connect to {url} failed: {e:?}");
self.close_current_and_notify(&scope, &key, generation)
.await;
}
}
}
pub fn write(self: &Arc<Self>, scope: Scope, key: SessionKey, bytes: Bytes) -> EffectEnqueue {
self.enqueue(scope, key, Outbound::Data(bytes))
}
pub fn shutdown(self: &Arc<Self>, scope: Scope, key: SessionKey) -> EffectEnqueue {
self.enqueue(scope, key, Outbound::Shutdown)
}
pub fn close_for_effect(&self, key: &SessionKey) {
let removed = self.lock_sessions().remove(key);
self.finish_close_without_feedback(removed);
}
async fn close_if_current(&self, scope: &Scope, key: &SessionKey, generation: u64) -> bool {
let removed = {
let mut map = self.lock_sessions();
let current = map.get(key).map(|handle| handle.generation());
(current == Some(generation))
.then(|| map.remove(key))
.flatten()
};
self.finish_close(scope, key, removed).await
}
async fn close_current_and_notify(&self, scope: &Scope, key: &SessionKey, generation: u64) {
if self.close_if_current(scope, key, generation).await {
let _ = send_frame(scope, key.peer, Frame::Close {
session: key.session,
from_opener: matches!(key.initiator, Initiator::Local),
})
.await;
}
}
async fn finish_close(
&self,
scope: &Scope,
key: &SessionKey,
removed: Option<SessionHandle>,
) -> bool {
let Some(handle) = removed else {
return false;
};
if let SessionHandle::Ready { transport, .. } = handle {
transport.close();
}
inject_untrack(scope, key).await;
true
}
fn finish_close_without_feedback(&self, removed: Option<SessionHandle>) -> bool {
let Some(handle) = removed else {
return false;
};
if let SessionHandle::Ready { transport, .. } = handle {
transport.close();
}
true
}
fn open_slot(&self, key: SessionKey) -> Option<u64> {
let generation = allocate_non_reusing(&self.generations)?;
self.insert(key, SessionHandle::Opening {
queue: VecDeque::new(),
budget: OutboundQueueBudget::default(),
generation,
});
Some(generation)
}
fn promote(
&self,
key: &SessionKey,
generation: u64,
writer: WritableStreamDefaultWriter,
transport: WebTransport,
) -> Option<bool> {
let mut map = self.lock_sessions();
let (queue, budget) = match map.get(key) {
Some(SessionHandle::Opening {
generation: current,
..
}) if *current == generation => match map.remove(key) {
Some(SessionHandle::Opening { queue, budget, .. }) => (queue, budget),
_ => return None,
},
_ => return None,
};
let mut drain = OutboundDrainState::Idle;
let start_drain = !queue.is_empty() && drain.claim();
map.insert(key.clone(), SessionHandle::Ready {
writer,
transport,
queue,
budget,
drain,
generation,
});
Some(start_drain)
}
fn enqueue(self: &Arc<Self>, scope: Scope, key: SessionKey, op: Outbound) -> EffectEnqueue {
let mut map = self.lock_sessions();
let Some(handle) = map.get_mut(&key) else {
return EffectEnqueue::Missing;
};
let data_bytes = op.data_bytes();
let (generation, start_drain, admitted) = match handle {
SessionHandle::Opening {
queue,
budget,
generation,
} => {
let admitted = budget.try_reserve(data_bytes);
if admitted {
queue.push_back(op);
}
(*generation, false, admitted)
}
SessionHandle::Ready {
queue,
budget,
drain,
generation,
..
} => {
let admitted = budget.try_reserve(data_bytes);
if admitted {
queue.push_back(op);
}
let start_drain = admitted && drain.claim();
(*generation, start_drain, admitted)
}
};
if admitted {
drop(map);
if start_drain {
self.spawn_writer_loop(scope, key, generation);
}
return EffectEnqueue::Enqueued;
}
let removed = (map.get(&key).map(SessionHandle::generation) == Some(generation))
.then(|| map.remove(&key))
.flatten();
drop(map);
self.finish_close_without_feedback(removed);
EffectEnqueue::Failed
}
fn spawn_writer_loop(self: &Arc<Self>, scope: Scope, key: SessionKey, generation: u64) {
let sessions = self.clone();
spawn_detached(async move {
loop {
let (writer, op) = match sessions.take_ready_outbound(&key, generation) {
OutboundDrainStep::Operation(writer, op) => (writer, op),
OutboundDrainStep::Complete | OutboundDrainStep::Superseded => break,
OutboundDrainStep::InvariantViolation => {
tracing::error!("WebTransport outbound queue budget diverged for {key:?}");
sessions
.close_current_and_notify(&scope, &key, generation)
.await;
break;
}
};
let promise = match op {
Outbound::Data(bytes) => {
let chunk = Uint8Array::from(bytes.as_ref());
writer.write_with_chunk(chunk.as_ref())
}
Outbound::Shutdown => writer.close(),
};
if JsFuture::from(promise).await.is_err() {
sessions
.close_current_and_notify(&scope, &key, generation)
.await;
break;
}
}
});
}
fn take_ready_outbound(&self, key: &SessionKey, generation: u64) -> OutboundDrainStep {
let mut map = self.lock_sessions();
let Some(SessionHandle::Ready {
writer,
queue,
budget,
drain,
generation: current,
..
}) = map.get_mut(key)
else {
return OutboundDrainStep::Superseded;
};
if *current != generation {
return OutboundDrainStep::Superseded;
}
let Some(op) = queue.pop_front() else {
drain.release();
return OutboundDrainStep::Complete;
};
if !budget.release(op.data_bytes()) {
return OutboundDrainStep::InvariantViolation;
}
OutboundDrainStep::Operation(writer.clone(), op)
}
fn insert(&self, key: SessionKey, handle: SessionHandle) {
let mut map = self.lock_sessions();
if let Some(SessionHandle::Ready { transport, .. }) = map.insert(key, handle) {
transport.close();
}
}
fn lock_sessions(&self) -> MutexGuard<'_, HashMap<SessionKey, SessionHandle>> {
self.map.lock().unwrap_or_else(|poisoned| {
tracing::error!("recovering poisoned WebTransport session table");
poisoned.into_inner()
})
}
fn spawn_read_loop(
self: &Arc<Self>,
scope: Scope,
key: SessionKey,
readable: ReadableStream,
generation: u64,
) {
let sessions = self.clone();
spawn_detached(async move {
let peer = key.peer;
let session = key.session;
let from_opener = matches!(key.initiator, Initiator::Local);
let reader: ReadableStreamDefaultReader = match readable.get_reader().dyn_into() {
Ok(reader) => reader,
Err(_) => return,
};
loop {
let result = match JsFuture::from(reader.read()).await {
Ok(result) => result,
Err(_) => break,
};
let done = Reflect::get(&result, &JsValue::from_str("done"))
.ok()
.and_then(|v| v.as_bool())
.unwrap_or(true);
if done {
break;
}
let value = match Reflect::get(&result, &JsValue::from_str("value")) {
Ok(value) => value,
Err(_) => break,
};
let bytes = Bytes::from(Uint8Array::new(&value).to_vec());
if send_frame(&scope, peer, Frame::Data {
session,
from_opener,
bytes,
})
.await
.is_err()
{
break;
}
}
sessions
.close_current_and_notify(&scope, &key, generation)
.await;
});
}
}
async fn open(
url: &str,
kind: TransportKind,
) -> std::result::Result<(WebTransport, ReadableStream, WritableStreamDefaultWriter), JsValue> {
let transport = WebTransport::new(url)?;
JsFuture::from(transport.ready()).await?;
let (readable, writable): (ReadableStream, WritableStream) = match kind {
TransportKind::Tcp => {
let bidi = JsFuture::from(transport.create_bidirectional_stream()).await?;
let bidi: web_sys::WebTransportBidirectionalStream = bidi.unchecked_into();
(
bidi.readable().unchecked_into(),
bidi.writable().unchecked_into(),
)
}
TransportKind::Udp => {
let datagrams = transport.datagrams();
(datagrams.readable(), datagrams.writable())
}
};
let writer = writable.get_writer()?;
Ok((transport, readable, writer))
}
async fn send_frame(scope: &Scope, peer: Did, frame: Frame) -> Result<()> {
let payload = rings_codec::serialize(&frame).map_err(|_| Error::EncodeError)?;
scope.send(peer, Bytes::from(payload)).await
}
async fn inject_untrack(scope: &Scope, key: &SessionKey) {
let command = RelayCommand::<String>::Untrack {
peer: key.peer,
session: key.session,
initiator: key.initiator,
};
if let Ok(bytes) = rings_codec::serialize(&command) {
if let Err(e) = scope.inject(Bytes::from(bytes)).await {
tracing::warn!(
"relay Untrack inject failed for {key:?}: {e:?}; pure state may still list \
this (now dropped) session"
);
}
}
}