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::confidential::{
32 clear_allowed, opened_request, sealed_request, stated, unsealed, Seal, StreamSeal,
33 CODE_SEALED_REFUSED, CODE_SEALED_REQUIRED,
34};
35use super::framing::{read_frame, FrameWriter, MAX_FRAME_BYTES};
36use super::serve::{bounded_detail, BoxFuture, StreamOffer, CODE_REQUEST_COPY};
37use super::{frame_type_of, now_ms, Inner, Link, LinkError};
38
39const STREAM_OPEN_BYTES: usize = 1024 * 1024;
40const STREAM_OPEN_WAIT: Duration = Duration::from_secs(10);
41const STREAM_INBOX: usize = 16 * 1024 * 1024;
42
43pub const DEFAULT_STREAM_DEADLINE: Duration = Duration::from_secs(30);
46
47const CODE_STREAM_NOT_FOUND: &str = "not_found";
50const CODE_MODE_MISMATCH: &str = "mode_mismatch";
51const CODE_TOO_MANY_SESSIONS: &str = "too_many_sessions";
52const CODE_STREAM_HANDLER_ERROR: &str = "error";
53
54pub type StreamHandler = Arc<dyn Fn(Stream) -> BoxFuture<Result<(), String>> + Send + Sync>;
59
60pub fn stream_handler<F, Fut>(f: F) -> StreamHandler
62where
63 F: Fn(Stream) -> Fut + Send + Sync + 'static,
64 Fut: Future<Output = Result<(), String>> + Send + 'static,
65{
66 Arc::new(move |s| Box::pin(f(s)))
67}
68
69#[derive(Debug, Clone, PartialEq)]
75pub struct StreamCall {
76 pub realm: [u8; 32],
77 pub procedure: String,
78 pub target: [u8; 32],
79 pub mode: StreamMode,
80 pub payload: Value,
81 pub deadline: Duration,
82 pub token: Option<Vec<u8>>,
83 pub proofs: Vec<Vec<u8>>,
84 pub seal: Option<Seal>,
85}
86
87impl Default for StreamCall {
88 fn default() -> Self {
89 StreamCall {
90 realm: [0; 32],
91 procedure: String::new(),
92 target: [0; 32],
93 mode: StreamMode::ServerStream,
94 payload: Value::Map(Vec::new()),
95 deadline: Duration::ZERO,
96 token: None,
97 proofs: Vec::new(),
98 seal: None,
99 }
100 }
101}
102
103#[derive(Debug, Clone, PartialEq)]
107pub enum StreamEvent {
108 Data {
109 encoding: StreamEncoding,
110 body: Value,
111 },
112 End {
113 role: StreamRole,
114 },
115 Reply {
116 payload: Value,
117 },
118}
119
120#[derive(Clone)]
122pub struct Stream {
123 pub(super) inner: Arc<StreamInner>,
124}
125
126struct Budget {
129 admission: Arc<Admission>,
130 caller: [u8; 32],
131 place: Option<SessionPlace>,
132}
133
134pub(super) struct StreamInner {
135 link: Arc<Inner>,
136 writer: FrameWriter,
137 pub(super) open: VerifiedRequest,
138 pub(super) caller: bool,
139 pub(super) sealing: Option<StreamSeal>,
141 send_seq: tokio::sync::Mutex<u64>,
143 state: Mutex<StreamSide>,
144 budget: Mutex<Option<Budget>>,
145 notify: Notify,
146 done_tx: watch::Sender<bool>,
147}
148
149#[derive(Default)]
150pub(super) struct StreamSide {
151 sent_end: bool,
153 peer_ended: bool,
155 inbox: VecDeque<(StreamEvent, usize)>,
156 held: usize,
157 ended: bool,
158 err: Option<LinkError>,
159 pub(super) settled: bool,
161}
162
163impl Stream {
164 pub fn request(&self) -> &VerifiedRequest {
168 &self.inner.open
169 }
170
171 pub fn sealed(&self) -> bool {
174 self.inner.sealing.is_some()
175 }
176
177 pub async fn send(&self, body: &[u8]) -> Result<(), LinkError> {
179 self.inner
180 .send(
181 |seq| StreamFields::Data {
182 seq,
183 encoding: StreamEncoding::Raw,
184 body: Value::Bytes(body.to_vec()),
185 },
186 false,
187 )
188 .await
189 }
190
191 pub async fn send_value(&self, v: Value) -> Result<(), LinkError> {
193 self.inner
194 .send(
195 |seq| StreamFields::Data {
196 seq,
197 encoding: StreamEncoding::Msgpack,
198 body: v.clone(),
199 },
200 false,
201 )
202 .await
203 }
204
205 pub async fn close_send(&self) -> Result<(), LinkError> {
207 self.inner
208 .send(
209 |seq| StreamFields::End {
210 seq,
211 role: StreamRole::Send,
212 },
213 true,
214 )
215 .await
216 }
217
218 pub async fn close(&self) -> Result<(), LinkError> {
220 let sent = self
221 .inner
222 .send(
223 |seq| StreamFields::End {
224 seq,
225 role: StreamRole::Both,
226 },
227 true,
228 )
229 .await;
230 StreamInner::end(&self.inner, None);
231 sent
232 }
233
234 pub async fn reply(&self, payload: Value) -> Result<(), LinkError> {
236 let sent = self
237 .inner
238 .send(
239 |seq| StreamFields::Reply {
240 seq,
241 payload: payload.clone(),
242 },
243 true,
244 )
245 .await;
246 StreamInner::end(&self.inner, None);
247 sent
248 }
249
250 pub async fn abort(&self, code: &str, message: &str) -> Result<(), LinkError> {
252 self.inner.abort(code, message).await
253 }
254
255 pub async fn recv(&self) -> Result<StreamEvent, LinkError> {
259 loop {
260 let notified = self.inner.notify.notified();
261 {
262 let mut side = self.inner.side();
263 if let Some((event, size)) = side.inbox.pop_front() {
264 side.held -= size;
265 drop(side);
266 self.inner.release_inbox(size);
267 return Ok(event);
268 }
269 if side.ended {
270 return Err(side.err.clone().unwrap_or(LinkError::EndOfStream));
271 }
272 }
273 notified.await;
274 }
275 }
276
277 pub async fn done(&self) -> Option<LinkError> {
280 let mut done = self.inner.done_tx.subscribe();
281 let _ = done.wait_for(|ended| *ended).await;
282 self.inner.side().err.clone()
283 }
284}
285
286impl StreamInner {
287 fn new(
288 link: Arc<Inner>,
289 send: quinn::SendStream,
290 open: VerifiedRequest,
291 caller: bool,
292 sealing: Option<StreamSeal>,
293 ) -> Arc<StreamInner> {
294 Arc::new(StreamInner {
295 link,
296 writer: FrameWriter::new(send),
297 open,
298 caller,
299 sealing,
300 send_seq: tokio::sync::Mutex::new(0),
301 state: Mutex::new(StreamSide::default()),
302 budget: Mutex::new(None),
303 notify: Notify::new(),
304 done_tx: watch::channel(false).0,
305 })
306 }
307
308 pub(super) fn side(&self) -> MutexGuard<'_, StreamSide> {
309 self.state.lock().unwrap_or_else(|p| p.into_inner())
310 }
311
312 fn budget(&self) -> MutexGuard<'_, Option<Budget>> {
313 self.budget.lock().unwrap_or_else(|p| p.into_inner())
314 }
315
316 async fn send(
323 self: &Arc<Self>,
324 at: impl FnOnce(u64) -> StreamFields,
325 last: bool,
326 ) -> Result<(), LinkError> {
327 let mut seq = self.send_seq.lock().await;
328 if self.side().sent_end {
329 return Err(LinkError::StreamClosed);
330 }
331 let fields = at(*seq);
332 let plain = match &self.sealing {
333 Some(sealing) => sealing.plain_of(&fields)?,
334 None => None,
335 };
336 let spent = plain.is_some();
337 let fields = match (&self.sealing, plain) {
338 (Some(sealing), Some(plain)) => {
339 *seq += 1;
340 let sealed = sealing.sealed(fields, plain);
341 sealed.inspect_err(|_| self.side().sent_end = true)?
342 }
343 _ => fields,
344 };
345 let encoded = match self.signed(&fields) {
346 Ok(encoded) => encoded,
347 Err(e) => {
348 if spent {
349 self.side().sent_end = true;
350 }
351 return Err(e);
352 }
353 };
354 if let Err(e) = self.writer.write(&encoded, MAX_FRAME_BYTES).await {
355 self.side().sent_end = true;
356 return Err(e);
357 }
358 if !spent {
359 *seq += 1;
360 }
361 if last {
362 let peer_ended = {
363 let mut side = self.side();
364 side.sent_end = true;
365 side.peer_ended
366 };
367 self.writer.finish().await;
368 if peer_ended {
369 StreamInner::end(self, None);
370 }
371 }
372 Ok(())
373 }
374
375 fn signed(&self, fields: &StreamFields) -> Result<Vec<u8>, LinkError> {
377 let signed = if self.caller {
378 frame::sign_caller_stream(fields, &self.open, &self.link.key)?
379 } else {
380 frame::sign_provider_stream(fields, &self.open, &self.link.key)?
381 };
382 cbor::encode(&signed)
383 .map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))
384 }
385
386 async fn abort(self: &Arc<Self>, code: &str, message: &str) -> Result<(), LinkError> {
387 let sent = self
388 .send(
389 |seq| StreamFields::Error {
390 seq,
391 code: code.to_string(),
392 message: message.to_string(),
393 },
394 true,
395 )
396 .await;
397 StreamInner::end(
398 self,
399 Some(LinkError::Stream {
400 code: code.to_string(),
401 message: message.to_string(),
402 relay: false,
403 }),
404 );
405 sent
406 }
407
408 fn deliver(&self, event: StreamEvent, size: usize) -> bool {
411 {
412 let mut side = self.side();
413 if side.held + size > STREAM_INBOX {
414 return false;
415 }
416 if let Some(budget) = &*self.budget() {
417 if !budget.admission.charge_inbox(budget.caller, size) {
418 return false;
419 }
420 }
421 side.inbox.push_back((event, size));
422 side.held += size;
423 }
424 self.notify.notify_one();
425 true
426 }
427
428 fn release_inbox(&self, size: usize) {
429 if let Some(budget) = &*self.budget() {
430 budget.admission.release_inbox(budget.caller, size);
431 }
432 }
433
434 fn peer_finished(self: &Arc<Self>, err: Option<LinkError>) {
437 self.side().peer_ended = true;
438 StreamInner::end(self, err);
439 }
440
441 async fn fail(self: &Arc<Self>, code: &str, cause: Option<String>) {
445 let message = cause
446 .as_deref()
447 .map(bounded_detail)
448 .unwrap_or("")
449 .to_string();
450 let _ = self.abort(code, &message).await;
451 }
452
453 pub(super) fn end(this: &Arc<StreamInner>, err: Option<LinkError>) {
457 let graceful = {
458 let mut side = this.side();
459 if side.ended {
460 return;
461 }
462 side.ended = true;
463 if side.err.is_none() {
464 side.err = err;
465 }
466 let graceful = side.sent_end;
467 side.sent_end = true;
468 graceful
469 };
470 if !graceful {
471 let released = this.clone();
472 if let Ok(runtime) = tokio::runtime::Handle::try_current() {
473 runtime.spawn(async move { released.writer.reset().await });
474 }
475 }
476 if let Some(budget) = this.budget().take() {
477 let held = std::mem::take(&mut this.side().held);
478 budget.admission.release_inbox(budget.caller, held);
479 drop(budget.place);
480 }
481 this.link
482 .lock()
483 .streams
484 .retain(|w| w.strong_count() > 0 && !std::ptr::eq(w.as_ptr(), Arc::as_ptr(this)));
485 let _ = this.done_tx.send_replace(true);
486 this.notify.notify_waiters();
487 this.notify.notify_one();
488 }
489}
490
491fn hold_stream(inner: &Inner, s: &Arc<StreamInner>) -> bool {
494 let mut state = inner.lock();
495 if state.ended.is_some() {
496 return false;
497 }
498 state.streams.push(Arc::downgrade(s));
499 true
500}
501
502fn abandon(mut send: quinn::SendStream, mut recv: quinn::RecvStream) {
504 let _ = send.reset(0u32.into());
505 let _ = recv.stop(0u32.into());
506}
507
508impl Link {
509 pub async fn open_stream(&self, c: StreamCall) -> Result<Stream, LinkError> {
513 let inner = &self.inner;
514 stated(&c.target, &inner.station.node_id, &c.seal)?;
515 let deadline = if c.deadline.is_zero() {
516 DEFAULT_STREAM_DEADLINE
517 } else {
518 c.deadline
519 };
520 let mut request_id = [0u8; 16];
521 aws_lc_rs::rand::fill(&mut request_id)
522 .map_err(|_| LinkError::Io("no randomness".into()))?;
523 let deadline = (now_ms() + deadline.as_millis() as i64) as u64;
524 let (sealed, sealing) = match &c.seal {
525 Some(Seal::To(key)) => {
526 let (sealed, s) = sealed_request(
527 inner.profile,
528 key,
529 crate::seal::FRAME_STREAM_OPEN,
530 c.realm,
531 &c.procedure,
532 inner.self_id,
533 c.target,
534 request_id,
535 deadline,
536 &c.payload,
537 )?;
538 (Some(sealed), Some(StreamSeal::caller(&s)))
539 }
540 _ => (None, None),
541 };
542 let signed = frame::sign_stream_open(
543 &RequestSpec {
544 request_id,
545 realm: c.realm,
546 procedure: c.procedure,
547 target: c.target,
548 deadline,
549 payload: c.payload,
550 sealed,
551 mode: Some(c.mode),
552 token: c.token,
553 proofs: c.proofs,
554 source_route: None,
555 retry_budget: None,
556 },
557 &inner.key,
558 )?;
559 let encoded = cbor::encode(&signed)
560 .map_err(|e| LinkError::Frame(frame::FrameError::Payload(e.to_string())))?;
561 if encoded.len() > STREAM_OPEN_BYTES {
562 return Err(LinkError::StreamOpenTooLarge(encoded.len()));
563 }
564 let open = frame::verify_request(&signed, inner.profile)?;
565 let state = frame::open_stream(&open)?;
566 let (send, recv) = inner
567 .connection
568 .open_bi()
569 .await
570 .map_err(|e| LinkError::Io(format!("open a stream: {e}")))?;
571 let s = StreamInner::new(inner.clone(), send, open, true, sealing);
572 let held = hold_stream(inner, &s);
573 let written = match held {
574 true => s.writer.write(&encoded, STREAM_OPEN_BYTES).await,
575 false => Err(inner.lock().ended.clone().unwrap_or(LinkError::Closed)),
576 };
577 if let Err(e) = written {
578 StreamInner::end(&s, Some(e.clone()));
579 let mut recv = recv;
580 let _ = recv.stop(0u32.into());
581 return Err(e);
582 }
583 tokio::spawn(read(s.clone(), recv, state));
584 Ok(Stream { inner: s })
585 }
586}
587
588pub(super) async fn accept_streams(link: Weak<Inner>) {
590 let Some(connection) = link.upgrade().map(|l| l.connection.clone()) else {
591 return;
592 };
593 while let Ok((send, recv)) = connection.accept_bi().await {
594 tokio::spawn(incoming(link.clone(), send, recv));
595 }
596}
597
598async fn incoming(link: Weak<Inner>, send: quinn::SendStream, mut recv: quinn::RecvStream) {
604 let Some(inner) = link.upgrade() else { return };
605 let payload = match tokio::time::timeout(
606 STREAM_OPEN_WAIT,
607 read_frame(&mut recv, STREAM_OPEN_BYTES),
608 )
609 .await
610 {
611 Ok(Ok(payload)) => payload,
612 _ => {
613 inner.count("stream_open_unread");
614 abandon(send, recv);
615 return;
616 }
617 };
618 let v = match cbor::decode(&payload) {
619 Ok(v) if frame_type_of(&v) == "stream_open" => v,
620 _ => {
621 inner.count("stream_open_malformed");
622 abandon(send, recv);
623 return;
624 }
625 };
626 let Ok(open) = frame::verify_request(&v, inner.profile) else {
627 inner.count("stream_open_unverified");
628 abandon(send, recv);
629 return;
630 };
631 if open.target != inner.self_id {
632 inner.count("stream_for_another_node");
633 abandon(send, recv);
634 return;
635 }
636 let Ok(state) = frame::open_stream(&open) else {
637 abandon(send, recv);
638 return;
639 };
640 let offer = inner
641 .lock()
642 .served
643 .get(&(open.realm, open.procedure.clone()))
644 .map(|s| s.offer.clone());
645 let refuse_clear = |code: &str, message: &str, send: quinn::SendStream, recv| {
646 let s = StreamInner::new(inner.clone(), send, open.clone(), false, None);
647 let (code, message) = (code.to_string(), message.to_string());
648 async move { refuse(&s, &code, &message, recv).await }
649 };
650 let place = match admit_stream(&inner, &open) {
653 Ok(place) => place,
654 Err(code) => return refuse_clear(code, "", send, recv).await,
655 };
656 let (session_open, sealing) = match &open.sealed {
657 Some(_) => match opened_request(inner.keyring.as_deref(), &open) {
658 Ok((payload, sealed)) => (
659 VerifiedRequest {
660 payload,
661 ..open.clone()
662 },
663 Some(StreamSeal::provider(&sealed)),
664 ),
665 Err(detail) => return refuse_clear(CODE_SEALED_REFUSED, &detail, send, recv).await,
666 },
667 None if offer
668 .as_ref()
669 .is_some_and(|o| !clear_allowed(o.confidential, inner.keyed_since(o), now_ms())) =>
670 {
671 let message = "this procedure takes sealed opens only";
672 return refuse_clear(CODE_SEALED_REQUIRED, message, send, recv).await;
673 }
674 None => (open.clone(), None),
675 };
676 let s = StreamInner::new(inner.clone(), send, session_open, false, sealing);
678 let Some(offer) = offer.and_then(|o| o.stream) else {
679 return refuse(&s, CODE_STREAM_NOT_FOUND, "", recv).await;
680 };
681 if Some(offer.mode) != open.mode {
682 return refuse(&s, CODE_MODE_MISMATCH, "", recv).await;
683 }
684 *s.budget() = Some(Budget {
685 admission: inner.admission.clone(),
686 caller: open.caller,
687 place: Some(place),
688 });
689 if !hold_stream(&inner, &s) {
690 StreamInner::end(&s, Some(LinkError::Closed));
691 let _ = recv.stop(0u32.into());
692 return;
693 }
694 tokio::spawn(read(s.clone(), recv, state));
695 tokio::spawn(serve(s, offer));
696}
697
698fn admit_stream(inner: &Inner, open: &VerifiedRequest) -> Result<SessionPlace, &'static str> {
702 match inner.admission.admit(open, &inner.share, now_ms()) {
703 Verdict::Refused(code) => return Err(code),
704 Verdict::Copy(_) => return Err(CODE_REQUEST_COPY),
705 Verdict::New => {}
706 }
707 inner
708 .admission
709 .open_session(open.caller)
710 .ok_or(CODE_TOO_MANY_SESSIONS)
711}
712
713async fn refuse(s: &Arc<StreamInner>, code: &str, message: &str, mut recv: quinn::RecvStream) {
716 s.link.count(&format!("stream_refused_{code}"));
717 let _ = s.abort(code, message).await;
718 let _ = recv.stop(0u32.into());
719}
720
721async fn serve(s: Arc<StreamInner>, offer: StreamOffer) {
725 let stream = Stream { inner: s.clone() };
726 let mut running = tokio::spawn((offer.handler)(stream.clone()));
727 let mut done = s.done_tx.subscribe();
728 let outcome = tokio::select! {
729 outcome = &mut running => outcome,
730 _ = done.wait_for(|ended| *ended) => {
731 running.abort();
732 return;
733 }
734 };
735 match outcome {
736 Ok(Ok(())) => {
737 let _ = stream.close().await;
738 }
739 Ok(Err(e)) => {
740 let _ = stream
741 .abort(CODE_STREAM_HANDLER_ERROR, bounded_detail(&e))
742 .await;
743 }
744 Err(panicked) => {
745 let _ = stream
746 .abort(
747 CODE_STREAM_HANDLER_ERROR,
748 bounded_detail(&panicked.to_string()),
749 )
750 .await;
751 }
752 }
753}
754
755async fn read(s: Arc<StreamInner>, mut recv: quinn::RecvStream, mut state: StreamState) {
759 let mut done = s.done_tx.subscribe();
760 loop {
761 let payload = tokio::select! {
762 _ = done.wait_for(|ended| *ended) => return,
763 payload = read_frame(&mut recv, MAX_FRAME_BYTES) => payload,
764 };
765 let payload = match payload {
766 Ok(payload) => payload,
767 Err(e) => return read_ended(&s, e),
768 };
769 match received(&s, &payload, &state).await {
770 Some(next) => state = next,
771 None => return,
772 }
773 }
774}
775
776fn read_ended(s: &Arc<StreamInner>, e: LinkError) {
779 if s.side().peer_ended {
780 return;
781 }
782 let err = s.link.lock().ended.clone().unwrap_or(e);
783 StreamInner::end(s, Some(err));
784}
785
786async fn received(
789 s: &Arc<StreamInner>,
790 payload: &[u8],
791 state: &StreamState,
792) -> Option<StreamState> {
793 let v = match cbor::decode(payload) {
794 Ok(v) => v,
795 Err(e) => {
796 s.fail("malformed_frame", Some(e.to_string())).await;
797 return None;
798 }
799 };
800 if s.caller && v.get("relay_error").is_some() {
801 match frame::verify_relay_error(&v, &s.open, s.link.profile, &s.link.station.node_id) {
802 Ok(relayed) => s.peer_finished(Some(LinkError::Stream {
803 code: relayed.code,
804 message: String::new(),
805 relay: true,
806 })),
807 Err(e) => s.fail("malformed_frame", Some(e.to_string())).await,
808 }
809 return None;
810 }
811 let verified = if s.caller {
812 frame::verify_provider_stream(&v, state, s.link.profile)
813 } else {
814 frame::verify_caller_stream(&v, state, s.link.profile)
815 };
816 let (verified, next) = match verified {
817 Ok(verified) => verified,
818 Err(e) => {
819 s.fail("malformed_frame", Some(e.to_string())).await;
820 return None;
821 }
822 };
823 let size = payload.len();
824 let fields = match unsealed(verified, s.sealing.as_ref()) {
825 Ok(fields) => fields,
826 Err(e) => {
827 StreamInner::end(s, Some(e));
828 return None;
829 }
830 };
831 if s.caller && settles(&fields, s.sealing.is_some()) {
834 s.side().settled = true;
835 }
836 match fields {
837 StreamFields::Error { code, message, .. } => {
838 s.peer_finished(Some(LinkError::Stream {
839 code,
840 message,
841 relay: false,
842 }));
843 None
844 }
845 StreamFields::Reply { payload, .. } => {
846 s.deliver(StreamEvent::Reply { payload }, size);
847 s.peer_finished(None);
848 None
849 }
850 StreamFields::End { role, .. } => {
851 s.deliver(StreamEvent::End { role }, size);
852 if role == StreamRole::Both {
853 s.peer_finished(None);
854 return None;
855 }
856 let mine = {
857 let mut side = s.side();
858 side.peer_ended = true;
859 side.sent_end
860 };
861 if mine {
862 StreamInner::end(s, None);
863 }
864 None
865 }
866 StreamFields::Data { encoding, body, .. } => {
867 if !s.deliver(StreamEvent::Data { encoding, body }, size) {
868 s.fail("resource_exhausted", None).await;
869 return None;
870 }
871 Some(next)
872 }
873 StreamFields::SealedData { .. }
875 | StreamFields::SealedError { .. }
876 | StreamFields::SealedReply { .. } => {
877 StreamInner::end(s, Some(LinkError::ClearAnswerToSealed));
878 None
879 }
880 }
881}
882
883fn settles(fields: &StreamFields, sealed: bool) -> bool {
889 match fields {
890 StreamFields::Data { .. } | StreamFields::Reply { .. } => true,
891 StreamFields::End { .. } => !sealed,
892 _ => false,
893 }
894}