1use std::sync::Arc;
16
17use dashmap::DashMap;
18use tokio::sync::{broadcast, mpsc, Mutex};
19use trustee_core::session::Session;
20use trustee_core::types::TuiMessage;
21
22use crate::auth::AuthState;
23
24pub struct UserSessionEntry {
30 pub session: Arc<Mutex<Session>>,
32 pub ws_tx: broadcast::Sender<String>,
34 pub created_at: chrono::DateTime<chrono::Utc>,
36 pub last_active: Arc<Mutex<chrono::DateTime<chrono::Utc>>>,
39}
40
41pub struct UserSessions {
43 pub sessions: DashMap<String, UserSessionEntry>,
45 pub token_store: Arc<pep::MemoryTokenStore>,
47 pub active_session_id: Mutex<String>,
49}
50
51#[derive(Debug, serde::Serialize)]
53pub struct SessionListItem {
54 pub session_id: String,
55 pub session_name: Option<String>,
56 pub workflow_state: String,
57 pub created_at: String,
58 pub last_active: String,
59 pub handoff_count: u32,
64}
65
66#[derive(Debug)]
68pub enum SessionError {
69 MaxSessionsReached(usize),
71 NotFound(String),
73 NotIdle(String),
75}
76
77impl std::fmt::Display for SessionError {
78 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
79 match self {
80 SessionError::MaxSessionsReached(n) => {
81 write!(f, "Maximum {} sessions per user reached", n)
82 }
83 SessionError::NotFound(id) => write!(f, "Session {} not found", id),
84 SessionError::NotIdle(state) => write!(f, "Session is not idle (state: {})", state),
85 }
86 }
87}
88
89impl std::error::Error for SessionError {}
90
91pub type SessionRegistry = Arc<DashMap<String, UserSessions>>;
93
94pub struct McpLoaderEntry {
105 pub loader: Option<std::sync::Arc<abk::agent::McpToolLoader>>,
107 pub fingerprint: u64,
110 pub built_at: chrono::DateTime<chrono::Utc>,
111 pub degraded: Option<String>,
114 pub failed_at: Option<chrono::DateTime<chrono::Utc>>,
117}
118
119pub const MCP_BUILD_RETRY_BACKOFF: std::time::Duration = std::time::Duration::from_secs(30);
121
122#[derive(Debug, Clone)]
129pub struct ThqDispatchEntry {
130 pub user_key: String,
135 pub service_token: Option<String>,
139}
140
141#[derive(Clone)]
143pub struct ServerState {
144 pub sessions: SessionRegistry,
146 pub ws_tx: broadcast::Sender<String>,
148 pub auth: Option<Arc<AuthState>>,
150 pub config_toml: Option<String>,
152 pub secrets: Option<std::collections::HashMap<String, String>>,
154 pub build_info: Option<trustee_core::types::BuildInfo>,
156 pub workflow_semaphore: Arc<tokio::sync::Semaphore>,
159 pub max_sessions_per_user: usize,
161 pub allow_llm_overlay: bool,
164 pub mcp_loaders: Arc<DashMap<String, McpLoaderEntry>>,
166 mcp_build_locks: Arc<DashMap<String, Arc<tokio::sync::Mutex<()>>>>,
168 pub thq_dispatch: Arc<DashMap<String, ThqDispatchEntry>>,
170 pub agent_dispatch_tokens: Arc<DashMap<String, (String, std::time::Instant)>>,
173}
174
175impl ServerState {
176 pub fn new(
178 session: Session,
179 ws_tx: broadcast::Sender<String>,
180 auth: Option<Arc<AuthState>>,
181 ) -> Self {
182 let sessions = Arc::new(DashMap::new());
183
184 let token_store = Arc::new(pep::MemoryTokenStore::new());
186 let (ws_tx_entry, _) = broadcast::channel::<String>(256);
187
188 let now = chrono::Utc::now();
189 let initial_entry = UserSessionEntry {
190 session: Arc::new(Mutex::new(session)),
191 ws_tx: ws_tx_entry,
192 created_at: now,
193 last_active: Arc::new(Mutex::new(now)),
194 };
195
196 let user_sessions = UserSessions {
197 sessions: DashMap::new(),
198 token_store,
199 active_session_id: Mutex::new(String::new()),
200 };
201 user_sessions.sessions.insert("default".to_string(), initial_entry);
202
203 sessions.insert("default".to_string(), user_sessions);
204
205 Self {
206 sessions,
207 ws_tx,
208 auth,
209 config_toml: None,
210 secrets: None,
211 build_info: None,
212 workflow_semaphore: Arc::new(tokio::sync::Semaphore::new(8)),
213 max_sessions_per_user: 4,
214 allow_llm_overlay: false,
215 mcp_loaders: Arc::new(DashMap::new()),
216 mcp_build_locks: Arc::new(DashMap::new()),
217 thq_dispatch: Arc::new(DashMap::new()),
218 agent_dispatch_tokens: Arc::new(DashMap::new()),
219 }
220 }
221
222 pub fn with_config_toml(mut self, config_toml: String) -> Self {
223 self.config_toml = Some(config_toml);
224 self
225 }
226
227 pub fn with_secrets(mut self, secrets: std::collections::HashMap<String, String>) -> Self {
228 self.secrets = Some(secrets);
229 self
230 }
231
232 pub fn with_build_info(mut self, build_info: trustee_core::types::BuildInfo) -> Self {
233 self.build_info = Some(build_info);
234 self
235 }
236
237 pub fn with_max_concurrent_workflows(mut self, max: usize) -> Self {
238 self.workflow_semaphore = Arc::new(tokio::sync::Semaphore::new(max));
239 self
240 }
241
242 pub fn with_max_sessions_per_user(mut self, max: usize) -> Self {
244 self.max_sessions_per_user = max;
245 self
246 }
247
248 pub fn with_allow_llm_overlay(mut self, allow: bool) -> Self {
251 self.allow_llm_overlay = allow;
252 self
253 }
254
255 pub async fn create_session(
265 &self,
266 user_key: &str,
267 session_name: Option<String>,
268 identity: Option<String>,
269 activate: bool,
270 ) -> Result<String, SessionError> {
271 let user_sessions = self
273 .sessions
274 .entry(user_key.to_string())
275 .or_insert_with(|| UserSessions {
276 sessions: DashMap::new(),
277 token_store: Arc::new(pep::MemoryTokenStore::new()),
278 active_session_id: Mutex::new(String::new()),
279 });
280
281 if user_sessions.sessions.len() >= self.max_sessions_per_user {
283 return Err(SessionError::MaxSessionsReached(self.max_sessions_per_user));
284 }
285
286 let (mut session, workflow_rx) = Session::new();
288
289 if let Some(ref config_toml) = self.config_toml {
291 session.config_toml = Some(config_toml.clone());
292 session.parse_auto_handoff_config();
293 if let Ok(table) = config_toml.parse::<toml::Value>() {
294 if let Some(name) = table
295 .get("agent")
296 .and_then(|a| a.get("name"))
297 .and_then(|n| n.as_str())
298 {
299 session.agent_name = name.to_string();
300 }
301 }
302 }
303
304 session.secrets = self.secrets.clone();
305 session.build_info = self.build_info.clone();
306
307 self.apply_user_isolation(&mut session, user_key);
309
310 session.session_name = session_name;
312
313 session.identity = identity;
315
316 let (ws_tx_entry, _) = broadcast::channel::<String>(256);
318
319 let session_id = format!(
321 "session_{}_{}",
322 chrono::Utc::now().format("%Y_%m_%d_%H_%M"),
323 &uuid::Uuid::new_v4().to_string()[..8]
324 );
325
326 let now = chrono::Utc::now();
327
328 user_sessions.sessions.insert(
330 session_id.clone(),
331 UserSessionEntry {
332 session: Arc::new(Mutex::new(session)),
333 ws_tx: ws_tx_entry.clone(),
334 created_at: now,
335 last_active: Arc::new(Mutex::new(now)),
336 },
337 );
338
339 if activate {
348 *user_sessions.active_session_id.lock().await = session_id.clone();
349 }
350
351 let session_arc = user_sessions
353 .sessions
354 .get(&session_id)
355 .map(|e| e.session.clone());
356 if let Some(session_arc) = session_arc {
357 self.spawn_user_drain_task(
358 session_id.clone(),
359 session_arc,
360 ws_tx_entry,
361 workflow_rx,
362 );
363 }
364
365 Ok(session_id)
366 }
367
368 pub async fn get_session(
371 &self,
372 user_key: &str,
373 session_id: &str,
374 ) -> Option<(Arc<Mutex<Session>>, broadcast::Sender<String>)> {
375 let user_sessions = self.sessions.get(user_key)?;
376 let entry = user_sessions.sessions.get(session_id)?;
377
378 let now = chrono::Utc::now();
380 *entry.last_active.lock().await = now;
381
382 Some((entry.session.clone(), entry.ws_tx.clone()))
383 }
384
385 pub async fn get_session_by_any_id(
403 &self,
404 user_key: &str,
405 id: &str,
406 ) -> Option<(String, Arc<Mutex<Session>>, broadcast::Sender<String>)> {
407 let user_sessions = self.sessions.get(user_key)?;
409 if let Some(entry) = user_sessions.sessions.get(id) {
410 let now = chrono::Utc::now();
412 *entry.last_active.lock().await = now;
413 return Some((id.to_string(), entry.session.clone(), entry.ws_tx.clone()));
414 }
415
416 for entry in user_sessions.sessions.iter() {
418 let session = entry.session.lock().await;
419 if session.session_id.as_deref() == Some(id) {
420 let key = entry.key().clone();
421 let ws_tx = entry.ws_tx.clone();
422 drop(session);
423 let now = chrono::Utc::now();
425 *entry.last_active.lock().await = now;
426 return Some((key, entry.session.clone(), ws_tx));
427 }
428 }
429
430 None
431 }
432
433 pub async fn list_sessions(&self, user_key: &str) -> Vec<SessionListItem> {
435 let Some(user_sessions) = self.sessions.get(user_key) else {
436 return Vec::new();
437 };
438
439 let mut items = Vec::new();
440 for entry in user_sessions.sessions.iter() {
441 let session = entry.session.lock().await;
442 let workflow_state = match session.workflow_state {
443 trustee_core::types::WorkflowState::Idle => "Idle",
444 trustee_core::types::WorkflowState::Running => "Running",
445 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
446 };
447 let last_active = entry.last_active.lock().await;
448 items.push(SessionListItem {
449 session_id: entry.key().clone(),
450 session_name: session.session_name.clone(),
451 workflow_state: workflow_state.to_string(),
452 created_at: entry.created_at.to_rfc3339(),
453 last_active: last_active.to_rfc3339(),
454 handoff_count: session.handoff_count,
455 });
456 }
457 drop(user_sessions);
458
459 items.sort_by(|a, b| b.last_active.cmp(&a.last_active));
461 items
462 }
463
464 pub async fn destroy_session(
466 &self,
467 user_key: &str,
468 session_id: &str,
469 ) -> Result<(), SessionError> {
470 let user_sessions = self
471 .sessions
472 .get(user_key)
473 .ok_or_else(|| SessionError::NotFound(session_id.to_string()))?;
474
475 {
477 let entry = user_sessions
478 .sessions
479 .get(session_id)
480 .ok_or_else(|| SessionError::NotFound(session_id.to_string()))?;
481 let session = entry.session.lock().await;
482 if session.workflow_state != trustee_core::types::WorkflowState::Idle {
483 let state_str = match session.workflow_state {
484 trustee_core::types::WorkflowState::Running => "Running",
485 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
486 _ => "Unknown",
487 };
488 return Err(SessionError::NotIdle(state_str.to_string()));
489 }
490 }
491
492 user_sessions.sessions.remove(session_id);
494
495 let mut active_id = user_sessions.active_session_id.lock().await;
497 if &*active_id == session_id {
498 let mut newest: Option<(String, chrono::DateTime<chrono::Utc>)> = None;
500 for entry in user_sessions.sessions.iter() {
501 let la = entry.last_active.lock().await;
502 if newest.as_ref().map_or(true, |(_, t)| *la > *t) {
503 newest = Some((entry.key().clone(), *la));
504 }
505 }
506 *active_id = newest.map(|(id, _)| id).unwrap_or_default();
507 }
508
509 Ok(())
510 }
511
512 pub async fn ensure_active_session(
521 &self,
522 user_key: &str,
523 ) -> (
524 String,
525 Arc<Mutex<Session>>,
526 broadcast::Sender<String>,
527 Arc<pep::MemoryTokenStore>,
528 ) {
529 let token_store = {
531 let user_sessions = self
532 .sessions
533 .entry(user_key.to_string())
534 .or_insert_with(|| UserSessions {
535 sessions: DashMap::new(),
536 token_store: Arc::new(pep::MemoryTokenStore::new()),
537 active_session_id: Mutex::new(String::new()),
538 });
539 user_sessions.token_store.clone()
540 };
541
542 let active_id = {
544 let user_sessions = self.sessions.get(user_key).unwrap();
545 let guard = user_sessions.active_session_id.lock().await;
546 guard.clone()
547 };
548
549 if !active_id.is_empty() {
550 if let Some((session, ws_tx)) = self.get_session(user_key, &active_id).await {
551 return (active_id, session, ws_tx, token_store);
552 }
553 }
555
556 let existing_session: Option<(String, Arc<Mutex<Session>>, broadcast::Sender<String>)> = {
560 let user_sessions = self.sessions.get(user_key).unwrap();
561 let result = user_sessions.sessions.iter().next().map(|first| {
562 (
563 first.key().clone(),
564 first.session.clone(),
565 first.ws_tx.clone(),
566 )
567 });
568 result
569 };
570 if let Some((id, session, ws_tx)) = existing_session {
571 let now = chrono::Utc::now();
572 if let Some(entry) = self.sessions.get(user_key) {
573 if let Some(e) = entry.sessions.get(&id) {
574 *e.last_active.lock().await = now;
575 }
576 *entry.active_session_id.lock().await = id.clone();
577 }
578
579 return (id, session, ws_tx, token_store);
580 }
581
582 let session_id = self
584 .create_session(user_key, None, None, true)
585 .await
586 .unwrap_or_else(|_| "default".to_string());
587
588 let (session, ws_tx) = self
589 .get_session(user_key, &session_id)
590 .await
591 .expect("just-created session must exist");
592
593 (session_id, session, ws_tx, token_store)
594 }
595
596 pub async fn ensure_user_session(
599 &self,
600 user_key: &str,
601 ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
602 let (_id, session, ws_tx, token_store) = self.ensure_active_session(user_key).await;
603 (session, ws_tx, token_store)
604 }
605
606 pub async fn set_active_session(&self, user_key: &str, session_id: &str) {
608 if let Some(user_sessions) = self.sessions.get(user_key) {
609 if user_sessions.sessions.contains_key(session_id) {
610 *user_sessions.active_session_id.lock().await = session_id.to_string();
611 }
612 }
613 }
614
615 pub fn get_user_home_dir(&self, user_key: &str) -> Option<std::path::PathBuf> {
626 let hash = trustee_core::user_hash(user_key);
627 dirs::home_dir().map(|home| home.join(".trustee").join("users").join(&hash))
628 }
629
630 pub fn get_user_config_and_home(&self, user_key: &str) -> (Option<String>, Option<std::path::PathBuf>) {
635 (self.config_toml.clone(), self.get_user_home_dir(user_key))
636 }
637
638 fn apply_user_isolation(&self, session: &mut Session, user_key: &str) {
647 let user_hash = trustee_core::user_hash(user_key);
648
649 let user_home = if let Some(home) = dirs::home_dir() {
651 let user_home = home.join(".trustee").join("users").join(&user_hash);
652 session.home_dir = Some(user_home.clone());
653 Some(user_home)
654 } else {
655 None
656 };
657
658 session.project_id = Some(format!("web{}", &user_hash[..16]));
659
660 let shared_secrets = session.secrets.clone().unwrap_or_default();
667 let mut merged_secrets = shared_secrets.clone();
668
669 if let Some(ref user_home) = user_home {
670 if let Ok(merged) = self.load_user_secrets(user_home, &merged_secrets) {
671 merged_secrets = merged;
672 }
673 }
674
675 if let Some(ref user_home) = user_home {
677 if let Some(merged) = self.merge_user_config(user_home, session.config_toml.as_deref()) {
678 session.config_toml = Some(merged);
679 tracing::debug!("Merged per-user config into session");
680 }
681 }
682
683 if let Some(ref mut config_toml) = session.config_toml {
685 substitute_env_vars(config_toml, &merged_secrets);
686 }
687
688 session.secrets = Some(shared_secrets);
690 }
691
692 fn load_user_secrets(
695 &self,
696 user_home: &std::path::Path,
697 base: &std::collections::HashMap<String, String>,
698 ) -> std::io::Result<std::collections::HashMap<String, String>> {
699 let user_env_path = user_home.join(".env");
700 if !user_env_path.exists() {
701 return Ok(base.clone());
702 }
703 let content = std::fs::read_to_string(&user_env_path)?;
704 let mut merged = base.clone();
705 for line in content.lines() {
706 let line = line.trim();
707 if line.is_empty() || line.starts_with('#') {
708 continue;
709 }
710 if let Some((key, value)) = line.split_once('=') {
711 let key = key.trim().to_string();
712 let value = value
713 .trim()
714 .trim_matches('"')
715 .trim_matches('\'')
716 .to_string();
717 merged.insert(key, value);
718 }
719 }
720 tracing::debug!("Loaded per-user secrets from {}", user_env_path.display());
721 Ok(merged)
722 }
723
724 fn merge_user_config(
739 &self,
740 user_home: &std::path::Path,
741 shared_config: Option<&str>,
742 ) -> Option<String> {
743 let user_config_path = user_home.join("config").join("trustee.toml");
744 if !user_config_path.exists() {
745 return None;
746 }
747 let user_config_toml = std::fs::read_to_string(&user_config_path).ok()?;
748 let shared = shared_config
749 .unwrap_or("")
750 .parse::<toml::Value>()
751 .ok()?;
752 let overlay = user_config_toml.parse::<toml::Value>().ok()?;
753
754 let allowed = |section: &str| {
755 section == "mcp"
756 || section == "thq"
760 || (self.allow_llm_overlay && section == "llm")
761 };
762 let dir_name = user_home
765 .file_name()
766 .and_then(|n| n.to_str())
767 .unwrap_or("<unknown>");
768 let masked_user = dir_name.get(..8).unwrap_or(dir_name);
769
770 let mut filtered_overlay = toml::map::Map::new();
771 if let Some(table) = overlay.as_table() {
772 for (section, value) in table {
773 if allowed(section) {
774 filtered_overlay.insert(section.clone(), value.clone());
775 } else {
776 tracing::warn!(
777 "user config overlay: dropping non-allowlisted section [{}] for user {}",
778 section,
779 masked_user
780 );
781 }
782 }
783 }
784
785 if filtered_overlay.is_empty() {
786 return None;
788 }
789 let overlay = toml::Value::Table(filtered_overlay);
790
791 let mut shared = shared;
792 deep_merge_toml(&mut shared, &overlay);
793 let merged = toml::to_string(&shared).ok()?;
794 tracing::debug!("Merged per-user config from {}", user_config_path.display());
795 Some(merged)
796 }
797
798 pub fn resolve_user_config(&self, user_key: &str) -> Option<String> {
807 let config_toml = self.config_toml.clone()?;
808
809 let user_home = self.get_user_home_dir(user_key)?;
811
812 let mut merged_secrets = self.secrets.clone().unwrap_or_default();
814 if let Ok(merged) = self.load_user_secrets(&user_home, &merged_secrets) {
815 merged_secrets = merged;
816 }
817
818 let mut resolved = config_toml;
820 if let Some(merged) = self.merge_user_config(&user_home, Some(&resolved)) {
821 resolved = merged;
822 }
823
824 substitute_env_vars(&mut resolved, &merged_secrets);
826
827 Some(resolved)
828 }
829
830 pub async fn get_or_build_mcp_loader(
846 &self,
847 user_key: &str,
848 token_store: &Arc<pep::MemoryTokenStore>,
849 ) -> Result<Option<std::sync::Arc<abk::agent::McpToolLoader>>, String> {
850 let user_hash = trustee_core::user_hash(user_key);
851
852 let resolved = self.resolve_user_config(user_key);
854 let fingerprint = fingerprint_mcp_section(resolved.as_deref());
855
856 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
858 if entry.degraded.is_none() {
859 if entry.fingerprint == fingerprint {
860 return Ok(entry.loader.clone());
861 }
862 } else if let (Some(err), Some(failed_at)) = (&entry.degraded, entry.failed_at) {
863 let backoff = chrono::Duration::from_std(MCP_BUILD_RETRY_BACKOFF)
864 .unwrap_or_else(|_| chrono::Duration::seconds(30));
865 if chrono::Utc::now() < failed_at + backoff {
866 return Err(err.clone());
867 }
868 }
869 }
870
871 let lock = self
873 .mcp_build_locks
874 .entry(user_hash.clone())
875 .or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
876 .clone();
877 let _guard = lock.lock().await;
878
879 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
881 if entry.degraded.is_none() && entry.fingerprint == fingerprint {
882 return Ok(entry.loader.clone());
883 }
884 }
885
886 match self
887 .build_mcp_loader(&user_hash, resolved.as_deref(), fingerprint, token_store)
888 .await
889 {
890 Ok(entry) => {
891 self.mcp_loaders.insert(user_hash.clone(), entry);
892 Ok(self.mcp_loaders.get(&user_hash).unwrap().loader.clone())
893 }
894 Err(err) => {
895 tracing::warn!(
896 "MCP loader build FAILED for user {}; dispatch fails loud, retry after {:?}",
897 &user_hash[..8.min(user_hash.len())],
898 MCP_BUILD_RETRY_BACKOFF
899 );
900 self.mcp_loaders.insert(
901 user_hash,
902 McpLoaderEntry {
903 loader: None,
904 fingerprint,
905 built_at: chrono::Utc::now(),
906 degraded: Some(err.clone()),
907 failed_at: Some(chrono::Utc::now()),
908 },
909 );
910 Err(err)
911 }
912 }
913 }
914
915 async fn build_mcp_loader(
918 &self,
919 user_hash: &str,
920 resolved: Option<&str>,
921 fingerprint: u64,
922 token_store: &Arc<pep::MemoryTokenStore>,
923 ) -> Result<McpLoaderEntry, String> {
924 let mcp_config: Option<abk::config::McpConfig> = match resolved {
925 Some(toml_str) => {
926 let value = toml_str
927 .parse::<toml::Value>()
928 .map_err(|e| format!("config parse failed: {}", e))?;
929 match value.get("mcp") {
930 Some(section) => {
931 use serde::Deserialize as _;
932 Some(
933 abk::config::McpConfig::deserialize(section.clone())
934 .map_err(|e| format!("invalid [mcp] config: {}", e))?,
935 )
936 }
937 None => None,
938 }
939 }
940 None => None,
941 };
942
943 let loader = match mcp_config {
944 Some(cfg) if cfg.enabled => {
945 let built = abk::agent::McpToolLoader::with_token_store(
946 &cfg,
947 Some(token_store.clone() as std::sync::Arc<dyn pep::token_store::TokenStore>),
948 )
949 .await
950 .map_err(|e| format!("MCP loader build failed: {}", e))?;
951
952 let servers: Vec<String> = built
956 .server_statuses
957 .iter()
958 .map(|s| {
959 if s.connected {
960 format!("{}(up,{}tools)", s.name, s.tool_count)
961 } else {
962 format!("{}(DOWN)", s.name)
963 }
964 })
965 .collect();
966 tracing::info!(
967 "MCP loader built for user {}: servers=[{}] total_tools={}",
968 &user_hash[..8.min(user_hash.len())],
969 servers.join(", "),
970 built.tool_count
971 );
972 Some(std::sync::Arc::new(built))
973 }
974 _ => None,
975 };
976
977 Ok(McpLoaderEntry {
978 loader,
979 fingerprint,
980 built_at: chrono::Utc::now(),
981 degraded: None,
982 failed_at: None,
983 })
984 }
985
986 fn spawn_user_drain_task(
988 &self,
989 session_id: String,
990 session: Arc<Mutex<Session>>,
991 ws_tx: broadcast::Sender<String>,
992 mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
993 ) {
994 let mut last_broadcast_state: Option<String> = None;
1000 tokio::spawn(async move {
1001 while let Some(msg) = workflow_rx.recv().await {
1002 {
1003 let mut session = session.lock().await;
1004 session.handle_workflow_message(msg.clone());
1005
1006 let state_str = match session.workflow_state {
1007 trustee_core::types::WorkflowState::Idle => "Idle",
1008 trustee_core::types::WorkflowState::Running => "Running",
1009 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
1010 };
1011 if last_broadcast_state.as_deref() != Some(state_str) {
1012 last_broadcast_state = Some(state_str.to_string());
1013 let state_msg = serde_json::json!({
1014 "type": "StateChanged",
1015 "state": state_str
1016 });
1017 let _ = ws_tx.send(state_msg.to_string());
1018 }
1019 }
1020
1021 let json =
1022 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
1023 let _ = ws_tx.send(json);
1024 }
1025 tracing::debug!("Drain task ended for session: {}", session_id);
1026 });
1027 }
1028
1029 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
1032 let default_user = self
1034 .sessions
1035 .get("default")
1036 .expect("default user must exist");
1037 let first_entry = default_user
1038 .sessions
1039 .iter()
1040 .next()
1041 .expect("default user must have at least one session");
1042 let session = first_entry.session.clone();
1043 let ws_tx = first_entry.ws_tx.clone();
1044 let session_id = first_entry.key().clone();
1045 drop(first_entry);
1046 drop(default_user);
1047
1048 let mut last_broadcast_state = Some("Running".to_string());
1052 tokio::spawn(async move {
1053 while let Some(msg) = workflow_rx.recv().await {
1054 {
1055 let mut session = session.lock().await;
1056 session.handle_workflow_message(msg.clone());
1057
1058 let state_str = match session.workflow_state {
1059 trustee_core::types::WorkflowState::Idle => "Idle",
1060 trustee_core::types::WorkflowState::Running => "Running",
1061 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
1062 };
1063 if last_broadcast_state.as_deref() != Some(state_str) {
1064 last_broadcast_state = Some(state_str.to_string());
1065 let state_msg = serde_json::json!({
1066 "type": "StateChanged",
1067 "state": state_str
1068 });
1069 let _ = ws_tx.send(state_msg.to_string());
1070 }
1071 }
1072
1073 let json =
1074 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
1075 let _ = ws_tx.send(json);
1076 }
1077 tracing::debug!("Drain task ended for session: {}", session_id);
1078 });
1079 }
1080
1081 pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
1083 let Some(ref auth) = self.auth else {
1084 return "default".to_string();
1085 };
1086
1087 if let Some(token) = headers
1089 .get(axum::http::header::AUTHORIZATION)
1090 .and_then(|v| v.to_str().ok())
1091 .and_then(|v| v.strip_prefix("Bearer "))
1092 .map(|s| s.to_string())
1093 {
1094 if token.starts_with("dev:") {
1095 let parts: Vec<&str> = token.splitn(4, ':').collect();
1096 if parts.len() >= 4 {
1097 return format!("dev:{}", parts[1]);
1098 }
1099 }
1100 if let Ok(claims) = auth.validate_token(&token).await {
1101 return claims.sub;
1102 }
1103 }
1104
1105 let cookie_session_id = headers
1107 .get(axum::http::header::COOKIE)
1108 .and_then(|v| v.to_str().ok())
1109 .and_then(|cookies| {
1110 cookies
1111 .split(';')
1112 .map(|c| c.trim())
1113 .find_map(|c| {
1114 c.strip_prefix(&format!("{}=", auth.config.cookie_name))
1115 .map(|s| s.to_string())
1116 })
1117 });
1118
1119 if let Some(session_id) = cookie_session_id {
1120 if session_id.starts_with("dev:") {
1121 let parts: Vec<&str> = session_id.splitn(4, ':').collect();
1122 if parts.len() >= 4 {
1123 return format!("dev:{}", parts[1]);
1124 }
1125 }
1126
1127 if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
1128 if let Ok(claims) = auth.validate_token(&access_token).await {
1129 return claims.sub;
1130 }
1131 }
1132 }
1133
1134 "default".to_string()
1135 }
1136}
1137
1138struct SerializableMessage<'a>(&'a TuiMessage);
1144
1145impl<'a> serde::Serialize for SerializableMessage<'a> {
1146 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1147 where
1148 S: serde::Serializer,
1149 {
1150 use serde::ser::SerializeStruct;
1151
1152 match self.0 {
1153 TuiMessage::OutputLine(line) => {
1154 let mut s = serializer.serialize_struct("msg", 2)?;
1155 s.serialize_field("type", "OutputLine")?;
1156 s.serialize_field("line", line)?;
1157 s.end()
1158 }
1159 TuiMessage::StreamDelta(delta) => {
1160 let mut s = serializer.serialize_struct("msg", 2)?;
1161 s.serialize_field("type", "StreamDelta")?;
1162 s.serialize_field("delta", delta)?;
1163 s.end()
1164 }
1165 TuiMessage::ReasoningDelta(delta) => {
1166 let mut s = serializer.serialize_struct("msg", 2)?;
1167 s.serialize_field("type", "ReasoningDelta")?;
1168 s.serialize_field("delta", delta)?;
1169 s.end()
1170 }
1171 TuiMessage::WorkflowCompleted => {
1172 let mut s = serializer.serialize_struct("msg", 2)?;
1173 s.serialize_field("type", "WorkflowCompleted")?;
1174 s.serialize_field("state", "Idle")?;
1175 s.end()
1176 }
1177 TuiMessage::WorkflowError(err) => {
1178 let mut s = serializer.serialize_struct("msg", 2)?;
1179 s.serialize_field("type", "WorkflowError")?;
1180 s.serialize_field("error", err)?;
1181 s.end()
1182 }
1183 TuiMessage::ResumeInfo(info) => match info {
1184 Some(ri) => {
1185 let mut s = serializer.serialize_struct("msg", 5)?;
1186 s.serialize_field("type", "ResumeInfo")?;
1187 s.serialize_field("state", "Idle")?;
1188 s.serialize_field("session_id", &ri.session_id)?;
1189 s.serialize_field("checkpoint_id", &ri.checkpoint_id)?;
1190 s.serialize_field("iteration", &ri.iteration)?;
1191 s.end()
1192 }
1193 None => {
1194 let mut s = serializer.serialize_struct("msg", 2)?;
1195 s.serialize_field("type", "ResumeInfo")?;
1196 s.serialize_field("state", "Idle")?;
1197 s.end()
1198 }
1199 },
1200 TuiMessage::TodoUpdate(content) => {
1201 let mut s = serializer.serialize_struct("msg", 2)?;
1202 s.serialize_field("type", "TodoUpdate")?;
1203 s.serialize_field("content", content)?;
1204 s.end()
1205 }
1206 TuiMessage::WorkflowCancelled => {
1207 let mut s = serializer.serialize_struct("msg", 2)?;
1208 s.serialize_field("type", "WorkflowCancelled")?;
1209 s.serialize_field("state", "Idle")?;
1210 s.end()
1211 }
1212 TuiMessage::HandoffReady(briefing) => {
1213 let mut s = serializer.serialize_struct("msg", 3)?;
1214 s.serialize_field("type", "HandoffReady")?;
1215 s.serialize_field("state", "Idle")?;
1216 s.serialize_field("briefing", briefing)?;
1217 s.end()
1218 }
1219 TuiMessage::SessionRotated { old, new } => {
1220 let mut s = serializer.serialize_struct("msg", 3)?;
1221 s.serialize_field("type", "SessionRotated")?;
1222 s.serialize_field("old", old)?;
1223 s.serialize_field("new", new)?;
1224 s.end()
1225 }
1226 TuiMessage::HandoffFailed => {
1227 let mut s = serializer.serialize_struct("msg", 2)?;
1228 s.serialize_field("type", "HandoffFailed")?;
1229 s.serialize_field("state", "Idle")?;
1230 s.end()
1231 }
1232 TuiMessage::ToolPending {
1233 tool_name,
1234 hint,
1235 } => {
1236 let mut s = serializer.serialize_struct("msg", 3)?;
1237 s.serialize_field("type", "ToolPending")?;
1238 s.serialize_field("tool_name", tool_name)?;
1239 s.serialize_field("hint", hint)?;
1240 s.end()
1241 }
1242 TuiMessage::ToolDone {
1243 tool_name,
1244 success,
1245 hint,
1246 } => {
1247 let mut s = serializer.serialize_struct("msg", 4)?;
1248 s.serialize_field("type", "ToolDone")?;
1249 s.serialize_field("tool_name", tool_name)?;
1250 s.serialize_field("success", success)?;
1251 s.serialize_field("hint", hint)?;
1252 s.end()
1253 }
1254 TuiMessage::ContextTokensUpdated(count) => {
1255 let mut s = serializer.serialize_struct("msg", 2)?;
1256 s.serialize_field("type", "ContextTokensUpdated")?;
1257 s.serialize_field("count", count)?;
1258 s.end()
1259 }
1260 TuiMessage::McpServerStatus {
1261 name,
1262 connected,
1263 tool_count,
1264 error,
1265 } => {
1266 let mut s = serializer.serialize_struct("msg", 5)?;
1267 s.serialize_field("type", "McpServerStatus")?;
1268 s.serialize_field("name", name)?;
1269 s.serialize_field("connected", connected)?;
1270 s.serialize_field("tool_count", tool_count)?;
1271 s.serialize_field("error", error)?;
1272 s.end()
1273 }
1274 TuiMessage::SessionTitleUpdated(title) => {
1275 let mut s = serializer.serialize_struct("msg", 2)?;
1276 s.serialize_field("type", "SessionTitleUpdated")?;
1277 s.serialize_field("title", title)?;
1278 s.end()
1279 }
1280 }
1281 }
1282}
1283
1284fn deep_merge_toml(base: &mut toml::Value, overlay: &toml::Value) {
1295 match (base, overlay) {
1296 (toml::Value::Table(base_table), toml::Value::Table(overlay_table)) => {
1297 for (key, overlay_val) in overlay_table {
1298 match base_table.get_mut(key) {
1299 Some(base_val) => {
1300 deep_merge_toml(base_val, overlay_val);
1302 }
1303 None => {
1304 base_table.insert(key.clone(), overlay_val.clone());
1306 }
1307 }
1308 }
1309 }
1310 (base, overlay) => {
1312 *base = overlay.clone();
1313 }
1314 }
1315}
1316
1317fn substitute_env_vars(s: &mut String, secrets: &std::collections::HashMap<String, String>) {
1322 let mut result = String::with_capacity(s.len());
1324 let bytes = s.as_bytes();
1325 let mut i = 0;
1326
1327 while i < bytes.len() {
1328 if i + 1 < bytes.len() && bytes[i] == b'$' && bytes[i + 1] == b'{' {
1329 if let Some(end) = s[i + 2..].find('}') {
1331 let var_name = &s[i + 2..i + 2 + end];
1332 if let Some(value) = secrets.get(var_name) {
1334 result.push_str(value);
1335 } else if let Ok(value) = std::env::var(var_name) {
1336 result.push_str(&value);
1337 } else {
1338 result.push_str(&s[i..i + 2 + end + 1]);
1340 }
1341 i = i + 2 + end + 1;
1342 } else {
1343 result.push('$');
1345 i += 1;
1346 }
1347 } else {
1348 result.push(bytes[i] as char);
1349 i += 1;
1350 }
1351 }
1352
1353 *s = result;
1354}
1355
1356fn fingerprint_mcp_section(resolved: Option<&str>) -> u64 {
1366 use sha2::{Digest, Sha256};
1367 let section = resolved
1368 .and_then(|s| s.parse::<toml::Value>().ok())
1369 .and_then(|v| v.get("mcp").cloned());
1370 let bytes = match section {
1371 Some(v) => v.to_string().into_bytes(),
1372 None => b"<no-mcp>".to_vec(),
1373 };
1374 let digest = Sha256::digest(&bytes);
1375 u64::from_be_bytes(digest[..8].try_into().expect("sha256 digest >= 8 bytes"))
1376}
1377
1378#[cfg(test)]
1379mod tests {
1380 use super::*;
1381
1382 fn test_state() -> ServerState {
1384 let (session, _rx) = Session::new();
1385 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1386 ServerState::new(session, ws_tx, None)
1387 }
1388
1389 fn temp_user_home(tag: &str) -> std::path::PathBuf {
1391 static COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1392 let n = COUNTER.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1393 let dir = std::env::temp_dir().join(format!(
1394 "trustee-state-test-{}-{}-{}",
1395 tag,
1396 std::process::id(),
1397 n
1398 ));
1399 std::fs::create_dir_all(dir.join("config")).expect("create temp user home");
1400 dir
1401 }
1402
1403 fn parse(toml_str: &str) -> toml::Value {
1404 toml_str.parse::<toml::Value>().expect("valid test TOML")
1405 }
1406
1407 #[test]
1410 fn overlay_allowlist_drops_non_allowlisted_sections() {
1411 let state = test_state();
1412 let home = temp_user_home("allowlist");
1413 std::fs::write(
1414 home.join("config").join("trustee.toml"),
1415 "[server]\nport = 1\n\n[auth]\nmode = \"kanidm\"\n\n[mcp]\nmode = \"user\"\n",
1416 )
1417 .expect("write overlay");
1418
1419 let shared = "[server]\nport = 8080\n\n[mcp]\nmode = \"shared\"\n";
1420 let merged = state
1421 .merge_user_config(&home, Some(shared))
1422 .expect("overlay has allowlisted content");
1423
1424 let merged_val = parse(&merged);
1425 let expected = parse("[server]\nport = 8080\n\n[mcp]\nmode = \"user\"\n");
1426 assert_eq!(merged_val, expected, "merged config must be shared + [mcp] overlay only");
1427 assert!(merged_val.get("auth").is_none(), "overlay [auth] must be dropped");
1428 }
1429
1430 #[test]
1432 fn no_overlay_file_returns_none() {
1433 let state = test_state();
1434 let home = temp_user_home("empty");
1435 assert!(state
1436 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1437 .is_none());
1438 }
1439
1440 #[test]
1442 fn overlay_with_no_allowlisted_sections_is_noop() {
1443 let state = test_state();
1444 let home = temp_user_home("all-dropped");
1445 std::fs::write(
1446 home.join("config").join("trustee.toml"),
1447 "[server]\nport = 1\n\n[storage]\npath = \"/tmp/x\"\n",
1448 )
1449 .expect("write overlay");
1450 assert!(state
1451 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1452 .is_none());
1453 }
1454
1455 #[test]
1458 fn llm_overlay_dropped_by_default_and_kept_when_enabled() {
1459 let shared = "[llm]\nprovider = \"openai\"\n\n[mcp]\nmode = \"shared\"\n";
1460 let overlay = "[llm]\nprovider = \"anthropic\"\n";
1461
1462 let state = test_state();
1464 assert!(!state.allow_llm_overlay);
1465 let home = temp_user_home("llm-off");
1466 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1467 assert!(state.merge_user_config(&home, Some(shared)).is_none());
1468
1469 let state = test_state().with_allow_llm_overlay(true);
1471 assert!(state.allow_llm_overlay);
1472 let home = temp_user_home("llm-on");
1473 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1474 let merged = state
1475 .merge_user_config(&home, Some(shared))
1476 .expect("[llm] overlay applies when opted in");
1477 let merged_val = parse(&merged);
1478 assert_eq!(merged_val["llm"]["provider"].as_str(), Some("anthropic"));
1479 assert_eq!(merged_val["mcp"]["mode"].as_str(), Some("shared"));
1480 }
1481
1482 #[test]
1486 fn user_home_dir_uses_consolidated_hash() {
1487 let state = test_state();
1488 if let Some(home) = state.get_user_home_dir("farzan@example.com") {
1489 assert_eq!(
1490 home.file_name().and_then(|n| n.to_str()),
1491 Some(trustee_core::user_hash("farzan@example.com")).as_deref()
1492 );
1493 let users_root = dirs::home_dir().unwrap().join(".trustee").join("users");
1494 assert_eq!(home.parent(), Some(&users_root).map(|p| p.as_path()));
1495 }
1496 }
1498
1499 #[tokio::test]
1504 async fn agent_principals_get_isolated_session_buckets() {
1505 let state = test_state();
1506 let key_a = "agent-farzan";
1507 let key_b = "agent-paydar";
1508
1509 let (sid_a, session_a, _tx_a, _ts_a) = state.ensure_active_session(key_a).await;
1510 let (_sid_b, _session_b, _tx_b, _ts_b) = state.ensure_active_session(key_b).await;
1511
1512 assert!(
1514 state.get_session_by_any_id(key_a, &sid_a).await.is_some(),
1515 "owner bucket resolves its own session"
1516 );
1517 assert!(
1518 state.get_session_by_any_id(key_b, &sid_a).await.is_none(),
1519 "cross-agent session access must be 404/None"
1520 );
1521 assert!(Arc::strong_count(&session_a) >= 1);
1522
1523 if let Some(home) = state.get_user_home_dir(key_a) {
1525 assert_eq!(
1526 home.file_name().and_then(|n| n.to_str()),
1527 Some(trustee_core::user_hash(key_a)).as_deref()
1528 );
1529 }
1530 }
1531}
1532
1533#[cfg(test)]
1538mod mcp_loader_cache_tests {
1539 use super::*;
1540
1541 fn state_with_shared(shared: &str) -> ServerState {
1542 let (session, _rx) = Session::new();
1543 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1544 let mut state = ServerState::new(session, ws_tx, None);
1545 state.config_toml = Some(shared.to_string());
1546 state
1547 }
1548
1549 struct TempUser {
1553 key: String,
1554 home: std::path::PathBuf,
1555 }
1556
1557 impl TempUser {
1558 fn new(tag: &str) -> Self {
1559 let key = format!("16c-{tag}-{}@test.invalid", std::process::id());
1560 let home = dirs::home_dir()
1561 .expect("HOME available in test env")
1562 .join(".trustee")
1563 .join("users")
1564 .join(trustee_core::user_hash(&key));
1565 std::fs::create_dir_all(home.join("config")).expect("create user home");
1566 Self { key, home }
1567 }
1568
1569 fn write_overlay(&self, toml_str: &str) {
1570 std::fs::write(self.home.join("config").join("trustee.toml"), toml_str)
1571 .expect("write overlay");
1572 }
1573 }
1574
1575 impl Drop for TempUser {
1576 fn drop(&mut self) {
1577 let _ = std::fs::remove_dir_all(&self.home);
1578 }
1579 }
1580
1581 #[tokio::test]
1582 async fn no_mcp_config_caches_disabled_marker() {
1583 let state = state_with_shared("[server]\nport = 8080\n");
1584 let user = TempUser::new("nomcp");
1585 let ts = Arc::new(pep::MemoryTokenStore::new());
1586
1587 let first = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1588 assert!(first.is_none(), "no [mcp] anywhere → disabled marker");
1589 assert_eq!(state.mcp_loaders.len(), 1, "exactly one cache entry");
1590
1591 let second = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1593 assert!(second.is_none());
1594 assert_eq!(state.mcp_loaders.len(), 1);
1595 }
1596
1597 #[tokio::test]
1598 async fn concurrent_cold_builds_single_flight() {
1599 let state = Arc::new(state_with_shared("[server]\nport = 8080\n"));
1600 let user = TempUser::new("singleflight");
1601 let ts = Arc::new(pep::MemoryTokenStore::new());
1602
1603 let mut handles = Vec::new();
1604 for _ in 0..5 {
1605 let state = state.clone();
1606 let key = user.key.clone();
1607 let ts = ts.clone();
1608 handles.push(tokio::spawn(async move {
1609 state.get_or_build_mcp_loader(&key, &ts).await
1610 }));
1611 }
1612 for h in handles {
1613 h.await.unwrap().expect("all five succeed");
1614 }
1615 assert_eq!(state.mcp_loaders.len(), 1, "single-flight → one entry");
1616 }
1617
1618 #[tokio::test]
1619 async fn fingerprint_change_triggers_rebuild() {
1620 let state = state_with_shared("[server]\nport = 8080\n");
1621 let user = TempUser::new("fpchange");
1622 let ts = Arc::new(pep::MemoryTokenStore::new());
1623
1624 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v1\"\nurl = \"http://127.0.0.1:9/sse\"\n");
1627 let v1 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1628 assert!(v1.is_some(), "enabled [mcp] → real loader");
1629 let fp1 = state
1630 .mcp_loaders
1631 .get(&trustee_core::user_hash(&user.key))
1632 .unwrap()
1633 .fingerprint;
1634
1635 std::thread::sleep(std::time::Duration::from_millis(5));
1637 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v2\"\nurl = \"http://127.0.0.1:9/other\"\n");
1638 let v2 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1639 assert!(v2.is_some());
1640 let entry = state
1641 .mcp_loaders
1642 .get(&trustee_core::user_hash(&user.key))
1643 .unwrap();
1644 assert_ne!(
1645 entry.fingerprint, fp1,
1646 "fingerprint must change with content"
1647 );
1648 assert!(entry.degraded.is_none());
1649
1650 let _still_usable = v1.as_ref().unwrap().tool_count;
1652 }
1653
1654 #[tokio::test]
1655 async fn degraded_entry_fails_loud_within_backoff_and_isolates_users() {
1656 let state = state_with_shared("[server]\nport = 8080\n");
1657 let bad = TempUser::new("degraded-bad");
1658 let good = TempUser::new("degraded-good");
1659 let ts = Arc::new(pep::MemoryTokenStore::new());
1660
1661 bad.write_overlay("[mcp]\nenabled = \"not-a-bool\"\n");
1663 let err = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1664 Ok(_) => panic!("invalid [mcp] must fail loud"),
1665 Err(e) => e,
1666 };
1667 assert!(
1668 err.contains("invalid [mcp]"),
1669 "surfaces the parse error: {err}"
1670 );
1671
1672 let entry = state
1673 .mcp_loaders
1674 .get(&trustee_core::user_hash(&bad.key))
1675 .unwrap();
1676 assert!(entry.degraded.is_some(), "poison entry recorded");
1677
1678 let err2 = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1680 Ok(_) => panic!("still within backoff"),
1681 Err(e) => e,
1682 };
1683 assert_eq!(err, err2, "same cached error");
1684
1685 let good_loader = state
1687 .get_or_build_mcp_loader(&good.key, &ts)
1688 .await
1689 .expect("other user unaffected");
1690 assert!(
1691 good_loader.is_none(),
1692 "good user has no [mcp] → disabled marker"
1693 );
1694 }
1695}