Skip to content

Commit ab63237

Browse files
committed
feat(sdk): reclaim session IO after completion
Retain the session driver result so native callers can recover the underlying stream and WASM callers can await exclusive reuse of their injected IoChannel. Exercise the handoff by exchanging application bytes after the SDK protocol completes. Assisted-by: GPT-5 Signed-off-by: Wondertan <hlibwondertan@gmail.com>
1 parent 0fe3c32 commit ab63237

8 files changed

Lines changed: 330 additions & 19 deletions

File tree

crates/harness/executor/src/provider.rs

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -83,10 +83,29 @@ mod wasm {
8383
use crate::io::Io;
8484
use anyhow::{Result, anyhow};
8585
use std::time::Duration;
86+
use wasm_bindgen::prelude::*;
87+
use wasm_bindgen_futures::JsFuture;
8688

8789
const CHECK_WS_OPEN_DELAY_MS: usize = 50;
8890
const MAX_RETRIES: usize = 50;
8991

92+
#[wasm_bindgen]
93+
extern "C" {
94+
type JsIoChannel;
95+
96+
#[wasm_bindgen(js_namespace = globalThis, js_name = connectIoChannel)]
97+
fn connect_io_channel(url: String) -> js_sys::Promise;
98+
99+
#[wasm_bindgen(method, js_name = isOpen)]
100+
fn is_open(this: &JsIoChannel) -> bool;
101+
}
102+
103+
async fn connect_js_io(url: String) -> Result<JsValue> {
104+
JsFuture::from(connect_io_channel(url))
105+
.await
106+
.map_err(|error| anyhow!("failed to connect JS IO: {error:?}"))
107+
}
108+
90109
impl IoProvider {
91110
/// Provides a connection to the server.
92111
pub async fn provide_server_io(&self) -> Result<impl Io> {
@@ -102,6 +121,18 @@ mod wasm {
102121
Ok(io.into_io())
103122
}
104123

124+
/// Provides a JavaScript `IoChannel` backed by a real WebSocket.
125+
pub async fn provide_server_js_io(&self) -> Result<JsValue> {
126+
connect_js_io(format!(
127+
"ws://{}:{}/tcp?addr={}%3A{}",
128+
&self.config.app_proxy.0,
129+
self.config.app_proxy.1,
130+
&self.config.app.0,
131+
self.config.app.1,
132+
))
133+
.await
134+
}
135+
105136
/// Provides a connection to the verifier.
106137
pub async fn provide_proto_io(&self) -> Result<impl Io> {
107138
let url = format!(
@@ -135,5 +166,30 @@ mod wasm {
135166

136167
Ok(io.into_io())
137168
}
169+
170+
/// Provides a JavaScript `IoChannel` backed by the protocol WebSocket.
171+
pub async fn provide_proto_js_io(&self) -> Result<JsValue> {
172+
let url = format!(
173+
"ws://{}:{}/tcp?addr={}%3A{}",
174+
&self.config.proto_proxy.0,
175+
self.config.proto_proxy.1,
176+
&self.config.proto_1.0,
177+
self.config.proto_1.1,
178+
);
179+
let mut retries = 0;
180+
181+
loop {
182+
let io = connect_js_io(url.clone()).await?;
183+
std::thread::sleep(Duration::from_millis(CHECK_WS_OPEN_DELAY_MS as u64));
184+
if io.unchecked_ref::<JsIoChannel>().is_open() {
185+
return Ok(io);
186+
}
187+
188+
retries += 1;
189+
if retries > MAX_RETRIES {
190+
return Err(anyhow!("verifier did not accept connection"));
191+
}
192+
}
193+
}
138194
}
139195
}

crates/harness/executor/test_plugins/sdk_core.rs

Lines changed: 113 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
1-
use tlsn_sdk_core::{
2-
HttpRequest, NetworkSetting, ProverConfig, Reveal, SdkProver, SdkVerifier, VerifierConfig,
3-
};
1+
use futures::io::{AsyncReadExt, AsyncWriteExt};
2+
#[cfg(not(target_arch = "wasm32"))]
3+
use tlsn_sdk_core::{HttpRequest, NetworkSetting, ProverConfig, Reveal, SdkProver};
4+
use tlsn_sdk_core::{SdkVerifier, VerifierConfig};
45
use tlsn_server_fixture_certs::{CA_CERT_DER, SERVER_DOMAIN};
56

67
use crate::IoProvider;
@@ -13,6 +14,15 @@ const MAX_RECV_DATA: usize = 1 << 11;
1314
crate::test!("sdk_core", prover, verifier);
1415

1516
async fn prover(provider: &IoProvider) {
17+
#[cfg(target_arch = "wasm32")]
18+
return prover_wasm(provider).await;
19+
20+
#[cfg(not(target_arch = "wasm32"))]
21+
prover_core(provider).await;
22+
}
23+
24+
#[cfg(not(target_arch = "wasm32"))]
25+
async fn prover_core(provider: &IoProvider) {
1626
let config = ProverConfig::builder(SERVER_DOMAIN)
1727
.max_sent_data(MAX_SENT_DATA)
1828
.max_recv_data(MAX_RECV_DATA)
@@ -60,6 +70,100 @@ async fn prover(provider: &IoProvider) {
6070
.unwrap();
6171

6272
assert!(prover.is_complete());
73+
74+
let mut io = prover.finish().await.unwrap();
75+
io.write_all(b"prover-finished").await.unwrap();
76+
let mut response = [0; 17];
77+
io.read_exact(&mut response).await.unwrap();
78+
assert_eq!(&response, b"verifier-finished");
79+
}
80+
81+
#[cfg(target_arch = "wasm32")]
82+
async fn prover_wasm(provider: &IoProvider) {
83+
use js_sys::{Promise, Uint8Array};
84+
use std::collections::HashMap;
85+
use tlsn_wasm::{
86+
prover::{JsProver, ProverConfig},
87+
types::{HttpRequest, Method, Reveal},
88+
};
89+
use wasm_bindgen::{JsCast, prelude::*};
90+
use wasm_bindgen_futures::JsFuture;
91+
92+
#[wasm_bindgen]
93+
extern "C" {
94+
type TestIoChannel;
95+
96+
#[wasm_bindgen(method)]
97+
fn read(this: &TestIoChannel) -> Promise;
98+
99+
#[wasm_bindgen(method)]
100+
fn write(this: &TestIoChannel, data: &Uint8Array) -> Promise;
101+
}
102+
103+
let config: ProverConfig = serde_json::from_value(serde_json::json!({
104+
"server_name": SERVER_DOMAIN,
105+
"mode": "Mpc",
106+
"max_sent_data": MAX_SENT_DATA,
107+
"max_sent_records": null,
108+
"max_recv_data_online": null,
109+
"max_recv_data": MAX_RECV_DATA,
110+
"max_recv_records_online": null,
111+
"defer_decryption_from_start": true,
112+
"network": "Latency",
113+
"client_auth": null,
114+
"root_certs": [CA_CERT_DER],
115+
}))
116+
.unwrap();
117+
let mut prover = JsProver::new(config).unwrap();
118+
119+
let proto_io = provider.provide_proto_js_io().await.unwrap();
120+
let retained_io = proto_io.unchecked_ref::<TestIoChannel>();
121+
prover
122+
.setup(proto_io.clone().unchecked_into())
123+
.await
124+
.unwrap();
125+
126+
let server_io = provider.provide_server_js_io().await.unwrap();
127+
let response = prover
128+
.send_request(
129+
Some(server_io.unchecked_into()),
130+
HttpRequest {
131+
uri: format!(
132+
"https://{}/bytes?size={}",
133+
SERVER_DOMAIN,
134+
MAX_RECV_DATA - 256
135+
),
136+
method: Method::GET,
137+
headers: HashMap::from([
138+
("Host".to_string(), SERVER_DOMAIN.as_bytes().to_vec()),
139+
("Connection".to_string(), b"close".to_vec()),
140+
]),
141+
body: None,
142+
},
143+
)
144+
.await
145+
.unwrap();
146+
assert_eq!(response.status, 200);
147+
148+
let transcript = prover.transcript().unwrap();
149+
prover
150+
.reveal(
151+
Reveal {
152+
sent: vec![0..transcript.sent.len() - 1],
153+
recv: vec![2..transcript.recv.len()],
154+
server_identity: true,
155+
},
156+
None,
157+
)
158+
.await
159+
.unwrap();
160+
prover.finish().await.unwrap();
161+
162+
JsFuture::from(retained_io.write(&Uint8Array::from(b"prover-finished".as_slice())))
163+
.await
164+
.unwrap();
165+
let response = JsFuture::from(retained_io.read()).await.unwrap();
166+
assert_eq!(Uint8Array::new(&response).to_vec(), b"verifier-finished");
63167
}
64168

65169
async fn verifier(provider: &IoProvider) {
@@ -86,4 +190,10 @@ async fn verifier(provider: &IoProvider) {
86190
assert_eq!(output.server_name.as_deref(), Some(SERVER_DOMAIN));
87191
assert!(output.transcript.is_some());
88192
assert!(verifier.is_complete());
193+
194+
let mut io = verifier.finish().await.unwrap();
195+
let mut request = [0; 15];
196+
io.read_exact(&mut request).await.unwrap();
197+
assert_eq!(&request, b"prover-finished");
198+
io.write_all(b"verifier-finished").await.unwrap();
89199
}

crates/harness/static/executor.js

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,67 @@
11
import * as Comlink from "./comlink.mjs";
22
import initWasm, * as wasm from "./generated/harness_executor.js";
33

4+
class WebSocketIoChannel {
5+
constructor(socket) {
6+
this.socket = socket;
7+
this.queue = [];
8+
this.reader = null;
9+
10+
socket.binaryType = "arraybuffer";
11+
socket.onmessage = ({ data }) => {
12+
const bytes = new Uint8Array(data);
13+
if (this.reader) {
14+
const reader = this.reader;
15+
this.reader = null;
16+
reader.resolve(bytes);
17+
} else {
18+
this.queue.push(bytes);
19+
}
20+
};
21+
socket.onclose = () => {
22+
if (this.reader) {
23+
const reader = this.reader;
24+
this.reader = null;
25+
reader.resolve(null);
26+
}
27+
};
28+
socket.onerror = () => {
29+
if (this.reader) {
30+
const reader = this.reader;
31+
this.reader = null;
32+
reader.reject(new Error("WebSocket error"));
33+
}
34+
};
35+
}
36+
37+
read() {
38+
if (this.queue.length) return Promise.resolve(this.queue.shift());
39+
if (this.socket.readyState === WebSocket.CLOSED) return Promise.resolve(null);
40+
if (this.reader) return Promise.reject(new Error("concurrent read"));
41+
return new Promise((resolve, reject) => { this.reader = { resolve, reject }; });
42+
}
43+
44+
write(data) {
45+
this.socket.send(data);
46+
return Promise.resolve();
47+
}
48+
49+
close() {
50+
this.socket.close();
51+
return Promise.resolve();
52+
}
53+
54+
isOpen() {
55+
return this.socket.readyState === WebSocket.OPEN;
56+
}
57+
}
58+
59+
globalThis.connectIoChannel = (url) => new Promise((resolve, reject) => {
60+
const socket = new WebSocket(url);
61+
socket.onopen = () => resolve(new WebSocketIoChannel(socket));
62+
socket.onerror = () => reject(new Error(`failed to connect to ${url}`));
63+
});
64+
465
class Executor {
566
executor;
667

crates/sdk-core/src/prover.rs

Lines changed: 21 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@ use crate::{
3333
pub struct SdkProver {
3434
config: ProverConfig,
3535
state: State,
36+
driver_task: Option<crate::spawn::DriverTask>,
3637
}
3738

3839
#[allow(clippy::large_enum_variant)]
@@ -91,6 +92,7 @@ impl SdkProver {
9192
Ok(SdkProver {
9293
config,
9394
state: State::Initialized,
95+
driver_task: None,
9496
})
9597
}
9698

@@ -111,15 +113,9 @@ impl SdkProver {
111113

112114
info!("connecting to verifier");
113115

114-
let session = Session::new(verifier_io);
116+
let session = Session::new(Box::new(verifier_io) as crate::spawn::BoxIo);
115117
let (driver, mut handle) = session.split();
116-
117-
crate::spawn::spawn(async move {
118-
match driver.await {
119-
Ok(_io) => tracing::warn!("session driver completed (mux closed)"),
120-
Err(e) => tracing::error!("session driver error: {e}"),
121-
}
122-
});
118+
self.driver_task = Some(crate::spawn::DriverTask::spawn(driver));
123119

124120
let prover_config = tlsn::config::prover::ProverConfig::builder().build()?;
125121
let prover = handle.new_prover(prover_config)?;
@@ -444,6 +440,23 @@ impl SdkProver {
444440
pub fn is_complete(&self) -> bool {
445441
matches!(self.state, State::Complete)
446442
}
443+
444+
/// Waits for the session driver to stop and returns the underlying IO.
445+
///
446+
/// This must be called after [`reveal`](Self::reveal). The returned stream
447+
/// is no longer read from or written to by TLSNotary and can be reused by
448+
/// the application.
449+
pub async fn finish(&mut self) -> Result<Box<dyn Io>> {
450+
if !self.is_complete() {
451+
return Err(SdkError::invalid_state("prover is not complete"));
452+
}
453+
454+
self.driver_task
455+
.take()
456+
.ok_or_else(|| SdkError::invalid_state("prover session already finished"))?
457+
.finish()
458+
.await
459+
}
447460
}
448461

449462
fn build_reveal_output(output: ProverOutput) -> Result<RevealOutput> {

crates/sdk-core/src/spawn.rs

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,40 @@
22
33
use std::future::Future;
44

5+
use futures::channel::oneshot;
6+
use tlsn::SessionDriver;
7+
8+
use crate::{
9+
error::{Result, SdkError},
10+
io::Io,
11+
};
12+
13+
pub(crate) type BoxIo = Box<dyn Io>;
14+
15+
pub(crate) struct DriverTask(oneshot::Receiver<tlsn::Result<BoxIo>>);
16+
17+
impl DriverTask {
18+
pub(crate) fn spawn(driver: SessionDriver<BoxIo>) -> Self {
19+
let (sender, receiver) = oneshot::channel();
20+
spawn(async move {
21+
let result = driver.await;
22+
match &result {
23+
Ok(_) => tracing::warn!("session driver completed (mux closed)"),
24+
Err(error) => tracing::error!("session driver error: {error}"),
25+
}
26+
let _ = sender.send(result);
27+
});
28+
Self(receiver)
29+
}
30+
31+
pub(crate) async fn finish(self) -> Result<BoxIo> {
32+
self.0
33+
.await
34+
.map_err(|_| SdkError::internal("session driver task dropped"))?
35+
.map_err(Into::into)
36+
}
37+
}
38+
539
/// Spawns a future on the appropriate runtime.
640
#[cfg(feature = "wasm")]
741
pub(crate) fn spawn(future: impl Future<Output = ()> + 'static) {

0 commit comments

Comments
 (0)