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 secrets: Option<std::collections::HashMap<String, String>>,
82 pub build_info: Option<trustee_core::types::BuildInfo>,
84 pub workflow_semaphore: Arc<tokio::sync::Semaphore>,
88}
89
90impl ServerState {
91 pub fn new(
96 session: Session,
97 ws_tx: broadcast::Sender<String>,
98 auth: Option<Arc<AuthState>>,
99 ) -> Self {
100 let sessions = Arc::new(DashMap::new());
101
102 let token_store = Arc::new(pep::MemoryTokenStore::new());
105 sessions.insert(
106 "default".to_string(),
107 UserSession {
108 session: Arc::new(Mutex::new(session)),
109 ws_tx: ws_tx.clone(),
110 token_store,
111 },
112 );
113
114 Self {
115 sessions,
116 ws_tx,
117 auth,
118 config_toml: None,
119 secrets: None,
120 build_info: None,
121 workflow_semaphore: Arc::new(tokio::sync::Semaphore::new(8)),
122 }
123 }
124
125 pub fn with_config_toml(mut self, config_toml: String) -> Self {
127 self.config_toml = Some(config_toml);
128 self
129 }
130
131 pub fn with_secrets(mut self, secrets: std::collections::HashMap<String, String>) -> Self {
133 self.secrets = Some(secrets);
134 self
135 }
136
137 pub fn with_build_info(mut self, build_info: trustee_core::types::BuildInfo) -> Self {
139 self.build_info = Some(build_info);
140 self
141 }
142
143 pub fn with_max_concurrent_workflows(mut self, max: usize) -> Self {
145 self.workflow_semaphore = Arc::new(tokio::sync::Semaphore::new(max));
146 self
147 }
148
149 pub async fn ensure_user_session(
155 &self,
156 user_key: &str,
157 ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
158 if let Some(entry) = self.sessions.get(user_key) {
160 return (
161 entry.session.clone(),
162 entry.ws_tx.clone(),
163 entry.token_store.clone(),
164 );
165 }
166
167 let (mut session, workflow_rx) = Session::new();
169
170 if let Some(ref config_toml) = self.config_toml {
172 session.config_toml = Some(config_toml.clone());
173 session.parse_auto_handoff_config();
174
175 if let Ok(table) = config_toml.parse::<toml::Value>() {
176 if let Some(name) = table.get("agent").and_then(|a| a.get("name")).and_then(|n| n.as_str()) {
177 session.agent_name = name.to_string();
178 }
179 }
180 }
181
182 session.secrets = self.secrets.clone();
184 session.build_info = self.build_info.clone();
185
186 if user_key != "default" {
191 use sha2::{Digest, Sha256};
193 let mut hasher = Sha256::new();
194 hasher.update(user_key.as_bytes());
195 let hash_bytes = hasher.finalize();
196 let user_hash = format!("{:016x}", u64::from_be_bytes(hash_bytes[..8].try_into().unwrap()));
197
198 if let Some(home) = dirs::home_dir() {
200 session.home_dir = Some(home.join(".trustee").join("users").join(&user_hash));
201 }
202
203 session.project_id = Some(uuid::Uuid::new_v4().to_string());
205 }
206
207 let user_session = UserSession::new(session);
208 let session_arc = user_session.session.clone();
209 let ws_tx = user_session.ws_tx.clone();
210 let token_store = user_session.token_store.clone();
211
212 self.sessions.insert(user_key.to_string(), user_session);
213
214 self.spawn_user_drain_task(
216 user_key.to_string(),
217 session_arc.clone(),
218 ws_tx.clone(),
219 workflow_rx,
220 );
221
222 (session_arc, ws_tx, token_store)
223 }
224
225 fn spawn_user_drain_task(
229 &self,
230 user_key: String,
231 session: Arc<Mutex<Session>>,
232 ws_tx: broadcast::Sender<String>,
233 mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
234 ) {
235 tokio::spawn(async move {
236 while let Some(msg) = workflow_rx.recv().await {
237 {
239 let mut session = session.lock().await;
240 session.handle_workflow_message(msg.clone());
241
242 let state_str = match session.workflow_state {
243 trustee_core::types::WorkflowState::Idle => "Idle",
244 trustee_core::types::WorkflowState::Running => "Running",
245 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
246 };
247 let state_msg = serde_json::json!({
248 "type": "StateChanged",
249 "state": state_str
250 });
251 let _ = ws_tx.send(state_msg.to_string());
252 }
253
254 let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
256 let _ = ws_tx.send(json);
257 }
258 tracing::debug!("Drain task ended for user: {}", user_key);
259 });
260 }
261
262 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
266 let default_entry = self.sessions.get("default").expect("default session must exist");
268 let session = default_entry.session.clone();
269 let ws_tx = default_entry.ws_tx.clone();
270 drop(default_entry);
271
272 tokio::spawn(async move {
273 while let Some(msg) = workflow_rx.recv().await {
274 {
275 let mut session = session.lock().await;
276 session.handle_workflow_message(msg.clone());
277
278 let state_str = match session.workflow_state {
279 trustee_core::types::WorkflowState::Idle => "Idle",
280 trustee_core::types::WorkflowState::Running => "Running",
281 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
282 };
283 let state_msg = serde_json::json!({
284 "type": "StateChanged",
285 "state": state_str
286 });
287 let _ = ws_tx.send(state_msg.to_string());
288 }
289
290 let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
291 let _ = ws_tx.send(json);
292 }
293 });
294 }
295
296 pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
301 let Some(ref auth) = self.auth else {
302 return "default".to_string();
303 };
304
305 if let Some(token) = headers
307 .get(axum::http::header::AUTHORIZATION)
308 .and_then(|v| v.to_str().ok())
309 .and_then(|v| v.strip_prefix("Bearer "))
310 .map(|s| s.to_string())
311 {
312 if token.starts_with("dev:") {
314 let parts: Vec<&str> = token.splitn(4, ':').collect();
315 if parts.len() >= 4 {
316 return format!("dev:{}", parts[1]);
317 }
318 }
319 if let Ok(claims) = auth.validate_token(&token).await {
321 return claims.sub;
322 }
323 }
324
325 let cookie_session_id = headers
327 .get(axum::http::header::COOKIE)
328 .and_then(|v| v.to_str().ok())
329 .and_then(|cookies| {
330 cookies
331 .split(';')
332 .map(|c| c.trim())
333 .find_map(|c| c.strip_prefix(&format!("{}=", auth.config.cookie_name)))
334 .map(|s| s.to_string())
335 });
336
337 if let Some(session_id) = cookie_session_id {
338 if session_id.starts_with("dev:") {
340 let parts: Vec<&str> = session_id.splitn(4, ':').collect();
341 if parts.len() >= 4 {
342 return format!("dev:{}", parts[1]);
343 }
344 }
345
346 if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
348 if let Ok(claims) = auth.validate_token(&access_token).await {
349 return claims.sub;
350 }
351 }
352 }
353
354 "default".to_string()
355 }
356}
357
358struct SerializableMessage<'a>(&'a TuiMessage);
360
361impl<'a> serde::Serialize for SerializableMessage<'a> {
362 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
363 where
364 S: serde::Serializer,
365 {
366 use serde::ser::SerializeStruct;
367
368 match self.0 {
369 TuiMessage::OutputLine(line) => {
370 let mut s = serializer.serialize_struct("msg", 2)?;
371 s.serialize_field("type", "OutputLine")?;
372 s.serialize_field("line", line)?;
373 s.end()
374 }
375 TuiMessage::StreamDelta(delta) => {
376 let mut s = serializer.serialize_struct("msg", 2)?;
377 s.serialize_field("type", "StreamDelta")?;
378 s.serialize_field("delta", delta)?;
379 s.end()
380 }
381 TuiMessage::ReasoningDelta(delta) => {
382 let mut s = serializer.serialize_struct("msg", 2)?;
383 s.serialize_field("type", "ReasoningDelta")?;
384 s.serialize_field("delta", delta)?;
385 s.end()
386 }
387 TuiMessage::WorkflowCompleted => {
388 let mut s = serializer.serialize_struct("msg", 2)?;
389 s.serialize_field("type", "WorkflowCompleted")?;
390 s.serialize_field("state", "Idle")?;
391 s.end()
392 }
393 TuiMessage::WorkflowError(err) => {
394 let mut s = serializer.serialize_struct("msg", 2)?;
395 s.serialize_field("type", "WorkflowError")?;
396 s.serialize_field("error", err)?;
397 s.end()
398 }
399 TuiMessage::ResumeInfo(info) => {
400 match info {
401 Some(ri) => {
402 let mut s = serializer.serialize_struct("msg", 5)?;
403 s.serialize_field("type", "ResumeInfo")?;
404 s.serialize_field("state", "Idle")?;
405 s.serialize_field("session_id", &ri.session_id)?;
406 s.serialize_field("checkpoint_id", &ri.checkpoint_id)?;
407 s.serialize_field("iteration", &ri.iteration)?;
408 s.end()
409 }
410 None => {
411 let mut s = serializer.serialize_struct("msg", 2)?;
412 s.serialize_field("type", "ResumeInfo")?;
413 s.serialize_field("state", "Idle")?;
414 s.end()
415 }
416 }
417 }
418 TuiMessage::TodoUpdate(content) => {
419 let mut s = serializer.serialize_struct("msg", 2)?;
420 s.serialize_field("type", "TodoUpdate")?;
421 s.serialize_field("content", content)?;
422 s.end()
423 }
424 TuiMessage::WorkflowCancelled => {
425 let mut s = serializer.serialize_struct("msg", 2)?;
426 s.serialize_field("type", "WorkflowCancelled")?;
427 s.serialize_field("state", "Idle")?;
428 s.end()
429 }
430 TuiMessage::HandoffReady(_) => {
431 let mut s = serializer.serialize_struct("msg", 2)?;
432 s.serialize_field("type", "HandoffReady")?;
433 s.serialize_field("state", "Idle")?;
434 s.end()
435 }
436 TuiMessage::ToolPending { tool_name, hint } => {
437 let mut s = serializer.serialize_struct("msg", 3)?;
438 s.serialize_field("type", "ToolPending")?;
439 s.serialize_field("tool_name", tool_name)?;
440 s.serialize_field("hint", hint)?;
441 s.end()
442 }
443 TuiMessage::ToolDone { tool_name, success, hint } => {
444 let mut s = serializer.serialize_struct("msg", 4)?;
445 s.serialize_field("type", "ToolDone")?;
446 s.serialize_field("tool_name", tool_name)?;
447 s.serialize_field("success", success)?;
448 s.serialize_field("hint", hint)?;
449 s.end()
450 }
451 TuiMessage::ContextTokensUpdated(count) => {
452 let mut s = serializer.serialize_struct("msg", 2)?;
453 s.serialize_field("type", "ContextTokensUpdated")?;
454 s.serialize_field("count", count)?;
455 s.end()
456 }
457 TuiMessage::McpServerStatus { name, connected, tool_count, error } => {
458 let mut s = serializer.serialize_struct("msg", 5)?;
459 s.serialize_field("type", "McpServerStatus")?;
460 s.serialize_field("name", name)?;
461 s.serialize_field("connected", connected)?;
462 s.serialize_field("tool_count", tool_count)?;
463 s.serialize_field("error", error)?;
464 s.end()
465 }
466 }
467 }
468}