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