1use std::collections::VecDeque;
18use std::future::Future;
19use std::sync::{Arc, Mutex, MutexGuard, Weak};
20use std::time::Duration;
21
22use tokio::sync::{watch, Notify};
23
24use crate::cbor::{self, Value};
25use crate::frame::{
26 self, RequestSpec, StreamEncoding, StreamFields, StreamMode, StreamRole, StreamState,
27 VerifiedRequest,
28};
29
30use super::admission::{Admission, SessionPlace, Verdict};
31use super::framing::{read_frame, FrameWriter, MAX_FRAME_BYTES};
32use super::serve::{bounded_detail, BoxFuture, StreamOffer, CODE_REQUEST_COPY};
33use super::{frame_type_of, now_ms, Inner, Link, LinkError};
34
35const STREAM_OPEN_BYTES: usize = 1024 * 1024;
36const STREAM_OPEN_WAIT: Duration = Duration::from_secs(10);
37const STREAM_INBOX: usize = 16 * 1024 * 1024;
38
39pub const DEFAULT_STREAM_DEADLINE: Duration = Duration::from_secs(30);
42
43const CODE_STREAM_NOT_FOUND: &str = "not_found";
46const CODE_MODE_MISMATCH: &str = "mode_mismatch";
47const CODE_TOO_MANY_SESSIONS: &str = "too_many_sessions";
48const CODE_STREAM_HANDLER_ERROR: &str = "error";
49
50pub type StreamHandler = Arc<dyn Fn(Stream) -> BoxFuture<Result<(), String>> + Send + Sync>;
55
56pub fn stream_handler<F, Fut>(f: F) -> StreamHandler
58where
59 F: Fn(Stream) -> Fut + Send + Sync + 'static,
60 Fut: Future<Output = Result<(), String>> + Send + 'static,
61{
62 Arc::new(move |s| Box::pin(f(s)))
63}
64
65#[derive(Debug, Clone, PartialEq)]
70pub struct StreamCall {
71 pub realm: [u8; 32],
72 pub procedure: String,
73 pub target: [u8; 32],
74 pub mode: StreamMode,
75 pub payload: Value,
76 pub deadline: Duration,
77 pub token: Option<Vec<u8>>,
78 pub proofs: Vec<Vec<u8>>,
79}
80
81impl Default for StreamCall {
82 fn default() -> Self {
83 StreamCall {
84 realm: [0; 32],
85 procedure: String::new(),
86 target: [0; 32],
87 mode: StreamMode::ServerStream,
88 payload: Value::Map(Vec::new()),
89 deadline: Duration::ZERO,
90 token: None,
91 proofs: Vec::new(),
92 }
93 }
94}
95
96#[derive(Debug, Clone, PartialEq)]
100pub enum StreamEvent {
101 Data {
102 encoding: StreamEncoding,
103 body: Value,
104 },
105 End {
106 role: StreamRole,
107 },
108 Reply {
109 payload: Value,
110 },
111}
112
113#[derive(Clone)]
115pub struct Stream {
116 inner: Arc<StreamInner>,
117}
118
119struct Budget {
122 admission: Arc<Admission>,
123 caller: [u8; 32],
124 place: Option<SessionPlace>,
125}
126
127pub(super) struct StreamInner {
128 link: Arc<Inner>,
129 writer: FrameWriter,
130 open: VerifiedRequest,
131 caller: bool,
132 send_seq: tokio::sync::Mutex<u64>,
134 state: Mutex<StreamSide>,
135 budget: Mutex<Option<Budget>>,
136 notify: Notify,
137 done_tx: watch::Sender<bool>,
138}
139
140#[derive(Default)]
141struct StreamSide {
142 sent_end: bool,
144 peer_ended: bool,
146 inbox: VecDeque<(StreamEvent, usize)>,
147 held: usize,
148 ended: bool,
149 err: Option<LinkError>,
150}
151
152impl Stream {
153 pub fn request(&self) -> &VerifiedRequest {
156 &self.inner.open
157 }
158
159 pub async fn send(&self, body: &[u8]) -> Result<(), LinkError> {
161 self.inner
162 .send(
163 |seq| StreamFields::Data {
164 seq,
165 encoding: StreamEncoding::Raw,
166 body: Value::Bytes(body.to_vec()),
167 },
168 false,
169 )
170 .await
171 }
172
173 pub async fn send_value(&self, v: Value) -> Result<(), LinkError> {
175 self.inner
176 .send(
177 |seq| StreamFields::Data {
178 seq,
179 encoding: StreamEncoding::Msgpack,
180 body: v.clone(),
181 },
182 false,
183 )
184 .await
185 }
186
187 pub async fn close_send(&self) -> Result<(), LinkError> {
189 self.inner
190 .send(
191 |seq| StreamFields::End {
192 seq,
193 role: StreamRole::Send,
194 },
195 true,
196 )
197 .await
198 }
199
200 pub async fn close(&self) -> Result<(), LinkError> {
202 let sent = self
203 .inner
204 .send(
205 |seq| StreamFields::End {
206 seq,
207 role: StreamRole::Both,
208 },
209 true,
210 )
211 .await;
212 StreamInner::end(&self.inner, None);
213 sent
214 }
215
216 pub async fn reply(&self, payload: Value) -> Result<(), LinkError> {
218 let sent = self
219 .inner
220 .send(
221 |seq| StreamFields::Reply {
222 seq,
223 payload: payload.clone(),
224 },
225 true,
226 )
227 .await;
228 StreamInner::end(&self.inner, None);
229 sent
230 }
231
232 pub async fn abort(&self, code: &str, message: &str) -> Result<(), LinkError> {
234 self.inner.abort(code, message).await
235 }
236
237 pub async fn recv(&self) -> Result<StreamEvent, LinkError> {
241 loop {
242 let notified = self.inner.notify.notified();
243 {
244 let mut side = self.inner.side();
245 if let Some((event, size)) = side.inbox.pop_front() {
246 side.held -= size;
247 drop(side);
248 self.inner.release_inbox(size);
249 return Ok(event);
250 }
251 if side.ended {
252 return Err(side.err.clone().unwrap_or(LinkError::EndOfStream));
253 }
254 }
255 notified.await;
256 }
257 }
258
259 pub async fn done(&self) -> Option<LinkError> {
262 let mut done = self.inner.done_tx.subscribe();
263 let _ = done.wait_for(|ended| *ended).await;
264 self.inner.side().err.clone()
265 }
266}
267
268impl StreamInner {
269 fn new(
270 link: Arc<Inner>,
271 send: quinn::SendStream,
272 open: VerifiedRequest,
273 caller: bool,
274 ) -> Arc<StreamInner> {
275 Arc::new(StreamInner {
276 link,
277 writer: FrameWriter::new(send),
278 open,
279 caller,
280 send_seq: tokio::sync::Mutex::new(0),
281 state: Mutex::new(StreamSide::default()),
282 budget: Mutex::new(None),
283 notify: Notify::new(),
284 done_tx: watch::channel(false).0,
285 })
286 }
287
288 fn side(&self) -> MutexGuard<'_, StreamSide> {
289 self.state.lock().unwrap_or_else(|p| p.into_inner())
290 }
291
292 fn budget(&self) -> MutexGuard<'_, Option<Budget>> {
293 self.budget.lock().unwrap_or_else(|p| p.into_inner())
294 }
295
296 async fn send(
300 self: &Arc<Self>,
301 at: impl FnOnce(u64) -> StreamFields,
302 last: bool,
303 ) -> Result<(), LinkError> {
304 let mut seq = self.send_seq.lock().await;
305 if self.side().sent_end {
306 return Err(LinkError::StreamClosed);
307 }
308 let fields = at(*seq);
309 let signed = if self.caller {
310 frame::sign_caller_stream(&fields, &self.open, &self.link.key)?
311 } else {
312 frame::sign_provider_stream(&fields, &self.open, &self.link.key)?
313 };
314 let encoded = cbor::encode(&signed)
315 .map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))?;
316 if let Err(e) = self.writer.write(&encoded, MAX_FRAME_BYTES).await {
317 self.side().sent_end = true;
318 return Err(e);
319 }
320 *seq += 1;
321 if last {
322 let peer_ended = {
323 let mut side = self.side();
324 side.sent_end = true;
325 side.peer_ended
326 };
327 self.writer.finish().await;
328 if peer_ended {
329 StreamInner::end(self, None);
330 }
331 }
332 Ok(())
333 }
334
335 async fn abort(self: &Arc<Self>, code: &str, message: &str) -> Result<(), LinkError> {
336 let sent = self
337 .send(
338 |seq| StreamFields::Error {
339 seq,
340 code: code.to_string(),
341 message: message.to_string(),
342 },
343 true,
344 )
345 .await;
346 StreamInner::end(
347 self,
348 Some(LinkError::Stream {
349 code: code.to_string(),
350 message: message.to_string(),
351 relay: false,
352 }),
353 );
354 sent
355 }
356
357 fn deliver(&self, event: StreamEvent, size: usize) -> bool {
360 {
361 let mut side = self.side();
362 if side.held + size > STREAM_INBOX {
363 return false;
364 }
365 if let Some(budget) = &*self.budget() {
366 if !budget.admission.charge_inbox(budget.caller, size) {
367 return false;
368 }
369 }
370 side.inbox.push_back((event, size));
371 side.held += size;
372 }
373 self.notify.notify_one();
374 true
375 }
376
377 fn release_inbox(&self, size: usize) {
378 if let Some(budget) = &*self.budget() {
379 budget.admission.release_inbox(budget.caller, size);
380 }
381 }
382
383 fn peer_finished(self: &Arc<Self>, err: Option<LinkError>) {
386 self.side().peer_ended = true;
387 StreamInner::end(self, err);
388 }
389
390 async fn fail(self: &Arc<Self>, code: &str, cause: Option<String>) {
394 let message = cause
395 .as_deref()
396 .map(bounded_detail)
397 .unwrap_or("")
398 .to_string();
399 let _ = self.abort(code, &message).await;
400 }
401
402 pub(super) fn end(this: &Arc<StreamInner>, err: Option<LinkError>) {
406 let graceful = {
407 let mut side = this.side();
408 if side.ended {
409 return;
410 }
411 side.ended = true;
412 if side.err.is_none() {
413 side.err = err;
414 }
415 let graceful = side.sent_end;
416 side.sent_end = true;
417 graceful
418 };
419 if !graceful {
420 let released = this.clone();
421 if let Ok(runtime) = tokio::runtime::Handle::try_current() {
422 runtime.spawn(async move { released.writer.reset().await });
423 }
424 }
425 if let Some(budget) = this.budget().take() {
426 let held = std::mem::take(&mut this.side().held);
427 budget.admission.release_inbox(budget.caller, held);
428 drop(budget.place);
429 }
430 this.link
431 .lock()
432 .streams
433 .retain(|w| w.strong_count() > 0 && !std::ptr::eq(w.as_ptr(), Arc::as_ptr(this)));
434 let _ = this.done_tx.send_replace(true);
435 this.notify.notify_waiters();
436 this.notify.notify_one();
437 }
438}
439
440fn hold_stream(inner: &Inner, s: &Arc<StreamInner>) -> bool {
443 let mut state = inner.lock();
444 if state.ended.is_some() {
445 return false;
446 }
447 state.streams.push(Arc::downgrade(s));
448 true
449}
450
451fn abandon(mut send: quinn::SendStream, mut recv: quinn::RecvStream) {
453 let _ = send.reset(0u32.into());
454 let _ = recv.stop(0u32.into());
455}
456
457impl Link {
458 pub async fn open_stream(&self, c: StreamCall) -> Result<Stream, LinkError> {
462 let inner = &self.inner;
463 let deadline = if c.deadline.is_zero() {
464 DEFAULT_STREAM_DEADLINE
465 } else {
466 c.deadline
467 };
468 let mut request_id = [0u8; 16];
469 aws_lc_rs::rand::fill(&mut request_id)
470 .map_err(|_| LinkError::Io("no randomness".into()))?;
471 let signed = frame::sign_stream_open(
472 &RequestSpec {
473 request_id,
474 realm: c.realm,
475 procedure: c.procedure,
476 target: c.target,
477 deadline: (now_ms() + deadline.as_millis() as i64) as u64,
478 payload: c.payload,
479 mode: Some(c.mode),
480 token: c.token,
481 proofs: c.proofs,
482 source_route: None,
483 retry_budget: None,
484 },
485 &inner.key,
486 )?;
487 let encoded = cbor::encode(&signed)
488 .map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))?;
489 if encoded.len() > STREAM_OPEN_BYTES {
490 return Err(LinkError::StreamOpenTooLarge(encoded.len()));
491 }
492 let open = frame::verify_request(&signed, inner.profile)?;
493 let state = frame::open_stream(&open)?;
494 let (send, recv) = inner
495 .connection
496 .open_bi()
497 .await
498 .map_err(|e| LinkError::Io(format!("open a stream: {e}")))?;
499 let s = StreamInner::new(inner.clone(), send, open, true);
500 let held = hold_stream(inner, &s);
501 let written = match held {
502 true => s.writer.write(&encoded, STREAM_OPEN_BYTES).await,
503 false => Err(inner.lock().ended.clone().unwrap_or(LinkError::Closed)),
504 };
505 if let Err(e) = written {
506 StreamInner::end(&s, Some(e.clone()));
507 let mut recv = recv;
508 let _ = recv.stop(0u32.into());
509 return Err(e);
510 }
511 tokio::spawn(read(s.clone(), recv, state));
512 Ok(Stream { inner: s })
513 }
514}
515
516pub(super) async fn accept_streams(link: Weak<Inner>) {
518 let Some(connection) = link.upgrade().map(|l| l.connection.clone()) else {
519 return;
520 };
521 while let Ok((send, recv)) = connection.accept_bi().await {
522 tokio::spawn(incoming(link.clone(), send, recv));
523 }
524}
525
526async fn incoming(link: Weak<Inner>, send: quinn::SendStream, mut recv: quinn::RecvStream) {
532 let Some(inner) = link.upgrade() else { return };
533 let payload = match tokio::time::timeout(
534 STREAM_OPEN_WAIT,
535 read_frame(&mut recv, STREAM_OPEN_BYTES),
536 )
537 .await
538 {
539 Ok(Ok(payload)) => payload,
540 _ => {
541 inner.count("stream_open_unread");
542 abandon(send, recv);
543 return;
544 }
545 };
546 let v = match cbor::decode(&payload) {
547 Ok(v) if frame_type_of(&v) == "stream_open" => v,
548 _ => {
549 inner.count("stream_open_malformed");
550 abandon(send, recv);
551 return;
552 }
553 };
554 let Ok(open) = frame::verify_request(&v, inner.profile) else {
555 inner.count("stream_open_unverified");
556 abandon(send, recv);
557 return;
558 };
559 if open.target != inner.self_id {
560 inner.count("stream_for_another_node");
561 abandon(send, recv);
562 return;
563 }
564 let Ok(state) = frame::open_stream(&open) else {
565 abandon(send, recv);
566 return;
567 };
568 let s = StreamInner::new(inner.clone(), send, open.clone(), false);
569 let offer = match admit_stream(&inner, &open) {
570 Ok(offer) => offer,
571 Err(code) => return refuse(&s, code, recv).await,
572 };
573 let Some(place) = inner.admission.open_session(open.caller) else {
574 return refuse(&s, CODE_TOO_MANY_SESSIONS, recv).await;
575 };
576 *s.budget() = Some(Budget {
577 admission: inner.admission.clone(),
578 caller: open.caller,
579 place: Some(place),
580 });
581 if !hold_stream(&inner, &s) {
582 StreamInner::end(&s, Some(LinkError::Closed));
583 let _ = recv.stop(0u32.into());
584 return;
585 }
586 tokio::spawn(read(s.clone(), recv, state));
587 tokio::spawn(serve(s, offer));
588}
589
590fn admit_stream(inner: &Inner, open: &VerifiedRequest) -> Result<StreamOffer, &'static str> {
594 match inner.admission.admit(open, &inner.share, now_ms()) {
595 Verdict::Refused(code) => return Err(code),
596 Verdict::Copy(_) => return Err(CODE_REQUEST_COPY),
597 Verdict::New => {}
598 }
599 let state = inner.lock();
600 let offer = state
601 .served
602 .get(&(open.realm, open.procedure.clone()))
603 .and_then(|s| s.offer.stream.clone())
604 .ok_or(CODE_STREAM_NOT_FOUND)?;
605 if Some(offer.mode) != open.mode {
606 return Err(CODE_MODE_MISMATCH);
607 }
608 Ok(offer)
609}
610
611async fn refuse(s: &Arc<StreamInner>, code: &str, mut recv: quinn::RecvStream) {
614 s.link.count(&format!("stream_refused_{code}"));
615 let _ = s.abort(code, "").await;
616 let _ = recv.stop(0u32.into());
617}
618
619async fn serve(s: Arc<StreamInner>, offer: StreamOffer) {
623 let stream = Stream { inner: s.clone() };
624 let mut running = tokio::spawn((offer.handler)(stream.clone()));
625 let mut done = s.done_tx.subscribe();
626 let outcome = tokio::select! {
627 outcome = &mut running => outcome,
628 _ = done.wait_for(|ended| *ended) => {
629 running.abort();
630 return;
631 }
632 };
633 match outcome {
634 Ok(Ok(())) => {
635 let _ = stream.close().await;
636 }
637 Ok(Err(e)) => {
638 let _ = stream
639 .abort(CODE_STREAM_HANDLER_ERROR, bounded_detail(&e))
640 .await;
641 }
642 Err(panicked) => {
643 let _ = stream
644 .abort(
645 CODE_STREAM_HANDLER_ERROR,
646 bounded_detail(&panicked.to_string()),
647 )
648 .await;
649 }
650 }
651}
652
653async fn read(s: Arc<StreamInner>, mut recv: quinn::RecvStream, mut state: StreamState) {
657 let mut done = s.done_tx.subscribe();
658 loop {
659 let payload = tokio::select! {
660 _ = done.wait_for(|ended| *ended) => return,
661 payload = read_frame(&mut recv, MAX_FRAME_BYTES) => payload,
662 };
663 let payload = match payload {
664 Ok(payload) => payload,
665 Err(e) => return read_ended(&s, e),
666 };
667 match received(&s, &payload, &state).await {
668 Some(next) => state = next,
669 None => return,
670 }
671 }
672}
673
674fn read_ended(s: &Arc<StreamInner>, e: LinkError) {
677 if s.side().peer_ended {
678 return;
679 }
680 let err = s.link.lock().ended.clone().unwrap_or(e);
681 StreamInner::end(s, Some(err));
682}
683
684async fn received(
687 s: &Arc<StreamInner>,
688 payload: &[u8],
689 state: &StreamState,
690) -> Option<StreamState> {
691 let v = match cbor::decode(payload) {
692 Ok(v) => v,
693 Err(e) => {
694 s.fail("malformed_frame", Some(e.to_string())).await;
695 return None;
696 }
697 };
698 if s.caller && v.get("relay_error").is_some() {
699 match frame::verify_relay_error(&v, &s.open, s.link.profile, &s.link.station.node_id) {
700 Ok(relayed) => s.peer_finished(Some(LinkError::Stream {
701 code: relayed.code,
702 message: String::new(),
703 relay: true,
704 })),
705 Err(e) => s.fail("malformed_frame", Some(e.to_string())).await,
706 }
707 return None;
708 }
709 let verified = if s.caller {
710 frame::verify_provider_stream(&v, state, s.link.profile)
711 } else {
712 frame::verify_caller_stream(&v, state, s.link.profile)
713 };
714 let (verified, next) = match verified {
715 Ok(verified) => verified,
716 Err(e) => {
717 s.fail("malformed_frame", Some(e.to_string())).await;
718 return None;
719 }
720 };
721 let size = payload.len();
722 match verified.fields {
723 StreamFields::Error { code, message, .. } => {
724 s.peer_finished(Some(LinkError::Stream {
725 code,
726 message,
727 relay: false,
728 }));
729 None
730 }
731 StreamFields::Reply { payload, .. } => {
732 s.deliver(StreamEvent::Reply { payload }, size);
733 s.peer_finished(None);
734 None
735 }
736 StreamFields::End { role, .. } => {
737 s.deliver(StreamEvent::End { role }, size);
738 if role == StreamRole::Both {
739 s.peer_finished(None);
740 return None;
741 }
742 let mine = {
743 let mut side = s.side();
744 side.peer_ended = true;
745 side.sent_end
746 };
747 if mine {
748 StreamInner::end(s, None);
749 }
750 None
751 }
752 StreamFields::Data { encoding, body, .. } => {
753 if !s.deliver(StreamEvent::Data { encoding, body }, size) {
754 s.fail("resource_exhausted", None).await;
755 return None;
756 }
757 Some(next)
758 }
759 }
760}