diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 6606d0cf00..cfafafb827 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -109,7 +109,6 @@ use crate::protocol::Submission; use crate::protocol::TaskCompleteEvent; use crate::protocol::TokenCountEvent; use crate::protocol::TokenUsage; -use crate::protocol::TokenUsageInfo; use crate::protocol::TurnDiffEvent; use crate::protocol::WebSearchBeginEvent; use crate::rollout::RolloutRecorder; @@ -117,6 +116,7 @@ use crate::rollout::RolloutRecorderParams; use crate::safety::SafetyCheck; use crate::safety::assess_command_safety; use crate::safety::assess_safety_for_untrusted_command; +use crate::services::SessionServices; use crate::shell; use crate::turn_diff_tracker::TurnDiffTracker; use crate::unified_exec::UnifiedExecSessionManager; @@ -259,21 +259,8 @@ use crate::state::SessionState; pub(crate) struct Session { conversation_id: ConversationId, tx_event: Sender, - - /// Manager for external MCP servers/tools. - mcp_connection_manager: McpConnectionManager, - session_manager: ExecSessionManager, - unified_exec_manager: UnifiedExecSessionManager, - - notifier: UserNotifier, - - /// Optional rollout recorder for persisting the conversation transcript so - /// sessions can be replayed or inspected later. - rollout: Mutex>, state: Mutex, - codex_linux_sandbox_exe: Option, - user_shell: shell::Shell, - show_raw_agent_reasoning: bool, + services: SessionServices, next_internal_sub_id: AtomicU64, } @@ -459,18 +446,22 @@ impl Session { is_review_mode: false, final_output_json_schema: None, }; - let sess = Arc::new(Session { - conversation_id, - tx_event: tx_event.clone(), + let services = SessionServices { mcp_connection_manager, session_manager: ExecSessionManager::default(), unified_exec_manager: UnifiedExecSessionManager::default(), notifier: notify, - state: Mutex::new(state), rollout: Mutex::new(Some(rollout_recorder)), codex_linux_sandbox_exe: config.codex_linux_sandbox_exe.clone(), user_shell: default_shell, show_raw_agent_reasoning: config.show_raw_agent_reasoning, + }; + + let sess = Arc::new(Session { + conversation_id, + tx_event: tx_event.clone(), + state: Mutex::new(state), + services, next_internal_sub_id: AtomicU64::new(0), }); @@ -577,7 +568,7 @@ impl Session { let event_id = sub_id.clone(); let prev_entry = { let mut state = self.state.lock().await; - state.pending_approvals.insert(sub_id, tx_approve) + state.insert_pending_approval(sub_id, tx_approve) }; if prev_entry.is_some() { warn!("Overwriting existing pending approval for sub_id: {event_id}"); @@ -609,7 +600,7 @@ impl Session { let event_id = sub_id.clone(); let prev_entry = { let mut state = self.state.lock().await; - state.pending_approvals.insert(sub_id, tx_approve) + state.insert_pending_approval(sub_id, tx_approve) }; if prev_entry.is_some() { warn!("Overwriting existing pending approval for sub_id: {event_id}"); @@ -631,7 +622,7 @@ impl Session { pub async fn notify_approval(&self, sub_id: &str, decision: ReviewDecision) { let entry = { let mut state = self.state.lock().await; - state.pending_approvals.remove(sub_id) + state.remove_pending_approval(sub_id) }; match entry { Some(tx_approve) => { @@ -645,7 +636,7 @@ impl Session { pub async fn add_approved_command(&self, cmd: Vec) { let mut state = self.state.lock().await; - state.approved_commands.insert(cmd); + state.add_approved_command(cmd); } /// Records input items: always append to conversation history and @@ -685,7 +676,7 @@ impl Session { /// Append ResponseItems to the in-memory conversation history only. async fn record_into_history(&self, items: &[ResponseItem]) { let mut state = self.state.lock().await; - state.history.record_items(items.iter()); + state.record_items(items.iter()); } async fn persist_rollout_response_items(&self, items: &[ResponseItem]) { @@ -706,14 +697,14 @@ impl Session { Some(turn_context.cwd.clone()), Some(turn_context.approval_policy), Some(turn_context.sandbox_policy.clone()), - Some(self.user_shell.clone()), + Some(self.services.user_shell.clone()), ))); items } async fn persist_rollout_items(&self, items: &[RolloutItem]) { let recorder = { - let guard = self.rollout.lock().await; + let guard = self.services.rollout.lock().await; guard.clone() }; if let Some(rec) = recorder @@ -732,12 +723,10 @@ impl Session { { let mut state = self.state.lock().await; if let Some(token_usage) = token_usage { - let info = TokenUsageInfo::new_or_append( - &state.token_info, - &Some(token_usage.clone()), + state.update_token_info_from_usage( + token_usage, turn_context.client.get_model_context_window(), ); - state.token_info = info; } } self.send_token_count_event(sub_id).await; @@ -746,7 +735,7 @@ impl Session { async fn update_rate_limits(&self, sub_id: &str, new_rate_limits: RateLimitSnapshot) { { let mut state = self.state.lock().await; - state.latest_rate_limits = Some(new_rate_limits); + state.set_rate_limits(new_rate_limits); } self.send_token_count_event(sub_id).await; } @@ -754,7 +743,7 @@ impl Session { async fn send_token_count_event(&self, sub_id: &str) { let (info, rate_limits) = { let state = self.state.lock().await; - (state.token_info.clone(), state.latest_rate_limits.clone()) + state.token_info_and_rate_limits() }; let event = Event { id: sub_id.to_string(), @@ -772,8 +761,10 @@ impl Session { .await; // Derive user message events and persist only UserMessage to rollout - let msgs = - map_response_item_to_event_messages(&response_item, self.show_raw_agent_reasoning); + let msgs = map_response_item_to_event_messages( + &response_item, + self.services.show_raw_agent_reasoning, + ); let user_msgs: Vec = msgs .into_iter() .filter_map(|m| match m { @@ -972,7 +963,7 @@ impl Session { pub async fn turn_input_with_history(&self, extra: Vec) -> Vec { let history = { let state = self.state.lock().await; - state.history.contents() + state.history_snapshot() }; [history, extra].concat() } @@ -981,7 +972,7 @@ impl Session { pub async fn inject_input(&self, input: Vec) -> Result<(), Vec> { let mut state = self.state.lock().await; if state.current_task.is_some() { - state.pending_input.push(input.into()); + state.push_pending_input(input.into()); Ok(()) } else { Err(input) @@ -990,13 +981,7 @@ impl Session { pub async fn get_pending_input(&self) -> Vec { let mut state = self.state.lock().await; - if state.pending_input.is_empty() { - Vec::with_capacity(0) - } else { - let mut ret = Vec::new(); - std::mem::swap(&mut ret, &mut state.pending_input); - ret - } + state.take_pending_input() } pub async fn call_tool( @@ -1005,7 +990,8 @@ impl Session { tool: &str, arguments: Option, ) -> anyhow::Result { - self.mcp_connection_manager + self.services + .mcp_connection_manager .call_tool(server, tool, arguments) .await } @@ -1013,8 +999,7 @@ impl Session { pub async fn interrupt_task(&self) { info!("interrupt received: abort current task, if any"); let mut state = self.state.lock().await; - state.pending_approvals.clear(); - state.pending_input.clear(); + state.clear_pending(); if let Some(task) = state.current_task.take() { task.abort(TurnAbortReason::Interrupted); } @@ -1022,8 +1007,7 @@ impl Session { fn interrupt_task_sync(&self) { if let Ok(mut state) = self.state.try_lock() { - state.pending_approvals.clear(); - state.pending_input.clear(); + state.clear_pending(); if let Some(task) = state.current_task.take() { task.abort(TurnAbortReason::Interrupted); } @@ -1031,7 +1015,7 @@ impl Session { } pub(crate) fn notifier(&self) -> &UserNotifier { - &self.notifier + &self.services.notifier } } @@ -1403,7 +1387,7 @@ async fn submission_loop( let sub_id = sub.id.clone(); // This is a cheap lookup from the connection manager's cache. - let tools = sess.mcp_connection_manager.list_all_tools(); + let tools = sess.services.mcp_connection_manager.list_all_tools(); let event = Event { id: sub_id, msg: EventMsg::McpListToolsResponse( @@ -1453,7 +1437,7 @@ async fn submission_loop( // Gracefully flush and shutdown rollout recorder on session end so tests // that inspect the rollout file do not race with the background writer. let recorder_opt = { - let mut guard = sess.rollout.lock().await; + let mut guard = sess.services.rollout.lock().await; guard.take() }; if let Some(rec) = recorder_opt @@ -1480,7 +1464,7 @@ async fn submission_loop( let sub_id = sub.id.clone(); // Flush rollout writes before returning the path so readers observe a consistent file. let (path, rec_opt) = { - let guard = sess.rollout.lock().await; + let guard = sess.services.rollout.lock().await; match guard.as_ref() { Some(rec) => (rec.get_rollout_path(), Some(rec.clone())), None => { @@ -1938,7 +1922,7 @@ async fn run_turn( ) -> CodexResult { let tools = get_openai_tools( &turn_context.tools_config, - Some(sess.mcp_connection_manager.list_all_tools()), + Some(sess.services.mcp_connection_manager.list_all_tools()), ); let prompt = Prompt { @@ -2188,7 +2172,7 @@ async fn try_run_turn( sess.send_event(event).await; } ResponseEvent::ReasoningContentDelta(delta) => { - if sess.show_raw_agent_reasoning { + if sess.services.show_raw_agent_reasoning { let event = Event { id: sub_id.to_string(), msg: EventMsg::AgentReasoningRawContentDelta( @@ -2310,7 +2294,10 @@ async fn handle_response_item( trace!("suppressing assistant Message in review mode"); Vec::new() } - _ => map_response_item_to_event_messages(&item, sess.show_raw_agent_reasoning), + _ => map_response_item_to_event_messages( + &item, + sess.services.show_raw_agent_reasoning, + ), }; for msg in msgs { let event = Event { @@ -2356,7 +2343,11 @@ async fn handle_unified_exec_tool_call( timeout_ms, }; - let result = sess.unified_exec_manager.handle_request(request).await; + let result = sess + .services + .unified_exec_manager + .handle_request(request) + .await; let output_payload = match result { Ok(value) => { @@ -2531,6 +2522,7 @@ async fn handle_function_call( } }; let result = sess + .services .session_manager .handle_exec_command_request(exec_params) .await; @@ -2554,6 +2546,7 @@ async fn handle_function_call( } }; let result = sess + .services .session_manager .handle_write_stdin_request(write_stdin_params) .await; @@ -2565,7 +2558,7 @@ async fn handle_function_call( } } _ => { - match sess.mcp_connection_manager.parse_tool_name(&name) { + match sess.services.mcp_connection_manager.parse_tool_name(&name) { Some((server, tool_name)) => { handle_mcp_tool_call(sess, &sub_id, call_id, server, tool_name, arguments).await } @@ -2683,11 +2676,12 @@ fn maybe_translate_shell_command( sess: &Session, turn_context: &TurnContext, ) -> ExecParams { - let should_translate = matches!(sess.user_shell, crate::shell::Shell::PowerShell(_)) + let should_translate = matches!(sess.services.user_shell, crate::shell::Shell::PowerShell(_)) || turn_context.shell_environment_policy.use_profile; if should_translate && let Some(command) = sess + .services .user_shell .format_default_shell_invocation(params.command.clone()) { @@ -2802,7 +2796,7 @@ async fn handle_container_exec_with_params( ¶ms.command, turn_context.approval_policy, &turn_context.sandbox_policy, - &state.approved_commands, + state.approved_commands_ref(), params.with_escalated_permissions.unwrap_or(false), ) }; @@ -2881,7 +2875,7 @@ async fn handle_container_exec_with_params( sandbox_type, sandbox_policy: &turn_context.sandbox_policy, sandbox_cwd: &turn_context.cwd, - codex_linux_sandbox_exe: &sess.codex_linux_sandbox_exe, + codex_linux_sandbox_exe: &sess.services.codex_linux_sandbox_exe, stdout_stream: if exec_command_context.apply_patch.is_some() { None } else { @@ -3016,7 +3010,7 @@ async fn handle_sandbox_error( sandbox_type: SandboxType::None, sandbox_policy: &turn_context.sandbox_policy, sandbox_cwd: &turn_context.cwd, - codex_linux_sandbox_exe: &sess.codex_linux_sandbox_exe, + codex_linux_sandbox_exe: &sess.services.codex_linux_sandbox_exe, stdout_stream: if exec_command_context.apply_patch.is_some() { None } else { @@ -3370,7 +3364,7 @@ mod tests { }), )); - let actual = tokio_test::block_on(async { session.state.lock().await.history.contents() }); + let actual = tokio_test::block_on(async { session.state.lock().await.history_snapshot() }); assert_eq!(expected, actual); } @@ -3383,7 +3377,7 @@ mod tests { session.record_initial_history(&turn_context, InitialHistory::Forked(rollout_items)), ); - let actual = tokio_test::block_on(async { session.state.lock().await.history.contents() }); + let actual = tokio_test::block_on(async { session.state.lock().await.history_snapshot() }); assert_eq!(expected, actual); } @@ -3609,18 +3603,21 @@ mod tests { is_review_mode: false, final_output_json_schema: None, }; - let session = Session { - conversation_id, - tx_event, + let services = SessionServices { mcp_connection_manager: McpConnectionManager::default(), session_manager: ExecSessionManager::default(), unified_exec_manager: UnifiedExecSessionManager::default(), notifier: UserNotifier::default(), rollout: Mutex::new(None), - state: Mutex::new(SessionState::new()), codex_linux_sandbox_exe: None, user_shell: shell::Shell::Unknown, show_raw_agent_reasoning: config.show_raw_agent_reasoning, + }; + let session = Session { + conversation_id, + tx_event, + state: Mutex::new(SessionState::new()), + services, next_internal_sub_id: AtomicU64::new(0), }; (session, turn_context) diff --git a/codex-rs/core/src/codex/compact.rs b/codex-rs/core/src/codex/compact.rs index 8f213d4e7e..88e8a8f433 100644 --- a/codex-rs/core/src/codex/compact.rs +++ b/codex-rs/core/src/codex/compact.rs @@ -153,7 +153,7 @@ async fn run_compact_task_inner( } let history_snapshot = { let state = sess.state.lock().await; - state.history.contents() + state.history_snapshot() }; let summary_text = get_last_assistant_message_from_turn(&history_snapshot).unwrap_or_default(); let user_messages = collect_user_messages(&history_snapshot); @@ -161,7 +161,7 @@ async fn run_compact_task_inner( let new_history = build_compacted_history(initial_context, &user_messages, &summary_text); { let mut state = sess.state.lock().await; - state.history.replace(new_history); + state.replace_history(new_history); } let rollout_item = RolloutItem::Compacted(CompactedItem { @@ -271,7 +271,7 @@ async fn drain_to_completed( match event { Ok(ResponseEvent::OutputItemDone(item)) => { let mut state = sess.state.lock().await; - state.history.record_items(std::slice::from_ref(&item)); + state.record_items(std::slice::from_ref(&item)); } Ok(ResponseEvent::Completed { .. }) => { return Ok(()); diff --git a/codex-rs/core/src/lib.rs b/codex-rs/core/src/lib.rs index e4f5709345..79203f1b37 100644 --- a/codex-rs/core/src/lib.rs +++ b/codex-rs/core/src/lib.rs @@ -47,6 +47,7 @@ pub use model_provider_info::create_oss_provider_with_base_url; mod conversation_manager; mod event_mapping; pub mod review_format; +mod services; pub use codex_protocol::protocol::InitialHistory; pub use conversation_manager::ConversationManager; pub use conversation_manager::NewConversation; diff --git a/codex-rs/core/src/services/mod.rs b/codex-rs/core/src/services/mod.rs new file mode 100644 index 0000000000..bad425aa12 --- /dev/null +++ b/codex-rs/core/src/services/mod.rs @@ -0,0 +1,23 @@ +//! Group long‑lived helpers/managers for a session. + +use std::path::PathBuf; + +use tokio::sync::Mutex; + +use crate::exec_command::ExecSessionManager; +use crate::mcp_connection_manager::McpConnectionManager; +use crate::rollout::RolloutRecorder; +use crate::unified_exec::UnifiedExecSessionManager; +use crate::user_notification::UserNotifier; + +/// Container for side‑effectful services and helpers used by `Session`. +pub(crate) struct SessionServices { + pub(crate) mcp_connection_manager: McpConnectionManager, + pub(crate) session_manager: ExecSessionManager, + pub(crate) unified_exec_manager: UnifiedExecSessionManager, + pub(crate) notifier: UserNotifier, + pub(crate) rollout: Mutex>, + pub(crate) codex_linux_sandbox_exe: Option, + pub(crate) user_shell: crate::shell::Shell, + pub(crate) show_raw_agent_reasoning: bool, +} diff --git a/codex-rs/core/src/state/session.rs b/codex-rs/core/src/state/session.rs index 8fc7cbb4fe..6b968769d2 100644 --- a/codex-rs/core/src/state/session.rs +++ b/codex-rs/core/src/state/session.rs @@ -4,12 +4,14 @@ use std::collections::HashMap; use std::collections::HashSet; use codex_protocol::models::ResponseInputItem; +use codex_protocol::models::ResponseItem; use tokio::sync::oneshot; use crate::codex::AgentTask; use crate::conversation_history::ConversationHistory; use crate::protocol::RateLimitSnapshot; use crate::protocol::ReviewDecision; +use crate::protocol::TokenUsage; use crate::protocol::TokenUsageInfo; /// Persistent, session-scoped state previously stored directly on `Session`. @@ -32,4 +34,88 @@ impl SessionState { ..Default::default() } } + + // History helpers + pub(crate) fn record_items(&mut self, items: I) + where + I: IntoIterator, + I::Item: std::ops::Deref, + { + self.history.record_items(items) + } + + pub(crate) fn history_snapshot(&self) -> Vec { + self.history.contents() + } + + pub(crate) fn replace_history(&mut self, items: Vec) { + self.history.replace(items); + } + + // Approved command helpers + pub(crate) fn add_approved_command(&mut self, cmd: Vec) { + self.approved_commands.insert(cmd); + } + + pub(crate) fn approved_commands_ref(&self) -> &HashSet> { + &self.approved_commands + } + + // Token/rate limit helpers + pub(crate) fn update_token_info_from_usage( + &mut self, + usage: &TokenUsage, + model_context_window: Option, + ) { + self.token_info = TokenUsageInfo::new_or_append( + &self.token_info, + &Some(usage.clone()), + model_context_window, + ); + } + + pub(crate) fn set_rate_limits(&mut self, snapshot: RateLimitSnapshot) { + self.latest_rate_limits = Some(snapshot); + } + + pub(crate) fn token_info_and_rate_limits( + &self, + ) -> (Option, Option) { + (self.token_info.clone(), self.latest_rate_limits.clone()) + } + + // Pending input/approval helpers + pub(crate) fn insert_pending_approval( + &mut self, + key: String, + tx: oneshot::Sender, + ) -> Option> { + self.pending_approvals.insert(key, tx) + } + + pub(crate) fn remove_pending_approval( + &mut self, + key: &str, + ) -> Option> { + self.pending_approvals.remove(key) + } + + pub(crate) fn clear_pending(&mut self) { + self.pending_approvals.clear(); + self.pending_input.clear(); + } + + pub(crate) fn push_pending_input(&mut self, input: ResponseInputItem) { + self.pending_input.push(input); + } + + pub(crate) fn take_pending_input(&mut self) -> Vec { + if self.pending_input.is_empty() { + Vec::with_capacity(0) + } else { + let mut ret = Vec::new(); + std::mem::swap(&mut ret, &mut self.pending_input); + ret + } + } }