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}
81
82impl ServerState {
83    /// Create new shared state from a default session, broadcast sender, and optional auth.
84    ///
85    /// The provided session becomes the `"default"` user's session. When auth is
86    /// enabled, authenticated users get their own sessions created on demand.
87    pub fn new(
88        session: Session,
89        ws_tx: broadcast::Sender<String>,
90        auth: Option<Arc<AuthState>>,
91    ) -> Self {
92        let sessions = Arc::new(DashMap::new());
93
94        // Store the default session under the "default" key
95        // Use the provided ws_tx as the default user's broadcast channel
96        let token_store = Arc::new(pep::MemoryTokenStore::new());
97        sessions.insert(
98            "default".to_string(),
99            UserSession {
100                session: Arc::new(Mutex::new(session)),
101                ws_tx: ws_tx.clone(),
102                token_store,
103            },
104        );
105
106        Self {
107            sessions,
108            ws_tx,
109            auth,
110            config_toml: None,
111        }
112    }
113
114    /// Set the shared config TOML.
115    pub fn with_config_toml(mut self, config_toml: String) -> Self {
116        self.config_toml = Some(config_toml);
117        self
118    }
119
120    /// Get or create a session for the given user key, returning the session + ws_tx.
121    ///
122    /// This is the main entry point for route handlers. It ensures the user
123    /// has a session, spawns a drain task if newly created, and returns
124    /// references to the session mutex, broadcast sender, and token store.
125    pub async fn ensure_user_session(
126        &self,
127        user_key: &str,
128    ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
129        // Fast path: user already has a session
130        if let Some(entry) = self.sessions.get(user_key) {
131            return (
132                entry.session.clone(),
133                entry.ws_tx.clone(),
134                entry.token_store.clone(),
135            );
136        }
137
138        // Slow path: create new session for this user
139        let (mut session, workflow_rx) = Session::new();
140
141        // Copy shared config
142        if let Some(ref config_toml) = self.config_toml {
143            session.config_toml = Some(config_toml.clone());
144            session.parse_auto_handoff_config();
145
146            if let Ok(table) = config_toml.parse::<toml::Value>() {
147                if let Some(name) = table.get("agent").and_then(|a| a.get("name")).and_then(|n| n.as_str()) {
148                    session.agent_name = name.to_string();
149                }
150            }
151        }
152
153        // Isolate checkpoint storage per user by setting a unique project_id.
154        // The project_id becomes the storage partition key in ABK's checkpoint
155        // system. By prefixing with the user_key, each user's checkpoints are
156        // stored in separate directories, preventing cross-user access.
157        // The "default" user (no auth) keeps the legacy behavior (no project_id).
158        if user_key != "default" {
159            session.project_id = Some(format!("user:{user_key}"));
160        }
161
162        let user_session = UserSession::new(session);
163        let session_arc = user_session.session.clone();
164        let ws_tx = user_session.ws_tx.clone();
165        let token_store = user_session.token_store.clone();
166
167        self.sessions.insert(user_key.to_string(), user_session);
168
169        // Spawn drain task for this user's workflow receiver
170        self.spawn_user_drain_task(
171            user_key.to_string(),
172            session_arc.clone(),
173            ws_tx.clone(),
174            workflow_rx,
175        );
176
177        (session_arc, ws_tx, token_store)
178    }
179
180    /// Spawn a background drain task for a specific user's workflow receiver.
181    ///
182    /// This replaces the old global drain task — each user gets their own.
183    fn spawn_user_drain_task(
184        &self,
185        user_key: String,
186        session: Arc<Mutex<Session>>,
187        ws_tx: broadcast::Sender<String>,
188        mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
189    ) {
190        tokio::spawn(async move {
191            while let Some(msg) = workflow_rx.recv().await {
192                // Process the message through Session's handler (updates state)
193                {
194                    let mut session = session.lock().await;
195                    session.handle_workflow_message(msg.clone());
196
197                    let state_str = match session.workflow_state {
198                        trustee_core::types::WorkflowState::Idle => "Idle",
199                        trustee_core::types::WorkflowState::Running => "Running",
200                        trustee_core::types::WorkflowState::Cancelling => "Cancelling",
201                    };
202                    let state_msg = serde_json::json!({
203                        "type": "StateChanged",
204                        "state": state_str
205                    });
206                    let _ = ws_tx.send(state_msg.to_string());
207                }
208
209                // Broadcast the raw message to WebSocket clients
210                let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
211                let _ = ws_tx.send(json);
212            }
213            tracing::debug!("Drain task ended for user: {}", user_key);
214        });
215    }
216
217    /// Spawn the default user's drain task (backward compatibility).
218    ///
219    /// Called during server startup for the initial session.
220    pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
221        // Get the default session's arc
222        let default_entry = self.sessions.get("default").expect("default session must exist");
223        let session = default_entry.session.clone();
224        let ws_tx = default_entry.ws_tx.clone();
225        drop(default_entry);
226
227        tokio::spawn(async move {
228            while let Some(msg) = workflow_rx.recv().await {
229                {
230                    let mut session = session.lock().await;
231                    session.handle_workflow_message(msg.clone());
232
233                    let state_str = match session.workflow_state {
234                        trustee_core::types::WorkflowState::Idle => "Idle",
235                        trustee_core::types::WorkflowState::Running => "Running",
236                        trustee_core::types::WorkflowState::Cancelling => "Cancelling",
237                    };
238                    let state_msg = serde_json::json!({
239                        "type": "StateChanged",
240                        "state": state_str
241                    });
242                    let _ = ws_tx.send(state_msg.to_string());
243                }
244
245                let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
246                let _ = ws_tx.send(json);
247            }
248        });
249    }
250
251    /// Resolve the user key from request headers.
252    ///
253    /// Returns `"default"` when auth is disabled.
254    /// Returns the JWT `sub` claim (or `dev:email` for dev mode) when auth is enabled.
255    pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
256        let Some(ref auth) = self.auth else {
257            return "default".to_string();
258        };
259
260        // Try Bearer header first
261        if let Some(token) = headers
262            .get(axum::http::header::AUTHORIZATION)
263            .and_then(|v| v.to_str().ok())
264            .and_then(|v| v.strip_prefix("Bearer "))
265            .map(|s| s.to_string())
266        {
267            // Dev token
268            if token.starts_with("dev:") {
269                let parts: Vec<&str> = token.splitn(4, ':').collect();
270                if parts.len() >= 4 {
271                    return format!("dev:{}", parts[1]);
272                }
273            }
274            // Real JWT — extract sub claim
275            if let Ok(claims) = auth.validate_token(&token).await {
276                return claims.sub;
277            }
278        }
279
280        // Try cookie
281        let cookie_session_id = headers
282            .get(axum::http::header::COOKIE)
283            .and_then(|v| v.to_str().ok())
284            .and_then(|cookies| {
285                cookies
286                    .split(';')
287                    .map(|c| c.trim())
288                    .find_map(|c| c.strip_prefix(&format!("{}=", auth.config.cookie_name)))
289                    .map(|s| s.to_string())
290            });
291
292        if let Some(session_id) = cookie_session_id {
293            // Dev token in cookie
294            if session_id.starts_with("dev:") {
295                let parts: Vec<&str> = session_id.splitn(4, ':').collect();
296                if parts.len() >= 4 {
297                    return format!("dev:{}", parts[1]);
298                }
299            }
300
301            // Resolve session_id → access_token → sub claim
302            if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
303                if let Ok(claims) = auth.validate_token(&access_token).await {
304                    return claims.sub;
305                }
306            }
307        }
308
309        "default".to_string()
310    }
311}
312
313/// Wrapper to serialize `TuiMessage` as JSON with a `type` discriminator.
314struct SerializableMessage<'a>(&'a TuiMessage);
315
316impl<'a> serde::Serialize for SerializableMessage<'a> {
317    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
318    where
319        S: serde::Serializer,
320    {
321        use serde::ser::SerializeStruct;
322
323        match self.0 {
324            TuiMessage::OutputLine(line) => {
325                let mut s = serializer.serialize_struct("msg", 2)?;
326                s.serialize_field("type", "OutputLine")?;
327                s.serialize_field("line", line)?;
328                s.end()
329            }
330            TuiMessage::StreamDelta(delta) => {
331                let mut s = serializer.serialize_struct("msg", 2)?;
332                s.serialize_field("type", "StreamDelta")?;
333                s.serialize_field("delta", delta)?;
334                s.end()
335            }
336            TuiMessage::ReasoningDelta(delta) => {
337                let mut s = serializer.serialize_struct("msg", 2)?;
338                s.serialize_field("type", "ReasoningDelta")?;
339                s.serialize_field("delta", delta)?;
340                s.end()
341            }
342            TuiMessage::WorkflowCompleted => {
343                let mut s = serializer.serialize_struct("msg", 2)?;
344                s.serialize_field("type", "WorkflowCompleted")?;
345                s.serialize_field("state", "Idle")?;
346                s.end()
347            }
348            TuiMessage::WorkflowError(err) => {
349                let mut s = serializer.serialize_struct("msg", 2)?;
350                s.serialize_field("type", "WorkflowError")?;
351                s.serialize_field("error", err)?;
352                s.end()
353            }
354            TuiMessage::ResumeInfo(_) => {
355                let mut s = serializer.serialize_struct("msg", 2)?;
356                s.serialize_field("type", "ResumeInfo")?;
357                s.serialize_field("state", "Idle")?;
358                s.end()
359            }
360            TuiMessage::TodoUpdate(content) => {
361                let mut s = serializer.serialize_struct("msg", 2)?;
362                s.serialize_field("type", "TodoUpdate")?;
363                s.serialize_field("content", content)?;
364                s.end()
365            }
366            TuiMessage::WorkflowCancelled => {
367                let mut s = serializer.serialize_struct("msg", 2)?;
368                s.serialize_field("type", "WorkflowCancelled")?;
369                s.serialize_field("state", "Idle")?;
370                s.end()
371            }
372            TuiMessage::HandoffReady(_) => {
373                let mut s = serializer.serialize_struct("msg", 2)?;
374                s.serialize_field("type", "HandoffReady")?;
375                s.serialize_field("state", "Idle")?;
376                s.end()
377            }
378            TuiMessage::ToolPending { tool_name, hint } => {
379                let mut s = serializer.serialize_struct("msg", 3)?;
380                s.serialize_field("type", "ToolPending")?;
381                s.serialize_field("tool_name", tool_name)?;
382                s.serialize_field("hint", hint)?;
383                s.end()
384            }
385            TuiMessage::ToolDone { tool_name, success, hint } => {
386                let mut s = serializer.serialize_struct("msg", 4)?;
387                s.serialize_field("type", "ToolDone")?;
388                s.serialize_field("tool_name", tool_name)?;
389                s.serialize_field("success", success)?;
390                s.serialize_field("hint", hint)?;
391                s.end()
392            }
393            TuiMessage::ContextTokensUpdated(count) => {
394                let mut s = serializer.serialize_struct("msg", 2)?;
395                s.serialize_field("type", "ContextTokensUpdated")?;
396                s.serialize_field("count", count)?;
397                s.end()
398            }
399            TuiMessage::McpServerStatus { name, connected, tool_count, error } => {
400                let mut s = serializer.serialize_struct("msg", 5)?;
401                s.serialize_field("type", "McpServerStatus")?;
402                s.serialize_field("name", name)?;
403                s.serialize_field("connected", connected)?;
404                s.serialize_field("tool_count", tool_count)?;
405                s.serialize_field("error", error)?;
406                s.end()
407            }
408        }
409    }
410}