From 1a92267a40e243da81f5d3e64bb4518155cc99a9 Mon Sep 17 00:00:00 2001 From: Ahmed Ibrahim Date: Thu, 28 Aug 2025 13:01:37 -0700 Subject: [PATCH] enable search config change --- codex-rs/core/src/codex.rs | 10 +- codex-rs/core/tests/suite/mod.rs | 1 + codex-rs/core/tests/suite/prompt_caching.rs | 2 + codex-rs/core/tests/suite/tools_web_search.rs | 170 ++++++++++++++++++ .../mcp-server/src/codex_message_processor.rs | 2 + .../suite/codex_message_processor_flow.rs | 1 + codex-rs/protocol/src/mcp_protocol.rs | 1 + codex-rs/protocol/src/protocol.rs | 7 + codex-rs/tui/src/chatwidget.rs | 2 + 9 files changed, 193 insertions(+), 3 deletions(-) create mode 100644 codex-rs/core/tests/suite/tools_web_search.rs diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index 365969ac02..3c8bfb4799 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -1058,6 +1058,7 @@ async fn submission_loop( model, effort, summary, + enable_web_search, } => { // Recalculate the persistent turn context with provided overrides. let prev = Arc::clone(&turn_context); @@ -1082,6 +1083,7 @@ async fn submission_loop( let mut updated_config = (*config).clone(); updated_config.model = effective_model.clone(); updated_config.model_family = effective_family.clone(); + updated_config.tools_web_search_request = enable_web_search.unwrap_or(false); if let Some(model_info) = get_model_info(&effective_family) { updated_config.model_context_window = Some(model_info.context_window); } @@ -1107,7 +1109,7 @@ async fn submission_loop( sandbox_policy: new_sandbox_policy.clone(), include_plan_tool: config.include_plan_tool, include_apply_patch_tool: config.include_apply_patch_tool, - include_web_search_request: config.tools_web_search_request, + include_web_search_request: enable_web_search.unwrap_or(false), use_streamable_shell_tool: config.use_experimental_streamable_shell_tool, include_view_image_tool: config.include_view_image_tool, }); @@ -1154,6 +1156,7 @@ async fn submission_loop( model, effort, summary, + enable_web_search, } => { // attempt to inject input into current task if let Err(items) = sess.inject_input(items) { @@ -1169,6 +1172,7 @@ async fn submission_loop( let mut per_turn_config = (*config).clone(); per_turn_config.model = model.clone(); per_turn_config.model_family = model_family.clone(); + per_turn_config.tools_web_search_request = enable_web_search; if let Some(model_info) = get_model_info(&model_family) { per_turn_config.model_context_window = Some(model_info.context_window); } @@ -1176,7 +1180,7 @@ async fn submission_loop( // Build a new client with per‑turn reasoning settings. // Reuse the same provider and session id; auth defaults to env/API key. let client = ModelClient::new( - Arc::new(per_turn_config), + Arc::new(per_turn_config.clone()), auth_manager, provider, effort, @@ -1192,7 +1196,7 @@ async fn submission_loop( sandbox_policy: sandbox_policy.clone(), include_plan_tool: config.include_plan_tool, include_apply_patch_tool: config.include_apply_patch_tool, - include_web_search_request: config.tools_web_search_request, + include_web_search_request: per_turn_config.tools_web_search_request, use_streamable_shell_tool: config .use_experimental_streamable_shell_tool, include_view_image_tool: config.include_view_image_tool, diff --git a/codex-rs/core/tests/suite/mod.rs b/codex-rs/core/tests/suite/mod.rs index 22aa826699..3569479c4b 100644 --- a/codex-rs/core/tests/suite/mod.rs +++ b/codex-rs/core/tests/suite/mod.rs @@ -10,3 +10,4 @@ mod prompt_caching; mod seatbelt; mod stream_error_allows_next_turn; mod stream_no_completed; +mod tools_web_search; diff --git a/codex-rs/core/tests/suite/prompt_caching.rs b/codex-rs/core/tests/suite/prompt_caching.rs index 999f807286..d88887c297 100644 --- a/codex-rs/core/tests/suite/prompt_caching.rs +++ b/codex-rs/core/tests/suite/prompt_caching.rs @@ -393,6 +393,7 @@ async fn overrides_turn_context_but_keeps_cached_prefix_and_key_constant() { model: Some("o3".to_string()), effort: Some(ReasoningEffort::High), summary: Some(ReasoningSummary::Detailed), + enable_web_search: Some(false), }) .await .unwrap(); @@ -521,6 +522,7 @@ async fn per_turn_overrides_keep_cached_prefix_and_key_constant() { model: "o3".to_string(), effort: ReasoningEffort::High, summary: ReasoningSummary::Detailed, + enable_web_search: false, }) .await .unwrap(); diff --git a/codex-rs/core/tests/suite/tools_web_search.rs b/codex-rs/core/tests/suite/tools_web_search.rs new file mode 100644 index 0000000000..2851c14a82 --- /dev/null +++ b/codex-rs/core/tests/suite/tools_web_search.rs @@ -0,0 +1,170 @@ +#![allow(clippy::unwrap_used)] + +use codex_core::ConversationManager; +use codex_core::ModelProviderInfo; +use codex_core::built_in_model_providers; +use codex_core::protocol::AskForApproval; +use codex_core::protocol::EventMsg; +use codex_core::protocol::InputItem; +use codex_core::protocol::Op; +use codex_core::protocol::SandboxPolicy; +use codex_core::protocol_config_types::ReasoningEffort; +use codex_core::protocol_config_types::ReasoningSummary; +use codex_login::CodexAuth; +use core_test_support::load_default_config_for_test; +use core_test_support::load_sse_fixture_with_id; +use core_test_support::wait_for_event; +use tempfile::TempDir; +use wiremock::Mock; +use wiremock::MockServer; +use wiremock::ResponseTemplate; +use wiremock::matchers::method; +use wiremock::matchers::path; + +/// Build minimal SSE stream with completed marker using the JSON fixture. +fn sse_completed(id: &str) -> String { + load_sse_fixture_with_id("tests/fixtures/completed_template.json", id) +} + +fn tools_include_web_search(body: &serde_json::Value) -> bool { + body["tools"] + .as_array() + .unwrap() + .iter() + .any(|t| t["type"].as_str() == Some("web_search")) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn web_search_tool_present_after_override_turn_context() { + // Mock server + let server = MockServer::start().await; + + let sse = sse_completed("resp"); + let template = ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_raw(sse, "text/event-stream"); + + // Expect one POST to /v1/responses + Mock::given(method("POST")) + .and(path("/v1/responses")) + .respond_with(template) + .expect(1) + .mount(&server) + .await; + + let model_provider = ModelProviderInfo { + base_url: Some(format!("{}/v1", server.uri())), + ..built_in_model_providers()["openai"].clone() + }; + + let cwd = TempDir::new().unwrap(); + let codex_home = TempDir::new().unwrap(); + let mut config = load_default_config_for_test(&codex_home); + config.cwd = cwd.path().to_path_buf(); + config.model_provider = model_provider; + + let conversation_manager = + ConversationManager::with_auth(CodexAuth::from_api_key("Test API Key")); + let codex = conversation_manager + .new_conversation(config) + .await + .expect("create new conversation") + .conversation; + + // Enable web search for subsequent turns + codex + .submit(Op::OverrideTurnContext { + cwd: None, + approval_policy: None, + sandbox_policy: None, + model: None, + effort: None, + summary: None, + enable_web_search: Some(true), + }) + .await + .unwrap(); + + // Trigger a user turn + codex + .submit(Op::UserInput { + items: vec![InputItem::Text { + text: "hello".into(), + }], + }) + .await + .unwrap(); + wait_for_event(&codex, |ev| matches!(ev, EventMsg::TaskComplete(_))).await; + + let requests = server.received_requests().await.unwrap(); + assert_eq!(requests.len(), 1, "expected one POST request"); + let body = requests[0].body_json::().unwrap(); + assert!( + tools_include_web_search(&body), + "tools should include web_search when enable_web_search is true via OverrideTurnContext", + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn web_search_tool_present_in_user_turn_when_enabled() { + // Mock server + let server = MockServer::start().await; + + let sse = sse_completed("resp"); + let template = ResponseTemplate::new(200) + .insert_header("content-type", "text/event-stream") + .set_body_raw(sse, "text/event-stream"); + + // Expect one POST to /v1/responses + Mock::given(method("POST")) + .and(path("/v1/responses")) + .respond_with(template) + .expect(1) + .mount(&server) + .await; + + let model_provider = ModelProviderInfo { + base_url: Some(format!("{}/v1", server.uri())), + ..built_in_model_providers()["openai"].clone() + }; + + let cwd = TempDir::new().unwrap(); + let codex_home = TempDir::new().unwrap(); + let mut config = load_default_config_for_test(&codex_home); + config.cwd = cwd.path().to_path_buf(); + config.model_provider = model_provider; + + let conversation_manager = + ConversationManager::with_auth(CodexAuth::from_api_key("Test API Key")); + let codex = conversation_manager + .new_conversation(config) + .await + .expect("create new conversation") + .conversation; + + // Submit a per-turn override with enable_web_search = true + codex + .submit(Op::UserTurn { + items: vec![InputItem::Text { + text: "hello".into(), + }], + cwd: cwd.path().to_path_buf(), + approval_policy: AskForApproval::Never, + sandbox_policy: SandboxPolicy::new_read_only_policy(), + model: "o3".to_string(), + effort: ReasoningEffort::High, + summary: ReasoningSummary::Detailed, + enable_web_search: true, + }) + .await + .unwrap(); + wait_for_event(&codex, |ev| matches!(ev, EventMsg::TaskComplete(_))).await; + + let requests = server.received_requests().await.unwrap(); + assert_eq!(requests.len(), 1, "expected one POST request"); + let body = requests[0].body_json::().unwrap(); + assert!( + tools_include_web_search(&body), + "tools should include web_search when enable_web_search is true in UserTurn", + ); +} diff --git a/codex-rs/mcp-server/src/codex_message_processor.rs b/codex-rs/mcp-server/src/codex_message_processor.rs index aae463ad92..e84e05eb81 100644 --- a/codex-rs/mcp-server/src/codex_message_processor.rs +++ b/codex-rs/mcp-server/src/codex_message_processor.rs @@ -505,6 +505,7 @@ impl CodexMessageProcessor { model, effort, summary, + enable_web_search, } = params; let Ok(conversation) = self @@ -539,6 +540,7 @@ impl CodexMessageProcessor { model, effort, summary, + enable_web_search, }) .await; diff --git a/codex-rs/mcp-server/tests/suite/codex_message_processor_flow.rs b/codex-rs/mcp-server/tests/suite/codex_message_processor_flow.rs index 092b135291..3315b20948 100644 --- a/codex-rs/mcp-server/tests/suite/codex_message_processor_flow.rs +++ b/codex-rs/mcp-server/tests/suite/codex_message_processor_flow.rs @@ -320,6 +320,7 @@ async fn test_send_user_turn_changes_approval_policy_behavior() { model: "mock-model".to_string(), effort: ReasoningEffort::Medium, summary: ReasoningSummary::Auto, + enable_web_search: false, }) .await .expect("send sendUserTurn"); diff --git a/codex-rs/protocol/src/mcp_protocol.rs b/codex-rs/protocol/src/mcp_protocol.rs index 9a45c1678d..e4fac113c7 100644 --- a/codex-rs/protocol/src/mcp_protocol.rs +++ b/codex-rs/protocol/src/mcp_protocol.rs @@ -266,6 +266,7 @@ pub struct SendUserTurnParams { pub model: String, pub effort: ReasoningEffort, pub summary: ReasoningSummary, + pub enable_web_search: bool, } #[derive(Serialize, Deserialize, Debug, Clone, PartialEq, TS)] diff --git a/codex-rs/protocol/src/protocol.rs b/codex-rs/protocol/src/protocol.rs index 7f317bbfba..2430b866a9 100644 --- a/codex-rs/protocol/src/protocol.rs +++ b/codex-rs/protocol/src/protocol.rs @@ -76,6 +76,9 @@ pub enum Op { /// Will only be honored if the model is configured to use reasoning. summary: ReasoningSummaryConfig, + + /// Whether to enable web search. + enable_web_search: bool, }, /// Override parts of the persistent turn context for subsequent turns. @@ -108,6 +111,10 @@ pub enum Op { /// Updated reasoning summary preference (honored only for reasoning-capable models). #[serde(skip_serializing_if = "Option::is_none")] summary: Option, + + /// Whether to enable web search. + #[serde(skip_serializing_if = "Option::is_none")] + enable_web_search: Option, }, /// Approve a command execution diff --git a/codex-rs/tui/src/chatwidget.rs b/codex-rs/tui/src/chatwidget.rs index e687fc038f..4bf8f3c07a 100644 --- a/codex-rs/tui/src/chatwidget.rs +++ b/codex-rs/tui/src/chatwidget.rs @@ -1058,6 +1058,7 @@ impl ChatWidget { model: Some(model_slug.clone()), effort: Some(effort), summary: None, + enable_web_search: None, })); tx.send(AppEvent::UpdateModel(model_slug.clone())); tx.send(AppEvent::UpdateReasoningEffort(effort)); @@ -1099,6 +1100,7 @@ impl ChatWidget { model: None, effort: None, summary: None, + enable_web_search: None, })); tx.send(AppEvent::UpdateAskForApprovalPolicy(approval)); tx.send(AppEvent::UpdateSandboxPolicy(sandbox.clone()));