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 outcome;
20pub mod workflow;
21
22const DEFAULT_TIMEOUT_MS: u64 = 35_000;
23
24type _WsStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
25
26struct PendingRequest {
27    tx: oneshot::Sender<Result<Value, String>>,
28}
29
30type PendingMap = Arc<Mutex<HashMap<String, PendingRequest>>>;
31
32/// Removing a local waiter never means the remote operation was cancelled.
33struct PendingGuard {
34    pending: PendingMap,
35    id: String,
36}
37
38impl Drop for PendingGuard {
39    fn drop(&mut self) {
40        self.pending.lock().unwrap().remove(&self.id);
41    }
42}
43
44fn reject_pending(pending: &PendingMap, reason: &str) {
45    for (id, req) in pending.lock().unwrap().drain() {
46        let _ = req.tx.send(Err(format!(
47            "outcome_unknown: {reason} (requestId: {id}); remote execution may continue"
48        )));
49    }
50}
51
52/// WebSocket client that communicates with tauri-plugin-connector.
53pub struct ConnectorClient {
54    write_tx: Option<mpsc::UnboundedSender<String>>,
55    pending: PendingMap,
56    _reader_handle: Option<tokio::task::JoinHandle<()>>,
57}
58
59impl ConnectorClient {
60    pub fn new() -> Self {
61        Self {
62            write_tx: None,
63            pending: Arc::new(Mutex::new(HashMap::new())),
64            _reader_handle: None,
65        }
66    }
67
68    /// Connect to the plugin's WebSocket server.
69    pub async fn connect(&mut self, host: &str, port: u16) -> Result<(), String> {
70        self.disconnect().await;
71        // An old I/O task finishing concurrently must not reject requests made
72        // on the replacement connection.
73        self.pending = Arc::new(Mutex::new(HashMap::new()));
74
75        let url = format!("ws://{host}:{port}");
76        let (ws, _) = tokio_tungstenite::connect_async(&url)
77            .await
78            .map_err(|e| format!("WebSocket connection failed: {e}"))?;
79
80        let (ws_write, ws_read) = ws.split();
81
82        // One task owns both halves: exiting either direction drops the other,
83        // closes the queue, and releases all waiters. No detached writer remains.
84        let (write_tx, mut write_rx) = mpsc::unbounded_channel::<String>();
85        let pending = self.pending.clone();
86        let reader_handle = tokio::spawn(async move {
87            let mut ws_write = ws_write;
88            let mut ws_read = ws_read;
89            loop {
90                tokio::select! {
91                    outbound = write_rx.recv() => {
92                        match outbound {
93                            Some(msg) => {
94                                if ws_write.send(Message::Text(msg.into())).await.is_err() {
95                                    break;
96                                }
97                            }
98                            None => break,
99                        }
100                    }
101                    inbound = ws_read.next() => {
102                        match inbound {
103                            Some(Ok(Message::Text(text))) => {
104                                if let Ok(response) = serde_json::from_str::<Value>(&text) {
105                                    let id = response.get("id").and_then(Value::as_str).unwrap_or("");
106                                    if let Some(req) = pending.lock().unwrap().remove(id) {
107                                        let result = if let Some(error) = response.get("error") {
108                                            if let Some(outcome) = response.get("outcome") {
109                                                Err(serde_json::json!({ "error": error, "outcome": outcome }).to_string())
110                                            } else {
111                                                Err(error.as_str().unwrap_or("Unknown error").to_string())
112                                            }
113                                        } else {
114                                            Ok(response.get("result").cloned().unwrap_or(Value::Null))
115                                        };
116                                        let _ = req.tx.send(result);
117                                    }
118                                }
119                            }
120                            Some(Ok(Message::Close(_))) | Some(Err(_)) | None => break,
121                            Some(Ok(_)) => {}
122                        }
123                    }
124                }
125            }
126            write_rx.close();
127            reject_pending(&pending, "Connection closed");
128        });
129
130        self.write_tx = Some(write_tx);
131        self._reader_handle = Some(reader_handle);
132
133        Ok(())
134    }
135
136    /// Disconnect from the WebSocket server.
137    pub async fn disconnect(&mut self) {
138        self.write_tx = None;
139        if let Some(handle) = self._reader_handle.take() {
140            handle.abort();
141        }
142        reject_pending(&self.pending, "Disconnected");
143    }
144
145    /// Check if connected.
146    pub fn is_connected(&self) -> bool {
147        self.write_tx.as_ref().is_some_and(|tx| !tx.is_closed())
148    }
149
150    /// Send a command and wait for a response.
151    pub async fn send(&self, command: Value) -> Result<Value, String> {
152        self.send_with_timeout(command, DEFAULT_TIMEOUT_MS).await
153    }
154
155    /// Send a command with a custom timeout.
156    pub async fn send_with_timeout(
157        &self,
158        command: Value,
159        timeout_ms: u64,
160    ) -> Result<Value, String> {
161        // Validate and serialize before registering any request state.
162        let mut msg = match command {
163            Value::Object(map) => map,
164            _ => return Err("Command must be a JSON object".to_string()),
165        };
166        let write_tx = self
167            .write_tx
168            .as_ref()
169            .ok_or_else(|| "Not connected".to_string())?;
170        let id = uuid::Uuid::new_v4().to_string();
171        msg.insert("id".to_string(), Value::String(id.clone()));
172        let json = serde_json::to_string(&msg).map_err(|e| e.to_string())?;
173        let (tx, rx) = oneshot::channel();
174        self.pending
175            .lock()
176            .unwrap()
177            .insert(id.clone(), PendingRequest { tx });
178        let _waiter = PendingGuard {
179            pending: self.pending.clone(),
180            id: id.clone(),
181        };
182        write_tx
183            .send(json)
184            .map_err(|_| "not_dispatched: Send failed: connection closed".to_string())?;
185
186        // Once queued, a timeout or connection loss cannot prove non-execution.
187        match tokio::time::timeout(Duration::from_millis(timeout_ms), rx).await {
188            Ok(Ok(result)) => result,
189            Ok(Err(_)) => Err(format!("outcome_unknown: Response channel closed (requestId: {id}); remote execution may continue")),
190            Err(_) => Err(format!("outcome_unknown: Request timeout (requestId: {id}); remote execution may continue")),
191        }
192    }
193}
194
195impl Drop for ConnectorClient {
196    fn drop(&mut self) {
197        self.write_tx = None;
198        if let Some(handle) = self._reader_handle.take() {
199            handle.abort();
200        }
201        reject_pending(&self.pending, "Client dropped");
202    }
203}
204
205impl Default for ConnectorClient {
206    fn default() -> Self {
207        Self::new()
208    }
209}
210
211#[cfg(test)]
212mod transport_tests {
213    use super::*;
214    use serde_json::json;
215
216    fn queued_client() -> (ConnectorClient, mpsc::UnboundedReceiver<String>) {
217        let (tx, rx) = mpsc::unbounded_channel();
218        let mut client = ConnectorClient::new();
219        client.write_tx = Some(tx);
220        (client, rx)
221    }
222
223    #[tokio::test]
224    async fn invalid_command_does_not_register_a_waiter() {
225        let (client, _rx) = queued_client();
226        assert!(client.send(Value::Null).await.is_err());
227        assert!(client.pending.lock().unwrap().is_empty());
228    }
229
230    #[tokio::test]
231    async fn failed_enqueue_removes_waiter() {
232        let (client, rx) = queued_client();
233        drop(rx);
234        assert!(client.send(json!({"command": "test"})).await.is_err());
235        assert!(client.pending.lock().unwrap().is_empty());
236    }
237
238    #[tokio::test]
239    async fn dropped_send_future_removes_waiter_without_remote_cancellation() {
240        let (client, mut rx) = queued_client();
241        let client = Arc::new(client);
242        let cloned = client.clone();
243        let task = tokio::spawn(async move { cloned.send(json!({"command": "test"})).await });
244        let dispatched = rx.recv().await.unwrap();
245        assert!(serde_json::from_str::<Value>(&dispatched).unwrap()["id"].is_string());
246        assert_eq!(client.pending.lock().unwrap().len(), 1);
247        task.abort();
248        assert!(task.await.unwrap_err().is_cancelled());
249        assert!(client.pending.lock().unwrap().is_empty());
250        assert!(
251            rx.try_recv().is_err(),
252            "dropping the waiter must not send another command"
253        );
254    }
255
256    #[tokio::test]
257    async fn timeout_retains_unknown_remote_outcome_and_request_id() {
258        let (client, mut rx) = queued_client();
259        let error = client
260            .send_with_timeout(json!({"command": "test"}), 1)
261            .await
262            .unwrap_err();
263        let sent: Value = serde_json::from_str(&rx.recv().await.unwrap()).unwrap();
264        assert!(error.contains("outcome_unknown"), "{error}");
265        assert!(error.contains(sent["id"].as_str().unwrap()), "{error}");
266        assert!(client.pending.lock().unwrap().is_empty());
267    }
268
269    async fn fixture_client(
270        response: Option<Value>,
271    ) -> (ConnectorClient, tokio::task::JoinHandle<()>) {
272        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
273        let port = listener.local_addr().unwrap().port();
274        let server = tokio::spawn(async move {
275            let (stream, _) = listener.accept().await.unwrap();
276            let mut socket = tokio_tungstenite::accept_async(stream).await.unwrap();
277            let command = socket.next().await.unwrap().unwrap();
278            let command: Value = serde_json::from_str(command.to_text().unwrap()).unwrap();
279            if let Some(mut response) = response {
280                response["id"] = command["id"].clone();
281                socket
282                    .send(Message::Text(response.to_string().into()))
283                    .await
284                    .unwrap();
285            }
286            socket.close(None).await.unwrap();
287        });
288        let mut client = ConnectorClient::new();
289        client.connect("127.0.0.1", port).await.unwrap();
290        (client, server)
291    }
292
293    #[tokio::test]
294    async fn closed_socket_cleans_pending_and_preserves_unknown_remote_state() {
295        let (client, server) = fixture_client(None).await;
296        let error = client.send(json!({"command":"write"})).await.unwrap_err();
297        assert!(error.contains("outcome_unknown"));
298        assert!(error.contains("requestId:"));
299        assert!(client.pending.lock().unwrap().is_empty());
300        server.await.unwrap();
301        assert!(!client.is_connected());
302    }
303
304    #[tokio::test]
305    async fn protocol_error_retains_the_explicit_outcome_envelope() {
306        let outcome =
307            json!({"execution":"completed", "verification":"failed", "effect":"possible"});
308        let (client, server) = fixture_client(Some(
309            json!({"error":"Postcondition failed", "outcome":outcome}),
310        ))
311        .await;
312        let error = client.send(json!({"command":"write"})).await.unwrap_err();
313        let envelope: Value = serde_json::from_str(&error).unwrap();
314        assert_eq!(envelope["outcome"], outcome);
315        assert_eq!(envelope["error"], "Postcondition failed");
316        assert!(client.pending.lock().unwrap().is_empty());
317        server.await.unwrap();
318    }
319
320    #[tokio::test]
321    async fn business_error_fields_remain_successful_data() {
322        let data = json!({"error":"user content", "found":false, "ok":false});
323        let (client, server) = fixture_client(Some(json!({"result":data}))).await;
324        assert_eq!(client.send(json!({"command":"read"})).await.unwrap(), data);
325        assert!(client.pending.lock().unwrap().is_empty());
326        server.await.unwrap();
327    }
328
329    #[tokio::test]
330    async fn explicit_disconnect_releases_every_waiter() {
331        let (mut client, _rx) = queued_client();
332        let (tx, result) = oneshot::channel();
333        client
334            .pending
335            .lock()
336            .unwrap()
337            .insert("queued-id".into(), PendingRequest { tx });
338        client.disconnect().await;
339        assert!(!client.is_connected());
340        let error = result.await.unwrap().unwrap_err();
341        assert!(error.contains("outcome_unknown"));
342        assert!(error.contains("queued-id"));
343        assert!(client.pending.lock().unwrap().is_empty());
344    }
345
346    #[tokio::test]
347    async fn dropping_client_closes_socket_without_detached_writer() {
348        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
349        let port = listener.local_addr().unwrap().port();
350        let server = tokio::spawn(async move {
351            let (stream, _) = listener.accept().await.unwrap();
352            let mut socket = tokio_tungstenite::accept_async(stream).await.unwrap();
353            match tokio::time::timeout(Duration::from_secs(1), socket.next()).await {
354                Ok(None | Some(Err(_)) | Some(Ok(Message::Close(_)))) => {}
355                other => panic!("client drop must close the connection: {other:?}"),
356            }
357        });
358        let mut client = ConnectorClient::new();
359        client.connect("127.0.0.1", port).await.unwrap();
360        drop(client);
361        server.await.unwrap();
362    }
363}