1use std::sync::Arc;
15
16use dashmap::DashMap;
17use tokio::sync::{broadcast, mpsc, Mutex};
18use trustee_core::session::Session;
19use trustee_core::types::TuiMessage;
20
21use crate::auth::AuthState;
22
23pub struct UserSession {
29 pub session: Arc<Mutex<Session>>,
31 pub ws_tx: broadcast::Sender<String>,
33 pub token_store: Arc<pep::MemoryTokenStore>,
39}
40
41impl UserSession {
42 pub fn new(session: Session) -> Self {
47 let (ws_tx, _ws_rx) = broadcast::channel::<String>(256);
48 let token_store = Arc::new(pep::MemoryTokenStore::new());
49 Self {
50 session: Arc::new(Mutex::new(session)),
51 ws_tx,
52 token_store,
53 }
54 }
55}
56
57pub type SessionRegistry = Arc<DashMap<String, UserSession>>;
67
68#[derive(Clone)]
70pub struct ServerState {
71 pub sessions: SessionRegistry,
73 pub ws_tx: broadcast::Sender<String>,
76 pub auth: Option<Arc<AuthState>>,
78 pub config_toml: Option<String>,
80 pub workflow_semaphore: Arc<tokio::sync::Semaphore>,
84}
85
86impl ServerState {
87 pub fn new(
92 session: Session,
93 ws_tx: broadcast::Sender<String>,
94 auth: Option<Arc<AuthState>>,
95 ) -> Self {
96 let sessions = Arc::new(DashMap::new());
97
98 let token_store = Arc::new(pep::MemoryTokenStore::new());
101 sessions.insert(
102 "default".to_string(),
103 UserSession {
104 session: Arc::new(Mutex::new(session)),
105 ws_tx: ws_tx.clone(),
106 token_store,
107 },
108 );
109
110 Self {
111 sessions,
112 ws_tx,
113 auth,
114 config_toml: None,
115 workflow_semaphore: Arc::new(tokio::sync::Semaphore::new(8)),
116 }
117 }
118
119 pub fn with_config_toml(mut self, config_toml: String) -> Self {
121 self.config_toml = Some(config_toml);
122 self
123 }
124
125 pub fn with_max_concurrent_workflows(mut self, max: usize) -> Self {
127 self.workflow_semaphore = Arc::new(tokio::sync::Semaphore::new(max));
128 self
129 }
130
131 pub async fn ensure_user_session(
137 &self,
138 user_key: &str,
139 ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
140 if let Some(entry) = self.sessions.get(user_key) {
142 return (
143 entry.session.clone(),
144 entry.ws_tx.clone(),
145 entry.token_store.clone(),
146 );
147 }
148
149 let (mut session, workflow_rx) = Session::new();
151
152 if let Some(ref config_toml) = self.config_toml {
154 session.config_toml = Some(config_toml.clone());
155 session.parse_auto_handoff_config();
156
157 if let Ok(table) = config_toml.parse::<toml::Value>() {
158 if let Some(name) = table.get("agent").and_then(|a| a.get("name")).and_then(|n| n.as_str()) {
159 session.agent_name = name.to_string();
160 }
161 }
162 }
163
164 if user_key != "default" {
170 session.project_id = Some(format!("user:{user_key}"));
171 }
172
173 let user_session = UserSession::new(session);
174 let session_arc = user_session.session.clone();
175 let ws_tx = user_session.ws_tx.clone();
176 let token_store = user_session.token_store.clone();
177
178 self.sessions.insert(user_key.to_string(), user_session);
179
180 self.spawn_user_drain_task(
182 user_key.to_string(),
183 session_arc.clone(),
184 ws_tx.clone(),
185 workflow_rx,
186 );
187
188 (session_arc, ws_tx, token_store)
189 }
190
191 fn spawn_user_drain_task(
195 &self,
196 user_key: String,
197 session: Arc<Mutex<Session>>,
198 ws_tx: broadcast::Sender<String>,
199 mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
200 ) {
201 tokio::spawn(async move {
202 while let Some(msg) = workflow_rx.recv().await {
203 {
205 let mut session = session.lock().await;
206 session.handle_workflow_message(msg.clone());
207
208 let state_str = match session.workflow_state {
209 trustee_core::types::WorkflowState::Idle => "Idle",
210 trustee_core::types::WorkflowState::Running => "Running",
211 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
212 };
213 let state_msg = serde_json::json!({
214 "type": "StateChanged",
215 "state": state_str
216 });
217 let _ = ws_tx.send(state_msg.to_string());
218 }
219
220 let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
222 let _ = ws_tx.send(json);
223 }
224 tracing::debug!("Drain task ended for user: {}", user_key);
225 });
226 }
227
228 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
232 let default_entry = self.sessions.get("default").expect("default session must exist");
234 let session = default_entry.session.clone();
235 let ws_tx = default_entry.ws_tx.clone();
236 drop(default_entry);
237
238 tokio::spawn(async move {
239 while let Some(msg) = workflow_rx.recv().await {
240 {
241 let mut session = session.lock().await;
242 session.handle_workflow_message(msg.clone());
243
244 let state_str = match session.workflow_state {
245 trustee_core::types::WorkflowState::Idle => "Idle",
246 trustee_core::types::WorkflowState::Running => "Running",
247 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
248 };
249 let state_msg = serde_json::json!({
250 "type": "StateChanged",
251 "state": state_str
252 });
253 let _ = ws_tx.send(state_msg.to_string());
254 }
255
256 let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
257 let _ = ws_tx.send(json);
258 }
259 });
260 }
261
262 pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
267 let Some(ref auth) = self.auth else {
268 return "default".to_string();
269 };
270
271 if let Some(token) = headers
273 .get(axum::http::header::AUTHORIZATION)
274 .and_then(|v| v.to_str().ok())
275 .and_then(|v| v.strip_prefix("Bearer "))
276 .map(|s| s.to_string())
277 {
278 if token.starts_with("dev:") {
280 let parts: Vec<&str> = token.splitn(4, ':').collect();
281 if parts.len() >= 4 {
282 return format!("dev:{}", parts[1]);
283 }
284 }
285 if let Ok(claims) = auth.validate_token(&token).await {
287 return claims.sub;
288 }
289 }
290
291 let cookie_session_id = headers
293 .get(axum::http::header::COOKIE)
294 .and_then(|v| v.to_str().ok())
295 .and_then(|cookies| {
296 cookies
297 .split(';')
298 .map(|c| c.trim())
299 .find_map(|c| c.strip_prefix(&format!("{}=", auth.config.cookie_name)))
300 .map(|s| s.to_string())
301 });
302
303 if let Some(session_id) = cookie_session_id {
304 if session_id.starts_with("dev:") {
306 let parts: Vec<&str> = session_id.splitn(4, ':').collect();
307 if parts.len() >= 4 {
308 return format!("dev:{}", parts[1]);
309 }
310 }
311
312 if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
314 if let Ok(claims) = auth.validate_token(&access_token).await {
315 return claims.sub;
316 }
317 }
318 }
319
320 "default".to_string()
321 }
322}
323
324struct SerializableMessage<'a>(&'a TuiMessage);
326
327impl<'a> serde::Serialize for SerializableMessage<'a> {
328 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
329 where
330 S: serde::Serializer,
331 {
332 use serde::ser::SerializeStruct;
333
334 match self.0 {
335 TuiMessage::OutputLine(line) => {
336 let mut s = serializer.serialize_struct("msg", 2)?;
337 s.serialize_field("type", "OutputLine")?;
338 s.serialize_field("line", line)?;
339 s.end()
340 }
341 TuiMessage::StreamDelta(delta) => {
342 let mut s = serializer.serialize_struct("msg", 2)?;
343 s.serialize_field("type", "StreamDelta")?;
344 s.serialize_field("delta", delta)?;
345 s.end()
346 }
347 TuiMessage::ReasoningDelta(delta) => {
348 let mut s = serializer.serialize_struct("msg", 2)?;
349 s.serialize_field("type", "ReasoningDelta")?;
350 s.serialize_field("delta", delta)?;
351 s.end()
352 }
353 TuiMessage::WorkflowCompleted => {
354 let mut s = serializer.serialize_struct("msg", 2)?;
355 s.serialize_field("type", "WorkflowCompleted")?;
356 s.serialize_field("state", "Idle")?;
357 s.end()
358 }
359 TuiMessage::WorkflowError(err) => {
360 let mut s = serializer.serialize_struct("msg", 2)?;
361 s.serialize_field("type", "WorkflowError")?;
362 s.serialize_field("error", err)?;
363 s.end()
364 }
365 TuiMessage::ResumeInfo(_) => {
366 let mut s = serializer.serialize_struct("msg", 2)?;
367 s.serialize_field("type", "ResumeInfo")?;
368 s.serialize_field("state", "Idle")?;
369 s.end()
370 }
371 TuiMessage::TodoUpdate(content) => {
372 let mut s = serializer.serialize_struct("msg", 2)?;
373 s.serialize_field("type", "TodoUpdate")?;
374 s.serialize_field("content", content)?;
375 s.end()
376 }
377 TuiMessage::WorkflowCancelled => {
378 let mut s = serializer.serialize_struct("msg", 2)?;
379 s.serialize_field("type", "WorkflowCancelled")?;
380 s.serialize_field("state", "Idle")?;
381 s.end()
382 }
383 TuiMessage::HandoffReady(_) => {
384 let mut s = serializer.serialize_struct("msg", 2)?;
385 s.serialize_field("type", "HandoffReady")?;
386 s.serialize_field("state", "Idle")?;
387 s.end()
388 }
389 TuiMessage::ToolPending { tool_name, hint } => {
390 let mut s = serializer.serialize_struct("msg", 3)?;
391 s.serialize_field("type", "ToolPending")?;
392 s.serialize_field("tool_name", tool_name)?;
393 s.serialize_field("hint", hint)?;
394 s.end()
395 }
396 TuiMessage::ToolDone { tool_name, success, hint } => {
397 let mut s = serializer.serialize_struct("msg", 4)?;
398 s.serialize_field("type", "ToolDone")?;
399 s.serialize_field("tool_name", tool_name)?;
400 s.serialize_field("success", success)?;
401 s.serialize_field("hint", hint)?;
402 s.end()
403 }
404 TuiMessage::ContextTokensUpdated(count) => {
405 let mut s = serializer.serialize_struct("msg", 2)?;
406 s.serialize_field("type", "ContextTokensUpdated")?;
407 s.serialize_field("count", count)?;
408 s.end()
409 }
410 TuiMessage::McpServerStatus { name, connected, tool_count, error } => {
411 let mut s = serializer.serialize_struct("msg", 5)?;
412 s.serialize_field("type", "McpServerStatus")?;
413 s.serialize_field("name", name)?;
414 s.serialize_field("connected", connected)?;
415 s.serialize_field("tool_count", tool_count)?;
416 s.serialize_field("error", error)?;
417 s.end()
418 }
419 }
420 }
421}