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 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 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 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 {
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 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 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
212 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 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 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 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 if let Ok(claims) = auth.validate_token(&token).await {
267 return claims.sub;
268 }
269 }
270
271 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 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 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
304struct 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}