graphwalker_restful/
session.rs1use 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}