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        let user_session = UserSession::new(session);
154        let session_arc = user_session.session.clone();
155        let ws_tx = user_session.ws_tx.clone();
156        let token_store = user_session.token_store.clone();
157
158        self.sessions.insert(user_key.to_string(), user_session);
159
160        // Spawn drain task for this user's workflow receiver
161        self.spawn_user_drain_task(
162            user_key.to_string(),
163            session_arc.clone(),
164            ws_tx.clone(),
165            workflow_rx,
166        );
167
168        (session_arc, ws_tx, token_store)
169    }
170
171    /// Spawn a background drain task for a specific user's workflow receiver.
172    ///
173    /// This replaces the old global drain task — each user gets their own.
174    fn spawn_user_drain_task(
175        &self,
176        user_key: String,
177        session: Arc<Mutex<Session>>,
178        ws_tx: broadcast::Sender<String>,
179        mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
180    ) {
181        tokio::spawn(async move {
182            while let Some(msg) = workflow_rx.recv().await {
183                // Process the message through Session's handler (updates state)
184                {
185                    let mut session = session.lock().await;
186                    session.handle_workflow_message(msg.clone());
187
188                    let state_str = match session.workflow_state {
189                        trustee_core::types::WorkflowState::Idle => "Idle",
190                        trustee_core::types::WorkflowState::Running => "Running",
191                        trustee_core::types::WorkflowState::Cancelling => "Cancelling",
192                    };
193                    let state_msg = serde_json::json!({
194                        "type": "StateChanged",
195                        "state": state_str
196                    });
197                    let _ = ws_tx.send(state_msg.to_string());
198                }
199
200                // Broadcast the raw message to WebSocket clients
201                let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
202                let _ = ws_tx.send(json);
203            }
204            tracing::debug!("Drain task ended for user: {}", user_key);
205        });
206    }
207
208    /// Spawn the default user's drain task (backward compatibility).
209    ///
210    /// Called during server startup for the initial session.
211    pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
212        // Get the default session's arc
213        let default_entry = self.sessions.get("default").expect("default session must exist");
214        let session = default_entry.session.clone();
215        let ws_tx = default_entry.ws_tx.clone();
216        drop(default_entry);
217
218        tokio::spawn(async move {
219            while let Some(msg) = workflow_rx.recv().await {
220                {
221                    let mut session = session.lock().await;
222                    session.handle_workflow_message(msg.clone());
223
224                    let state_str = match session.workflow_state {
225                        trustee_core::types::WorkflowState::Idle => "Idle",
226                        trustee_core::types::WorkflowState::Running => "Running",
227                        trustee_core::types::WorkflowState::Cancelling => "Cancelling",
228                    };
229                    let state_msg = serde_json::json!({
230                        "type": "StateChanged",
231                        "state": state_str
232                    });
233                    let _ = ws_tx.send(state_msg.to_string());
234                }
235
236                let json = serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
237                let _ = ws_tx.send(json);
238            }
239        });
240    }
241
242    /// Resolve the user key from request headers.
243    ///
244    /// Returns `"default"` when auth is disabled.
245    /// Returns the JWT `sub` claim (or `dev:email` for dev mode) when auth is enabled.
246    pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
247        let Some(ref auth) = self.auth else {
248            return "default".to_string();
249        };
250
251        // Try Bearer header first
252        if let Some(token) = headers
253            .get(axum::http::header::AUTHORIZATION)
254            .and_then(|v| v.to_str().ok())
255            .and_then(|v| v.strip_prefix("Bearer "))
256            .map(|s| s.to_string())
257        {
258            // Dev token
259            if token.starts_with("dev:") {
260                let parts: Vec<&str> = token.splitn(4, ':').collect();
261                if parts.len() >= 4 {
262                    return format!("dev:{}", parts[1]);
263                }
264            }
265            // Real JWT — extract sub claim
266            if let Ok(claims) = auth.validate_token(&token).await {
267                return claims.sub;
268            }
269        }
270
271        // Try cookie
272        let cookie_session_id = headers
273            .get(axum::http::header::COOKIE)
274            .and_then(|v| v.to_str().ok())
275            .and_then(|cookies| {
276                cookies
277                    .split(';')
278                    .map(|c| c.trim())
279                    .find_map(|c| c.strip_prefix(&format!("{}=", auth.config.cookie_name)))
280                    .map(|s| s.to_string())
281            });
282
283        if let Some(session_id) = cookie_session_id {
284            // Dev token in cookie
285            if session_id.starts_with("dev:") {
286                let parts: Vec<&str> = session_id.splitn(4, ':').collect();
287                if parts.len() >= 4 {
288                    return format!("dev:{}", parts[1]);
289                }
290            }
291
292            // Resolve session_id → access_token → sub claim
293            if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
294                if let Ok(claims) = auth.validate_token(&access_token).await {
295                    return claims.sub;
296                }
297            }
298        }
299
300        "default".to_string()
301    }
302}
303
304/// Wrapper to serialize `TuiMessage` as JSON with a `type` discriminator.
305struct SerializableMessage<'a>(&'a TuiMessage);
306
307impl<'a> serde::Serialize for SerializableMessage<'a> {
308    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
309    where
310        S: serde::Serializer,
311    {
312        use serde::ser::SerializeStruct;
313
314        match self.0 {
315            TuiMessage::OutputLine(line) => {
316                let mut s = serializer.serialize_struct("msg", 2)?;
317                s.serialize_field("type", "OutputLine")?;
318                s.serialize_field("line", line)?;
319                s.end()
320            }
321            TuiMessage::StreamDelta(delta) => {
322                let mut s = serializer.serialize_struct("msg", 2)?;
323                s.serialize_field("type", "StreamDelta")?;
324                s.serialize_field("delta", delta)?;
325                s.end()
326            }
327            TuiMessage::ReasoningDelta(delta) => {
328                let mut s = serializer.serialize_struct("msg", 2)?;
329                s.serialize_field("type", "ReasoningDelta")?;
330                s.serialize_field("delta", delta)?;
331                s.end()
332            }
333            TuiMessage::WorkflowCompleted => {
334                let mut s = serializer.serialize_struct("msg", 2)?;
335                s.serialize_field("type", "WorkflowCompleted")?;
336                s.serialize_field("state", "Idle")?;
337                s.end()
338            }
339            TuiMessage::WorkflowError(err) => {
340                let mut s = serializer.serialize_struct("msg", 2)?;
341                s.serialize_field("type", "WorkflowError")?;
342                s.serialize_field("error", err)?;
343                s.end()
344            }
345            TuiMessage::ResumeInfo(_) => {
346                let mut s = serializer.serialize_struct("msg", 2)?;
347                s.serialize_field("type", "ResumeInfo")?;
348                s.serialize_field("state", "Idle")?;
349                s.end()
350            }
351            TuiMessage::TodoUpdate(content) => {
352                let mut s = serializer.serialize_struct("msg", 2)?;
353                s.serialize_field("type", "TodoUpdate")?;
354                s.serialize_field("content", content)?;
355                s.end()
356            }
357            TuiMessage::WorkflowCancelled => {
358                let mut s = serializer.serialize_struct("msg", 2)?;
359                s.serialize_field("type", "WorkflowCancelled")?;
360                s.serialize_field("state", "Idle")?;
361                s.end()
362            }
363            TuiMessage::HandoffReady(_) => {
364                let mut s = serializer.serialize_struct("msg", 2)?;
365                s.serialize_field("type", "HandoffReady")?;
366                s.serialize_field("state", "Idle")?;
367                s.end()
368            }
369            TuiMessage::ToolPending { tool_name, hint } => {
370                let mut s = serializer.serialize_struct("msg", 3)?;
371                s.serialize_field("type", "ToolPending")?;
372                s.serialize_field("tool_name", tool_name)?;
373                s.serialize_field("hint", hint)?;
374                s.end()
375            }
376            TuiMessage::ToolDone { tool_name, success, hint } => {
377                let mut s = serializer.serialize_struct("msg", 4)?;
378                s.serialize_field("type", "ToolDone")?;
379                s.serialize_field("tool_name", tool_name)?;
380                s.serialize_field("success", success)?;
381                s.serialize_field("hint", hint)?;
382                s.end()
383            }
384            TuiMessage::ContextTokensUpdated(count) => {
385                let mut s = serializer.serialize_struct("msg", 2)?;
386                s.serialize_field("type", "ContextTokensUpdated")?;
387                s.serialize_field("count", count)?;
388                s.end()
389            }
390            TuiMessage::McpServerStatus { name, connected, tool_count, error } => {
391                let mut s = serializer.serialize_struct("msg", 5)?;
392                s.serialize_field("type", "McpServerStatus")?;
393                s.serialize_field("name", name)?;
394                s.serialize_field("connected", connected)?;
395                s.serialize_field("tool_count", tool_count)?;
396                s.serialize_field("error", error)?;
397                s.end()
398            }
399        }
400    }
401}