Skip to main content

graphwalker_restful/
websocket.rs

1use std::collections::HashSet;
2use std::time::Duration;
3
4use axum::extract::ws::{Message, WebSocket};
5use futures::{SinkExt, StreamExt};
6use serde_json::{json, Value};
7use tokio::sync::{broadcast, oneshot};
8use tracing::{debug, warn};
9
10use crate::actor::{self, spawn_machine_thread, Command};
11use crate::session::{SessionHandle, SessionManager};
12
13pub async fn handle_socket(socket: WebSocket, session_mgr: SessionManager) {
14    let (mut ws_write, mut ws_read) = socket.split();
15
16    let mut own_session: Option<SessionHandle> = None;
17    let mut subscribed_rx: Option<broadcast::Receiver<Value>> = None;
18    let mut subscribed_session_id: Option<String> = None;
19    let mut change_rx: Option<broadcast::Receiver<Value>> = None;
20
21    debug!("new websocket connection");
22
23    loop {
24        tokio::select! {
25            msg = ws_read.next() => {
26                let text = match msg {
27                    Some(Ok(Message::Text(t))) => t.to_string(),
28                    Some(Ok(Message::Close(_))) | None => {
29                        debug!("websocket closed by client");
30                        break;
31                    }
32                    _ => continue,
33                };
34
35                let response = match serde_json::from_str::<Value>(&text) {
36                    Ok(request) => {
37                        let cmd = request.get("command").and_then(|c| c.as_str()).unwrap_or("?");
38                        debug!(command = cmd, "recv command");
39                        process_command(
40                            &session_mgr,
41                            &mut own_session,
42                            &mut subscribed_rx,
43                            &mut subscribed_session_id,
44                            &mut change_rx,
45                            &request,
46                        )
47                        .await
48                    }
49                    Err(e) => json!({
50                        "command": "unknown",
51                        "success": false,
52                        "message": e.to_string(),
53                    }),
54                };
55
56                let resp_cmd = response.get("command").and_then(|c| c.as_str()).unwrap_or("?");
57                let resp_ok = response.get("success").and_then(|v| v.as_bool()).unwrap_or(false);
58                debug!(command = resp_cmd, success = resp_ok, "send response");
59
60                if ws_write
61                    .send(Message::Text(response.to_string().into()))
62                    .await
63                    .is_err()
64                {
65                    warn!("ws_write failed sending response");
66                    break;
67                }
68            }
69
70            event = async {
71                match &mut subscribed_rx {
72                    Some(rx) => rx.recv().await,
73                    None => std::future::pending().await,
74                }
75            } => {
76                match event {
77                    Ok(value) => {
78                        let cmd = value.get("command").and_then(|c| c.as_str()).unwrap_or("?");
79                        debug!(command = cmd, "forward broadcast to subscriber");
80                        if ws_write
81                            .send(Message::Text(value.to_string().into()))
82                            .await
83                            .is_err()
84                        {
85                            warn!("ws_write failed sending broadcast");
86                            break;
87                        }
88                    }
89                    Err(broadcast::error::RecvError::Lagged(n)) => {
90                        warn!(skipped = n, "broadcast receiver lagged");
91                    }
92                    Err(broadcast::error::RecvError::Closed) => {
93                        debug!("broadcast channel closed, clearing subscription");
94                        subscribed_rx = None;
95                    }
96                }
97            }
98
99            event = async {
100                match &mut change_rx {
101                    Some(rx) => rx.recv().await,
102                    None => std::future::pending().await,
103                }
104            } => {
105                match event {
106                    Ok(value) => {
107                        let cmd = value.get("command").and_then(|c| c.as_str()).unwrap_or("?");
108                        debug!(command = cmd, "forward change event");
109                        if ws_write
110                            .send(Message::Text(value.to_string().into()))
111                            .await
112                            .is_err()
113                        {
114                            warn!("ws_write failed sending change event");
115                            break;
116                        }
117                    }
118                    Err(_) => {}
119                }
120            }
121        }
122    }
123
124    debug!("websocket handler cleaning up");
125
126    if let Some(sid) = subscribed_session_id {
127        debug!(session_id = %sid, "resetting control for watched session");
128        if let Some(session) = session_mgr.get_session(&sid) {
129            session.control.reset();
130        }
131    }
132
133    if let Some(session) = own_session {
134        debug!(session_id = %session.id, "removing owned session");
135        let _ = session.broadcast_tx.send(json!({
136            "command": "sessionEnded",
137            "sessionId": session.id,
138        }));
139        session_mgr.remove_session(&session.id);
140    }
141}
142
143async fn send_command(
144    tx: &std::sync::mpsc::Sender<Command>,
145    build: impl FnOnce(oneshot::Sender<Result<Value, String>>) -> Command,
146) -> Result<Value, String> {
147    let (reply_tx, reply_rx) = oneshot::channel();
148    tx.send(build(reply_tx))
149        .map_err(|_| "Machine thread unavailable".to_string())?;
150    match reply_rx.await {
151        Ok(result) => result,
152        Err(_) => Err("Machine thread dropped".to_string()),
153    }
154}
155
156fn get_target_session(
157    session_mgr: &SessionManager,
158    request: &Value,
159) -> Result<SessionHandle, String> {
160    let session_id = request
161        .get("sessionId")
162        .and_then(|v| v.as_str())
163        .ok_or("Missing 'sessionId' field")?;
164    session_mgr
165        .get_session(session_id)
166        .ok_or_else(|| format!("Session '{}' not found", session_id))
167}
168
169async fn process_command(
170    session_mgr: &SessionManager,
171    own_session: &mut Option<SessionHandle>,
172    subscribed_rx: &mut Option<broadcast::Receiver<Value>>,
173    subscribed_session_id: &mut Option<String>,
174    change_rx: &mut Option<broadcast::Receiver<Value>>,
175    request: &Value,
176) -> Value {
177    let command = request
178        .get("command")
179        .and_then(|c| c.as_str())
180        .unwrap_or("")
181        .to_uppercase();
182
183    let result = match command.as_str() {
184        "MODE" => Ok(json!({
185            "command": "mode",
186            "mode": "EDITOR",
187            "success": true,
188        })),
189        "CHECK" => cmd_check(request),
190        "CONVERTGRAPHML" => cmd_convert_graphml(request),
191        "START" => cmd_start(session_mgr, own_session, request).await,
192        "GETNEXT" => cmd_get_next(own_session).await,
193        "HASNEXT" => cmd_has_next(own_session).await,
194        "GETDATA" => cmd_get_data(own_session).await,
195        "SETDATA" => cmd_set_data(own_session, request).await,
196        "GETMODEL" => cmd_get_model(own_session).await,
197        "UPDATEALLELEMENTS" => cmd_update_all_elements(own_session).await,
198        "LISTSESSIONS" => cmd_list_sessions(session_mgr, change_rx),
199        "SUBSCRIBESESSION" => {
200            cmd_subscribe_session(session_mgr, subscribed_rx, subscribed_session_id, request).await
201        }
202        "UNSUBSCRIBESESSION" => {
203            cmd_unsubscribe_session(session_mgr, subscribed_rx, subscribed_session_id)
204        }
205        "PAUSESESSION" => cmd_pause_session(session_mgr, request),
206        "RESUMESESSION" => cmd_resume_session(session_mgr, request),
207        "STEPSESSION" => cmd_step_session(session_mgr, request),
208        "SETDELAY" => cmd_set_delay(session_mgr, request),
209        "SETBREAKPOINTS" => cmd_set_breakpoints(session_mgr, request),
210        _ => Ok(json!({
211            "command": command.to_lowercase(),
212            "success": false,
213            "message": format!("Unknown command: {}", command),
214        })),
215    };
216
217    match result {
218        Ok(val) => val,
219        Err(msg) => {
220            warn!(command = %command, error = %msg, "command failed");
221            json!({
222                "command": command.to_lowercase(),
223                "success": false,
224                "message": msg,
225            })
226        }
227    }
228}
229
230fn cmd_check(request: &Value) -> Result<Value, String> {
231    let gw = request.get("gw").ok_or("Missing 'gw' field")?;
232    let val = actor::handle_check(&gw.to_string())?;
233    Ok(json!({
234        "command": "check",
235        "issues": val.get("issues").unwrap_or(&json!([])),
236        "success": true,
237    }))
238}
239
240fn cmd_convert_graphml(request: &Value) -> Result<Value, String> {
241    let graphml = request
242        .get("graphml")
243        .and_then(|g| g.as_str())
244        .ok_or("Missing 'graphml' field")?;
245    let val = actor::handle_convert_graphml(graphml)?;
246    Ok(json!({
247        "command": "convertGraphml",
248        "models": val.get("models").unwrap_or(&json!("")),
249        "success": true,
250    }))
251}
252
253async fn cmd_start(
254    session_mgr: &SessionManager,
255    own_session: &mut Option<SessionHandle>,
256    request: &Value,
257) -> Result<Value, String> {
258    if let Some(prev) = own_session.take() {
259        debug!(session_id = %prev.id, "removing previous session");
260        session_mgr.remove_session(&prev.id);
261    }
262
263    let gw = request.get("gw").ok_or("Missing 'gw' field")?;
264    let json_body = gw.to_string();
265    let seed = request.get("seed").and_then(|v| v.as_u64());
266    let global_data = request
267        .get("globalData")
268        .and_then(|v| v.as_str())
269        .map(|s| s.to_string());
270    let session_name = request
271        .get("name")
272        .and_then(|v| v.as_str())
273        .map(|s| s.to_string())
274        .unwrap_or_else(|| {
275            gw.get("models")
276                .and_then(|m| m.as_array())
277                .and_then(|arr| arr.first())
278                .and_then(|m| m.get("name"))
279                .and_then(|n| n.as_str())
280                .unwrap_or("Unnamed")
281                .to_string()
282        });
283
284    let machine_tx = spawn_machine_thread();
285
286    let val = send_command(&machine_tx, |reply| Command::Load {
287        json_body: json_body.clone(),
288        seed,
289        global_data,
290        reply,
291    })
292    .await?;
293
294    let actual_seed = val.get("seed").and_then(|v| v.as_u64());
295    let handle =
296        session_mgr.create_session(session_name.clone(), json_body, actual_seed, machine_tx);
297    let session_id = handle.id.clone();
298    *own_session = Some(handle);
299
300    debug!(session_id = %session_id, name = %session_name, seed = ?actual_seed, "session started");
301
302    Ok(json!({
303        "command": "start",
304        "success": true,
305        "seed": val.get("seed").unwrap_or(&json!(0)),
306        "sessionId": session_id,
307    }))
308}
309
310async fn cmd_get_next(own_session: &Option<SessionHandle>) -> Result<Value, String> {
311    let session = own_session.as_ref().ok_or("No active session")?;
312
313    debug!(session_id = %session.id, paused = session.control.is_paused(), "getNext: entering gate");
314    session.control.gate().await;
315    debug!(session_id = %session.id, "getNext: gate passed");
316
317    let val = send_command(&session.machine_tx, |reply| Command::GetNext {
318        verbose: true,
319        reply,
320    })
321    .await?;
322
323    let model_id = val
324        .get("modelId")
325        .and_then(|v| v.as_str())
326        .unwrap_or("");
327    let element_id = val
328        .get("currentElementID")
329        .and_then(|v| v.as_str())
330        .unwrap_or("");
331    let name = val
332        .get("currentElementName")
333        .and_then(|v| v.as_str())
334        .unwrap_or("");
335
336    debug!(session_id = %session.id, element = %name, "getNext: got element");
337
338    let response = json!({
339        "command": "visitedElement",
340        "modelId": val.get("modelId").unwrap_or(&json!("")),
341        "elementId": val.get("currentElementID").unwrap_or(&json!("")),
342        "name": val.get("currentElementName").unwrap_or(&json!("")),
343        "visitedCount": val.get("visitedCount").unwrap_or(&json!(0)),
344        "totalCount": val.get("totalCount").unwrap_or(&json!(0)),
345        "stopConditionFulfillment": val.get("stopConditionFulfillment").unwrap_or(&json!(0.0)),
346        "data": val.get("data").unwrap_or(&json!("")),
347        "success": true,
348    });
349
350    let subscribers = session.broadcast_tx.send(response.clone()).unwrap_or(0);
351    debug!(session_id = %session.id, subscribers, "getNext: broadcast sent");
352
353    if session
354        .control
355        .check_and_pause_if_breakpoint(model_id, element_id)
356    {
357        debug!(session_id = %session.id, element = %name, "breakpoint hit, pausing");
358        let _ = session.broadcast_tx.send(json!({
359            "command": "sessionPaused",
360            "reason": "breakpoint",
361            "modelId": model_id,
362            "elementId": element_id,
363        }));
364    }
365
366    let delay_ms = session.control.delay_ms();
367    if delay_ms > 0 {
368        debug!(session_id = %session.id, delay_ms, "getNext: sleeping");
369        tokio::time::sleep(Duration::from_millis(delay_ms)).await;
370    }
371
372    Ok(response)
373}
374
375async fn cmd_has_next(own_session: &Option<SessionHandle>) -> Result<Value, String> {
376    let session = own_session.as_ref().ok_or("No active session")?;
377    let val = send_command(&session.machine_tx, |reply| Command::HasNext { reply }).await?;
378    let has_next = val
379        .get("hasNext")
380        .and_then(|v| v.as_str())
381        .unwrap_or("false");
382
383    Ok(json!({
384        "command": "hasNext",
385        "hasNext": has_next == "true",
386        "success": true,
387    }))
388}
389
390async fn cmd_get_data(own_session: &Option<SessionHandle>) -> Result<Value, String> {
391    let session = own_session.as_ref().ok_or("No active session")?;
392    let val = send_command(&session.machine_tx, |reply| Command::GetData { reply }).await?;
393    Ok(json!({
394        "command": "getData",
395        "data": val.get("data").unwrap_or(&json!("")),
396        "success": true,
397    }))
398}
399
400async fn cmd_set_data(
401    own_session: &Option<SessionHandle>,
402    request: &Value,
403) -> Result<Value, String> {
404    let session = own_session.as_ref().ok_or("No active session")?;
405    let script = request
406        .get("action")
407        .and_then(|a| a.as_str())
408        .ok_or("Missing 'action' field")?
409        .to_string();
410    send_command(&session.machine_tx, |reply| Command::SetData { script, reply }).await?;
411    Ok(json!({"command": "setData", "success": true}))
412}
413
414async fn cmd_get_model(own_session: &Option<SessionHandle>) -> Result<Value, String> {
415    let session = own_session.as_ref().ok_or("No active session")?;
416    let val = send_command(&session.machine_tx, |reply| Command::GetModel { reply }).await?;
417    Ok(json!({
418        "command": "getModel",
419        "models": val.get("models").unwrap_or(&json!("")),
420        "success": true,
421    }))
422}
423
424async fn cmd_update_all_elements(own_session: &Option<SessionHandle>) -> Result<Value, String> {
425    let session = own_session.as_ref().ok_or("No active session")?;
426    let val =
427        send_command(&session.machine_tx, |reply| Command::UpdateAllElements { reply }).await?;
428    Ok(json!({
429        "command": "updateAllElements",
430        "elements": val.get("elements").unwrap_or(&json!([])),
431        "success": true,
432    }))
433}
434
435fn cmd_list_sessions(
436    session_mgr: &SessionManager,
437    change_rx: &mut Option<broadcast::Receiver<Value>>,
438) -> Result<Value, String> {
439    if change_rx.is_none() {
440        *change_rx = Some(session_mgr.subscribe_changes());
441    }
442
443    let sessions = session_mgr.list_sessions();
444    let list: Vec<Value> = sessions
445        .iter()
446        .map(|s| {
447            json!({
448                "id": s.id,
449                "name": s.name,
450            })
451        })
452        .collect();
453
454    debug!(count = sessions.len(), "listSessions");
455
456    Ok(json!({
457        "command": "sessions",
458        "sessions": list,
459        "success": true,
460    }))
461}
462
463async fn cmd_subscribe_session(
464    session_mgr: &SessionManager,
465    subscribed_rx: &mut Option<broadcast::Receiver<Value>>,
466    subscribed_session_id: &mut Option<String>,
467    request: &Value,
468) -> Result<Value, String> {
469    if let Some(prev_id) = subscribed_session_id.take() {
470        debug!(prev_session = %prev_id, "resetting previous subscription");
471        if let Some(prev_session) = session_mgr.get_session(&prev_id) {
472            prev_session.control.reset();
473        }
474    }
475
476    let session_id = request
477        .get("sessionId")
478        .and_then(|v| v.as_str())
479        .ok_or("Missing 'sessionId' field")?;
480
481    let session = session_mgr
482        .get_session(session_id)
483        .ok_or_else(|| format!("Session '{}' not found", session_id))?;
484
485    debug!(session_id = %session_id, "fetching element snapshot");
486
487    let elements = send_command(&session.machine_tx, |reply| Command::UpdateAllElements {
488        reply,
489    })
490    .await
491    .ok()
492    .and_then(|v| v.get("elements").cloned())
493    .unwrap_or(json!([]));
494
495    let models_value = send_command(&session.machine_tx, |reply| Command::GetModel { reply })
496        .await
497        .ok()
498        .and_then(|v| v.get("models").and_then(|m| m.as_str()).map(String::from))
499        .and_then(|s| serde_json::from_str::<Value>(&s).ok())
500        .unwrap_or_else(|| serde_json::from_str::<Value>(&session.model_json).unwrap_or(json!(null)));
501
502    let elem_count = elements.as_array().map(|a| a.len()).unwrap_or(0);
503    debug!(session_id = %session_id, elements = elem_count, "subscribing to broadcast");
504
505    *subscribed_rx = Some(session.broadcast_tx.subscribe());
506    *subscribed_session_id = Some(session_id.to_string());
507
508    Ok(json!({
509        "command": "subscribeSession",
510        "sessionId": session.id,
511        "name": session.name,
512        "models": models_value,
513        "elements": elements,
514        "seed": session.seed,
515        "paused": session.control.is_paused(),
516        "success": true,
517    }))
518}
519
520fn cmd_unsubscribe_session(
521    session_mgr: &SessionManager,
522    subscribed_rx: &mut Option<broadcast::Receiver<Value>>,
523    subscribed_session_id: &mut Option<String>,
524) -> Result<Value, String> {
525    if let Some(sid) = subscribed_session_id.take() {
526        debug!(session_id = %sid, "unsubscribing, resetting control");
527        if let Some(session) = session_mgr.get_session(&sid) {
528            session.control.reset();
529        }
530    }
531    *subscribed_rx = None;
532    Ok(json!({
533        "command": "unsubscribeSession",
534        "success": true,
535    }))
536}
537
538fn cmd_pause_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
539    let session = get_target_session(session_mgr, request)?;
540    debug!(session_id = %session.id, "pauseSession");
541    session.control.pause();
542    let _ = session.broadcast_tx.send(json!({
543        "command": "sessionPaused",
544        "sessionId": session.id,
545    }));
546    Ok(json!({"command": "pauseSession", "success": true}))
547}
548
549fn cmd_resume_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
550    let session = get_target_session(session_mgr, request)?;
551    debug!(session_id = %session.id, "resumeSession");
552    session.control.resume();
553    let _ = session.broadcast_tx.send(json!({
554        "command": "sessionResumed",
555        "sessionId": session.id,
556    }));
557    Ok(json!({"command": "resumeSession", "success": true}))
558}
559
560fn cmd_step_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
561    let session = get_target_session(session_mgr, request)?;
562    debug!(session_id = %session.id, "stepSession (pause + step)");
563    session.control.pause();
564    session.control.step();
565    Ok(json!({"command": "stepSession", "success": true}))
566}
567
568fn cmd_set_delay(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
569    let session = get_target_session(session_mgr, request)?;
570    let ms = request
571        .get("value")
572        .and_then(|v| v.as_u64())
573        .unwrap_or(0);
574    debug!(session_id = %session.id, delay_ms = ms, "setDelay");
575    session.control.set_delay(ms);
576    Ok(json!({"command": "setDelay", "success": true}))
577}
578
579fn cmd_set_breakpoints(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
580    let session = get_target_session(session_mgr, request)?;
581    let bps: HashSet<String> = request
582        .get("breakpoints")
583        .and_then(|v| v.as_array())
584        .map(|arr| {
585            arr.iter()
586                .filter_map(|v| v.as_str().map(|s| s.to_string()))
587                .collect()
588        })
589        .unwrap_or_default();
590    debug!(session_id = %session.id, count = bps.len(), "setBreakpoints");
591    session.control.set_breakpoints(bps);
592    Ok(json!({"command": "setBreakpoints", "success": true}))
593}