use std::collections::VecDeque;
use std::future::Future;
use std::sync::{Arc, Mutex, MutexGuard, Weak};
use std::time::Duration;
use tokio::sync::{watch, Notify};
use crate::cbor::{self, Value};
use crate::frame::{
self, RequestSpec, StreamEncoding, StreamFields, StreamMode, StreamRole, StreamState,
VerifiedRequest,
};
use super::admission::{Admission, SessionPlace, Verdict};
use super::framing::{read_frame, FrameWriter, MAX_FRAME_BYTES};
use super::serve::{bounded_detail, BoxFuture, StreamOffer, CODE_REQUEST_COPY};
use super::{frame_type_of, now_ms, Inner, Link, LinkError};
const STREAM_OPEN_BYTES: usize = 1024 * 1024;
const STREAM_OPEN_WAIT: Duration = Duration::from_secs(10);
const STREAM_INBOX: usize = 16 * 1024 * 1024;
pub const DEFAULT_STREAM_DEADLINE: Duration = Duration::from_secs(30);
const CODE_STREAM_NOT_FOUND: &str = "not_found";
const CODE_MODE_MISMATCH: &str = "mode_mismatch";
const CODE_TOO_MANY_SESSIONS: &str = "too_many_sessions";
const CODE_STREAM_HANDLER_ERROR: &str = "error";
pub type StreamHandler = Arc<dyn Fn(Stream) -> BoxFuture<Result<(), String>> + Send + Sync>;
pub fn stream_handler<F, Fut>(f: F) -> StreamHandler
where
F: Fn(Stream) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), String>> + Send + 'static,
{
Arc::new(move |s| Box::pin(f(s)))
}
#[derive(Debug, Clone, PartialEq)]
pub struct StreamCall {
pub realm: [u8; 32],
pub procedure: String,
pub target: [u8; 32],
pub mode: StreamMode,
pub payload: Value,
pub deadline: Duration,
pub token: Option<Vec<u8>>,
pub proofs: Vec<Vec<u8>>,
}
impl Default for StreamCall {
fn default() -> Self {
StreamCall {
realm: [0; 32],
procedure: String::new(),
target: [0; 32],
mode: StreamMode::ServerStream,
payload: Value::Map(Vec::new()),
deadline: Duration::ZERO,
token: None,
proofs: Vec::new(),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum StreamEvent {
Data {
encoding: StreamEncoding,
body: Value,
},
End {
role: StreamRole,
},
Reply {
payload: Value,
},
}
#[derive(Clone)]
pub struct Stream {
inner: Arc<StreamInner>,
}
struct Budget {
admission: Arc<Admission>,
caller: [u8; 32],
place: Option<SessionPlace>,
}
pub(super) struct StreamInner {
link: Arc<Inner>,
writer: FrameWriter,
open: VerifiedRequest,
caller: bool,
send_seq: tokio::sync::Mutex<u64>,
state: Mutex<StreamSide>,
budget: Mutex<Option<Budget>>,
notify: Notify,
done_tx: watch::Sender<bool>,
}
#[derive(Default)]
struct StreamSide {
sent_end: bool,
peer_ended: bool,
inbox: VecDeque<(StreamEvent, usize)>,
held: usize,
ended: bool,
err: Option<LinkError>,
}
impl Stream {
pub fn request(&self) -> &VerifiedRequest {
&self.inner.open
}
pub async fn send(&self, body: &[u8]) -> Result<(), LinkError> {
self.inner
.send(
|seq| StreamFields::Data {
seq,
encoding: StreamEncoding::Raw,
body: Value::Bytes(body.to_vec()),
},
false,
)
.await
}
pub async fn send_value(&self, v: Value) -> Result<(), LinkError> {
self.inner
.send(
|seq| StreamFields::Data {
seq,
encoding: StreamEncoding::Msgpack,
body: v.clone(),
},
false,
)
.await
}
pub async fn close_send(&self) -> Result<(), LinkError> {
self.inner
.send(
|seq| StreamFields::End {
seq,
role: StreamRole::Send,
},
true,
)
.await
}
pub async fn close(&self) -> Result<(), LinkError> {
let sent = self
.inner
.send(
|seq| StreamFields::End {
seq,
role: StreamRole::Both,
},
true,
)
.await;
StreamInner::end(&self.inner, None);
sent
}
pub async fn reply(&self, payload: Value) -> Result<(), LinkError> {
let sent = self
.inner
.send(
|seq| StreamFields::Reply {
seq,
payload: payload.clone(),
},
true,
)
.await;
StreamInner::end(&self.inner, None);
sent
}
pub async fn abort(&self, code: &str, message: &str) -> Result<(), LinkError> {
self.inner.abort(code, message).await
}
pub async fn recv(&self) -> Result<StreamEvent, LinkError> {
loop {
let notified = self.inner.notify.notified();
{
let mut side = self.inner.side();
if let Some((event, size)) = side.inbox.pop_front() {
side.held -= size;
drop(side);
self.inner.release_inbox(size);
return Ok(event);
}
if side.ended {
return Err(side.err.clone().unwrap_or(LinkError::EndOfStream));
}
}
notified.await;
}
}
pub async fn done(&self) -> Option<LinkError> {
let mut done = self.inner.done_tx.subscribe();
let _ = done.wait_for(|ended| *ended).await;
self.inner.side().err.clone()
}
}
impl StreamInner {
fn new(
link: Arc<Inner>,
send: quinn::SendStream,
open: VerifiedRequest,
caller: bool,
) -> Arc<StreamInner> {
Arc::new(StreamInner {
link,
writer: FrameWriter::new(send),
open,
caller,
send_seq: tokio::sync::Mutex::new(0),
state: Mutex::new(StreamSide::default()),
budget: Mutex::new(None),
notify: Notify::new(),
done_tx: watch::channel(false).0,
})
}
fn side(&self) -> MutexGuard<'_, StreamSide> {
self.state.lock().unwrap_or_else(|p| p.into_inner())
}
fn budget(&self) -> MutexGuard<'_, Option<Budget>> {
self.budget.lock().unwrap_or_else(|p| p.into_inner())
}
async fn send(
self: &Arc<Self>,
at: impl FnOnce(u64) -> StreamFields,
last: bool,
) -> Result<(), LinkError> {
let mut seq = self.send_seq.lock().await;
if self.side().sent_end {
return Err(LinkError::StreamClosed);
}
let fields = at(*seq);
let signed = if self.caller {
frame::sign_caller_stream(&fields, &self.open, &self.link.key)?
} else {
frame::sign_provider_stream(&fields, &self.open, &self.link.key)?
};
let encoded = cbor::encode(&signed)
.map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))?;
if let Err(e) = self.writer.write(&encoded, MAX_FRAME_BYTES).await {
self.side().sent_end = true;
return Err(e);
}
*seq += 1;
if last {
let peer_ended = {
let mut side = self.side();
side.sent_end = true;
side.peer_ended
};
self.writer.finish().await;
if peer_ended {
StreamInner::end(self, None);
}
}
Ok(())
}
async fn abort(self: &Arc<Self>, code: &str, message: &str) -> Result<(), LinkError> {
let sent = self
.send(
|seq| StreamFields::Error {
seq,
code: code.to_string(),
message: message.to_string(),
},
true,
)
.await;
StreamInner::end(
self,
Some(LinkError::Stream {
code: code.to_string(),
message: message.to_string(),
relay: false,
}),
);
sent
}
fn deliver(&self, event: StreamEvent, size: usize) -> bool {
{
let mut side = self.side();
if side.held + size > STREAM_INBOX {
return false;
}
if let Some(budget) = &*self.budget() {
if !budget.admission.charge_inbox(budget.caller, size) {
return false;
}
}
side.inbox.push_back((event, size));
side.held += size;
}
self.notify.notify_one();
true
}
fn release_inbox(&self, size: usize) {
if let Some(budget) = &*self.budget() {
budget.admission.release_inbox(budget.caller, size);
}
}
fn peer_finished(self: &Arc<Self>, err: Option<LinkError>) {
self.side().peer_ended = true;
StreamInner::end(self, err);
}
async fn fail(self: &Arc<Self>, code: &str, cause: Option<String>) {
let message = cause
.as_deref()
.map(bounded_detail)
.unwrap_or("")
.to_string();
let _ = self.abort(code, &message).await;
}
pub(super) fn end(this: &Arc<StreamInner>, err: Option<LinkError>) {
let graceful = {
let mut side = this.side();
if side.ended {
return;
}
side.ended = true;
if side.err.is_none() {
side.err = err;
}
let graceful = side.sent_end;
side.sent_end = true;
graceful
};
if !graceful {
let released = this.clone();
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(async move { released.writer.reset().await });
}
}
if let Some(budget) = this.budget().take() {
let held = std::mem::take(&mut this.side().held);
budget.admission.release_inbox(budget.caller, held);
drop(budget.place);
}
this.link
.lock()
.streams
.retain(|w| w.strong_count() > 0 && !std::ptr::eq(w.as_ptr(), Arc::as_ptr(this)));
let _ = this.done_tx.send_replace(true);
this.notify.notify_waiters();
this.notify.notify_one();
}
}
fn hold_stream(inner: &Inner, s: &Arc<StreamInner>) -> bool {
let mut state = inner.lock();
if state.ended.is_some() {
return false;
}
state.streams.push(Arc::downgrade(s));
true
}
fn abandon(mut send: quinn::SendStream, mut recv: quinn::RecvStream) {
let _ = send.reset(0u32.into());
let _ = recv.stop(0u32.into());
}
impl Link {
pub async fn open_stream(&self, c: StreamCall) -> Result<Stream, LinkError> {
let inner = &self.inner;
let deadline = if c.deadline.is_zero() {
DEFAULT_STREAM_DEADLINE
} else {
c.deadline
};
let mut request_id = [0u8; 16];
aws_lc_rs::rand::fill(&mut request_id)
.map_err(|_| LinkError::Io("no randomness".into()))?;
let signed = frame::sign_stream_open(
&RequestSpec {
request_id,
realm: c.realm,
procedure: c.procedure,
target: c.target,
deadline: (now_ms() + deadline.as_millis() as i64) as u64,
payload: c.payload,
mode: Some(c.mode),
token: c.token,
proofs: c.proofs,
source_route: None,
retry_budget: None,
},
&inner.key,
)?;
let encoded = cbor::encode(&signed)
.map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))?;
if encoded.len() > STREAM_OPEN_BYTES {
return Err(LinkError::StreamOpenTooLarge(encoded.len()));
}
let open = frame::verify_request(&signed, inner.profile)?;
let state = frame::open_stream(&open)?;
let (send, recv) = inner
.connection
.open_bi()
.await
.map_err(|e| LinkError::Io(format!("open a stream: {e}")))?;
let s = StreamInner::new(inner.clone(), send, open, true);
let held = hold_stream(inner, &s);
let written = match held {
true => s.writer.write(&encoded, STREAM_OPEN_BYTES).await,
false => Err(inner.lock().ended.clone().unwrap_or(LinkError::Closed)),
};
if let Err(e) = written {
StreamInner::end(&s, Some(e.clone()));
let mut recv = recv;
let _ = recv.stop(0u32.into());
return Err(e);
}
tokio::spawn(read(s.clone(), recv, state));
Ok(Stream { inner: s })
}
}
pub(super) async fn accept_streams(link: Weak<Inner>) {
let Some(connection) = link.upgrade().map(|l| l.connection.clone()) else {
return;
};
while let Ok((send, recv)) = connection.accept_bi().await {
tokio::spawn(incoming(link.clone(), send, recv));
}
}
async fn incoming(link: Weak<Inner>, send: quinn::SendStream, mut recv: quinn::RecvStream) {
let Some(inner) = link.upgrade() else { return };
let payload = match tokio::time::timeout(
STREAM_OPEN_WAIT,
read_frame(&mut recv, STREAM_OPEN_BYTES),
)
.await
{
Ok(Ok(payload)) => payload,
_ => {
inner.count("stream_open_unread");
abandon(send, recv);
return;
}
};
let v = match cbor::decode(&payload) {
Ok(v) if frame_type_of(&v) == "stream_open" => v,
_ => {
inner.count("stream_open_malformed");
abandon(send, recv);
return;
}
};
let Ok(open) = frame::verify_request(&v, inner.profile) else {
inner.count("stream_open_unverified");
abandon(send, recv);
return;
};
if open.target != inner.self_id {
inner.count("stream_for_another_node");
abandon(send, recv);
return;
}
let Ok(state) = frame::open_stream(&open) else {
abandon(send, recv);
return;
};
let s = StreamInner::new(inner.clone(), send, open.clone(), false);
let offer = match admit_stream(&inner, &open) {
Ok(offer) => offer,
Err(code) => return refuse(&s, code, recv).await,
};
let Some(place) = inner.admission.open_session(open.caller) else {
return refuse(&s, CODE_TOO_MANY_SESSIONS, recv).await;
};
*s.budget() = Some(Budget {
admission: inner.admission.clone(),
caller: open.caller,
place: Some(place),
});
if !hold_stream(&inner, &s) {
StreamInner::end(&s, Some(LinkError::Closed));
let _ = recv.stop(0u32.into());
return;
}
tokio::spawn(read(s.clone(), recv, state));
tokio::spawn(serve(s, offer));
}
fn admit_stream(inner: &Inner, open: &VerifiedRequest) -> Result<StreamOffer, &'static str> {
match inner.admission.admit(open, &inner.share, now_ms()) {
Verdict::Refused(code) => return Err(code),
Verdict::Copy(_) => return Err(CODE_REQUEST_COPY),
Verdict::New => {}
}
let state = inner.lock();
let offer = state
.served
.get(&(open.realm, open.procedure.clone()))
.and_then(|s| s.offer.stream.clone())
.ok_or(CODE_STREAM_NOT_FOUND)?;
if Some(offer.mode) != open.mode {
return Err(CODE_MODE_MISMATCH);
}
Ok(offer)
}
async fn refuse(s: &Arc<StreamInner>, code: &str, mut recv: quinn::RecvStream) {
s.link.count(&format!("stream_refused_{code}"));
let _ = s.abort(code, "").await;
let _ = recv.stop(0u32.into());
}
async fn serve(s: Arc<StreamInner>, offer: StreamOffer) {
let stream = Stream { inner: s.clone() };
let mut running = tokio::spawn((offer.handler)(stream.clone()));
let mut done = s.done_tx.subscribe();
let outcome = tokio::select! {
outcome = &mut running => outcome,
_ = done.wait_for(|ended| *ended) => {
running.abort();
return;
}
};
match outcome {
Ok(Ok(())) => {
let _ = stream.close().await;
}
Ok(Err(e)) => {
let _ = stream
.abort(CODE_STREAM_HANDLER_ERROR, bounded_detail(&e))
.await;
}
Err(panicked) => {
let _ = stream
.abort(
CODE_STREAM_HANDLER_ERROR,
bounded_detail(&panicked.to_string()),
)
.await;
}
}
}
async fn read(s: Arc<StreamInner>, mut recv: quinn::RecvStream, mut state: StreamState) {
let mut done = s.done_tx.subscribe();
loop {
let payload = tokio::select! {
_ = done.wait_for(|ended| *ended) => return,
payload = read_frame(&mut recv, MAX_FRAME_BYTES) => payload,
};
let payload = match payload {
Ok(payload) => payload,
Err(e) => return read_ended(&s, e),
};
match received(&s, &payload, &state).await {
Some(next) => state = next,
None => return,
}
}
}
fn read_ended(s: &Arc<StreamInner>, e: LinkError) {
if s.side().peer_ended {
return;
}
let err = s.link.lock().ended.clone().unwrap_or(e);
StreamInner::end(s, Some(err));
}
async fn received(
s: &Arc<StreamInner>,
payload: &[u8],
state: &StreamState,
) -> Option<StreamState> {
let v = match cbor::decode(payload) {
Ok(v) => v,
Err(e) => {
s.fail("malformed_frame", Some(e.to_string())).await;
return None;
}
};
if s.caller && v.get("relay_error").is_some() {
match frame::verify_relay_error(&v, &s.open, s.link.profile, &s.link.station.node_id) {
Ok(relayed) => s.peer_finished(Some(LinkError::Stream {
code: relayed.code,
message: String::new(),
relay: true,
})),
Err(e) => s.fail("malformed_frame", Some(e.to_string())).await,
}
return None;
}
let verified = if s.caller {
frame::verify_provider_stream(&v, state, s.link.profile)
} else {
frame::verify_caller_stream(&v, state, s.link.profile)
};
let (verified, next) = match verified {
Ok(verified) => verified,
Err(e) => {
s.fail("malformed_frame", Some(e.to_string())).await;
return None;
}
};
let size = payload.len();
match verified.fields {
StreamFields::Error { code, message, .. } => {
s.peer_finished(Some(LinkError::Stream {
code,
message,
relay: false,
}));
None
}
StreamFields::Reply { payload, .. } => {
s.deliver(StreamEvent::Reply { payload }, size);
s.peer_finished(None);
None
}
StreamFields::End { role, .. } => {
s.deliver(StreamEvent::End { role }, size);
if role == StreamRole::Both {
s.peer_finished(None);
return None;
}
let mine = {
let mut side = s.side();
side.peer_ended = true;
side.sent_end
};
if mine {
StreamInner::end(s, None);
}
None
}
StreamFields::Data { encoding, body, .. } => {
if !s.deliver(StreamEvent::Data { encoding, body }, size) {
s.fail("resource_exhausted", None).await;
return None;
}
Some(next)
}
}
}