1use std::collections::HashMap;
7use std::sync::{Arc, Mutex};
8use std::time::Duration;
9
10use futures_util::{SinkExt, StreamExt};
11use serde_json::Value;
12use tokio::net::TcpStream;
13use tokio::sync::{mpsc, oneshot};
14use tokio_tungstenite::tungstenite::Message;
15use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
16
17pub mod batch;
18pub mod discovery;
19pub mod identity;
20pub mod inspection;
21pub mod outcome;
22pub mod workflow;
23
24const DEFAULT_TIMEOUT_MS: u64 = 35_000;
25
26type _WsStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
27
28struct PendingRequest {
29 tx: oneshot::Sender<Result<Value, String>>,
30}
31
32type PendingMap = Arc<Mutex<HashMap<String, PendingRequest>>>;
33
34struct PendingGuard {
36 pending: PendingMap,
37 id: String,
38}
39
40impl Drop for PendingGuard {
41 fn drop(&mut self) {
42 self.pending.lock().unwrap().remove(&self.id);
43 }
44}
45
46fn reject_pending(pending: &PendingMap, reason: &str) {
47 for (id, req) in pending.lock().unwrap().drain() {
48 let _ = req.tx.send(Err(format!(
49 "outcome_unknown: {reason} (requestId: {id}); remote execution may continue"
50 )));
51 }
52}
53
54pub struct ConnectorClient {
56 write_tx: Option<mpsc::UnboundedSender<String>>,
57 pending: PendingMap,
58 _reader_handle: Option<tokio::task::JoinHandle<()>>,
59 expected_instance: Mutex<Option<String>>,
60}
61
62impl ConnectorClient {
63 pub fn new() -> Self {
64 Self {
65 write_tx: None,
66 pending: Arc::new(Mutex::new(HashMap::new())),
67 _reader_handle: None,
68 expected_instance: Mutex::new(None),
69 }
70 }
71
72 pub async fn connect(&mut self, host: &str, port: u16) -> Result<(), String> {
74 self.disconnect().await;
75 self.pending = Arc::new(Mutex::new(HashMap::new()));
78
79 let url = format!("ws://{host}:{port}");
80 let (ws, _) = tokio_tungstenite::connect_async(&url)
81 .await
82 .map_err(|e| format!("WebSocket connection failed: {e}"))?;
83
84 let (ws_write, ws_read) = ws.split();
85
86 let (write_tx, mut write_rx) = mpsc::unbounded_channel::<String>();
89 let pending = self.pending.clone();
90 let reader_handle = tokio::spawn(async move {
91 let mut ws_write = ws_write;
92 let mut ws_read = ws_read;
93 loop {
94 tokio::select! {
95 outbound = write_rx.recv() => {
96 match outbound {
97 Some(msg) => {
98 if ws_write.send(Message::Text(msg.into())).await.is_err() {
99 break;
100 }
101 }
102 None => break,
103 }
104 }
105 inbound = ws_read.next() => {
106 match inbound {
107 Some(Ok(Message::Text(text))) => {
108 if let Ok(response) = serde_json::from_str::<Value>(&text) {
109 let id = response.get("id").and_then(Value::as_str).unwrap_or("");
110 if let Some(req) = pending.lock().unwrap().remove(id) {
111 let result = if let Some(error) = response.get("error") {
112 if let Some(outcome) = response.get("outcome") {
113 Err(serde_json::json!({ "error": error, "outcome": outcome }).to_string())
114 } else {
115 Err(error.as_str().unwrap_or("Unknown error").to_string())
116 }
117 } else {
118 Ok(response.get("result").cloned().unwrap_or(Value::Null))
119 };
120 let _ = req.tx.send(result);
121 }
122 }
123 }
124 Some(Ok(Message::Close(_))) | Some(Err(_)) | None => break,
125 Some(Ok(_)) => {}
126 }
127 }
128 }
129 }
130 write_rx.close();
131 reject_pending(&pending, "Connection closed");
132 });
133
134 self.write_tx = Some(write_tx);
135 self._reader_handle = Some(reader_handle);
136 let pinned = self.expected_instance.lock().unwrap().is_some();
137 if pinned {
138 let token = std::env::var("TAURI_CONNECTOR_WORKFLOW_TOKEN").ok();
139 if let Err(error) = self.verify_app_identity(token.as_deref()).await {
140 self.disconnect().await;
141 return Err(error);
142 }
143 }
144 Ok(())
145 }
146
147 pub fn bind_instance(&self, instance: &str) -> Result<(), String> {
150 let mut expected = self.expected_instance.lock().unwrap();
151 if expected.as_deref().is_some_and(|old| old != instance) {
152 return Err("app_identity_mismatch: client is bound to another app instance".into());
153 }
154 *expected = Some(instance.to_owned());
155 Ok(())
156 }
157
158 pub async fn verify_app_identity(
159 &self,
160 auth_token: Option<&str>,
161 ) -> Result<identity::AppIdentity, String> {
162 let mut args = serde_json::json!({});
163 if let Some(token) = auth_token {
164 args["authToken"] = serde_json::json!(token);
165 }
166 let identity = identity::AppIdentity::parse(
167 self.send_with_timeout(
168 serde_json::json!({"type":"inspection","operation":"app_identity","args":args}),
169 2000,
170 )
171 .await?,
172 )?;
173 self.bind_instance(&identity.app_instance_id)?;
174 Ok(identity)
175 }
176
177 pub async fn inspect(&self, operation: &str, arguments: &Value) -> Result<Value, String> {
180 let mut args = arguments.clone();
181 if !args.is_object() {
182 return Err("invalid_arguments: arguments must be an object".into());
183 }
184 if args.get("authToken").is_none() {
185 if let Ok(token) = std::env::var("TAURI_CONNECTOR_WORKFLOW_TOKEN") {
186 args["authToken"] = serde_json::json!(token);
187 }
188 }
189 for field in [
190 "pickerId",
191 "captureSessionId",
192 "artifactId",
193 "artifact",
194 "before",
195 "after",
196 "baselineId",
197 "currentId",
198 ] {
199 if let Some(instance) = args
200 .get(field)
201 .and_then(Value::as_str)
202 .and_then(identity::instance_from_handle)
203 {
204 self.bind_instance(instance)?;
205 }
206 }
207 if operation == "webview_select_element" {
208 let request = inspection::PickerRequest::parse(&args).map_err(|e| e.to_string())?;
209 if request.action == "start" && request.request_key.is_none() {
210 args["requestKey"] = serde_json::json!(inspection::new_request_key());
211 }
212 }
213 let capabilities = self
214 .send_with_timeout(serde_json::json!({"type":"bridge_status"}), 2000)
215 .await?;
216 if capabilities
217 .get("inspectionProtocolVersion")
218 .and_then(Value::as_u64)
219 != Some(inspection::INSPECTION_PROTOCOL_VERSION)
220 {
221 return Err("capability_unavailable: connected app does not support inspection protocol v1; upgrade the plugin".into());
222 }
223 let identity = self
224 .verify_app_identity(args.get("authToken").and_then(Value::as_str))
225 .await?;
226 if operation == "app_identity" {
227 return Ok(serde_json::json!(identity));
228 }
229 let wait_ms = args
230 .get("waitMs")
231 .and_then(Value::as_u64)
232 .unwrap_or(10000)
233 .min(10000);
234 let transport_timeout = if operation == "webview_screenshot" {
235 args.get("timeoutMs")
236 .and_then(Value::as_u64)
237 .unwrap_or(10000)
238 .min(30000)
239 + 5000
240 } else {
241 wait_ms + 15000
242 };
243 let recovery_key = args.get("requestKey").cloned();
244 self.send_with_timeout(
245 serde_json::json!({"type":"inspection","operation":operation,"args":args}),
246 transport_timeout,
247 )
248 .await
249 .map_err(|error| {
250 if operation != "webview_select_element" {
251 return error;
252 }
253 if let Some(key) = recovery_key {
254 let mut report = serde_json::from_str::<Value>(&error)
255 .ok()
256 .filter(Value::is_object)
257 .unwrap_or_else(|| serde_json::json!({"error":error}));
258 report["requestKey"] = key;
259 report.to_string()
260 } else {
261 error
262 }
263 })
264 }
265
266 pub async fn disconnect(&mut self) {
268 self.write_tx = None;
269 if let Some(handle) = self._reader_handle.take() {
270 handle.abort();
271 }
272 reject_pending(&self.pending, "Disconnected");
273 }
274
275 pub fn is_connected(&self) -> bool {
277 self.write_tx.as_ref().is_some_and(|tx| !tx.is_closed())
278 }
279
280 pub async fn send(&self, command: Value) -> Result<Value, String> {
282 self.send_with_timeout(command, DEFAULT_TIMEOUT_MS).await
283 }
284
285 pub async fn send_with_timeout(
287 &self,
288 command: Value,
289 timeout_ms: u64,
290 ) -> Result<Value, String> {
291 let mut msg = match command {
293 Value::Object(map) => map,
294 _ => return Err("Command must be a JSON object".to_string()),
295 };
296 let write_tx = self
297 .write_tx
298 .as_ref()
299 .ok_or_else(|| "Not connected".to_string())?;
300 let id = uuid::Uuid::new_v4().to_string();
301 msg.insert("id".to_string(), Value::String(id.clone()));
302 let json = serde_json::to_string(&msg).map_err(|e| e.to_string())?;
303 let (tx, rx) = oneshot::channel();
304 self.pending
305 .lock()
306 .unwrap()
307 .insert(id.clone(), PendingRequest { tx });
308 let _waiter = PendingGuard {
309 pending: self.pending.clone(),
310 id: id.clone(),
311 };
312 write_tx
313 .send(json)
314 .map_err(|_| "not_dispatched: Send failed: connection closed".to_string())?;
315
316 match tokio::time::timeout(Duration::from_millis(timeout_ms), rx).await {
318 Ok(Ok(result)) => result,
319 Ok(Err(_)) => Err(format!("outcome_unknown: Response channel closed (requestId: {id}); remote execution may continue")),
320 Err(_) => Err(format!("outcome_unknown: Request timeout (requestId: {id}); remote execution may continue")),
321 }
322 }
323}
324
325impl Drop for ConnectorClient {
326 fn drop(&mut self) {
327 self.write_tx = None;
328 if let Some(handle) = self._reader_handle.take() {
329 handle.abort();
330 }
331 reject_pending(&self.pending, "Client dropped");
332 }
333}
334
335impl Default for ConnectorClient {
336 fn default() -> Self {
337 Self::new()
338 }
339}
340
341#[cfg(test)]
342mod transport_tests {
343 use super::*;
344 use serde_json::json;
345
346 fn queued_client() -> (ConnectorClient, mpsc::UnboundedReceiver<String>) {
347 let (tx, rx) = mpsc::unbounded_channel();
348 let mut client = ConnectorClient::new();
349 client.write_tx = Some(tx);
350 (client, rx)
351 }
352
353 #[tokio::test]
354 async fn invalid_command_does_not_register_a_waiter() {
355 let (client, _rx) = queued_client();
356 assert!(client.send(Value::Null).await.is_err());
357 assert!(client.pending.lock().unwrap().is_empty());
358 }
359
360 #[tokio::test]
361 async fn failed_enqueue_removes_waiter() {
362 let (client, rx) = queued_client();
363 drop(rx);
364 assert!(client.send(json!({"command": "test"})).await.is_err());
365 assert!(client.pending.lock().unwrap().is_empty());
366 }
367
368 #[tokio::test]
369 async fn dropped_send_future_removes_waiter_without_remote_cancellation() {
370 let (client, mut rx) = queued_client();
371 let client = Arc::new(client);
372 let cloned = client.clone();
373 let task = tokio::spawn(async move { cloned.send(json!({"command": "test"})).await });
374 let dispatched = rx.recv().await.unwrap();
375 assert!(serde_json::from_str::<Value>(&dispatched).unwrap()["id"].is_string());
376 assert_eq!(client.pending.lock().unwrap().len(), 1);
377 task.abort();
378 assert!(task.await.unwrap_err().is_cancelled());
379 assert!(client.pending.lock().unwrap().is_empty());
380 assert!(
381 rx.try_recv().is_err(),
382 "dropping the waiter must not send another command"
383 );
384 }
385
386 #[tokio::test]
387 async fn timeout_retains_unknown_remote_outcome_and_request_id() {
388 let (client, mut rx) = queued_client();
389 let error = client
390 .send_with_timeout(json!({"command": "test"}), 1)
391 .await
392 .unwrap_err();
393 let sent: Value = serde_json::from_str(&rx.recv().await.unwrap()).unwrap();
394 assert!(error.contains("outcome_unknown"), "{error}");
395 assert!(error.contains(sent["id"].as_str().unwrap()), "{error}");
396 assert!(client.pending.lock().unwrap().is_empty());
397 }
398
399 async fn fixture_client(
400 response: Option<Value>,
401 ) -> (ConnectorClient, tokio::task::JoinHandle<()>) {
402 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
403 let port = listener.local_addr().unwrap().port();
404 let server = tokio::spawn(async move {
405 let (stream, _) = listener.accept().await.unwrap();
406 let mut socket = tokio_tungstenite::accept_async(stream).await.unwrap();
407 let command = socket.next().await.unwrap().unwrap();
408 let command: Value = serde_json::from_str(command.to_text().unwrap()).unwrap();
409 if let Some(mut response) = response {
410 response["id"] = command["id"].clone();
411 socket
412 .send(Message::Text(response.to_string().into()))
413 .await
414 .unwrap();
415 }
416 socket.close(None).await.unwrap();
417 });
418 let mut client = ConnectorClient::new();
419 client.connect("127.0.0.1", port).await.unwrap();
420 (client, server)
421 }
422
423 #[tokio::test]
424 async fn closed_socket_cleans_pending_and_preserves_unknown_remote_state() {
425 let (client, server) = fixture_client(None).await;
426 let error = client.send(json!({"command":"write"})).await.unwrap_err();
427 assert!(error.contains("outcome_unknown"));
428 assert!(error.contains("requestId:"));
429 assert!(client.pending.lock().unwrap().is_empty());
430 server.await.unwrap();
431 assert!(!client.is_connected());
432 }
433
434 #[tokio::test]
435 async fn protocol_error_retains_the_explicit_outcome_envelope() {
436 let outcome =
437 json!({"execution":"completed", "verification":"failed", "effect":"possible"});
438 let (client, server) = fixture_client(Some(
439 json!({"error":"Postcondition failed", "outcome":outcome}),
440 ))
441 .await;
442 let error = client.send(json!({"command":"write"})).await.unwrap_err();
443 let envelope: Value = serde_json::from_str(&error).unwrap();
444 assert_eq!(envelope["outcome"], outcome);
445 assert_eq!(envelope["error"], "Postcondition failed");
446 assert!(client.pending.lock().unwrap().is_empty());
447 server.await.unwrap();
448 }
449
450 #[tokio::test]
451 async fn business_error_fields_remain_successful_data() {
452 let data = json!({"error":"user content", "found":false, "ok":false});
453 let (client, server) = fixture_client(Some(json!({"result":data}))).await;
454 assert_eq!(client.send(json!({"command":"read"})).await.unwrap(), data);
455 assert!(client.pending.lock().unwrap().is_empty());
456 server.await.unwrap();
457 }
458
459 #[tokio::test]
460 async fn explicit_disconnect_releases_every_waiter() {
461 let (mut client, _rx) = queued_client();
462 let (tx, result) = oneshot::channel();
463 client
464 .pending
465 .lock()
466 .unwrap()
467 .insert("queued-id".into(), PendingRequest { tx });
468 client.disconnect().await;
469 assert!(!client.is_connected());
470 let error = result.await.unwrap().unwrap_err();
471 assert!(error.contains("outcome_unknown"));
472 assert!(error.contains("queued-id"));
473 assert!(client.pending.lock().unwrap().is_empty());
474 }
475
476 #[tokio::test]
477 async fn dropping_client_closes_socket_without_detached_writer() {
478 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
479 let port = listener.local_addr().unwrap().port();
480 let server = tokio::spawn(async move {
481 let (stream, _) = listener.accept().await.unwrap();
482 let mut socket = tokio_tungstenite::accept_async(stream).await.unwrap();
483 match tokio::time::timeout(Duration::from_secs(1), socket.next()).await {
484 Ok(None | Some(Err(_)) | Some(Ok(Message::Close(_)))) => {}
485 other => panic!("client drop must close the connection: {other:?}"),
486 }
487 });
488 let mut client = ConnectorClient::new();
489 client.connect("127.0.0.1", port).await.unwrap();
490 drop(client);
491 server.await.unwrap();
492 }
493}