1use crate::p2::bindings::http::types;
4use crate::{Error, FieldMap};
5use bytes::Bytes;
6use http_body::{Body, Frame};
7use http_body_util::BodyExt;
8use http_body_util::combinators::UnsyncBoxBody;
9use std::future::Future;
10use std::mem;
11use std::task::{Context, Poll};
12use std::{pin::Pin, sync::Arc};
13use tokio::sync::{mpsc, oneshot};
14use wasmtime::format_err;
15use wasmtime_wasi::p2::{InputStream, OutputStream, Pollable, StreamError};
16use wasmtime_wasi::runtime::{AbortOnDropJoinHandle, poll_noop};
17
18pub type HyperIncomingBody = UnsyncBoxBody<Bytes, Error>;
20
21pub type HyperOutgoingBody = UnsyncBoxBody<Bytes, Error>;
23
24#[derive(Debug)]
26pub struct HostIncomingBody {
27 body: IncomingBodyState,
28 worker: Option<AbortOnDropJoinHandle<()>>,
32}
33
34impl HostIncomingBody {
35 pub fn new(body: HyperIncomingBody) -> HostIncomingBody {
37 HostIncomingBody {
38 body: IncomingBodyState::Start(body),
39 worker: None,
40 }
41 }
42
43 pub fn retain_worker(&mut self, worker: AbortOnDropJoinHandle<()>) {
45 assert!(self.worker.is_none());
46 self.worker = Some(worker);
47 }
48
49 pub fn take_stream(&mut self) -> Option<HostIncomingBodyStream> {
51 match &mut self.body {
52 IncomingBodyState::Start(_) => {}
53 IncomingBodyState::InBodyStream(_) => return None,
54 }
55 let (tx, rx) = oneshot::channel();
56 let body = match mem::replace(&mut self.body, IncomingBodyState::InBodyStream(rx)) {
57 IncomingBodyState::Start(b) => b,
58 IncomingBodyState::InBodyStream(_) => unreachable!(),
59 };
60 Some(HostIncomingBodyStream {
61 state: IncomingBodyStreamState::Open { body, tx },
62 buffer: Bytes::new(),
63 error: None,
64 })
65 }
66
67 pub fn into_future_trailers(self) -> HostFutureTrailers {
69 HostFutureTrailers::Waiting(self)
70 }
71}
72
73#[derive(Debug)]
75enum IncomingBodyState {
76 Start(HyperIncomingBody),
79
80 InBodyStream(oneshot::Receiver<StreamEnd>),
84}
85
86#[derive(Debug)]
89enum StreamEnd {
90 Remaining(HyperIncomingBody),
93
94 Trailers(Option<http::HeaderMap>),
97}
98
99#[derive(Debug)]
102pub struct HostIncomingBodyStream {
103 state: IncomingBodyStreamState,
104 buffer: Bytes,
105 error: Option<Error>,
106}
107
108impl HostIncomingBodyStream {
109 fn record_frame(&mut self, frame: Option<Result<Frame<Bytes>, Error>>) {
110 match frame {
111 Some(Ok(frame)) => match frame.into_data() {
112 Ok(bytes) => {
115 assert!(self.buffer.is_empty());
116 self.buffer = bytes;
117 }
118
119 Err(trailers) => {
123 let trailers = trailers.into_trailers().unwrap();
124 let tx = match mem::replace(&mut self.state, IncomingBodyStreamState::Closed) {
125 IncomingBodyStreamState::Open { body: _, tx } => tx,
126 IncomingBodyStreamState::Closed => unreachable!(),
127 };
128
129 let _ = tx.send(StreamEnd::Trailers(Some(trailers)));
132 }
133 },
134
135 Some(Err(e)) => {
139 self.error = Some(e);
140 self.state = IncomingBodyStreamState::Closed;
141 }
142
143 None => {
147 self.state = IncomingBodyStreamState::Closed;
148 }
149 }
150 }
151}
152
153#[derive(Debug)]
154enum IncomingBodyStreamState {
155 Open {
163 body: HyperIncomingBody,
164 tx: oneshot::Sender<StreamEnd>,
165 },
166
167 Closed,
170}
171
172#[async_trait::async_trait]
173impl InputStream for HostIncomingBodyStream {
174 fn read(&mut self, size: usize) -> Result<Bytes, StreamError> {
175 loop {
176 if !self.buffer.is_empty() {
178 let len = size.min(self.buffer.len());
179 let chunk = self.buffer.split_to(len);
180 return Ok(chunk);
181 }
182
183 if let Some(e) = self.error.take() {
184 return Err(StreamError::LastOperationFailed(e.into()));
185 }
186
187 let body = match &mut self.state {
193 IncomingBodyStreamState::Open { body, .. } => body,
194 IncomingBodyStreamState::Closed => return Err(StreamError::Closed),
195 };
196
197 let future = body.frame();
198 futures::pin_mut!(future);
199 match poll_noop(future) {
200 Some(result) => {
201 self.record_frame(result);
202 }
203 None => return Ok(Bytes::new()),
204 }
205 }
206 }
207}
208
209#[async_trait::async_trait]
210impl Pollable for HostIncomingBodyStream {
211 async fn ready(&mut self) {
212 if !self.buffer.is_empty() || self.error.is_some() {
213 return;
214 }
215
216 if let IncomingBodyStreamState::Open { body, .. } = &mut self.state {
217 let frame = body.frame().await;
218 self.record_frame(frame);
219 }
220 }
221}
222
223impl Drop for HostIncomingBodyStream {
224 fn drop(&mut self) {
225 let prev = mem::replace(&mut self.state, IncomingBodyStreamState::Closed);
231 if let IncomingBodyStreamState::Open { body, tx } = prev {
232 let _ = tx.send(StreamEnd::Remaining(body));
233 }
234 }
235}
236
237#[derive(Debug)]
239pub enum HostFutureTrailers {
240 Waiting(HostIncomingBody),
254
255 Done(Result<Option<http::HeaderMap>, Error>),
260
261 Consumed,
263}
264
265#[async_trait::async_trait]
266impl Pollable for HostFutureTrailers {
267 async fn ready(&mut self) {
268 let body = match self {
269 HostFutureTrailers::Waiting(body) => body,
270 HostFutureTrailers::Done(_) => return,
271 HostFutureTrailers::Consumed => return,
272 };
273
274 if let IncomingBodyState::InBodyStream(rx) = &mut body.body {
277 match rx.await {
278 Ok(StreamEnd::Trailers(Some(t))) => {
281 *self = Self::Done(Ok(Some(t)));
282 }
283 Ok(StreamEnd::Remaining(b)) => body.body = IncomingBodyState::Start(b),
286
287 Ok(StreamEnd::Trailers(None)) | Err(_) => {
289 *self = HostFutureTrailers::Done(Ok(None));
290 }
291 }
292 }
293
294 let body = match self {
297 HostFutureTrailers::Waiting(body) => body,
298 HostFutureTrailers::Done(_) => return,
299 HostFutureTrailers::Consumed => return,
300 };
301 let hyper_body = match &mut body.body {
302 IncomingBodyState::Start(body) => body,
303 IncomingBodyState::InBodyStream(_) => unreachable!(),
304 };
305 let result = loop {
306 match hyper_body.frame().await {
307 None => break Ok(None),
308 Some(Err(e)) => break Err(e),
309 Some(Ok(frame)) => {
310 if let Ok(header_map) = frame.into_trailers() {
313 break Ok(Some(header_map));
314 }
315 }
316 }
317 };
318 *self = HostFutureTrailers::Done(result);
319 }
320}
321
322#[derive(Debug, Clone)]
323struct WrittenState {
324 expected: u64,
325 written: Arc<std::sync::atomic::AtomicU64>,
326}
327
328impl WrittenState {
329 fn new(expected_size: u64) -> Self {
330 Self {
331 expected: expected_size,
332 written: Arc::new(std::sync::atomic::AtomicU64::new(0)),
333 }
334 }
335
336 fn written(&self) -> u64 {
338 self.written.load(std::sync::atomic::Ordering::Relaxed)
339 }
340
341 fn update(&self, len: usize) -> bool {
344 let len = len as u64;
345 let old = self
346 .written
347 .fetch_add(len, std::sync::atomic::Ordering::Relaxed);
348 old + len <= self.expected
349 }
350}
351
352pub struct HostOutgoingBody {
354 body_output_stream: Option<Box<dyn OutputStream>>,
356 context: StreamContext,
357 written: Option<WrittenState>,
358 finish_sender: Option<tokio::sync::oneshot::Sender<FinishMessage>>,
359}
360
361impl HostOutgoingBody {
362 pub fn new(
364 context: StreamContext,
365 size: Option<u64>,
366 buffer_chunks: usize,
367 chunk_size: usize,
368 ) -> (Self, HyperOutgoingBody) {
369 assert!(buffer_chunks >= 1);
370
371 let written = size.map(WrittenState::new);
372
373 use tokio::sync::oneshot::error::RecvError;
374 struct BodyImpl {
375 body_receiver: mpsc::Receiver<Bytes>,
376 finish_receiver: Option<oneshot::Receiver<FinishMessage>>,
377 }
378 impl Body for BodyImpl {
379 type Data = Bytes;
380 type Error = Error;
381 fn poll_frame(
382 mut self: Pin<&mut Self>,
383 cx: &mut Context<'_>,
384 ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
385 match self.as_mut().body_receiver.poll_recv(cx) {
386 Poll::Pending => Poll::Pending,
387 Poll::Ready(Some(frame)) => Poll::Ready(Some(Ok(Frame::data(frame)))),
388
389 Poll::Ready(None) => {
391 if let Some(mut finish_receiver) = self.as_mut().finish_receiver.take() {
392 match Pin::new(&mut finish_receiver).poll(cx) {
393 Poll::Pending => {
394 self.as_mut().finish_receiver = Some(finish_receiver);
395 Poll::Pending
396 }
397 Poll::Ready(Ok(message)) => match message {
398 FinishMessage::Finished => Poll::Ready(None),
399 FinishMessage::Trailers(trailers) => {
400 Poll::Ready(Some(Ok(Frame::trailers(trailers))))
401 }
402 FinishMessage::Abort => {
403 Poll::Ready(Some(Err(Error::HttpProtocolError)))
404 }
405 },
406 Poll::Ready(Err(RecvError { .. })) => Poll::Ready(None),
407 }
408 } else {
409 Poll::Ready(None)
410 }
411 }
412 }
413 }
414 }
415
416 let (body_sender, body_receiver) = mpsc::channel(buffer_chunks + 1);
418 let (finish_sender, finish_receiver) = oneshot::channel();
419 let body_impl = BodyImpl {
420 body_receiver,
421 finish_receiver: Some(finish_receiver),
422 }
423 .boxed_unsync();
424
425 let output_stream = BodyWriteStream::new(context, chunk_size, body_sender, written.clone());
426
427 (
428 Self {
429 body_output_stream: Some(Box::new(output_stream)),
430 context,
431 written,
432 finish_sender: Some(finish_sender),
433 },
434 body_impl,
435 )
436 }
437
438 pub fn take_output_stream(&mut self) -> Option<Box<dyn OutputStream>> {
440 self.body_output_stream.take()
441 }
442
443 pub fn finish(mut self, trailers: Option<FieldMap>) -> Result<(), types::ErrorCode> {
445 drop(self.body_output_stream);
448
449 let sender = self
450 .finish_sender
451 .take()
452 .expect("outgoing-body trailer_sender consumed by a non-owning function");
453
454 if let Some(w) = self.written {
455 let written = w.written();
456 if written != w.expected {
457 let _ = sender.send(FinishMessage::Abort);
458 return Err(self.context.as_body_size_error(written));
459 }
460 }
461
462 let message = if let Some(ts) = trailers {
463 FinishMessage::Trailers(ts.into())
464 } else {
465 FinishMessage::Finished
466 };
467
468 let _ = sender.send(message);
470
471 Ok(())
472 }
473
474 pub fn abort(mut self) {
476 drop(self.body_output_stream);
479
480 let sender = self
481 .finish_sender
482 .take()
483 .expect("outgoing-body trailer_sender consumed by a non-owning function");
484
485 let _ = sender.send(FinishMessage::Abort);
486 }
487}
488
489#[derive(Debug)]
491enum FinishMessage {
492 Finished,
493 Trailers(hyper::HeaderMap),
494 Abort,
495}
496
497#[derive(Clone, Copy, Debug, Eq, PartialEq)]
499pub enum StreamContext {
500 Request,
502 Response,
504}
505
506impl StreamContext {
507 pub fn as_body_size_error(&self, size: u64) -> types::ErrorCode {
509 match self {
510 StreamContext::Request => types::ErrorCode::HttpRequestBodySize(Some(size)),
511 StreamContext::Response => types::ErrorCode::HttpResponseBodySize(Some(size)),
512 }
513 }
514}
515
516#[derive(Debug)]
518struct BodyWriteStream {
519 context: StreamContext,
520 writer: mpsc::Sender<Bytes>,
521 write_budget: usize,
522 written: Option<WrittenState>,
523}
524
525impl BodyWriteStream {
526 fn new(
528 context: StreamContext,
529 write_budget: usize,
530 writer: mpsc::Sender<Bytes>,
531 written: Option<WrittenState>,
532 ) -> Self {
533 assert!(writer.max_capacity() >= 1);
535 BodyWriteStream {
536 context,
537 writer,
538 write_budget,
539 written,
540 }
541 }
542}
543
544#[async_trait::async_trait]
545impl OutputStream for BodyWriteStream {
546 fn write(&mut self, bytes: Bytes) -> Result<(), StreamError> {
547 let len = bytes.len();
548 match self.writer.try_send(bytes) {
549 Ok(()) => {
552 if let Some(written) = self.written.as_ref() {
553 if !written.update(len) {
554 let total = written.written();
555 return Err(StreamError::LastOperationFailed(format_err!(
556 self.context.as_body_size_error(total)
557 )));
558 }
559 }
560
561 Ok(())
562 }
563
564 Err(mpsc::error::TrySendError::Full(_)) => {
568 Err(StreamError::Trap(format_err!("write exceeded budget")))
569 }
570
571 Err(mpsc::error::TrySendError::Closed(_)) => Err(StreamError::Closed),
573 }
574 }
575
576 fn flush(&mut self) -> Result<(), StreamError> {
577 if self.writer.is_closed() {
580 Err(StreamError::Closed)
581 } else {
582 Ok(())
583 }
584 }
585
586 fn check_write(&mut self) -> Result<usize, StreamError> {
587 if self.writer.is_closed() {
588 Err(StreamError::Closed)
589 } else if self.writer.capacity() == 0 {
590 Ok(0)
598 } else {
599 Ok(self.write_budget)
600 }
601 }
602}
603
604#[async_trait::async_trait]
605impl Pollable for BodyWriteStream {
606 async fn ready(&mut self) {
607 let _ = self.writer.reserve().await;
611 }
612}