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(Clone)]
124pub struct ServerState {
125 pub sessions: SessionRegistry,
127 pub ws_tx: broadcast::Sender<String>,
129 pub auth: Option<Arc<AuthState>>,
131 pub config_toml: Option<String>,
133 pub secrets: Option<std::collections::HashMap<String, String>>,
135 pub build_info: Option<trustee_core::types::BuildInfo>,
137 pub workflow_semaphore: Arc<tokio::sync::Semaphore>,
140 pub max_sessions_per_user: usize,
142 pub allow_llm_overlay: bool,
145 pub mcp_loaders: Arc<DashMap<String, McpLoaderEntry>>,
147 mcp_build_locks: Arc<DashMap<String, Arc<tokio::sync::Mutex<()>>>>,
149}
150
151impl ServerState {
152 pub fn new(
154 session: Session,
155 ws_tx: broadcast::Sender<String>,
156 auth: Option<Arc<AuthState>>,
157 ) -> Self {
158 let sessions = Arc::new(DashMap::new());
159
160 let token_store = Arc::new(pep::MemoryTokenStore::new());
162 let (ws_tx_entry, _) = broadcast::channel::<String>(256);
163
164 let now = chrono::Utc::now();
165 let initial_entry = UserSessionEntry {
166 session: Arc::new(Mutex::new(session)),
167 ws_tx: ws_tx_entry,
168 created_at: now,
169 last_active: Arc::new(Mutex::new(now)),
170 };
171
172 let user_sessions = UserSessions {
173 sessions: DashMap::new(),
174 token_store,
175 active_session_id: Mutex::new(String::new()),
176 };
177 user_sessions.sessions.insert("default".to_string(), initial_entry);
178
179 sessions.insert("default".to_string(), user_sessions);
180
181 Self {
182 sessions,
183 ws_tx,
184 auth,
185 config_toml: None,
186 secrets: None,
187 build_info: None,
188 workflow_semaphore: Arc::new(tokio::sync::Semaphore::new(8)),
189 max_sessions_per_user: 4,
190 allow_llm_overlay: false,
191 mcp_loaders: Arc::new(DashMap::new()),
192 mcp_build_locks: Arc::new(DashMap::new()),
193 }
194 }
195
196 pub fn with_config_toml(mut self, config_toml: String) -> Self {
197 self.config_toml = Some(config_toml);
198 self
199 }
200
201 pub fn with_secrets(mut self, secrets: std::collections::HashMap<String, String>) -> Self {
202 self.secrets = Some(secrets);
203 self
204 }
205
206 pub fn with_build_info(mut self, build_info: trustee_core::types::BuildInfo) -> Self {
207 self.build_info = Some(build_info);
208 self
209 }
210
211 pub fn with_max_concurrent_workflows(mut self, max: usize) -> Self {
212 self.workflow_semaphore = Arc::new(tokio::sync::Semaphore::new(max));
213 self
214 }
215
216 pub fn with_max_sessions_per_user(mut self, max: usize) -> Self {
218 self.max_sessions_per_user = max;
219 self
220 }
221
222 pub fn with_allow_llm_overlay(mut self, allow: bool) -> Self {
225 self.allow_llm_overlay = allow;
226 self
227 }
228
229 pub async fn create_session(
239 &self,
240 user_key: &str,
241 session_name: Option<String>,
242 identity: Option<String>,
243 activate: bool,
244 ) -> Result<String, SessionError> {
245 let user_sessions = self
247 .sessions
248 .entry(user_key.to_string())
249 .or_insert_with(|| UserSessions {
250 sessions: DashMap::new(),
251 token_store: Arc::new(pep::MemoryTokenStore::new()),
252 active_session_id: Mutex::new(String::new()),
253 });
254
255 if user_sessions.sessions.len() >= self.max_sessions_per_user {
257 return Err(SessionError::MaxSessionsReached(self.max_sessions_per_user));
258 }
259
260 let (mut session, workflow_rx) = Session::new();
262
263 if let Some(ref config_toml) = self.config_toml {
265 session.config_toml = Some(config_toml.clone());
266 session.parse_auto_handoff_config();
267 if let Ok(table) = config_toml.parse::<toml::Value>() {
268 if let Some(name) = table
269 .get("agent")
270 .and_then(|a| a.get("name"))
271 .and_then(|n| n.as_str())
272 {
273 session.agent_name = name.to_string();
274 }
275 }
276 }
277
278 session.secrets = self.secrets.clone();
279 session.build_info = self.build_info.clone();
280
281 self.apply_user_isolation(&mut session, user_key);
283
284 session.session_name = session_name;
286
287 session.identity = identity;
289
290 let (ws_tx_entry, _) = broadcast::channel::<String>(256);
292
293 let session_id = format!(
295 "session_{}_{}",
296 chrono::Utc::now().format("%Y_%m_%d_%H_%M"),
297 &uuid::Uuid::new_v4().to_string()[..8]
298 );
299
300 let now = chrono::Utc::now();
301
302 user_sessions.sessions.insert(
304 session_id.clone(),
305 UserSessionEntry {
306 session: Arc::new(Mutex::new(session)),
307 ws_tx: ws_tx_entry.clone(),
308 created_at: now,
309 last_active: Arc::new(Mutex::new(now)),
310 },
311 );
312
313 if activate {
322 *user_sessions.active_session_id.lock().await = session_id.clone();
323 }
324
325 let session_arc = user_sessions
327 .sessions
328 .get(&session_id)
329 .map(|e| e.session.clone());
330 if let Some(session_arc) = session_arc {
331 self.spawn_user_drain_task(
332 session_id.clone(),
333 session_arc,
334 ws_tx_entry,
335 workflow_rx,
336 );
337 }
338
339 Ok(session_id)
340 }
341
342 pub async fn get_session(
345 &self,
346 user_key: &str,
347 session_id: &str,
348 ) -> Option<(Arc<Mutex<Session>>, broadcast::Sender<String>)> {
349 let user_sessions = self.sessions.get(user_key)?;
350 let entry = user_sessions.sessions.get(session_id)?;
351
352 let now = chrono::Utc::now();
354 *entry.last_active.lock().await = now;
355
356 Some((entry.session.clone(), entry.ws_tx.clone()))
357 }
358
359 pub async fn get_session_by_any_id(
377 &self,
378 user_key: &str,
379 id: &str,
380 ) -> Option<(String, Arc<Mutex<Session>>, broadcast::Sender<String>)> {
381 let user_sessions = self.sessions.get(user_key)?;
383 if let Some(entry) = user_sessions.sessions.get(id) {
384 let now = chrono::Utc::now();
386 *entry.last_active.lock().await = now;
387 return Some((id.to_string(), entry.session.clone(), entry.ws_tx.clone()));
388 }
389
390 for entry in user_sessions.sessions.iter() {
392 let session = entry.session.lock().await;
393 if session.session_id.as_deref() == Some(id) {
394 let key = entry.key().clone();
395 let ws_tx = entry.ws_tx.clone();
396 drop(session);
397 let now = chrono::Utc::now();
399 *entry.last_active.lock().await = now;
400 return Some((key, entry.session.clone(), ws_tx));
401 }
402 }
403
404 None
405 }
406
407 pub async fn list_sessions(&self, user_key: &str) -> Vec<SessionListItem> {
409 let Some(user_sessions) = self.sessions.get(user_key) else {
410 return Vec::new();
411 };
412
413 let mut items = Vec::new();
414 for entry in user_sessions.sessions.iter() {
415 let session = entry.session.lock().await;
416 let workflow_state = match session.workflow_state {
417 trustee_core::types::WorkflowState::Idle => "Idle",
418 trustee_core::types::WorkflowState::Running => "Running",
419 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
420 };
421 let last_active = entry.last_active.lock().await;
422 items.push(SessionListItem {
423 session_id: entry.key().clone(),
424 session_name: session.session_name.clone(),
425 workflow_state: workflow_state.to_string(),
426 created_at: entry.created_at.to_rfc3339(),
427 last_active: last_active.to_rfc3339(),
428 handoff_count: session.handoff_count,
429 });
430 }
431 drop(user_sessions);
432
433 items.sort_by(|a, b| b.last_active.cmp(&a.last_active));
435 items
436 }
437
438 pub async fn destroy_session(
440 &self,
441 user_key: &str,
442 session_id: &str,
443 ) -> Result<(), SessionError> {
444 let user_sessions = self
445 .sessions
446 .get(user_key)
447 .ok_or_else(|| SessionError::NotFound(session_id.to_string()))?;
448
449 {
451 let entry = user_sessions
452 .sessions
453 .get(session_id)
454 .ok_or_else(|| SessionError::NotFound(session_id.to_string()))?;
455 let session = entry.session.lock().await;
456 if session.workflow_state != trustee_core::types::WorkflowState::Idle {
457 let state_str = match session.workflow_state {
458 trustee_core::types::WorkflowState::Running => "Running",
459 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
460 _ => "Unknown",
461 };
462 return Err(SessionError::NotIdle(state_str.to_string()));
463 }
464 }
465
466 user_sessions.sessions.remove(session_id);
468
469 let mut active_id = user_sessions.active_session_id.lock().await;
471 if &*active_id == session_id {
472 let mut newest: Option<(String, chrono::DateTime<chrono::Utc>)> = None;
474 for entry in user_sessions.sessions.iter() {
475 let la = entry.last_active.lock().await;
476 if newest.as_ref().map_or(true, |(_, t)| *la > *t) {
477 newest = Some((entry.key().clone(), *la));
478 }
479 }
480 *active_id = newest.map(|(id, _)| id).unwrap_or_default();
481 }
482
483 Ok(())
484 }
485
486 pub async fn ensure_active_session(
495 &self,
496 user_key: &str,
497 ) -> (
498 String,
499 Arc<Mutex<Session>>,
500 broadcast::Sender<String>,
501 Arc<pep::MemoryTokenStore>,
502 ) {
503 let token_store = {
505 let user_sessions = self
506 .sessions
507 .entry(user_key.to_string())
508 .or_insert_with(|| UserSessions {
509 sessions: DashMap::new(),
510 token_store: Arc::new(pep::MemoryTokenStore::new()),
511 active_session_id: Mutex::new(String::new()),
512 });
513 user_sessions.token_store.clone()
514 };
515
516 let active_id = {
518 let user_sessions = self.sessions.get(user_key).unwrap();
519 let guard = user_sessions.active_session_id.lock().await;
520 guard.clone()
521 };
522
523 if !active_id.is_empty() {
524 if let Some((session, ws_tx)) = self.get_session(user_key, &active_id).await {
525 return (active_id, session, ws_tx, token_store);
526 }
527 }
529
530 let existing_session: Option<(String, Arc<Mutex<Session>>, broadcast::Sender<String>)> = {
534 let user_sessions = self.sessions.get(user_key).unwrap();
535 let result = user_sessions.sessions.iter().next().map(|first| {
536 (
537 first.key().clone(),
538 first.session.clone(),
539 first.ws_tx.clone(),
540 )
541 });
542 result
543 };
544 if let Some((id, session, ws_tx)) = existing_session {
545 let now = chrono::Utc::now();
546 if let Some(entry) = self.sessions.get(user_key) {
547 if let Some(e) = entry.sessions.get(&id) {
548 *e.last_active.lock().await = now;
549 }
550 *entry.active_session_id.lock().await = id.clone();
551 }
552
553 return (id, session, ws_tx, token_store);
554 }
555
556 let session_id = self
558 .create_session(user_key, None, None, true)
559 .await
560 .unwrap_or_else(|_| "default".to_string());
561
562 let (session, ws_tx) = self
563 .get_session(user_key, &session_id)
564 .await
565 .expect("just-created session must exist");
566
567 (session_id, session, ws_tx, token_store)
568 }
569
570 pub async fn ensure_user_session(
573 &self,
574 user_key: &str,
575 ) -> (Arc<Mutex<Session>>, broadcast::Sender<String>, Arc<pep::MemoryTokenStore>) {
576 let (_id, session, ws_tx, token_store) = self.ensure_active_session(user_key).await;
577 (session, ws_tx, token_store)
578 }
579
580 pub async fn set_active_session(&self, user_key: &str, session_id: &str) {
582 if let Some(user_sessions) = self.sessions.get(user_key) {
583 if user_sessions.sessions.contains_key(session_id) {
584 *user_sessions.active_session_id.lock().await = session_id.to_string();
585 }
586 }
587 }
588
589 pub fn get_user_home_dir(&self, user_key: &str) -> Option<std::path::PathBuf> {
600 let hash = trustee_core::user_hash(user_key);
601 dirs::home_dir().map(|home| home.join(".trustee").join("users").join(&hash))
602 }
603
604 pub fn get_user_config_and_home(&self, user_key: &str) -> (Option<String>, Option<std::path::PathBuf>) {
609 (self.config_toml.clone(), self.get_user_home_dir(user_key))
610 }
611
612 fn apply_user_isolation(&self, session: &mut Session, user_key: &str) {
621 let user_hash = trustee_core::user_hash(user_key);
622
623 let user_home = if let Some(home) = dirs::home_dir() {
625 let user_home = home.join(".trustee").join("users").join(&user_hash);
626 session.home_dir = Some(user_home.clone());
627 Some(user_home)
628 } else {
629 None
630 };
631
632 session.project_id = Some(format!("web{}", &user_hash[..16]));
633
634 let shared_secrets = session.secrets.clone().unwrap_or_default();
641 let mut merged_secrets = shared_secrets.clone();
642
643 if let Some(ref user_home) = user_home {
644 if let Ok(merged) = self.load_user_secrets(user_home, &merged_secrets) {
645 merged_secrets = merged;
646 }
647 }
648
649 if let Some(ref user_home) = user_home {
651 if let Some(merged) = self.merge_user_config(user_home, session.config_toml.as_deref()) {
652 session.config_toml = Some(merged);
653 tracing::debug!("Merged per-user config into session");
654 }
655 }
656
657 if let Some(ref mut config_toml) = session.config_toml {
659 substitute_env_vars(config_toml, &merged_secrets);
660 }
661
662 session.secrets = Some(shared_secrets);
664 }
665
666 fn load_user_secrets(
669 &self,
670 user_home: &std::path::Path,
671 base: &std::collections::HashMap<String, String>,
672 ) -> std::io::Result<std::collections::HashMap<String, String>> {
673 let user_env_path = user_home.join(".env");
674 if !user_env_path.exists() {
675 return Ok(base.clone());
676 }
677 let content = std::fs::read_to_string(&user_env_path)?;
678 let mut merged = base.clone();
679 for line in content.lines() {
680 let line = line.trim();
681 if line.is_empty() || line.starts_with('#') {
682 continue;
683 }
684 if let Some((key, value)) = line.split_once('=') {
685 let key = key.trim().to_string();
686 let value = value
687 .trim()
688 .trim_matches('"')
689 .trim_matches('\'')
690 .to_string();
691 merged.insert(key, value);
692 }
693 }
694 tracing::debug!("Loaded per-user secrets from {}", user_env_path.display());
695 Ok(merged)
696 }
697
698 fn merge_user_config(
713 &self,
714 user_home: &std::path::Path,
715 shared_config: Option<&str>,
716 ) -> Option<String> {
717 let user_config_path = user_home.join("config").join("trustee.toml");
718 if !user_config_path.exists() {
719 return None;
720 }
721 let user_config_toml = std::fs::read_to_string(&user_config_path).ok()?;
722 let shared = shared_config
723 .unwrap_or("")
724 .parse::<toml::Value>()
725 .ok()?;
726 let overlay = user_config_toml.parse::<toml::Value>().ok()?;
727
728 let allowed = |section: &str| {
729 section == "mcp"
730 || section == "thq"
734 || (self.allow_llm_overlay && section == "llm")
735 };
736 let dir_name = user_home
739 .file_name()
740 .and_then(|n| n.to_str())
741 .unwrap_or("<unknown>");
742 let masked_user = dir_name.get(..8).unwrap_or(dir_name);
743
744 let mut filtered_overlay = toml::map::Map::new();
745 if let Some(table) = overlay.as_table() {
746 for (section, value) in table {
747 if allowed(section) {
748 filtered_overlay.insert(section.clone(), value.clone());
749 } else {
750 tracing::warn!(
751 "user config overlay: dropping non-allowlisted section [{}] for user {}",
752 section,
753 masked_user
754 );
755 }
756 }
757 }
758
759 if filtered_overlay.is_empty() {
760 return None;
762 }
763 let overlay = toml::Value::Table(filtered_overlay);
764
765 let mut shared = shared;
766 deep_merge_toml(&mut shared, &overlay);
767 let merged = toml::to_string(&shared).ok()?;
768 tracing::debug!("Merged per-user config from {}", user_config_path.display());
769 Some(merged)
770 }
771
772 pub fn resolve_user_config(&self, user_key: &str) -> Option<String> {
781 let config_toml = self.config_toml.clone()?;
782
783 let user_home = self.get_user_home_dir(user_key)?;
785
786 let mut merged_secrets = self.secrets.clone().unwrap_or_default();
788 if let Ok(merged) = self.load_user_secrets(&user_home, &merged_secrets) {
789 merged_secrets = merged;
790 }
791
792 let mut resolved = config_toml;
794 if let Some(merged) = self.merge_user_config(&user_home, Some(&resolved)) {
795 resolved = merged;
796 }
797
798 substitute_env_vars(&mut resolved, &merged_secrets);
800
801 Some(resolved)
802 }
803
804 pub async fn get_or_build_mcp_loader(
820 &self,
821 user_key: &str,
822 token_store: &Arc<pep::MemoryTokenStore>,
823 ) -> Result<Option<std::sync::Arc<abk::agent::McpToolLoader>>, String> {
824 let user_hash = trustee_core::user_hash(user_key);
825
826 let resolved = self.resolve_user_config(user_key);
828 let fingerprint = fingerprint_mcp_section(resolved.as_deref());
829
830 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
832 if entry.degraded.is_none() {
833 if entry.fingerprint == fingerprint {
834 return Ok(entry.loader.clone());
835 }
836 } else if let (Some(err), Some(failed_at)) = (&entry.degraded, entry.failed_at) {
837 let backoff = chrono::Duration::from_std(MCP_BUILD_RETRY_BACKOFF)
838 .unwrap_or_else(|_| chrono::Duration::seconds(30));
839 if chrono::Utc::now() < failed_at + backoff {
840 return Err(err.clone());
841 }
842 }
843 }
844
845 let lock = self
847 .mcp_build_locks
848 .entry(user_hash.clone())
849 .or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
850 .clone();
851 let _guard = lock.lock().await;
852
853 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
855 if entry.degraded.is_none() && entry.fingerprint == fingerprint {
856 return Ok(entry.loader.clone());
857 }
858 }
859
860 match self
861 .build_mcp_loader(&user_hash, resolved.as_deref(), fingerprint, token_store)
862 .await
863 {
864 Ok(entry) => {
865 self.mcp_loaders.insert(user_hash.clone(), entry);
866 Ok(self.mcp_loaders.get(&user_hash).unwrap().loader.clone())
867 }
868 Err(err) => {
869 tracing::warn!(
870 "MCP loader build FAILED for user {}; dispatch fails loud, retry after {:?}",
871 &user_hash[..8.min(user_hash.len())],
872 MCP_BUILD_RETRY_BACKOFF
873 );
874 self.mcp_loaders.insert(
875 user_hash,
876 McpLoaderEntry {
877 loader: None,
878 fingerprint,
879 built_at: chrono::Utc::now(),
880 degraded: Some(err.clone()),
881 failed_at: Some(chrono::Utc::now()),
882 },
883 );
884 Err(err)
885 }
886 }
887 }
888
889 async fn build_mcp_loader(
892 &self,
893 user_hash: &str,
894 resolved: Option<&str>,
895 fingerprint: u64,
896 token_store: &Arc<pep::MemoryTokenStore>,
897 ) -> Result<McpLoaderEntry, String> {
898 let mcp_config: Option<abk::config::McpConfig> = match resolved {
899 Some(toml_str) => {
900 let value = toml_str
901 .parse::<toml::Value>()
902 .map_err(|e| format!("config parse failed: {}", e))?;
903 match value.get("mcp") {
904 Some(section) => {
905 use serde::Deserialize as _;
906 Some(
907 abk::config::McpConfig::deserialize(section.clone())
908 .map_err(|e| format!("invalid [mcp] config: {}", e))?,
909 )
910 }
911 None => None,
912 }
913 }
914 None => None,
915 };
916
917 let loader = match mcp_config {
918 Some(cfg) if cfg.enabled => {
919 let built = abk::agent::McpToolLoader::with_token_store(
920 &cfg,
921 Some(token_store.clone() as std::sync::Arc<dyn pep::token_store::TokenStore>),
922 )
923 .await
924 .map_err(|e| format!("MCP loader build failed: {}", e))?;
925
926 let servers: Vec<String> = built
930 .server_statuses
931 .iter()
932 .map(|s| {
933 if s.connected {
934 format!("{}(up,{}tools)", s.name, s.tool_count)
935 } else {
936 format!("{}(DOWN)", s.name)
937 }
938 })
939 .collect();
940 tracing::info!(
941 "MCP loader built for user {}: servers=[{}] total_tools={}",
942 &user_hash[..8.min(user_hash.len())],
943 servers.join(", "),
944 built.tool_count
945 );
946 Some(std::sync::Arc::new(built))
947 }
948 _ => None,
949 };
950
951 Ok(McpLoaderEntry {
952 loader,
953 fingerprint,
954 built_at: chrono::Utc::now(),
955 degraded: None,
956 failed_at: None,
957 })
958 }
959
960 fn spawn_user_drain_task(
962 &self,
963 session_id: String,
964 session: Arc<Mutex<Session>>,
965 ws_tx: broadcast::Sender<String>,
966 mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
967 ) {
968 let mut last_broadcast_state: Option<String> = None;
974 tokio::spawn(async move {
975 while let Some(msg) = workflow_rx.recv().await {
976 {
977 let mut session = session.lock().await;
978 session.handle_workflow_message(msg.clone());
979
980 let state_str = match session.workflow_state {
981 trustee_core::types::WorkflowState::Idle => "Idle",
982 trustee_core::types::WorkflowState::Running => "Running",
983 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
984 };
985 if last_broadcast_state.as_deref() != Some(state_str) {
986 last_broadcast_state = Some(state_str.to_string());
987 let state_msg = serde_json::json!({
988 "type": "StateChanged",
989 "state": state_str
990 });
991 let _ = ws_tx.send(state_msg.to_string());
992 }
993 }
994
995 let json =
996 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
997 let _ = ws_tx.send(json);
998 }
999 tracing::debug!("Drain task ended for session: {}", session_id);
1000 });
1001 }
1002
1003 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
1006 let default_user = self
1008 .sessions
1009 .get("default")
1010 .expect("default user must exist");
1011 let first_entry = default_user
1012 .sessions
1013 .iter()
1014 .next()
1015 .expect("default user must have at least one session");
1016 let session = first_entry.session.clone();
1017 let ws_tx = first_entry.ws_tx.clone();
1018 let session_id = first_entry.key().clone();
1019 drop(first_entry);
1020 drop(default_user);
1021
1022 let mut last_broadcast_state = Some("Running".to_string());
1026 tokio::spawn(async move {
1027 while let Some(msg) = workflow_rx.recv().await {
1028 {
1029 let mut session = session.lock().await;
1030 session.handle_workflow_message(msg.clone());
1031
1032 let state_str = match session.workflow_state {
1033 trustee_core::types::WorkflowState::Idle => "Idle",
1034 trustee_core::types::WorkflowState::Running => "Running",
1035 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
1036 };
1037 if last_broadcast_state.as_deref() != Some(state_str) {
1038 last_broadcast_state = Some(state_str.to_string());
1039 let state_msg = serde_json::json!({
1040 "type": "StateChanged",
1041 "state": state_str
1042 });
1043 let _ = ws_tx.send(state_msg.to_string());
1044 }
1045 }
1046
1047 let json =
1048 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
1049 let _ = ws_tx.send(json);
1050 }
1051 tracing::debug!("Drain task ended for session: {}", session_id);
1052 });
1053 }
1054
1055 pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
1057 let Some(ref auth) = self.auth else {
1058 return "default".to_string();
1059 };
1060
1061 if let Some(token) = headers
1063 .get(axum::http::header::AUTHORIZATION)
1064 .and_then(|v| v.to_str().ok())
1065 .and_then(|v| v.strip_prefix("Bearer "))
1066 .map(|s| s.to_string())
1067 {
1068 if token.starts_with("dev:") {
1069 let parts: Vec<&str> = token.splitn(4, ':').collect();
1070 if parts.len() >= 4 {
1071 return format!("dev:{}", parts[1]);
1072 }
1073 }
1074 if let Ok(claims) = auth.validate_token(&token).await {
1075 return claims.sub;
1076 }
1077 }
1078
1079 let cookie_session_id = headers
1081 .get(axum::http::header::COOKIE)
1082 .and_then(|v| v.to_str().ok())
1083 .and_then(|cookies| {
1084 cookies
1085 .split(';')
1086 .map(|c| c.trim())
1087 .find_map(|c| {
1088 c.strip_prefix(&format!("{}=", auth.config.cookie_name))
1089 .map(|s| s.to_string())
1090 })
1091 });
1092
1093 if let Some(session_id) = cookie_session_id {
1094 if session_id.starts_with("dev:") {
1095 let parts: Vec<&str> = session_id.splitn(4, ':').collect();
1096 if parts.len() >= 4 {
1097 return format!("dev:{}", parts[1]);
1098 }
1099 }
1100
1101 if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
1102 if let Ok(claims) = auth.validate_token(&access_token).await {
1103 return claims.sub;
1104 }
1105 }
1106 }
1107
1108 "default".to_string()
1109 }
1110}
1111
1112struct SerializableMessage<'a>(&'a TuiMessage);
1118
1119impl<'a> serde::Serialize for SerializableMessage<'a> {
1120 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1121 where
1122 S: serde::Serializer,
1123 {
1124 use serde::ser::SerializeStruct;
1125
1126 match self.0 {
1127 TuiMessage::OutputLine(line) => {
1128 let mut s = serializer.serialize_struct("msg", 2)?;
1129 s.serialize_field("type", "OutputLine")?;
1130 s.serialize_field("line", line)?;
1131 s.end()
1132 }
1133 TuiMessage::StreamDelta(delta) => {
1134 let mut s = serializer.serialize_struct("msg", 2)?;
1135 s.serialize_field("type", "StreamDelta")?;
1136 s.serialize_field("delta", delta)?;
1137 s.end()
1138 }
1139 TuiMessage::ReasoningDelta(delta) => {
1140 let mut s = serializer.serialize_struct("msg", 2)?;
1141 s.serialize_field("type", "ReasoningDelta")?;
1142 s.serialize_field("delta", delta)?;
1143 s.end()
1144 }
1145 TuiMessage::WorkflowCompleted => {
1146 let mut s = serializer.serialize_struct("msg", 2)?;
1147 s.serialize_field("type", "WorkflowCompleted")?;
1148 s.serialize_field("state", "Idle")?;
1149 s.end()
1150 }
1151 TuiMessage::WorkflowError(err) => {
1152 let mut s = serializer.serialize_struct("msg", 2)?;
1153 s.serialize_field("type", "WorkflowError")?;
1154 s.serialize_field("error", err)?;
1155 s.end()
1156 }
1157 TuiMessage::ResumeInfo(info) => match info {
1158 Some(ri) => {
1159 let mut s = serializer.serialize_struct("msg", 5)?;
1160 s.serialize_field("type", "ResumeInfo")?;
1161 s.serialize_field("state", "Idle")?;
1162 s.serialize_field("session_id", &ri.session_id)?;
1163 s.serialize_field("checkpoint_id", &ri.checkpoint_id)?;
1164 s.serialize_field("iteration", &ri.iteration)?;
1165 s.end()
1166 }
1167 None => {
1168 let mut s = serializer.serialize_struct("msg", 2)?;
1169 s.serialize_field("type", "ResumeInfo")?;
1170 s.serialize_field("state", "Idle")?;
1171 s.end()
1172 }
1173 },
1174 TuiMessage::TodoUpdate(content) => {
1175 let mut s = serializer.serialize_struct("msg", 2)?;
1176 s.serialize_field("type", "TodoUpdate")?;
1177 s.serialize_field("content", content)?;
1178 s.end()
1179 }
1180 TuiMessage::WorkflowCancelled => {
1181 let mut s = serializer.serialize_struct("msg", 2)?;
1182 s.serialize_field("type", "WorkflowCancelled")?;
1183 s.serialize_field("state", "Idle")?;
1184 s.end()
1185 }
1186 TuiMessage::HandoffReady(briefing) => {
1187 let mut s = serializer.serialize_struct("msg", 3)?;
1188 s.serialize_field("type", "HandoffReady")?;
1189 s.serialize_field("state", "Idle")?;
1190 s.serialize_field("briefing", briefing)?;
1191 s.end()
1192 }
1193 TuiMessage::SessionRotated { old, new } => {
1194 let mut s = serializer.serialize_struct("msg", 3)?;
1195 s.serialize_field("type", "SessionRotated")?;
1196 s.serialize_field("old", old)?;
1197 s.serialize_field("new", new)?;
1198 s.end()
1199 }
1200 TuiMessage::HandoffFailed => {
1201 let mut s = serializer.serialize_struct("msg", 2)?;
1202 s.serialize_field("type", "HandoffFailed")?;
1203 s.serialize_field("state", "Idle")?;
1204 s.end()
1205 }
1206 TuiMessage::ToolPending {
1207 tool_name,
1208 hint,
1209 } => {
1210 let mut s = serializer.serialize_struct("msg", 3)?;
1211 s.serialize_field("type", "ToolPending")?;
1212 s.serialize_field("tool_name", tool_name)?;
1213 s.serialize_field("hint", hint)?;
1214 s.end()
1215 }
1216 TuiMessage::ToolDone {
1217 tool_name,
1218 success,
1219 hint,
1220 } => {
1221 let mut s = serializer.serialize_struct("msg", 4)?;
1222 s.serialize_field("type", "ToolDone")?;
1223 s.serialize_field("tool_name", tool_name)?;
1224 s.serialize_field("success", success)?;
1225 s.serialize_field("hint", hint)?;
1226 s.end()
1227 }
1228 TuiMessage::ContextTokensUpdated(count) => {
1229 let mut s = serializer.serialize_struct("msg", 2)?;
1230 s.serialize_field("type", "ContextTokensUpdated")?;
1231 s.serialize_field("count", count)?;
1232 s.end()
1233 }
1234 TuiMessage::McpServerStatus {
1235 name,
1236 connected,
1237 tool_count,
1238 error,
1239 } => {
1240 let mut s = serializer.serialize_struct("msg", 5)?;
1241 s.serialize_field("type", "McpServerStatus")?;
1242 s.serialize_field("name", name)?;
1243 s.serialize_field("connected", connected)?;
1244 s.serialize_field("tool_count", tool_count)?;
1245 s.serialize_field("error", error)?;
1246 s.end()
1247 }
1248 TuiMessage::SessionTitleUpdated(title) => {
1249 let mut s = serializer.serialize_struct("msg", 2)?;
1250 s.serialize_field("type", "SessionTitleUpdated")?;
1251 s.serialize_field("title", title)?;
1252 s.end()
1253 }
1254 }
1255 }
1256}
1257
1258fn deep_merge_toml(base: &mut toml::Value, overlay: &toml::Value) {
1269 match (base, overlay) {
1270 (toml::Value::Table(base_table), toml::Value::Table(overlay_table)) => {
1271 for (key, overlay_val) in overlay_table {
1272 match base_table.get_mut(key) {
1273 Some(base_val) => {
1274 deep_merge_toml(base_val, overlay_val);
1276 }
1277 None => {
1278 base_table.insert(key.clone(), overlay_val.clone());
1280 }
1281 }
1282 }
1283 }
1284 (base, overlay) => {
1286 *base = overlay.clone();
1287 }
1288 }
1289}
1290
1291fn substitute_env_vars(s: &mut String, secrets: &std::collections::HashMap<String, String>) {
1296 let mut result = String::with_capacity(s.len());
1298 let bytes = s.as_bytes();
1299 let mut i = 0;
1300
1301 while i < bytes.len() {
1302 if i + 1 < bytes.len() && bytes[i] == b'$' && bytes[i + 1] == b'{' {
1303 if let Some(end) = s[i + 2..].find('}') {
1305 let var_name = &s[i + 2..i + 2 + end];
1306 if let Some(value) = secrets.get(var_name) {
1308 result.push_str(value);
1309 } else if let Ok(value) = std::env::var(var_name) {
1310 result.push_str(&value);
1311 } else {
1312 result.push_str(&s[i..i + 2 + end + 1]);
1314 }
1315 i = i + 2 + end + 1;
1316 } else {
1317 result.push('$');
1319 i += 1;
1320 }
1321 } else {
1322 result.push(bytes[i] as char);
1323 i += 1;
1324 }
1325 }
1326
1327 *s = result;
1328}
1329
1330fn fingerprint_mcp_section(resolved: Option<&str>) -> u64 {
1340 use sha2::{Digest, Sha256};
1341 let section = resolved
1342 .and_then(|s| s.parse::<toml::Value>().ok())
1343 .and_then(|v| v.get("mcp").cloned());
1344 let bytes = match section {
1345 Some(v) => v.to_string().into_bytes(),
1346 None => b"<no-mcp>".to_vec(),
1347 };
1348 let digest = Sha256::digest(&bytes);
1349 u64::from_be_bytes(digest[..8].try_into().expect("sha256 digest >= 8 bytes"))
1350}
1351
1352#[cfg(test)]
1353mod tests {
1354 use super::*;
1355
1356 fn test_state() -> ServerState {
1358 let (session, _rx) = Session::new();
1359 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1360 ServerState::new(session, ws_tx, None)
1361 }
1362
1363 fn temp_user_home(tag: &str) -> std::path::PathBuf {
1365 static COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1366 let n = COUNTER.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1367 let dir = std::env::temp_dir().join(format!(
1368 "trustee-state-test-{}-{}-{}",
1369 tag,
1370 std::process::id(),
1371 n
1372 ));
1373 std::fs::create_dir_all(dir.join("config")).expect("create temp user home");
1374 dir
1375 }
1376
1377 fn parse(toml_str: &str) -> toml::Value {
1378 toml_str.parse::<toml::Value>().expect("valid test TOML")
1379 }
1380
1381 #[test]
1384 fn overlay_allowlist_drops_non_allowlisted_sections() {
1385 let state = test_state();
1386 let home = temp_user_home("allowlist");
1387 std::fs::write(
1388 home.join("config").join("trustee.toml"),
1389 "[server]\nport = 1\n\n[auth]\nmode = \"kanidm\"\n\n[mcp]\nmode = \"user\"\n",
1390 )
1391 .expect("write overlay");
1392
1393 let shared = "[server]\nport = 8080\n\n[mcp]\nmode = \"shared\"\n";
1394 let merged = state
1395 .merge_user_config(&home, Some(shared))
1396 .expect("overlay has allowlisted content");
1397
1398 let merged_val = parse(&merged);
1399 let expected = parse("[server]\nport = 8080\n\n[mcp]\nmode = \"user\"\n");
1400 assert_eq!(merged_val, expected, "merged config must be shared + [mcp] overlay only");
1401 assert!(merged_val.get("auth").is_none(), "overlay [auth] must be dropped");
1402 }
1403
1404 #[test]
1406 fn no_overlay_file_returns_none() {
1407 let state = test_state();
1408 let home = temp_user_home("empty");
1409 assert!(state
1410 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1411 .is_none());
1412 }
1413
1414 #[test]
1416 fn overlay_with_no_allowlisted_sections_is_noop() {
1417 let state = test_state();
1418 let home = temp_user_home("all-dropped");
1419 std::fs::write(
1420 home.join("config").join("trustee.toml"),
1421 "[server]\nport = 1\n\n[storage]\npath = \"/tmp/x\"\n",
1422 )
1423 .expect("write overlay");
1424 assert!(state
1425 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1426 .is_none());
1427 }
1428
1429 #[test]
1432 fn llm_overlay_dropped_by_default_and_kept_when_enabled() {
1433 let shared = "[llm]\nprovider = \"openai\"\n\n[mcp]\nmode = \"shared\"\n";
1434 let overlay = "[llm]\nprovider = \"anthropic\"\n";
1435
1436 let state = test_state();
1438 assert!(!state.allow_llm_overlay);
1439 let home = temp_user_home("llm-off");
1440 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1441 assert!(state.merge_user_config(&home, Some(shared)).is_none());
1442
1443 let state = test_state().with_allow_llm_overlay(true);
1445 assert!(state.allow_llm_overlay);
1446 let home = temp_user_home("llm-on");
1447 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1448 let merged = state
1449 .merge_user_config(&home, Some(shared))
1450 .expect("[llm] overlay applies when opted in");
1451 let merged_val = parse(&merged);
1452 assert_eq!(merged_val["llm"]["provider"].as_str(), Some("anthropic"));
1453 assert_eq!(merged_val["mcp"]["mode"].as_str(), Some("shared"));
1454 }
1455
1456 #[test]
1460 fn user_home_dir_uses_consolidated_hash() {
1461 let state = test_state();
1462 if let Some(home) = state.get_user_home_dir("farzan@example.com") {
1463 assert_eq!(
1464 home.file_name().and_then(|n| n.to_str()),
1465 Some(trustee_core::user_hash("farzan@example.com")).as_deref()
1466 );
1467 let users_root = dirs::home_dir().unwrap().join(".trustee").join("users");
1468 assert_eq!(home.parent(), Some(&users_root).map(|p| p.as_path()));
1469 }
1470 }
1472
1473 #[tokio::test]
1478 async fn agent_principals_get_isolated_session_buckets() {
1479 let state = test_state();
1480 let key_a = "agent-farzan";
1481 let key_b = "agent-paydar";
1482
1483 let (sid_a, session_a, _tx_a, _ts_a) = state.ensure_active_session(key_a).await;
1484 let (_sid_b, _session_b, _tx_b, _ts_b) = state.ensure_active_session(key_b).await;
1485
1486 assert!(
1488 state.get_session_by_any_id(key_a, &sid_a).await.is_some(),
1489 "owner bucket resolves its own session"
1490 );
1491 assert!(
1492 state.get_session_by_any_id(key_b, &sid_a).await.is_none(),
1493 "cross-agent session access must be 404/None"
1494 );
1495 assert!(Arc::strong_count(&session_a) >= 1);
1496
1497 if let Some(home) = state.get_user_home_dir(key_a) {
1499 assert_eq!(
1500 home.file_name().and_then(|n| n.to_str()),
1501 Some(trustee_core::user_hash(key_a)).as_deref()
1502 );
1503 }
1504 }
1505}
1506
1507#[cfg(test)]
1512mod mcp_loader_cache_tests {
1513 use super::*;
1514
1515 fn state_with_shared(shared: &str) -> ServerState {
1516 let (session, _rx) = Session::new();
1517 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1518 let mut state = ServerState::new(session, ws_tx, None);
1519 state.config_toml = Some(shared.to_string());
1520 state
1521 }
1522
1523 struct TempUser {
1527 key: String,
1528 home: std::path::PathBuf,
1529 }
1530
1531 impl TempUser {
1532 fn new(tag: &str) -> Self {
1533 let key = format!("16c-{tag}-{}@test.invalid", std::process::id());
1534 let home = dirs::home_dir()
1535 .expect("HOME available in test env")
1536 .join(".trustee")
1537 .join("users")
1538 .join(trustee_core::user_hash(&key));
1539 std::fs::create_dir_all(home.join("config")).expect("create user home");
1540 Self { key, home }
1541 }
1542
1543 fn write_overlay(&self, toml_str: &str) {
1544 std::fs::write(self.home.join("config").join("trustee.toml"), toml_str)
1545 .expect("write overlay");
1546 }
1547 }
1548
1549 impl Drop for TempUser {
1550 fn drop(&mut self) {
1551 let _ = std::fs::remove_dir_all(&self.home);
1552 }
1553 }
1554
1555 #[tokio::test]
1556 async fn no_mcp_config_caches_disabled_marker() {
1557 let state = state_with_shared("[server]\nport = 8080\n");
1558 let user = TempUser::new("nomcp");
1559 let ts = Arc::new(pep::MemoryTokenStore::new());
1560
1561 let first = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1562 assert!(first.is_none(), "no [mcp] anywhere → disabled marker");
1563 assert_eq!(state.mcp_loaders.len(), 1, "exactly one cache entry");
1564
1565 let second = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1567 assert!(second.is_none());
1568 assert_eq!(state.mcp_loaders.len(), 1);
1569 }
1570
1571 #[tokio::test]
1572 async fn concurrent_cold_builds_single_flight() {
1573 let state = Arc::new(state_with_shared("[server]\nport = 8080\n"));
1574 let user = TempUser::new("singleflight");
1575 let ts = Arc::new(pep::MemoryTokenStore::new());
1576
1577 let mut handles = Vec::new();
1578 for _ in 0..5 {
1579 let state = state.clone();
1580 let key = user.key.clone();
1581 let ts = ts.clone();
1582 handles.push(tokio::spawn(async move {
1583 state.get_or_build_mcp_loader(&key, &ts).await
1584 }));
1585 }
1586 for h in handles {
1587 h.await.unwrap().expect("all five succeed");
1588 }
1589 assert_eq!(state.mcp_loaders.len(), 1, "single-flight → one entry");
1590 }
1591
1592 #[tokio::test]
1593 async fn fingerprint_change_triggers_rebuild() {
1594 let state = state_with_shared("[server]\nport = 8080\n");
1595 let user = TempUser::new("fpchange");
1596 let ts = Arc::new(pep::MemoryTokenStore::new());
1597
1598 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v1\"\nurl = \"http://127.0.0.1:9/sse\"\n");
1601 let v1 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1602 assert!(v1.is_some(), "enabled [mcp] → real loader");
1603 let fp1 = state
1604 .mcp_loaders
1605 .get(&trustee_core::user_hash(&user.key))
1606 .unwrap()
1607 .fingerprint;
1608
1609 std::thread::sleep(std::time::Duration::from_millis(5));
1611 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v2\"\nurl = \"http://127.0.0.1:9/other\"\n");
1612 let v2 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1613 assert!(v2.is_some());
1614 let entry = state
1615 .mcp_loaders
1616 .get(&trustee_core::user_hash(&user.key))
1617 .unwrap();
1618 assert_ne!(
1619 entry.fingerprint, fp1,
1620 "fingerprint must change with content"
1621 );
1622 assert!(entry.degraded.is_none());
1623
1624 let _still_usable = v1.as_ref().unwrap().tool_count;
1626 }
1627
1628 #[tokio::test]
1629 async fn degraded_entry_fails_loud_within_backoff_and_isolates_users() {
1630 let state = state_with_shared("[server]\nport = 8080\n");
1631 let bad = TempUser::new("degraded-bad");
1632 let good = TempUser::new("degraded-good");
1633 let ts = Arc::new(pep::MemoryTokenStore::new());
1634
1635 bad.write_overlay("[mcp]\nenabled = \"not-a-bool\"\n");
1637 let err = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1638 Ok(_) => panic!("invalid [mcp] must fail loud"),
1639 Err(e) => e,
1640 };
1641 assert!(
1642 err.contains("invalid [mcp]"),
1643 "surfaces the parse error: {err}"
1644 );
1645
1646 let entry = state
1647 .mcp_loaders
1648 .get(&trustee_core::user_hash(&bad.key))
1649 .unwrap();
1650 assert!(entry.degraded.is_some(), "poison entry recorded");
1651
1652 let err2 = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1654 Ok(_) => panic!("still within backoff"),
1655 Err(e) => e,
1656 };
1657 assert_eq!(err, err2, "same cached error");
1658
1659 let good_loader = state
1661 .get_or_build_mcp_loader(&good.key, &ts)
1662 .await
1663 .expect("other user unaffected");
1664 assert!(
1665 good_loader.is_none(),
1666 "good user has no [mcp] → disabled marker"
1667 );
1668 }
1669}