Skip to main content

connector_client/
lib.rs

1//! Shared WebSocket client for connecting to tauri-plugin-connector.
2//!
3//! Both the MCP server and CLI use this crate to communicate with the
4//! running Tauri app's connector plugin over WebSocket.
5
6use 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
34/// Removing a local waiter never means the remote operation was cancelled.
35struct 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
54/// WebSocket client that communicates with tauri-plugin-connector.
55pub 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    /// Connect to the plugin's WebSocket server.
73    pub async fn connect(&mut self, host: &str, port: u16) -> Result<(), String> {
74        self.disconnect().await;
75        // An old I/O task finishing concurrently must not reject requests made
76        // on the replacement connection.
77        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        // One task owns both halves: exiting either direction drops the other,
87        // closes the queue, and releases all waiters. No detached writer remains.
88        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    /// Pin the application expected by discovery or a previously returned handle.
148    /// Reconnecting does not clear this binding and cannot silently switch applications.
149    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    /// Forward to the application-owned inspection service after capability and identity checks.
178    /// A failed response is never retried and no arbitrary JS fallback exists.
179    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    /// Disconnect from the WebSocket server.
267    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    /// Check if connected.
276    pub fn is_connected(&self) -> bool {
277        self.write_tx.as_ref().is_some_and(|tx| !tx.is_closed())
278    }
279
280    /// Send a command and wait for a response.
281    pub async fn send(&self, command: Value) -> Result<Value, String> {
282        self.send_with_timeout(command, DEFAULT_TIMEOUT_MS).await
283    }
284
285    /// Send a command with a custom timeout.
286    pub async fn send_with_timeout(
287        &self,
288        command: Value,
289        timeout_ms: u64,
290    ) -> Result<Value, String> {
291        // Validate and serialize before registering any request state.
292        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        // Once queued, a timeout or connection loss cannot prove non-execution.
317        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}