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 =
729 |section: &str| section == "mcp" || (self.allow_llm_overlay && section == "llm");
730 let dir_name = user_home
733 .file_name()
734 .and_then(|n| n.to_str())
735 .unwrap_or("<unknown>");
736 let masked_user = dir_name.get(..8).unwrap_or(dir_name);
737
738 let mut filtered_overlay = toml::map::Map::new();
739 if let Some(table) = overlay.as_table() {
740 for (section, value) in table {
741 if allowed(section) {
742 filtered_overlay.insert(section.clone(), value.clone());
743 } else {
744 tracing::warn!(
745 "user config overlay: dropping non-allowlisted section [{}] for user {}",
746 section,
747 masked_user
748 );
749 }
750 }
751 }
752
753 if filtered_overlay.is_empty() {
754 return None;
756 }
757 let overlay = toml::Value::Table(filtered_overlay);
758
759 let mut shared = shared;
760 deep_merge_toml(&mut shared, &overlay);
761 let merged = toml::to_string(&shared).ok()?;
762 tracing::debug!("Merged per-user config from {}", user_config_path.display());
763 Some(merged)
764 }
765
766 pub fn resolve_user_config(&self, user_key: &str) -> Option<String> {
775 let config_toml = self.config_toml.clone()?;
776
777 let user_home = self.get_user_home_dir(user_key)?;
779
780 let mut merged_secrets = self.secrets.clone().unwrap_or_default();
782 if let Ok(merged) = self.load_user_secrets(&user_home, &merged_secrets) {
783 merged_secrets = merged;
784 }
785
786 let mut resolved = config_toml;
788 if let Some(merged) = self.merge_user_config(&user_home, Some(&resolved)) {
789 resolved = merged;
790 }
791
792 substitute_env_vars(&mut resolved, &merged_secrets);
794
795 Some(resolved)
796 }
797
798 pub async fn get_or_build_mcp_loader(
814 &self,
815 user_key: &str,
816 token_store: &Arc<pep::MemoryTokenStore>,
817 ) -> Result<Option<std::sync::Arc<abk::agent::McpToolLoader>>, String> {
818 let user_hash = trustee_core::user_hash(user_key);
819
820 let resolved = self.resolve_user_config(user_key);
822 let fingerprint = fingerprint_mcp_section(resolved.as_deref());
823
824 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
826 if entry.degraded.is_none() {
827 if entry.fingerprint == fingerprint {
828 return Ok(entry.loader.clone());
829 }
830 } else if let (Some(err), Some(failed_at)) = (&entry.degraded, entry.failed_at) {
831 let backoff = chrono::Duration::from_std(MCP_BUILD_RETRY_BACKOFF)
832 .unwrap_or_else(|_| chrono::Duration::seconds(30));
833 if chrono::Utc::now() < failed_at + backoff {
834 return Err(err.clone());
835 }
836 }
837 }
838
839 let lock = self
841 .mcp_build_locks
842 .entry(user_hash.clone())
843 .or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
844 .clone();
845 let _guard = lock.lock().await;
846
847 if let Some(entry) = self.mcp_loaders.get(&user_hash) {
849 if entry.degraded.is_none() && entry.fingerprint == fingerprint {
850 return Ok(entry.loader.clone());
851 }
852 }
853
854 match self
855 .build_mcp_loader(&user_hash, resolved.as_deref(), fingerprint, token_store)
856 .await
857 {
858 Ok(entry) => {
859 self.mcp_loaders.insert(user_hash.clone(), entry);
860 Ok(self.mcp_loaders.get(&user_hash).unwrap().loader.clone())
861 }
862 Err(err) => {
863 tracing::warn!(
864 "MCP loader build FAILED for user {}; dispatch fails loud, retry after {:?}",
865 &user_hash[..8.min(user_hash.len())],
866 MCP_BUILD_RETRY_BACKOFF
867 );
868 self.mcp_loaders.insert(
869 user_hash,
870 McpLoaderEntry {
871 loader: None,
872 fingerprint,
873 built_at: chrono::Utc::now(),
874 degraded: Some(err.clone()),
875 failed_at: Some(chrono::Utc::now()),
876 },
877 );
878 Err(err)
879 }
880 }
881 }
882
883 async fn build_mcp_loader(
886 &self,
887 user_hash: &str,
888 resolved: Option<&str>,
889 fingerprint: u64,
890 token_store: &Arc<pep::MemoryTokenStore>,
891 ) -> Result<McpLoaderEntry, String> {
892 let mcp_config: Option<abk::config::McpConfig> = match resolved {
893 Some(toml_str) => {
894 let value = toml_str
895 .parse::<toml::Value>()
896 .map_err(|e| format!("config parse failed: {}", e))?;
897 match value.get("mcp") {
898 Some(section) => {
899 use serde::Deserialize as _;
900 Some(
901 abk::config::McpConfig::deserialize(section.clone())
902 .map_err(|e| format!("invalid [mcp] config: {}", e))?,
903 )
904 }
905 None => None,
906 }
907 }
908 None => None,
909 };
910
911 let loader = match mcp_config {
912 Some(cfg) if cfg.enabled => {
913 let built = abk::agent::McpToolLoader::with_token_store(
914 &cfg,
915 Some(token_store.clone() as std::sync::Arc<dyn pep::token_store::TokenStore>),
916 )
917 .await
918 .map_err(|e| format!("MCP loader build failed: {}", e))?;
919
920 let servers: Vec<String> = built
924 .server_statuses
925 .iter()
926 .map(|s| {
927 if s.connected {
928 format!("{}(up,{}tools)", s.name, s.tool_count)
929 } else {
930 format!("{}(DOWN)", s.name)
931 }
932 })
933 .collect();
934 tracing::info!(
935 "MCP loader built for user {}: servers=[{}] total_tools={}",
936 &user_hash[..8.min(user_hash.len())],
937 servers.join(", "),
938 built.tool_count
939 );
940 Some(std::sync::Arc::new(built))
941 }
942 _ => None,
943 };
944
945 Ok(McpLoaderEntry {
946 loader,
947 fingerprint,
948 built_at: chrono::Utc::now(),
949 degraded: None,
950 failed_at: None,
951 })
952 }
953
954 fn spawn_user_drain_task(
956 &self,
957 session_id: String,
958 session: Arc<Mutex<Session>>,
959 ws_tx: broadcast::Sender<String>,
960 mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>,
961 ) {
962 let mut last_broadcast_state: Option<String> = None;
968 tokio::spawn(async move {
969 while let Some(msg) = workflow_rx.recv().await {
970 {
971 let mut session = session.lock().await;
972 session.handle_workflow_message(msg.clone());
973
974 let state_str = match session.workflow_state {
975 trustee_core::types::WorkflowState::Idle => "Idle",
976 trustee_core::types::WorkflowState::Running => "Running",
977 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
978 };
979 if last_broadcast_state.as_deref() != Some(state_str) {
980 last_broadcast_state = Some(state_str.to_string());
981 let state_msg = serde_json::json!({
982 "type": "StateChanged",
983 "state": state_str
984 });
985 let _ = ws_tx.send(state_msg.to_string());
986 }
987 }
988
989 let json =
990 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
991 let _ = ws_tx.send(json);
992 }
993 tracing::debug!("Drain task ended for session: {}", session_id);
994 });
995 }
996
997 pub fn spawn_drain_task(self, mut workflow_rx: mpsc::UnboundedReceiver<TuiMessage>) {
1000 let default_user = self
1002 .sessions
1003 .get("default")
1004 .expect("default user must exist");
1005 let first_entry = default_user
1006 .sessions
1007 .iter()
1008 .next()
1009 .expect("default user must have at least one session");
1010 let session = first_entry.session.clone();
1011 let ws_tx = first_entry.ws_tx.clone();
1012 let session_id = first_entry.key().clone();
1013 drop(first_entry);
1014 drop(default_user);
1015
1016 let mut last_broadcast_state = Some("Running".to_string());
1020 tokio::spawn(async move {
1021 while let Some(msg) = workflow_rx.recv().await {
1022 {
1023 let mut session = session.lock().await;
1024 session.handle_workflow_message(msg.clone());
1025
1026 let state_str = match session.workflow_state {
1027 trustee_core::types::WorkflowState::Idle => "Idle",
1028 trustee_core::types::WorkflowState::Running => "Running",
1029 trustee_core::types::WorkflowState::Cancelling => "Cancelling",
1030 };
1031 if last_broadcast_state.as_deref() != Some(state_str) {
1032 last_broadcast_state = Some(state_str.to_string());
1033 let state_msg = serde_json::json!({
1034 "type": "StateChanged",
1035 "state": state_str
1036 });
1037 let _ = ws_tx.send(state_msg.to_string());
1038 }
1039 }
1040
1041 let json =
1042 serde_json::to_string(&SerializableMessage(&msg)).unwrap_or_default();
1043 let _ = ws_tx.send(json);
1044 }
1045 tracing::debug!("Drain task ended for session: {}", session_id);
1046 });
1047 }
1048
1049 pub async fn resolve_user_key(&self, headers: &axum::http::HeaderMap) -> String {
1051 let Some(ref auth) = self.auth else {
1052 return "default".to_string();
1053 };
1054
1055 if let Some(token) = headers
1057 .get(axum::http::header::AUTHORIZATION)
1058 .and_then(|v| v.to_str().ok())
1059 .and_then(|v| v.strip_prefix("Bearer "))
1060 .map(|s| s.to_string())
1061 {
1062 if token.starts_with("dev:") {
1063 let parts: Vec<&str> = token.splitn(4, ':').collect();
1064 if parts.len() >= 4 {
1065 return format!("dev:{}", parts[1]);
1066 }
1067 }
1068 if let Ok(claims) = auth.validate_token(&token).await {
1069 return claims.sub;
1070 }
1071 }
1072
1073 let cookie_session_id = headers
1075 .get(axum::http::header::COOKIE)
1076 .and_then(|v| v.to_str().ok())
1077 .and_then(|cookies| {
1078 cookies
1079 .split(';')
1080 .map(|c| c.trim())
1081 .find_map(|c| {
1082 c.strip_prefix(&format!("{}=", auth.config.cookie_name))
1083 .map(|s| s.to_string())
1084 })
1085 });
1086
1087 if let Some(session_id) = cookie_session_id {
1088 if session_id.starts_with("dev:") {
1089 let parts: Vec<&str> = session_id.splitn(4, ':').collect();
1090 if parts.len() >= 4 {
1091 return format!("dev:{}", parts[1]);
1092 }
1093 }
1094
1095 if let Ok(access_token) = auth.session_manager.get_token(&session_id).await {
1096 if let Ok(claims) = auth.validate_token(&access_token).await {
1097 return claims.sub;
1098 }
1099 }
1100 }
1101
1102 "default".to_string()
1103 }
1104}
1105
1106struct SerializableMessage<'a>(&'a TuiMessage);
1112
1113impl<'a> serde::Serialize for SerializableMessage<'a> {
1114 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
1115 where
1116 S: serde::Serializer,
1117 {
1118 use serde::ser::SerializeStruct;
1119
1120 match self.0 {
1121 TuiMessage::OutputLine(line) => {
1122 let mut s = serializer.serialize_struct("msg", 2)?;
1123 s.serialize_field("type", "OutputLine")?;
1124 s.serialize_field("line", line)?;
1125 s.end()
1126 }
1127 TuiMessage::StreamDelta(delta) => {
1128 let mut s = serializer.serialize_struct("msg", 2)?;
1129 s.serialize_field("type", "StreamDelta")?;
1130 s.serialize_field("delta", delta)?;
1131 s.end()
1132 }
1133 TuiMessage::ReasoningDelta(delta) => {
1134 let mut s = serializer.serialize_struct("msg", 2)?;
1135 s.serialize_field("type", "ReasoningDelta")?;
1136 s.serialize_field("delta", delta)?;
1137 s.end()
1138 }
1139 TuiMessage::WorkflowCompleted => {
1140 let mut s = serializer.serialize_struct("msg", 2)?;
1141 s.serialize_field("type", "WorkflowCompleted")?;
1142 s.serialize_field("state", "Idle")?;
1143 s.end()
1144 }
1145 TuiMessage::WorkflowError(err) => {
1146 let mut s = serializer.serialize_struct("msg", 2)?;
1147 s.serialize_field("type", "WorkflowError")?;
1148 s.serialize_field("error", err)?;
1149 s.end()
1150 }
1151 TuiMessage::ResumeInfo(info) => match info {
1152 Some(ri) => {
1153 let mut s = serializer.serialize_struct("msg", 5)?;
1154 s.serialize_field("type", "ResumeInfo")?;
1155 s.serialize_field("state", "Idle")?;
1156 s.serialize_field("session_id", &ri.session_id)?;
1157 s.serialize_field("checkpoint_id", &ri.checkpoint_id)?;
1158 s.serialize_field("iteration", &ri.iteration)?;
1159 s.end()
1160 }
1161 None => {
1162 let mut s = serializer.serialize_struct("msg", 2)?;
1163 s.serialize_field("type", "ResumeInfo")?;
1164 s.serialize_field("state", "Idle")?;
1165 s.end()
1166 }
1167 },
1168 TuiMessage::TodoUpdate(content) => {
1169 let mut s = serializer.serialize_struct("msg", 2)?;
1170 s.serialize_field("type", "TodoUpdate")?;
1171 s.serialize_field("content", content)?;
1172 s.end()
1173 }
1174 TuiMessage::WorkflowCancelled => {
1175 let mut s = serializer.serialize_struct("msg", 2)?;
1176 s.serialize_field("type", "WorkflowCancelled")?;
1177 s.serialize_field("state", "Idle")?;
1178 s.end()
1179 }
1180 TuiMessage::HandoffReady(briefing) => {
1181 let mut s = serializer.serialize_struct("msg", 3)?;
1182 s.serialize_field("type", "HandoffReady")?;
1183 s.serialize_field("state", "Idle")?;
1184 s.serialize_field("briefing", briefing)?;
1185 s.end()
1186 }
1187 TuiMessage::SessionRotated { old, new } => {
1188 let mut s = serializer.serialize_struct("msg", 3)?;
1189 s.serialize_field("type", "SessionRotated")?;
1190 s.serialize_field("old", old)?;
1191 s.serialize_field("new", new)?;
1192 s.end()
1193 }
1194 TuiMessage::HandoffFailed => {
1195 let mut s = serializer.serialize_struct("msg", 2)?;
1196 s.serialize_field("type", "HandoffFailed")?;
1197 s.serialize_field("state", "Idle")?;
1198 s.end()
1199 }
1200 TuiMessage::ToolPending {
1201 tool_name,
1202 hint,
1203 } => {
1204 let mut s = serializer.serialize_struct("msg", 3)?;
1205 s.serialize_field("type", "ToolPending")?;
1206 s.serialize_field("tool_name", tool_name)?;
1207 s.serialize_field("hint", hint)?;
1208 s.end()
1209 }
1210 TuiMessage::ToolDone {
1211 tool_name,
1212 success,
1213 hint,
1214 } => {
1215 let mut s = serializer.serialize_struct("msg", 4)?;
1216 s.serialize_field("type", "ToolDone")?;
1217 s.serialize_field("tool_name", tool_name)?;
1218 s.serialize_field("success", success)?;
1219 s.serialize_field("hint", hint)?;
1220 s.end()
1221 }
1222 TuiMessage::ContextTokensUpdated(count) => {
1223 let mut s = serializer.serialize_struct("msg", 2)?;
1224 s.serialize_field("type", "ContextTokensUpdated")?;
1225 s.serialize_field("count", count)?;
1226 s.end()
1227 }
1228 TuiMessage::McpServerStatus {
1229 name,
1230 connected,
1231 tool_count,
1232 error,
1233 } => {
1234 let mut s = serializer.serialize_struct("msg", 5)?;
1235 s.serialize_field("type", "McpServerStatus")?;
1236 s.serialize_field("name", name)?;
1237 s.serialize_field("connected", connected)?;
1238 s.serialize_field("tool_count", tool_count)?;
1239 s.serialize_field("error", error)?;
1240 s.end()
1241 }
1242 TuiMessage::SessionTitleUpdated(title) => {
1243 let mut s = serializer.serialize_struct("msg", 2)?;
1244 s.serialize_field("type", "SessionTitleUpdated")?;
1245 s.serialize_field("title", title)?;
1246 s.end()
1247 }
1248 }
1249 }
1250}
1251
1252fn deep_merge_toml(base: &mut toml::Value, overlay: &toml::Value) {
1263 match (base, overlay) {
1264 (toml::Value::Table(base_table), toml::Value::Table(overlay_table)) => {
1265 for (key, overlay_val) in overlay_table {
1266 match base_table.get_mut(key) {
1267 Some(base_val) => {
1268 deep_merge_toml(base_val, overlay_val);
1270 }
1271 None => {
1272 base_table.insert(key.clone(), overlay_val.clone());
1274 }
1275 }
1276 }
1277 }
1278 (base, overlay) => {
1280 *base = overlay.clone();
1281 }
1282 }
1283}
1284
1285fn substitute_env_vars(s: &mut String, secrets: &std::collections::HashMap<String, String>) {
1290 let mut result = String::with_capacity(s.len());
1292 let bytes = s.as_bytes();
1293 let mut i = 0;
1294
1295 while i < bytes.len() {
1296 if i + 1 < bytes.len() && bytes[i] == b'$' && bytes[i + 1] == b'{' {
1297 if let Some(end) = s[i + 2..].find('}') {
1299 let var_name = &s[i + 2..i + 2 + end];
1300 if let Some(value) = secrets.get(var_name) {
1302 result.push_str(value);
1303 } else if let Ok(value) = std::env::var(var_name) {
1304 result.push_str(&value);
1305 } else {
1306 result.push_str(&s[i..i + 2 + end + 1]);
1308 }
1309 i = i + 2 + end + 1;
1310 } else {
1311 result.push('$');
1313 i += 1;
1314 }
1315 } else {
1316 result.push(bytes[i] as char);
1317 i += 1;
1318 }
1319 }
1320
1321 *s = result;
1322}
1323
1324fn fingerprint_mcp_section(resolved: Option<&str>) -> u64 {
1334 use sha2::{Digest, Sha256};
1335 let section = resolved
1336 .and_then(|s| s.parse::<toml::Value>().ok())
1337 .and_then(|v| v.get("mcp").cloned());
1338 let bytes = match section {
1339 Some(v) => v.to_string().into_bytes(),
1340 None => b"<no-mcp>".to_vec(),
1341 };
1342 let digest = Sha256::digest(&bytes);
1343 u64::from_be_bytes(digest[..8].try_into().expect("sha256 digest >= 8 bytes"))
1344}
1345
1346#[cfg(test)]
1347mod tests {
1348 use super::*;
1349
1350 fn test_state() -> ServerState {
1352 let (session, _rx) = Session::new();
1353 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1354 ServerState::new(session, ws_tx, None)
1355 }
1356
1357 fn temp_user_home(tag: &str) -> std::path::PathBuf {
1359 static COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
1360 let n = COUNTER.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1361 let dir = std::env::temp_dir().join(format!(
1362 "trustee-state-test-{}-{}-{}",
1363 tag,
1364 std::process::id(),
1365 n
1366 ));
1367 std::fs::create_dir_all(dir.join("config")).expect("create temp user home");
1368 dir
1369 }
1370
1371 fn parse(toml_str: &str) -> toml::Value {
1372 toml_str.parse::<toml::Value>().expect("valid test TOML")
1373 }
1374
1375 #[test]
1378 fn overlay_allowlist_drops_non_allowlisted_sections() {
1379 let state = test_state();
1380 let home = temp_user_home("allowlist");
1381 std::fs::write(
1382 home.join("config").join("trustee.toml"),
1383 "[server]\nport = 1\n\n[auth]\nmode = \"kanidm\"\n\n[mcp]\nmode = \"user\"\n",
1384 )
1385 .expect("write overlay");
1386
1387 let shared = "[server]\nport = 8080\n\n[mcp]\nmode = \"shared\"\n";
1388 let merged = state
1389 .merge_user_config(&home, Some(shared))
1390 .expect("overlay has allowlisted content");
1391
1392 let merged_val = parse(&merged);
1393 let expected = parse("[server]\nport = 8080\n\n[mcp]\nmode = \"user\"\n");
1394 assert_eq!(merged_val, expected, "merged config must be shared + [mcp] overlay only");
1395 assert!(merged_val.get("auth").is_none(), "overlay [auth] must be dropped");
1396 }
1397
1398 #[test]
1400 fn no_overlay_file_returns_none() {
1401 let state = test_state();
1402 let home = temp_user_home("empty");
1403 assert!(state
1404 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1405 .is_none());
1406 }
1407
1408 #[test]
1410 fn overlay_with_no_allowlisted_sections_is_noop() {
1411 let state = test_state();
1412 let home = temp_user_home("all-dropped");
1413 std::fs::write(
1414 home.join("config").join("trustee.toml"),
1415 "[server]\nport = 1\n\n[storage]\npath = \"/tmp/x\"\n",
1416 )
1417 .expect("write overlay");
1418 assert!(state
1419 .merge_user_config(&home, Some("[server]\nport = 8080\n"))
1420 .is_none());
1421 }
1422
1423 #[test]
1426 fn llm_overlay_dropped_by_default_and_kept_when_enabled() {
1427 let shared = "[llm]\nprovider = \"openai\"\n\n[mcp]\nmode = \"shared\"\n";
1428 let overlay = "[llm]\nprovider = \"anthropic\"\n";
1429
1430 let state = test_state();
1432 assert!(!state.allow_llm_overlay);
1433 let home = temp_user_home("llm-off");
1434 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1435 assert!(state.merge_user_config(&home, Some(shared)).is_none());
1436
1437 let state = test_state().with_allow_llm_overlay(true);
1439 assert!(state.allow_llm_overlay);
1440 let home = temp_user_home("llm-on");
1441 std::fs::write(home.join("config").join("trustee.toml"), overlay).expect("write overlay");
1442 let merged = state
1443 .merge_user_config(&home, Some(shared))
1444 .expect("[llm] overlay applies when opted in");
1445 let merged_val = parse(&merged);
1446 assert_eq!(merged_val["llm"]["provider"].as_str(), Some("anthropic"));
1447 assert_eq!(merged_val["mcp"]["mode"].as_str(), Some("shared"));
1448 }
1449
1450 #[test]
1454 fn user_home_dir_uses_consolidated_hash() {
1455 let state = test_state();
1456 if let Some(home) = state.get_user_home_dir("farzan@example.com") {
1457 assert_eq!(
1458 home.file_name().and_then(|n| n.to_str()),
1459 Some(trustee_core::user_hash("farzan@example.com")).as_deref()
1460 );
1461 let users_root = dirs::home_dir().unwrap().join(".trustee").join("users");
1462 assert_eq!(home.parent(), Some(&users_root).map(|p| p.as_path()));
1463 }
1464 }
1466
1467 #[tokio::test]
1472 async fn agent_principals_get_isolated_session_buckets() {
1473 let state = test_state();
1474 let key_a = "agent-farzan";
1475 let key_b = "agent-paydar";
1476
1477 let (sid_a, session_a, _tx_a, _ts_a) = state.ensure_active_session(key_a).await;
1478 let (_sid_b, _session_b, _tx_b, _ts_b) = state.ensure_active_session(key_b).await;
1479
1480 assert!(
1482 state.get_session_by_any_id(key_a, &sid_a).await.is_some(),
1483 "owner bucket resolves its own session"
1484 );
1485 assert!(
1486 state.get_session_by_any_id(key_b, &sid_a).await.is_none(),
1487 "cross-agent session access must be 404/None"
1488 );
1489 assert!(Arc::strong_count(&session_a) >= 1);
1490
1491 if let Some(home) = state.get_user_home_dir(key_a) {
1493 assert_eq!(
1494 home.file_name().and_then(|n| n.to_str()),
1495 Some(trustee_core::user_hash(key_a)).as_deref()
1496 );
1497 }
1498 }
1499}
1500
1501#[cfg(test)]
1506mod mcp_loader_cache_tests {
1507 use super::*;
1508
1509 fn state_with_shared(shared: &str) -> ServerState {
1510 let (session, _rx) = Session::new();
1511 let (ws_tx, _ws_rx) = tokio::sync::broadcast::channel::<String>(16);
1512 let mut state = ServerState::new(session, ws_tx, None);
1513 state.config_toml = Some(shared.to_string());
1514 state
1515 }
1516
1517 struct TempUser {
1521 key: String,
1522 home: std::path::PathBuf,
1523 }
1524
1525 impl TempUser {
1526 fn new(tag: &str) -> Self {
1527 let key = format!("16c-{tag}-{}@test.invalid", std::process::id());
1528 let home = dirs::home_dir()
1529 .expect("HOME available in test env")
1530 .join(".trustee")
1531 .join("users")
1532 .join(trustee_core::user_hash(&key));
1533 std::fs::create_dir_all(home.join("config")).expect("create user home");
1534 Self { key, home }
1535 }
1536
1537 fn write_overlay(&self, toml_str: &str) {
1538 std::fs::write(self.home.join("config").join("trustee.toml"), toml_str)
1539 .expect("write overlay");
1540 }
1541 }
1542
1543 impl Drop for TempUser {
1544 fn drop(&mut self) {
1545 let _ = std::fs::remove_dir_all(&self.home);
1546 }
1547 }
1548
1549 #[tokio::test]
1550 async fn no_mcp_config_caches_disabled_marker() {
1551 let state = state_with_shared("[server]\nport = 8080\n");
1552 let user = TempUser::new("nomcp");
1553 let ts = Arc::new(pep::MemoryTokenStore::new());
1554
1555 let first = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1556 assert!(first.is_none(), "no [mcp] anywhere → disabled marker");
1557 assert_eq!(state.mcp_loaders.len(), 1, "exactly one cache entry");
1558
1559 let second = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1561 assert!(second.is_none());
1562 assert_eq!(state.mcp_loaders.len(), 1);
1563 }
1564
1565 #[tokio::test]
1566 async fn concurrent_cold_builds_single_flight() {
1567 let state = Arc::new(state_with_shared("[server]\nport = 8080\n"));
1568 let user = TempUser::new("singleflight");
1569 let ts = Arc::new(pep::MemoryTokenStore::new());
1570
1571 let mut handles = Vec::new();
1572 for _ in 0..5 {
1573 let state = state.clone();
1574 let key = user.key.clone();
1575 let ts = ts.clone();
1576 handles.push(tokio::spawn(async move {
1577 state.get_or_build_mcp_loader(&key, &ts).await
1578 }));
1579 }
1580 for h in handles {
1581 h.await.unwrap().expect("all five succeed");
1582 }
1583 assert_eq!(state.mcp_loaders.len(), 1, "single-flight → one entry");
1584 }
1585
1586 #[tokio::test]
1587 async fn fingerprint_change_triggers_rebuild() {
1588 let state = state_with_shared("[server]\nport = 8080\n");
1589 let user = TempUser::new("fpchange");
1590 let ts = Arc::new(pep::MemoryTokenStore::new());
1591
1592 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v1\"\nurl = \"http://127.0.0.1:9/sse\"\n");
1595 let v1 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1596 assert!(v1.is_some(), "enabled [mcp] → real loader");
1597 let fp1 = state
1598 .mcp_loaders
1599 .get(&trustee_core::user_hash(&user.key))
1600 .unwrap()
1601 .fingerprint;
1602
1603 std::thread::sleep(std::time::Duration::from_millis(5));
1605 user.write_overlay("[mcp]\nenabled = true\n\n[[mcp.servers]]\nname = \"v2\"\nurl = \"http://127.0.0.1:9/other\"\n");
1606 let v2 = state.get_or_build_mcp_loader(&user.key, &ts).await.unwrap();
1607 assert!(v2.is_some());
1608 let entry = state
1609 .mcp_loaders
1610 .get(&trustee_core::user_hash(&user.key))
1611 .unwrap();
1612 assert_ne!(
1613 entry.fingerprint, fp1,
1614 "fingerprint must change with content"
1615 );
1616 assert!(entry.degraded.is_none());
1617
1618 let _still_usable = v1.as_ref().unwrap().tool_count;
1620 }
1621
1622 #[tokio::test]
1623 async fn degraded_entry_fails_loud_within_backoff_and_isolates_users() {
1624 let state = state_with_shared("[server]\nport = 8080\n");
1625 let bad = TempUser::new("degraded-bad");
1626 let good = TempUser::new("degraded-good");
1627 let ts = Arc::new(pep::MemoryTokenStore::new());
1628
1629 bad.write_overlay("[mcp]\nenabled = \"not-a-bool\"\n");
1631 let err = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1632 Ok(_) => panic!("invalid [mcp] must fail loud"),
1633 Err(e) => e,
1634 };
1635 assert!(
1636 err.contains("invalid [mcp]"),
1637 "surfaces the parse error: {err}"
1638 );
1639
1640 let entry = state
1641 .mcp_loaders
1642 .get(&trustee_core::user_hash(&bad.key))
1643 .unwrap();
1644 assert!(entry.degraded.is_some(), "poison entry recorded");
1645
1646 let err2 = match state.get_or_build_mcp_loader(&bad.key, &ts).await {
1648 Ok(_) => panic!("still within backoff"),
1649 Err(e) => e,
1650 };
1651 assert_eq!(err, err2, "same cached error");
1652
1653 let good_loader = state
1655 .get_or_build_mcp_loader(&good.key, &ts)
1656 .await
1657 .expect("other user unaffected");
1658 assert!(
1659 good_loader.is_none(),
1660 "good user has no [mcp] → disabled marker"
1661 );
1662 }
1663}