Skip to main content

wasmtime_wasi_tls/p3/
host.rs

1//! p3 host implementation for `wasi:tls`.
2
3use crate::p3::util::{AsyncReadProducer, AsyncWriteConsumer, Closed, Deferred, Shared, pipe};
4use crate::p3::{WasiTls, bindings};
5use crate::{BoxFutureTlsStream, Error, TlsStream, WasiTlsCtxView};
6use std::pin::Pin;
7use std::task::{Context, Poll};
8use tokio::{io::AsyncWriteExt as _, sync::oneshot};
9use wasmtime::StoreContextMut;
10use wasmtime::component::{
11    Access, Accessor, AccessorTask, FutureProducer, FutureReader, HasData, Resource, StreamReader,
12};
13
14/// Host-side state stored for `wasi:tls/client` `connector` resources.
15pub struct Connector {
16    connection: Shared<Deferred<Box<dyn TlsStream>>>,
17    send: Option<pipe::Writer>,
18    recv: Option<pipe::Reader>,
19}
20
21impl<'a> bindings::tls::client::Host for WasiTlsCtxView<'a> {}
22impl<'a> bindings::tls::types::Host for WasiTlsCtxView<'a> {}
23
24impl<'a> bindings::tls::types::HostError for WasiTlsCtxView<'a> {
25    fn to_debug_string(&mut self, this: Resource<Error>) -> wasmtime::Result<String> {
26        Ok(self.table.get(&this)?.to_string())
27    }
28
29    fn drop(&mut self, rep: Resource<Error>) -> wasmtime::Result<()> {
30        self.table.delete(rep)?;
31        Ok(())
32    }
33}
34
35impl<'a> bindings::tls::client::HostConnector for WasiTlsCtxView<'a> {
36    fn new(&mut self) -> wasmtime::Result<Resource<Connector>> {
37        Ok(self.table.push(Connector {
38            connection: Shared::new(Deferred::pending()),
39            send: None,
40            recv: None,
41        })?)
42    }
43
44    fn drop(&mut self, rep: Resource<Connector>) -> wasmtime::Result<()> {
45        self.table.delete(rep)?;
46        Ok(())
47    }
48}
49
50impl<T> bindings::tls::client::HostConnectorWithStore<T> for WasiTls {
51    fn send(
52        mut store: Access<'_, T, Self>,
53        this: Resource<Connector>,
54        mut cleartext: StreamReader<u8>,
55    ) -> wasmtime::Result<(StreamReader<u8>, FutureReader<Result<(), Resource<Error>>>)> {
56        let getter = store.getter();
57
58        {
59            let ctx = store.get();
60            let connector = ctx.table.get(&this)?;
61            if connector.send.is_some() {
62                cleartext.close(&mut store)?;
63
64                let err = Error::msg("send already configured");
65                let ciphertext = Closed(err.clone());
66                let result = ResultProducer::ready(getter, Err(err));
67
68                return Ok((
69                    StreamReader::new(&mut store, ciphertext)?,
70                    FutureReader::new(&mut store, result)?,
71                ));
72            }
73        }
74
75        let (ciphertext_reader, ciphertext_writer) = pipe::pipe();
76        let (ciphertext_result_tx, ciphertext_result_rx) = oneshot::channel();
77        let (cleartext_result_tx, cleartext_result_rx) = oneshot::channel();
78        let (send_result_tx, send_result_rx) = oneshot::channel();
79
80        let connection = {
81            let ctx = store.get();
82            let connector = ctx.table.get_mut(&this)?;
83            connector.send = Some(ciphertext_writer);
84            connector.connection.clone()
85        };
86
87        cleartext.pipe(
88            &mut store,
89            AsyncWriteConsumer::new(connection, cleartext_result_tx),
90        )?;
91
92        let ciphertext = AsyncReadProducer::new(ciphertext_reader, ciphertext_result_tx);
93        store.spawn(FnTask(async move || {
94            let cleartext_result = match cleartext_result_rx.await? {
95                Ok(mut inner) => inner.shutdown().await, // Drive the close_notify sequence
96                Err(e) => Err(e),
97            };
98            let ciphertext_result = ciphertext_result_rx.await?.map(drop);
99            let combined_result = cleartext_result
100                .and(ciphertext_result)
101                .map_err(|e| Error::from(e));
102            _ = send_result_tx.send(combined_result);
103            Ok(())
104        }))?;
105        let result = ResultProducer::new(getter, send_result_rx);
106
107        Ok((
108            StreamReader::new(&mut store, ciphertext)?,
109            FutureReader::new(&mut store, result)?,
110        ))
111    }
112
113    fn receive(
114        mut store: Access<'_, T, Self>,
115        this: Resource<Connector>,
116        mut ciphertext: StreamReader<u8>,
117    ) -> wasmtime::Result<(StreamReader<u8>, FutureReader<Result<(), Resource<Error>>>)> {
118        let getter = store.getter();
119
120        {
121            let ctx = store.get();
122            let connector = ctx.table.get(&this)?;
123            if connector.recv.is_some() {
124                ciphertext.close(&mut store)?;
125
126                let err = Error::msg("receive already configured");
127                let cleartext = Closed(err.clone());
128                let result = ResultProducer::ready(getter, Err(err));
129
130                return Ok((
131                    StreamReader::new(&mut store, cleartext)?,
132                    FutureReader::new(&mut store, result)?,
133                ));
134            }
135        }
136
137        let (ciphertext_reader, ciphertext_writer) = pipe::pipe();
138        let (ciphertext_result_tx, ciphertext_result_rx) = oneshot::channel();
139        let (cleartext_result_tx, cleartext_result_rx) = oneshot::channel();
140        let (recv_result_tx, recv_result_rx) = oneshot::channel();
141
142        let connection = {
143            let ctx = store.get();
144            let connector = ctx.table.get_mut(&this)?;
145            connector.recv = Some(ciphertext_reader);
146            connector.connection.clone()
147        };
148
149        ciphertext.pipe(
150            &mut store,
151            AsyncWriteConsumer::new(ciphertext_writer, ciphertext_result_tx),
152        )?;
153
154        let cleartext = AsyncReadProducer::new(connection, cleartext_result_tx);
155        store.spawn(FnTask(async move || {
156            let ciphertext_result = match ciphertext_result_rx.await? {
157                // Let the TLS implementation know the transport is closed.
158                // Most likely, `shutdown` will be entirely synchronous and
159                // complete immediately, but awaiting it anyway to adhere to
160                // the AsyncWrite contract:
161                Ok(mut inner) => inner.shutdown().await,
162                Err(e) => Err(e),
163            };
164            let cleartext_result = cleartext_result_rx.await?.map(drop);
165            let combined_result = cleartext_result
166                .and(ciphertext_result)
167                .map_err(|e| Error::from(e));
168            _ = recv_result_tx.send(combined_result);
169            Ok(())
170        }))?;
171        let result = ResultProducer::new(getter, recv_result_rx);
172
173        Ok((
174            StreamReader::new(&mut store, cleartext)?,
175            FutureReader::new(&mut store, result)?,
176        ))
177    }
178
179    async fn connect(
180        accessor: &Accessor<T, Self>,
181        this: Resource<Connector>,
182        server_name: String,
183    ) -> wasmtime::Result<Result<(), Resource<Error>>> {
184        fn connect_err(msg: &'static str) -> BoxFutureTlsStream {
185            Box::pin(async move { Err(Error::msg(msg)) })
186        }
187        let (fut, connection) = accessor.with(
188            move |mut access| -> wasmtime::Result<(BoxFutureTlsStream, _)> {
189                let WasiTlsCtxView { table, ctx } = access.get();
190                let connector = table.delete(this)?;
191                let connection = connector.connection;
192
193                let Some(ciphertext_writer) = connector.send else {
194                    return Ok((
195                        connect_err("send() must be called before connect()"),
196                        connection,
197                    ));
198                };
199                let Some(ciphertext_reader) = connector.recv else {
200                    return Ok((
201                        connect_err("receive() must be called before connect()"),
202                        connection,
203                    ));
204                };
205
206                let transport = Box::new(tokio::io::join(ciphertext_reader, ciphertext_writer));
207                let fut = ctx.provider.connect(server_name, transport);
208
209                Ok((fut, connection))
210            },
211        )?;
212
213        match fut.await {
214            Ok(tls_stream) => {
215                connection.lock().resolve(tls_stream);
216                Ok(Ok(()))
217            }
218            Err(e) => {
219                connection.lock().resolve(Box::new(Closed(e.clone())));
220                let resource = accessor.with(|mut access| access.get().table.push(e))?;
221                Ok(Err(resource))
222            }
223        }
224    }
225}
226
227pub(crate) struct FnTask<Fn>(pub(crate) Fn);
228impl<Fn, Fut, T, D> AccessorTask<T, D> for FnTask<Fn>
229where
230    Fn: FnOnce() -> Fut + Send + 'static,
231    Fut: Future<Output = wasmtime::Result<()>> + Send + 'static,
232    D: HasData + ?Sized,
233{
234    fn run(
235        self,
236        _accessor: &wasmtime::component::Accessor<T, D>,
237    ) -> impl Future<Output = wasmtime::Result<()>> + Send {
238        self.0()
239    }
240}
241
242pub(crate) struct ResultProducer<D> {
243    result: oneshot::Receiver<Result<(), Error>>,
244    getter: for<'a> fn(&'a mut D) -> WasiTlsCtxView<'a>,
245}
246impl<D> ResultProducer<D> {
247    pub(crate) fn new(
248        getter: for<'a> fn(&'a mut D) -> WasiTlsCtxView<'a>,
249        result: oneshot::Receiver<Result<(), Error>>,
250    ) -> Self {
251        Self { result, getter }
252    }
253
254    pub(crate) fn ready(
255        getter: for<'a> fn(&'a mut D) -> WasiTlsCtxView<'a>,
256        result: Result<(), Error>,
257    ) -> Self {
258        let (sender, receiver) = oneshot::channel();
259        sender.send(result).expect("receiver dropped");
260        Self {
261            result: receiver,
262            getter,
263        }
264    }
265}
266impl<D> FutureProducer<D> for ResultProducer<D>
267where
268    D: 'static,
269{
270    type Item = Result<(), Resource<Error>>;
271
272    fn poll_produce(
273        mut self: Pin<&mut Self>,
274        cx: &mut Context<'_>,
275        mut store: StoreContextMut<D>,
276        finish: bool,
277    ) -> Poll<wasmtime::error::Result<Option<Self::Item>>> {
278        match Pin::new(&mut self.result).poll(cx) {
279            Poll::Ready(Ok(Ok(()))) => Poll::Ready(Ok(Some(Ok(())))),
280            Poll::Ready(Ok(Err(err))) => {
281                let WasiTlsCtxView { table, .. } = (self.getter)(store.data_mut());
282                let err = table.push(err)?;
283                Poll::Ready(Ok(Some(Err(err))))
284            }
285            Poll::Ready(Err(_)) => Poll::Ready(Err(wasmtime::format_err!("sender dropped"))),
286            Poll::Pending if finish => Poll::Ready(Ok(None)),
287            Poll::Pending => Poll::Pending,
288        }
289    }
290}