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}
81
82impl ServerState {
83 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 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 pub fn with_config_toml(mut self, config_toml: String) -> Self {
116 self.config_toml = Some(config_toml);
117 self
118 }
119
120 pub async fn ensure_user_session(
126 &self,
127 user_key: &str,
128 ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
129 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 let (mut session, workflow_rx) = Session::new();
140
141 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 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 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 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 {
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 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 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
221 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 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 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 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 if let Ok(claims) = auth.validate_token(&token).await {
276 return claims.sub;
277 }
278 }
279
280 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 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 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
313struct 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}