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 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
32struct 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
52pub 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 pub async fn connect(&mut self, host: &str, port: u16) -> Result<(), String> {
70 self.disconnect().await;
71 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 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 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 pub fn is_connected(&self) -> bool {
147 self.write_tx.as_ref().is_some_and(|tx| !tx.is_closed())
148 }
149
150 pub async fn send(&self, command: Value) -> Result<Value, String> {
152 self.send_with_timeout(command, DEFAULT_TIMEOUT_MS).await
153 }
154
155 pub async fn send_with_timeout(
157 &self,
158 command: Value,
159 timeout_ms: u64,
160 ) -> Result<Value, String> {
161 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 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}