Skip to main content

graphwalker_restful/
session.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::atomic::{AtomicU64, Ordering};
3use std::sync::{Arc, Mutex, RwLock};
4
5use serde_json::{json, Value};
6use tokio::sync::{broadcast, Notify};
7
8use crate::actor::Command;
9
10pub struct ExecutionControl {
11    inner: Mutex<ControlState>,
12    notify: Notify,
13}
14
15struct ControlState {
16    paused: bool,
17    step_once: bool,
18    delay_ms: u64,
19    breakpoints: HashSet<String>,
20}
21
22impl ExecutionControl {
23    pub fn new() -> Self {
24        Self {
25            inner: Mutex::new(ControlState {
26                paused: false,
27                step_once: false,
28                delay_ms: 0,
29                breakpoints: HashSet::new(),
30            }),
31            notify: Notify::new(),
32        }
33    }
34
35    pub async fn gate(&self) {
36        loop {
37            {
38                let mut state = self.inner.lock().unwrap();
39                if !state.paused {
40                    return;
41                }
42                if state.step_once {
43                    state.step_once = false;
44                    return;
45                }
46            }
47            self.notify.notified().await;
48        }
49    }
50
51    pub fn pause(&self) {
52        self.inner.lock().unwrap().paused = true;
53    }
54
55    pub fn resume(&self) {
56        self.inner.lock().unwrap().paused = false;
57        self.notify.notify_waiters();
58    }
59
60    pub fn step(&self) {
61        self.inner.lock().unwrap().step_once = true;
62        self.notify.notify_waiters();
63    }
64
65    pub fn set_delay(&self, ms: u64) {
66        self.inner.lock().unwrap().delay_ms = ms;
67    }
68
69    pub fn delay_ms(&self) -> u64 {
70        self.inner.lock().unwrap().delay_ms
71    }
72
73    pub fn set_breakpoints(&self, bps: HashSet<String>) {
74        self.inner.lock().unwrap().breakpoints = bps;
75    }
76
77    pub fn check_and_pause_if_breakpoint(&self, model_id: &str, element_id: &str) -> bool {
78        let key = format!("{},{}", model_id, element_id);
79        let mut state = self.inner.lock().unwrap();
80        if state.breakpoints.contains(&key) {
81            state.paused = true;
82            true
83        } else {
84            false
85        }
86    }
87
88    pub fn is_paused(&self) -> bool {
89        self.inner.lock().unwrap().paused
90    }
91
92    pub fn reset(&self) {
93        let mut state = self.inner.lock().unwrap();
94        state.paused = false;
95        state.step_once = false;
96        state.delay_ms = 0;
97        state.breakpoints.clear();
98        drop(state);
99        self.notify.notify_waiters();
100    }
101}
102
103#[derive(Clone)]
104pub struct SessionHandle {
105    pub id: String,
106    pub name: String,
107    pub model_json: String,
108    pub seed: Option<u64>,
109    pub machine_tx: std::sync::mpsc::Sender<Command>,
110    pub broadcast_tx: broadcast::Sender<Value>,
111    pub control: Arc<ExecutionControl>,
112}
113
114#[derive(Clone)]
115pub struct SessionManager {
116    sessions: Arc<RwLock<HashMap<String, SessionHandle>>>,
117    counter: Arc<AtomicU64>,
118    change_tx: broadcast::Sender<Value>,
119}
120
121impl SessionManager {
122    pub fn new() -> Self {
123        let (change_tx, _) = broadcast::channel(64);
124        Self {
125            sessions: Arc::new(RwLock::new(HashMap::new())),
126            counter: Arc::new(AtomicU64::new(1)),
127            change_tx,
128        }
129    }
130
131    pub fn create_session(
132        &self,
133        name: String,
134        model_json: String,
135        seed: Option<u64>,
136        machine_tx: std::sync::mpsc::Sender<Command>,
137    ) -> SessionHandle {
138        let id = format!("session-{}", self.counter.fetch_add(1, Ordering::Relaxed));
139        let (broadcast_tx, _) = broadcast::channel(256);
140        let handle = SessionHandle {
141            id: id.clone(),
142            name: name.clone(),
143            model_json,
144            seed,
145            machine_tx,
146            broadcast_tx,
147            control: Arc::new(ExecutionControl::new()),
148        };
149        self.sessions
150            .write()
151            .unwrap()
152            .insert(id.clone(), handle.clone());
153        let _ = self.change_tx.send(json!({
154            "command": "sessionCreated",
155            "sessionId": id,
156            "name": name,
157        }));
158        handle
159    }
160
161    pub fn remove_session(&self, id: &str) {
162        if let Some(session) = self.sessions.write().unwrap().remove(id) {
163            session.control.reset();
164            let _ = self.change_tx.send(json!({
165                "command": "sessionEnded",
166                "sessionId": id,
167            }));
168        }
169    }
170
171    pub fn list_sessions(&self) -> Vec<SessionHandle> {
172        self.sessions.read().unwrap().values().cloned().collect()
173    }
174
175    pub fn get_session(&self, id: &str) -> Option<SessionHandle> {
176        self.sessions.read().unwrap().get(id).cloned()
177    }
178
179    pub fn subscribe_changes(&self) -> broadcast::Receiver<Value> {
180        self.change_tx.subscribe()
181    }
182}