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.get("modelId").and_then(|v| v.as_str()).unwrap_or("");
324 let element_id = val
325 .get("currentElementID")
326 .and_then(|v| v.as_str())
327 .unwrap_or("");
328 let name = val
329 .get("currentElementName")
330 .and_then(|v| v.as_str())
331 .unwrap_or("");
332
333 debug!(session_id = %session.id, element = %name, "getNext: got element");
334
335 let response = json!({
336 "command": "visitedElement",
337 "modelId": val.get("modelId").unwrap_or(&json!("")),
338 "elementId": val.get("currentElementID").unwrap_or(&json!("")),
339 "name": val.get("currentElementName").unwrap_or(&json!("")),
340 "visitedCount": val.get("visitedCount").unwrap_or(&json!(0)),
341 "totalCount": val.get("totalCount").unwrap_or(&json!(0)),
342 "stopConditionFulfillment": val.get("stopConditionFulfillment").unwrap_or(&json!(0.0)),
343 "data": val.get("data").unwrap_or(&json!("")),
344 "success": true,
345 });
346
347 let subscribers = session.broadcast_tx.send(response.clone()).unwrap_or(0);
348 debug!(session_id = %session.id, subscribers, "getNext: broadcast sent");
349
350 if session
351 .control
352 .check_and_pause_if_breakpoint(model_id, element_id)
353 {
354 debug!(session_id = %session.id, element = %name, "breakpoint hit, pausing");
355 let _ = session.broadcast_tx.send(json!({
356 "command": "sessionPaused",
357 "reason": "breakpoint",
358 "modelId": model_id,
359 "elementId": element_id,
360 }));
361 }
362
363 let delay_ms = session.control.delay_ms();
364 if delay_ms > 0 {
365 debug!(session_id = %session.id, delay_ms, "getNext: sleeping");
366 tokio::time::sleep(Duration::from_millis(delay_ms)).await;
367 }
368
369 Ok(response)
370}
371
372async fn cmd_has_next(own_session: &Option<SessionHandle>) -> Result<Value, String> {
373 let session = own_session.as_ref().ok_or("No active session")?;
374 let val = send_command(&session.machine_tx, |reply| Command::HasNext { reply }).await?;
375 let has_next = val
376 .get("hasNext")
377 .and_then(|v| v.as_str())
378 .unwrap_or("false");
379
380 Ok(json!({
381 "command": "hasNext",
382 "hasNext": has_next == "true",
383 "success": true,
384 }))
385}
386
387async fn cmd_get_data(own_session: &Option<SessionHandle>) -> Result<Value, String> {
388 let session = own_session.as_ref().ok_or("No active session")?;
389 let val = send_command(&session.machine_tx, |reply| Command::GetData { reply }).await?;
390 Ok(json!({
391 "command": "getData",
392 "data": val.get("data").unwrap_or(&json!("")),
393 "success": true,
394 }))
395}
396
397async fn cmd_set_data(
398 own_session: &Option<SessionHandle>,
399 request: &Value,
400) -> Result<Value, String> {
401 let session = own_session.as_ref().ok_or("No active session")?;
402 let script = request
403 .get("action")
404 .and_then(|a| a.as_str())
405 .ok_or("Missing 'action' field")?
406 .to_string();
407 send_command(&session.machine_tx, |reply| Command::SetData {
408 script,
409 reply,
410 })
411 .await?;
412 Ok(json!({"command": "setData", "success": true}))
413}
414
415async fn cmd_get_model(own_session: &Option<SessionHandle>) -> Result<Value, String> {
416 let session = own_session.as_ref().ok_or("No active session")?;
417 let val = send_command(&session.machine_tx, |reply| Command::GetModel { reply }).await?;
418 Ok(json!({
419 "command": "getModel",
420 "models": val.get("models").unwrap_or(&json!("")),
421 "success": true,
422 }))
423}
424
425async fn cmd_update_all_elements(own_session: &Option<SessionHandle>) -> Result<Value, String> {
426 let session = own_session.as_ref().ok_or("No active session")?;
427 let val = send_command(&session.machine_tx, |reply| Command::UpdateAllElements {
428 reply,
429 })
430 .await?;
431 Ok(json!({
432 "command": "updateAllElements",
433 "elements": val.get("elements").unwrap_or(&json!([])),
434 "success": true,
435 }))
436}
437
438fn cmd_list_sessions(
439 session_mgr: &SessionManager,
440 change_rx: &mut Option<broadcast::Receiver<Value>>,
441) -> Result<Value, String> {
442 if change_rx.is_none() {
443 *change_rx = Some(session_mgr.subscribe_changes());
444 }
445
446 let sessions = session_mgr.list_sessions();
447 let list: Vec<Value> = sessions
448 .iter()
449 .map(|s| {
450 json!({
451 "id": s.id,
452 "name": s.name,
453 })
454 })
455 .collect();
456
457 debug!(count = sessions.len(), "listSessions");
458
459 Ok(json!({
460 "command": "sessions",
461 "sessions": list,
462 "success": true,
463 }))
464}
465
466async fn cmd_subscribe_session(
467 session_mgr: &SessionManager,
468 subscribed_rx: &mut Option<broadcast::Receiver<Value>>,
469 subscribed_session_id: &mut Option<String>,
470 request: &Value,
471) -> Result<Value, String> {
472 if let Some(prev_id) = subscribed_session_id.take() {
473 debug!(prev_session = %prev_id, "resetting previous subscription");
474 if let Some(prev_session) = session_mgr.get_session(&prev_id) {
475 prev_session.control.reset();
476 }
477 }
478
479 let session_id = request
480 .get("sessionId")
481 .and_then(|v| v.as_str())
482 .ok_or("Missing 'sessionId' field")?;
483
484 let session = session_mgr
485 .get_session(session_id)
486 .ok_or_else(|| format!("Session '{}' not found", session_id))?;
487
488 debug!(session_id = %session_id, "fetching element snapshot");
489
490 let elements = send_command(&session.machine_tx, |reply| Command::UpdateAllElements {
491 reply,
492 })
493 .await
494 .ok()
495 .and_then(|v| v.get("elements").cloned())
496 .unwrap_or(json!([]));
497
498 let models_value = send_command(&session.machine_tx, |reply| Command::GetModel { reply })
499 .await
500 .ok()
501 .and_then(|v| v.get("models").and_then(|m| m.as_str()).map(String::from))
502 .and_then(|s| serde_json::from_str::<Value>(&s).ok())
503 .unwrap_or_else(|| {
504 serde_json::from_str::<Value>(&session.model_json).unwrap_or(json!(null))
505 });
506
507 let elem_count = elements.as_array().map(|a| a.len()).unwrap_or(0);
508 debug!(session_id = %session_id, elements = elem_count, "subscribing to broadcast");
509
510 *subscribed_rx = Some(session.broadcast_tx.subscribe());
511 *subscribed_session_id = Some(session_id.to_string());
512
513 Ok(json!({
514 "command": "subscribeSession",
515 "sessionId": session.id,
516 "name": session.name,
517 "models": models_value,
518 "elements": elements,
519 "seed": session.seed,
520 "paused": session.control.is_paused(),
521 "success": true,
522 }))
523}
524
525fn cmd_unsubscribe_session(
526 session_mgr: &SessionManager,
527 subscribed_rx: &mut Option<broadcast::Receiver<Value>>,
528 subscribed_session_id: &mut Option<String>,
529) -> Result<Value, String> {
530 if let Some(sid) = subscribed_session_id.take() {
531 debug!(session_id = %sid, "unsubscribing, resetting control");
532 if let Some(session) = session_mgr.get_session(&sid) {
533 session.control.reset();
534 }
535 }
536 *subscribed_rx = None;
537 Ok(json!({
538 "command": "unsubscribeSession",
539 "success": true,
540 }))
541}
542
543fn cmd_pause_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
544 let session = get_target_session(session_mgr, request)?;
545 debug!(session_id = %session.id, "pauseSession");
546 session.control.pause();
547 let _ = session.broadcast_tx.send(json!({
548 "command": "sessionPaused",
549 "sessionId": session.id,
550 }));
551 Ok(json!({"command": "pauseSession", "success": true}))
552}
553
554fn cmd_resume_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
555 let session = get_target_session(session_mgr, request)?;
556 debug!(session_id = %session.id, "resumeSession");
557 session.control.resume();
558 let _ = session.broadcast_tx.send(json!({
559 "command": "sessionResumed",
560 "sessionId": session.id,
561 }));
562 Ok(json!({"command": "resumeSession", "success": true}))
563}
564
565fn cmd_step_session(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
566 let session = get_target_session(session_mgr, request)?;
567 debug!(session_id = %session.id, "stepSession (pause + step)");
568 session.control.pause();
569 session.control.step();
570 Ok(json!({"command": "stepSession", "success": true}))
571}
572
573fn cmd_set_delay(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
574 let session = get_target_session(session_mgr, request)?;
575 let ms = request.get("value").and_then(|v| v.as_u64()).unwrap_or(0);
576 debug!(session_id = %session.id, delay_ms = ms, "setDelay");
577 session.control.set_delay(ms);
578 Ok(json!({"command": "setDelay", "success": true}))
579}
580
581fn cmd_set_breakpoints(session_mgr: &SessionManager, request: &Value) -> Result<Value, String> {
582 let session = get_target_session(session_mgr, request)?;
583 let bps: HashSet<String> = request
584 .get("breakpoints")
585 .and_then(|v| v.as_array())
586 .map(|arr| {
587 arr.iter()
588 .filter_map(|v| v.as_str().map(|s| s.to_string()))
589 .collect()
590 })
591 .unwrap_or_default();
592 debug!(session_id = %session.id, count = bps.len(), "setBreakpoints");
593 session.control.set_breakpoints(bps);
594 Ok(json!({"command": "setBreakpoints", "success": true}))
595}