1use 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
14pub 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, 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 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}