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}