1 use {
2     super::util::{config, make_component},
3     component_async_tests::{
4         Ctx, closed_streams,
5         util::{OneshotConsumer, OneshotProducer, PipeConsumer, PipeProducer, yield_times},
6     },
7     futures::{
8         FutureExt, Sink, SinkExt, Stream, StreamExt,
9         channel::{mpsc, oneshot},
10         future,
11     },
12     std::{
13         mem,
14         ops::DerefMut,
15         pin::Pin,
16         sync::{Arc, Mutex},
17         task::{self, Context, Poll},
18     },
19     wasmtime::{
20         Engine, Result, Store, StoreContextMut,
21         component::{
22             Destination, FutureReader, Lift, Linker, ResourceTable, Source, StreamConsumer,
23             StreamProducer, StreamReader, StreamResult, VecBuffer,
24         },
25     },
26     wasmtime_wasi::WasiCtxBuilder,
27 };
28 
29 pub struct DirectPipeProducer<S>(S);
30 
31 impl<D, S: Stream<Item = u8> + Send + 'static> StreamProducer<D> for DirectPipeProducer<S> {
32     type Item = u8;
33     type Buffer = Option<u8>;
34 
poll_produce<'a>( self: Pin<&mut Self>, cx: &mut Context<'_>, store: StoreContextMut<D>, destination: Destination<'a, Self::Item, Self::Buffer>, finish: bool, ) -> Poll<Result<StreamResult>>35     fn poll_produce<'a>(
36         self: Pin<&mut Self>,
37         cx: &mut Context<'_>,
38         store: StoreContextMut<D>,
39         destination: Destination<'a, Self::Item, Self::Buffer>,
40         finish: bool,
41     ) -> Poll<Result<StreamResult>> {
42         // SAFETY: This is a standard pin-projection, and we never move
43         // out of `self`.
44         let stream = unsafe { self.map_unchecked_mut(|v| &mut v.0) };
45 
46         match stream.poll_next(cx) {
47             Poll::Pending => {
48                 if finish {
49                     Poll::Ready(Ok(StreamResult::Cancelled))
50                 } else {
51                     Poll::Pending
52                 }
53             }
54             Poll::Ready(Some(item)) => {
55                 let mut destination = destination.as_direct(store, 1);
56                 destination.remaining()[0] = item;
57                 destination.mark_written(1);
58                 Poll::Ready(Ok(StreamResult::Completed))
59             }
60             Poll::Ready(None) => Poll::Ready(Ok(StreamResult::Dropped)),
61         }
62     }
63 }
64 
65 pub struct DirectPipeConsumer<S>(S);
66 
67 impl<D, S: Sink<u8, Error: std::error::Error + Send + Sync> + Send + 'static> StreamConsumer<D>
68     for DirectPipeConsumer<S>
69 {
70     type Item = u8;
71 
poll_consume( self: Pin<&mut Self>, cx: &mut Context<'_>, store: StoreContextMut<D>, source: Source<Self::Item>, finish: bool, ) -> Poll<Result<StreamResult>>72     fn poll_consume(
73         self: Pin<&mut Self>,
74         cx: &mut Context<'_>,
75         store: StoreContextMut<D>,
76         source: Source<Self::Item>,
77         finish: bool,
78     ) -> Poll<Result<StreamResult>> {
79         // SAFETY: This is a standard pin-projection, and we never move
80         // out of `self`.
81         let mut sink = unsafe { self.map_unchecked_mut(|v| &mut v.0) };
82 
83         let on_pending = || {
84             if finish {
85                 Poll::Ready(Ok(StreamResult::Cancelled))
86             } else {
87                 Poll::Pending
88             }
89         };
90 
91         match sink.as_mut().poll_flush(cx) {
92             Poll::Pending => on_pending(),
93             Poll::Ready(result) => {
94                 result?;
95                 match sink.as_mut().poll_ready(cx) {
96                     Poll::Pending => on_pending(),
97                     Poll::Ready(result) => {
98                         result?;
99                         let mut source = source.as_direct(store);
100                         let item = source.remaining()[0];
101                         source.mark_read(1);
102                         sink.start_send(item)?;
103                         Poll::Ready(Ok(StreamResult::Completed))
104                     }
105                 }
106             }
107         }
108     }
109 }
110 
111 #[tokio::test]
async_closed_streams() -> Result<()>112 pub async fn async_closed_streams() -> Result<()> {
113     let engine = Engine::new(&config())?;
114 
115     let mut store = Store::new(
116         &engine,
117         Ctx {
118             wasi: WasiCtxBuilder::new().inherit_stdio().build(),
119             table: ResourceTable::default(),
120             continue_: false,
121         },
122     );
123 
124     let mut linker = Linker::new(&engine);
125 
126     wasmtime_wasi::p2::add_to_linker_async(&mut linker)?;
127 
128     let component = make_component(
129         &engine,
130         &[test_programs_artifacts::ASYNC_CLOSED_STREAMS_COMPONENT],
131     )
132     .await?;
133 
134     let instance = linker.instantiate_async(&mut store, &component).await?;
135 
136     let values = vec![42_u8, 43, 44];
137 
138     let value = 42_u8;
139 
140     // First, test stream host->host
141     for direct_producer in [true, false] {
142         for direct_consumer in [true, false] {
143             let (mut input_tx, input_rx) = mpsc::channel(1);
144             let (output_tx, mut output_rx) = mpsc::channel(1);
145             let reader = if direct_producer {
146                 StreamReader::new(&mut store, DirectPipeProducer(input_rx))?
147             } else {
148                 StreamReader::new(&mut store, PipeProducer::new(input_rx))?
149             };
150             if direct_consumer {
151                 reader.pipe(&mut store, DirectPipeConsumer(output_tx))?;
152             } else {
153                 reader.pipe(&mut store, PipeConsumer::new(output_tx))?;
154             }
155 
156             store
157                 .run_concurrent(async |_| {
158                     let (a, b) = future::join(
159                         async {
160                             for &value in &values {
161                                 input_tx.send(value).await?;
162                             }
163                             drop(input_tx);
164                             wasmtime::error::Ok(())
165                         },
166                         async {
167                             for &value in &values {
168                                 assert_eq!(Some(value), output_rx.next().await);
169                             }
170                             assert!(output_rx.next().await.is_none());
171                             Ok(())
172                         },
173                     )
174                     .await;
175 
176                     a.and(b)
177                 })
178                 .await??;
179         }
180     }
181 
182     // Next, test futures host->host
183     {
184         let (input_tx, input_rx) = oneshot::channel();
185         let (output_tx, output_rx) = oneshot::channel();
186         FutureReader::new(&mut store, OneshotProducer::new(input_rx))?
187             .pipe(&mut store, OneshotConsumer::new(output_tx))?;
188 
189         store
190             .run_concurrent(async |_| {
191                 _ = input_tx.send(value);
192                 assert_eq!(value, output_rx.await?);
193                 wasmtime::error::Ok(())
194             })
195             .await??;
196     }
197 
198     // Next, test stream host->guest
199     {
200         let (mut tx, rx) = mpsc::channel(1);
201         let rx = StreamReader::new(&mut store, PipeProducer::new(rx))?;
202 
203         let closed_streams = closed_streams::bindings::ClosedStreams::new(&mut store, &instance)?;
204 
205         let values = values.clone();
206 
207         store
208             .run_concurrent(async move |accessor| {
209                 let (a, b) = future::join(
210                     async {
211                         for &value in &values {
212                             tx.send(value).await?;
213                         }
214                         drop(tx);
215                         Ok(())
216                     },
217                     closed_streams.local_local_closed().call_read_stream(
218                         accessor,
219                         rx,
220                         values.clone(),
221                     ),
222                 )
223                 .await;
224 
225                 a.and(b)
226             })
227             .await??;
228     }
229 
230     // Next, test futures host->guest
231     {
232         let (tx, rx) = oneshot::channel();
233         let rx = FutureReader::new(&mut store, OneshotProducer::new(rx))?;
234         let (_, rx_ignored) = oneshot::channel();
235         let rx_ignored = FutureReader::new(&mut store, OneshotProducer::new(rx_ignored))?;
236 
237         let closed_streams = closed_streams::bindings::ClosedStreams::new(&mut store, &instance)?;
238 
239         store
240             .run_concurrent(async move |accessor| {
241                 _ = tx.send(value);
242                 closed_streams
243                     .local_local_closed()
244                     .call_read_future(accessor, rx, value, rx_ignored)
245                     .await
246             })
247             .await??;
248     }
249 
250     Ok(())
251 }
252 
253 mod closed_stream {
254     wasmtime::component::bindgen!({
255         path: "wit",
256         world: "closed-stream-guest",
257         exports: { default: store | async },
258     });
259 }
260 
261 #[tokio::test]
async_closed_stream() -> Result<()>262 pub async fn async_closed_stream() -> Result<()> {
263     let engine = Engine::new(&config())?;
264 
265     let component = make_component(
266         &engine,
267         &[test_programs_artifacts::ASYNC_CLOSED_STREAM_COMPONENT],
268     )
269     .await?;
270 
271     let mut linker = Linker::new(&engine);
272 
273     wasmtime_wasi::p2::add_to_linker_async(&mut linker)?;
274 
275     let mut store = Store::new(
276         &engine,
277         Ctx {
278             wasi: WasiCtxBuilder::new().inherit_stdio().build(),
279             table: ResourceTable::default(),
280             continue_: false,
281         },
282     );
283 
284     let instance = linker.instantiate_async(&mut store, &component).await?;
285     let guest = closed_stream::ClosedStreamGuest::new(&mut store, &instance)?;
286     store
287         .run_concurrent(async move |accessor| {
288             let stream = guest.local_local_closed_stream().call_get(accessor).await?;
289 
290             let (tx, mut rx) = mpsc::channel(1);
291             accessor.with(move |store| stream.pipe(store, PipeConsumer::new(tx)))?;
292             assert!(rx.next().await.is_none());
293 
294             Ok(())
295         })
296         .await?
297 }
298 
299 struct VecProducer<T> {
300     source: Vec<T>,
301     maybe_yield: Pin<Box<dyn Future<Output = ()> + Send>>,
302 }
303 
304 impl<T> VecProducer<T> {
new(source: Vec<T>, delay: bool) -> Self305     fn new(source: Vec<T>, delay: bool) -> Self {
306         Self {
307             source,
308             maybe_yield: if delay {
309                 yield_times(5).boxed()
310             } else {
311                 async {}.boxed()
312             },
313         }
314     }
315 }
316 
317 impl<D, T: Lift + Unpin + 'static> StreamProducer<D> for VecProducer<T> {
318     type Item = T;
319     type Buffer = VecBuffer<T>;
320 
poll_produce( mut self: Pin<&mut Self>, cx: &mut Context<'_>, _: StoreContextMut<D>, mut destination: Destination<Self::Item, Self::Buffer>, _: bool, ) -> Poll<Result<StreamResult>>321     fn poll_produce(
322         mut self: Pin<&mut Self>,
323         cx: &mut Context<'_>,
324         _: StoreContextMut<D>,
325         mut destination: Destination<Self::Item, Self::Buffer>,
326         _: bool,
327     ) -> Poll<Result<StreamResult>> {
328         let maybe_yield = &mut self.as_mut().get_mut().maybe_yield;
329         task::ready!(maybe_yield.as_mut().poll(cx));
330         *maybe_yield = async {}.boxed();
331 
332         destination.set_buffer(mem::take(&mut self.get_mut().source).into());
333         Poll::Ready(Ok(StreamResult::Dropped))
334     }
335 }
336 
337 struct OneAtATime<T> {
338     destination: Arc<Mutex<Vec<T>>>,
339     maybe_yield: Pin<Box<dyn Future<Output = ()> + Send>>,
340 }
341 
342 impl<T> OneAtATime<T> {
new(destination: Arc<Mutex<Vec<T>>>, delay: bool) -> Self343     fn new(destination: Arc<Mutex<Vec<T>>>, delay: bool) -> Self {
344         Self {
345             destination,
346             maybe_yield: if delay {
347                 yield_times(5).boxed()
348             } else {
349                 async {}.boxed()
350             },
351         }
352     }
353 }
354 
355 impl<D, T: Lift + 'static> StreamConsumer<D> for OneAtATime<T> {
356     type Item = T;
357 
poll_consume( mut self: Pin<&mut Self>, cx: &mut Context<'_>, store: StoreContextMut<D>, mut source: Source<Self::Item>, _: bool, ) -> Poll<Result<StreamResult>>358     fn poll_consume(
359         mut self: Pin<&mut Self>,
360         cx: &mut Context<'_>,
361         store: StoreContextMut<D>,
362         mut source: Source<Self::Item>,
363         _: bool,
364     ) -> Poll<Result<StreamResult>> {
365         let maybe_yield = &mut self.as_mut().get_mut().maybe_yield;
366         task::ready!(maybe_yield.as_mut().poll(cx));
367         *maybe_yield = async {}.boxed();
368 
369         let value = &mut None;
370         source.read(store, value)?;
371         self.destination.lock().unwrap().push(value.take().unwrap());
372         Poll::Ready(Ok(StreamResult::Completed))
373     }
374 }
375 
376 mod short_reads {
377     wasmtime::component::bindgen!({
378         path: "wit",
379         world: "short-reads-guest",
380         exports: { default: async },
381     });
382 }
383 
384 #[tokio::test]
async_short_reads() -> Result<()>385 pub async fn async_short_reads() -> Result<()> {
386     test_async_short_reads(false).await
387 }
388 
389 #[tokio::test]
async_short_reads_with_delay() -> Result<()>390 async fn async_short_reads_with_delay() -> Result<()> {
391     test_async_short_reads(true).await
392 }
393 
test_async_short_reads(delay: bool) -> Result<()>394 async fn test_async_short_reads(delay: bool) -> Result<()> {
395     use short_reads::exports::local::local::short_reads::Thing;
396 
397     let engine = Engine::new(&config())?;
398 
399     let component = make_component(
400         &engine,
401         &[test_programs_artifacts::ASYNC_SHORT_READS_COMPONENT],
402     )
403     .await?;
404 
405     let mut linker = Linker::new(&engine);
406 
407     wasmtime_wasi::p2::add_to_linker_async(&mut linker)?;
408 
409     let mut store = Store::new(
410         &engine,
411         Ctx {
412             wasi: WasiCtxBuilder::new().inherit_stdio().build(),
413             table: ResourceTable::default(),
414             continue_: false,
415         },
416     );
417 
418     let guest =
419         short_reads::ShortReadsGuest::instantiate_async(&mut store, &component, &linker).await?;
420     let thing = guest.local_local_short_reads().thing();
421 
422     let strings = ["a", "b", "c", "d", "e"];
423     let mut things = Vec::with_capacity(strings.len());
424     for string in strings {
425         things.push(thing.call_constructor(&mut store, string).await?);
426     }
427 
428     store
429         .run_concurrent(async |store| {
430             let count = things.len();
431             let stream =
432                 store.with(|store| StreamReader::new(store, VecProducer::new(things, delay)))?;
433 
434             let stream = guest
435                 .local_local_short_reads()
436                 .call_short_reads(store, stream)
437                 .await?;
438 
439             let received_things = Arc::new(Mutex::new(Vec::<Thing>::with_capacity(count)));
440             // Read just one item at a time from the guest, forcing it to
441             // re-take ownership of any unwritten items.
442             store.with(|store| {
443                 stream.pipe(store, OneAtATime::new(received_things.clone(), delay))
444             })?;
445 
446             for i in 0.. {
447                 assert!(i < 1000);
448                 if count == received_things.lock().unwrap().len() {
449                     break;
450                 }
451                 tokio::task::yield_now().await;
452             }
453 
454             let mut received_strings = Vec::with_capacity(strings.len());
455             let received_things = mem::take(received_things.lock().unwrap().deref_mut());
456             for it in received_things {
457                 received_strings.push(thing.call_get(store, it).await?);
458             }
459 
460             assert_eq!(
461                 &strings[..],
462                 &received_strings
463                     .iter()
464                     .map(|s| s.as_str())
465                     .collect::<Vec<_>>()
466             );
467 
468             wasmtime::error::Ok(())
469         })
470         .await?
471 }
472