Skip to main content

trustee_api/
state.rs

1//! Shared server state: per-user session registry, broadcast channels, and auth state.
2//!
3//! ## Multi-User Architecture (TMU Phase 2)
4//!
5//! Each authenticated user gets their own [`UserSession`] containing:
6//! - An independent `Session` (workflow state, output, etc.)
7//! - A dedicated broadcast channel for WebSocket fan-out
8//! - A per-user token store for MCP credential isolation
9//!
10//! Sessions are keyed by user identity (`sub` claim from JWT, or `dev:email` for
11//! dev mode). Unauthenticated deployments use a single `"default"` key, preserving
12//! backward compatibility with single-user CLI operation.
13
14use 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
23/// Per-user session bundle.
24///
25/// Each user gets their own Session instance, broadcast channel, and
26/// token store. This struct is stored in the [`SessionRegistry`]
27/// and accessed via the user's identity key.
28pub struct UserSession {
29    /// The agent session, protected by a mutex.
30    pub session: Arc<Mutex<Session>>,
31    /// Broadcast sender for this user's WebSocket fan-out.
32    pub ws_tx: broadcast::Sender<String>,
33    /// Per-user in-memory token store for MCP credential isolation.
34    ///
35    /// Replaces the process-wide FileTokenStore that was vulnerable to
36    /// cross-user token leakage via __web_session.json. Each user's
37    /// MCP `web-session` tokens are stored here, isolated from other users.
38    pub token_store: Arc<pep::MemoryTokenStore>,
39}
40
41impl UserSession {
42    /// Create a new per-user session from an existing Session.
43    ///
44    /// Creates a fresh broadcast channel (256 capacity) for WebSocket fan-out
45    /// and a per-user MemoryTokenStore for MCP credential isolation.
46    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
57/// Concurrent registry of per-user sessions.
58///
59/// Keyed by user identity string:
60/// - Authenticated: JWT `sub` claim (e.g., Kanidm UUID)
61/// - Dev mode: `dev:{email}`
62/// - No auth: `"default"`
63///
64/// Falls back to the `"default"` entry when no user key is provided,
65/// preserving backward compatibility.
66pub type SessionRegistry = Arc<DashMap<String, UserSession>>;
67
68/// Shared state accessible by all axum handlers.
69#[derive(Clone)]
70pub struct ServerState {
71    /// Per-user session registry (TMU Phase 2).
72    pub sessions: SessionRegistry,
73    /// Broadcast sender for backward compat — delegates to the default user's channel.
74    /// New code should use `user_ws_tx(user_key)` instead.
75    pub ws_tx: broadcast::Sender<String>,
76    /// Auth state (None = auth disabled, all endpoints open).
77    pub auth: Option<Arc<AuthState>>,
78    /// Shared config TOML (all users share the same agent config).
79    pub config_toml: Option<String>,
80    /// Shared secrets (injected into every per-user session).
81    pub secrets: Option<std::collections::HashMap<String, String>>,
82    /// Shared build info (injected into every per-user session).
83    pub build_info: Option<trustee_core::types::BuildInfo>,
84    /// Global concurrency limiter — limits the number of simultaneous workflows
85    /// across all users. Prevents resource exhaustion on shared infrastructure.
86    /// Default: 8 concurrent workflows.
87    pub workflow_semaphore: Arc<tokio::sync::Semaphore>,
88}
89
90impl ServerState {
91    /// Create new shared state from a default session, broadcast sender, and optional auth.
92    ///
93    /// The provided session becomes the `"default"` user's session. When auth is
94    /// enabled, authenticated users get their own sessions created on demand.
95    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        // Store the default session under the "default" key
103        // Use the provided ws_tx as the default user's broadcast channel
104        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    /// Set the shared config TOML.
126    pub fn with_config_toml(mut self, config_toml: String) -> Self {
127        self.config_toml = Some(config_toml);
128        self
129    }
130
131    /// Set the shared secrets.
132    pub fn with_secrets(mut self, secrets: std::collections::HashMap<String, String>) -> Self {
133        self.secrets = Some(secrets);
134        self
135    }
136
137    /// Set the shared build info.
138    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    /// Set the max concurrent workflows.
144    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    /// Get or create a session for the given user key, returning the session + ws_tx.
150    ///
151    /// This is the main entry point for route handlers. It ensures the user
152    /// has a session, spawns a drain task if newly created, and returns
153    /// references to the session mutex, broadcast sender, and token store.
154    pub async fn ensure_user_session(
155        &self,
156        user_key: &str,
157    ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
158        // Fast path: user already has a session
159        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        // Slow path: create new session for this user
168        let (mut session, workflow_rx) = Session::new();
169
170        // Copy shared config
171        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        // Copy shared secrets and build info
183        session.secrets = self.secrets.clone();
184        session.build_info = self.build_info.clone();
185
186        // Isolate checkpoint storage per user by setting a per-user home_dir.
187        // Each user gets their own checkpoint directory tree at
188        // ~/.trustee/users/{user_hash}/, preventing cross-user access.
189        // The "default" user (no auth) keeps the legacy behavior (no home_dir override).
190        if user_key != "default" {
191            // Compute SHA-256 of the user_key, take first 16 hex chars
192            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            // Set per-user home_dir for checkpoint isolation
199            if let Some(home) = dirs::home_dir() {
200                session.home_dir = Some(home.join(".trustee").join("users").join(&user_hash));
201            }
202
203            // Use a UUID for project_id (web mode doesn't have a stable working dir)
204            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        // Spawn drain task for this user's workflow receiver
215        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    /// Spawn a background drain task for a specific user's workflow receiver.
226    ///
227    /// This replaces the old global drain task — each user gets their own.
228    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                // Process the message through Session's handler (updates state)
238                {
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                // Broadcast the raw message to WebSocket clients
255                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    /// Spawn the default user's drain task (backward compatibility).
263    ///
264    /// Called during server startup for the initial session.
265    pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
266        // Get the default session's arc
267        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    /// Resolve the user key from request headers.
297    ///
298    /// Returns `"default"` when auth is disabled.
299    /// Returns the JWT `sub` claim (or `dev:email` for dev mode) when auth is enabled.
300    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        // Try Bearer header first
306        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            // Dev token
313            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            // Real JWT — extract sub claim
320            if let Ok(claims) = auth.validate_token(&token).await {
321                return claims.sub;
322            }
323        }
324
325        // Try cookie
326        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            // Dev token in cookie
339            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            // Resolve session_id → access_token → sub claim
347            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
358/// Wrapper to serialize `TuiMessage` as JSON with a `type` discriminator.
359struct 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}