diff --git a/codex-rs/tui/src/app.rs b/codex-rs/tui/src/app.rs index 9085fbb163a5..1a5823ef7dcb 100644 --- a/codex-rs/tui/src/app.rs +++ b/codex-rs/tui/src/app.rs @@ -212,6 +212,7 @@ mod activity_groups; mod app_server_event_targets; mod app_server_events; pub(crate) mod app_server_requests; +mod app_server_thread_ownership; mod backend_banner_fallback; mod background_requests; mod composer_hints; diff --git a/codex-rs/tui/src/app/app_server_events.rs b/codex-rs/tui/src/app/app_server_events.rs index 5a0e798a6035..6d79b7e0c390 100644 --- a/codex-rs/tui/src/app/app_server_events.rs +++ b/codex-rs/tui/src/app/app_server_events.rs @@ -460,25 +460,12 @@ impl App { { return; } - if self.primary_thread_id.is_some() - && self.primary_thread_id != Some(thread_id) - && !self.thread_event_channels.contains_key(&thread_id) - && self.agent_navigation.get(&thread_id).is_none() - && !self.side_threads.contains_key(&thread_id) - && !matches!(¬ification, ServerNotification::McpServerStatusUpdated(_)) - && !matches!( - ¬ification, - ServerNotification::ThreadStarted(started) - if matches!( - &started.thread.source, - SessionSource::SubAgent(SubAgentSource::ThreadSpawn { - parent_thread_id, - .. - }) if self.primary_thread_id == Some(*parent_thread_id) - || self.thread_event_channels.contains_key(parent_thread_id) - || self.agent_navigation.get(parent_thread_id).is_some() - ) - ) + let untracked_thread = + self.primary_thread_id.is_some() && !self.owns_thread_for_routing(thread_id); + if untracked_thread + && !self + .owns_untracked_notification(app_server_client, thread_id, ¬ification) + .await { return; } @@ -754,10 +741,7 @@ impl App { else { return; }; - if self.primary_thread_id != Some(parent_thread_id) - && !self.thread_event_channels.contains_key(&parent_thread_id) - && self.agent_navigation.get(&parent_thread_id).is_none() - { + if !self.owns_thread_for_routing(parent_thread_id) { if self .agents_overview .dispatched_requests diff --git a/codex-rs/tui/src/app/app_server_thread_ownership.rs b/codex-rs/tui/src/app/app_server_thread_ownership.rs new file mode 100644 index 000000000000..a5259e619f8b --- /dev/null +++ b/codex-rs/tui/src/app/app_server_thread_ownership.rs @@ -0,0 +1,83 @@ +//! Ownership checks for routing app-server events to a TUI root. + +use super::App; +use crate::app_server_session::AppServerSession; +use codex_app_server_client::TypedRequestError; +use codex_app_server_protocol::ClientRequest; +use codex_app_server_protocol::RequestId; +use codex_app_server_protocol::ServerNotification; +use codex_app_server_protocol::SessionSource; +use codex_app_server_protocol::ThreadReadParams; +use codex_app_server_protocol::ThreadReadResponse; +use codex_protocol::ThreadId; +use codex_protocol::protocol::SubAgentSource; +use std::time::Duration; + +async fn read_mcp_subagent_parent_thread_id( + app_server_client: &AppServerSession, + thread_id: ThreadId, +) -> Option { + let response = tokio::time::timeout(Duration::from_secs(/*secs*/ 1), async { + let mut retry_delay = Duration::from_millis(50); + loop { + match app_server_client + .request_handle() + .request_typed::(ClientRequest::ThreadRead { + request_id: RequestId::String(uuid::Uuid::new_v4().to_string()), + params: ThreadReadParams { + thread_id: thread_id.to_string(), + include_turns: false, + }, + }) + .await + { + Ok(response) => return Some(response), + Err(TypedRequestError::Server { source, .. }) + if source.message.starts_with("thread not found:") + || source.message.starts_with("thread not loaded:") => + { + tokio::time::sleep(retry_delay).await; + retry_delay = retry_delay.saturating_mul(2).min(Duration::from_secs(1)); + } + Err(_) => return None, + } + } + }) + .await + .ok()??; + subagent_parent(&response.thread.source) +} + +fn subagent_parent(source: &SessionSource) -> Option { + match source { + SessionSource::SubAgent(SubAgentSource::ThreadSpawn { + parent_thread_id, .. + }) => Some(*parent_thread_id), + _ => None, + } +} + +impl App { + pub(super) fn owns_thread_for_routing(&self, thread_id: ThreadId) -> bool { + self.primary_thread_id == Some(thread_id) + || self.thread_event_channels.contains_key(&thread_id) + || self.side_threads.contains_key(&thread_id) + || self.agent_navigation.get(&thread_id).is_some() + } + + pub(super) async fn owns_untracked_notification( + &self, + app_server_client: &AppServerSession, + thread_id: ThreadId, + notification: &ServerNotification, + ) -> bool { + let parent_thread_id = match notification { + ServerNotification::ThreadStarted(started) => subagent_parent(&started.thread.source), + ServerNotification::McpServerStatusUpdated(_) => { + read_mcp_subagent_parent_thread_id(app_server_client, thread_id).await + } + _ => None, + }; + parent_thread_id.is_some_and(|thread_id| self.owns_thread_for_routing(thread_id)) + } +} diff --git a/codex-rs/tui/src/app/tests.rs b/codex-rs/tui/src/app/tests.rs index a76ac39d9cd6..33b0977ffdcb 100644 --- a/codex-rs/tui/src/app/tests.rs +++ b/codex-rs/tui/src/app/tests.rs @@ -5436,6 +5436,12 @@ async fn primary_thread_ignores_child_mcp_startup_notifications() { let child_thread_id = ThreadId::new(); app.primary_thread_id = Some(parent_thread_id); app.active_thread_id = Some(parent_thread_id); + app.upsert_agent_picker_thread( + child_thread_id, + /*agent_nickname*/ None, + /*agent_role*/ None, + /*is_closed*/ false, + ); app.handle_app_server_event( &app_server, diff --git a/codex-rs/tui/src/app/tests/startup.rs b/codex-rs/tui/src/app/tests/startup.rs index c5f43238be5d..c4c0fee68e71 100644 --- a/codex-rs/tui/src/app/tests/startup.rs +++ b/codex-rs/tui/src/app/tests/startup.rs @@ -1394,14 +1394,35 @@ async fn startup_thread_started_does_not_replay_resolved_approval() -> Result<() } #[tokio::test] -async fn owned_subagent_approval_before_thread_started_is_preserved() -> Result<()> { +async fn subagent_approval_respects_root_ownership() -> Result<()> { + check_subagent_approval_routing(ApprovalRouting::OwnedApprovalFirst).await?; + check_subagent_approval_routing(ApprovalRouting::OwnedAfterMcp).await?; + check_subagent_approval_routing(ApprovalRouting::ForeignAfterMcp).await +} + +enum ApprovalRouting { + OwnedApprovalFirst, + OwnedAfterMcp, + ForeignAfterMcp, +} + +async fn check_subagent_approval_routing(routing: ApprovalRouting) -> Result<()> { + let owned_parent = !matches!(&routing, ApprovalRouting::ForeignAfterMcp); let (mut app, _app_event_rx, _op_rx) = make_test_app_with_channels().await; let codex_home = tempdir()?; app.config.codex_home = codex_home.path().to_path_buf().abs(); app.config.sqlite = codex_state::SqliteConfig::new_for_testing(codex_home.path().abs()); let mut app_server = Box::pin(crate::start_embedded_app_server_for_picker(&app.config)).await?; let parent = app_server.start_thread(&app.config).await?; - let parent_thread_id = parent.session.thread_id; + let parent_thread_id = if owned_parent { + parent.session.thread_id + } else { + app_server + .start_thread(&app.config) + .await? + .session + .thread_id + }; app.enqueue_primary_thread_session(parent.session, parent.turns) .await?; let child_thread_id = ThreadId::from_string( @@ -1439,20 +1460,54 @@ async fn owned_subagent_approval_before_thread_started_is_preserved() -> Result< /*approval_id*/ None, ); + if matches!(&routing, ApprovalRouting::ForeignAfterMcp) { + send_failed_mcp_startup(&mut app, &app_server, parent_thread_id).await; + } + if !matches!(&routing, ApprovalRouting::OwnedApprovalFirst) { + send_failed_mcp_startup(&mut app, &app_server, child_thread_id).await; + let has_channel = app.thread_event_channels.contains_key(&child_thread_id); + assert_eq!(has_channel, owned_parent); + } + app.handle_app_server_event( &app_server, codex_app_server_client::AppServerEvent::ServerRequest(Box::new(request.clone())), ) .await; - assert!( - app.pending_app_server_requests - .contains_server_request(&request) - ); - assert!(app.thread_event_channels.contains_key(&child_thread_id)); + let requests = &app.pending_app_server_requests; + assert_eq!(requests.contains_server_request(&request), owned_parent); + assert_eq!(app.chat_widget.has_active_view(), owned_parent); + if !owned_parent { + let popup = render_bottom_popup(&app.chat_widget, /*width*/ 80); + let approval = popup + .lines() + .find(|line| line.contains("Would you like to run")); + insta::assert_snapshot!(approval.unwrap_or_default(), @""); + } Ok(()) } +async fn send_failed_mcp_startup( + app: &mut App, + app_server: &AppServerSession, + thread_id: ThreadId, +) { + app.handle_app_server_event( + app_server, + codex_app_server_client::AppServerEvent::ServerNotification(Box::new( + ServerNotification::McpServerStatusUpdated(McpServerStatusUpdatedNotification { + thread_id: Some(thread_id.to_string()), + name: "fixture".to_string(), + status: McpServerStartupState::Failed, + error: Some("fixture failed".to_string()), + failure_reason: None, + }), + )), + ) + .await; +} + #[tokio::test] async fn startup_thread_start_failure_returns_error() { let (mut app, _app_event_rx, _op_rx) = make_test_app_with_channels().await;