diff --git a/codex-rs/app-server/tests/common/models_cache.rs b/codex-rs/app-server/tests/common/models_cache.rs index 14b4e8d45..00359b28d 100644 --- a/codex-rs/app-server/tests/common/models_cache.rs +++ b/codex-rs/app-server/tests/common/models_cache.rs @@ -40,6 +40,7 @@ fn preset_to_info(preset: &ModelPreset, priority: i32) -> ModelInfo { effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, } } diff --git a/codex-rs/codex-api/tests/models_integration.rs b/codex-rs/codex-api/tests/models_integration.rs index 8442133b4..b33f8b308 100644 --- a/codex-rs/codex-api/tests/models_integration.rs +++ b/codex-rs/codex-api/tests/models_integration.rs @@ -88,6 +88,7 @@ async fn models_client_hits_models_endpoint() { effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, }], }; diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index ee4bc6f42..0999c3291 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -340,8 +340,9 @@ impl ModelClient { /// /// This combines provider capability and feature gating; both must be true for websocket paths /// to be eligible. - fn responses_websocket_enabled(&self) -> bool { - self.state.provider.supports_websockets && self.state.enable_responses_websockets + fn responses_websocket_enabled(&self, model_info: &ModelInfo) -> bool { + self.state.provider.supports_websockets + && (self.state.enable_responses_websockets || model_info.prefer_websockets) } fn responses_websockets_v2_enabled(&self) -> bool { @@ -612,9 +613,11 @@ impl ModelClientSession { pub async fn prewarm_websocket( &mut self, otel_manager: &OtelManager, + model_info: &ModelInfo, turn_metadata_header: Option<&str>, ) -> std::result::Result<(), ApiError> { - if !self.client.responses_websocket_enabled() || self.client.disable_websockets() { + if !self.client.responses_websocket_enabled(model_info) || self.client.disable_websockets() + { return Ok(()); } if self.connection.is_some() { @@ -881,8 +884,8 @@ impl ModelClientSession { let wire_api = self.client.state.provider.wire_api; match wire_api { WireApi::Responses => { - let websocket_enabled = - self.client.responses_websocket_enabled() && !self.client.disable_websockets(); + let websocket_enabled = self.client.responses_websocket_enabled(model_info) + && !self.client.disable_websockets(); if websocket_enabled { match self @@ -898,7 +901,7 @@ impl ModelClientSession { { WebsocketStreamOutcome::Stream(stream) => return Ok(stream), WebsocketStreamOutcome::FallbackToHttp => { - self.try_switch_fallback_transport(otel_manager); + self.try_switch_fallback_transport(otel_manager, model_info); } } } @@ -922,8 +925,12 @@ impl ModelClientSession { /// the HTTP transport. /// /// Returns `true` if this call activated fallback, or `false` if fallback was already active. - pub(crate) fn try_switch_fallback_transport(&mut self, otel_manager: &OtelManager) -> bool { - let websocket_enabled = self.client.responses_websocket_enabled(); + pub(crate) fn try_switch_fallback_transport( + &mut self, + otel_manager: &OtelManager, + model_info: &ModelInfo, + ) -> bool { + let websocket_enabled = self.client.responses_websocket_enabled(model_info); let activated = self.activate_http_fallback(websocket_enabled); if activated { warn!("falling back to HTTP"); diff --git a/codex-rs/core/src/codex.rs b/codex-rs/core/src/codex.rs index cd653793a..7ee094b56 100644 --- a/codex-rs/core/src/codex.rs +++ b/codex-rs/core/src/codex.rs @@ -1128,6 +1128,9 @@ impl Session { ), }; + let prewarm_model_info = models_manager + .get_model_info(session_configuration.collaboration_mode.model(), &config) + .await; let prewarm_cwd = session_configuration.cwd.clone(); let turn_metadata_header = resolve_turn_metadata_header_with_timeout( async move { build_turn_metadata_header(prewarm_cwd.as_path(), None).await }, @@ -1137,6 +1140,7 @@ impl Session { let startup_regular_task = RegularTask::with_startup_prewarm( services.model_client.clone(), services.otel_manager.clone(), + prewarm_model_info, turn_metadata_header, ); state.set_startup_regular_task(startup_regular_task); @@ -4293,7 +4297,8 @@ async fn run_sampling_request( // Use the configured provider-specific stream retry budget. let max_retries = turn_context.provider.stream_max_retries(); if retries >= max_retries - && client_session.try_switch_fallback_transport(&turn_context.otel_manager) + && client_session + .try_switch_fallback_transport(&turn_context.otel_manager, &turn_context.model_info) { sess.send_event( &turn_context, @@ -6261,7 +6266,8 @@ mod tests { session_configuration.provider.clone(), session_configuration.session_source.clone(), config.model_verbosity, - config.features.enabled(Feature::ResponsesWebsockets) + model_info.prefer_websockets + || config.features.enabled(Feature::ResponsesWebsockets) || config.features.enabled(Feature::ResponsesWebsocketsV2), config.features.enabled(Feature::ResponsesWebsocketsV2), config.features.enabled(Feature::EnableRequestCompression), @@ -6396,7 +6402,8 @@ mod tests { session_configuration.provider.clone(), session_configuration.session_source.clone(), config.model_verbosity, - config.features.enabled(Feature::ResponsesWebsockets) + model_info.prefer_websockets + || config.features.enabled(Feature::ResponsesWebsockets) || config.features.enabled(Feature::ResponsesWebsocketsV2), config.features.enabled(Feature::ResponsesWebsocketsV2), config.features.enabled(Feature::EnableRequestCompression), diff --git a/codex-rs/core/src/models_manager/model_info.rs b/codex-rs/core/src/models_manager/model_info.rs index b062a93a1..6a29cd96f 100644 --- a/codex-rs/core/src/models_manager/model_info.rs +++ b/codex-rs/core/src/models_manager/model_info.rs @@ -80,6 +80,7 @@ pub(crate) fn model_info_from_slug(slug: &str) -> ModelInfo { effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, } } diff --git a/codex-rs/core/src/tasks/regular.rs b/codex-rs/core/src/tasks/regular.rs index 2bc9002b5..8782d7ced 100644 --- a/codex-rs/core/src/tasks/regular.rs +++ b/codex-rs/core/src/tasks/regular.rs @@ -8,6 +8,7 @@ use crate::codex::run_turn; use crate::state::TaskKind; use async_trait::async_trait; use codex_otel::OtelManager; +use codex_protocol::openai_models::ModelInfo; use codex_protocol::user_input::UserInput; use futures::future::BoxFuture; use tokio::task::JoinHandle; @@ -37,13 +38,14 @@ impl RegularTask { pub(crate) fn with_startup_prewarm( model_client: ModelClient, otel_manager: OtelManager, + model_info: ModelInfo, turn_metadata_header: BoxFuture<'static, Option>, ) -> Self { let prewarmed_session_task = tokio::spawn(async move { let mut client_session = model_client.new_session(); let turn_metadata_header = turn_metadata_header.await; match client_session - .prewarm_websocket(&otel_manager, turn_metadata_header.as_deref()) + .prewarm_websocket(&otel_manager, &model_info, turn_metadata_header.as_deref()) .await { Ok(()) => Some(client_session), diff --git a/codex-rs/core/tests/suite/client_websockets.rs b/codex-rs/core/tests/suite/client_websockets.rs index d1ebc6377..c4f63f844 100755 --- a/codex-rs/core/tests/suite/client_websockets.rs +++ b/codex-rs/core/tests/suite/client_websockets.rs @@ -105,7 +105,7 @@ async fn responses_websocket_preconnect_reuses_connection() { let harness = websocket_harness(&server).await; let mut client_session = harness.client.new_session(); client_session - .prewarm_websocket(&harness.otel_manager, None) + .prewarm_websocket(&harness.otel_manager, &harness.model_info, None) .await .expect("websocket prewarm failed"); let prompt = prompt_with_input(vec![message_item("hello")]); @@ -130,7 +130,7 @@ async fn responses_websocket_preconnect_is_reused_even_with_header_changes() { let harness = websocket_harness(&server).await; let mut client_session = harness.client.new_session(); client_session - .prewarm_websocket(&harness.otel_manager, None) + .prewarm_websocket(&harness.otel_manager, &harness.model_info, None) .await .expect("websocket prewarm failed"); let prompt = prompt_with_input(vec![message_item("hello")]); @@ -158,6 +158,36 @@ async fn responses_websocket_preconnect_is_reused_even_with_header_changes() { server.shutdown().await; } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn responses_websocket_prewarm_uses_model_preference_when_feature_disabled() { + skip_if_no_network!(); + + let server = start_websocket_server(vec![vec![vec![ + ev_response_created("resp-1"), + ev_completed("resp-1"), + ]]]) + .await; + + let harness = websocket_harness_with_options(&server, false, false, false, true).await; + let mut client_session = harness.client.new_session(); + client_session + .prewarm_websocket(&harness.otel_manager, &harness.model_info, None) + .await + .expect("websocket prewarm failed"); + + // Prewarm should only perform the handshake, not send response.create. + assert_eq!(server.handshakes().len(), 1); + assert_eq!(server.single_connection().len(), 0); + + let prompt = prompt_with_input(vec![message_item("hello")]); + stream_until_complete(&mut client_session, &harness, &prompt).await; + + assert_eq!(server.handshakes().len(), 1); + assert_eq!(server.single_connection().len(), 1); + + server.shutdown().await; +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[traced_test] async fn responses_websocket_emits_websocket_telemetry_events() { @@ -887,26 +917,32 @@ async fn websocket_harness_with_runtime_metrics( server: &WebSocketTestServer, runtime_metrics_enabled: bool, ) -> WebsocketTestHarness { - websocket_harness_with_options(server, runtime_metrics_enabled, false).await + websocket_harness_with_options(server, runtime_metrics_enabled, true, false, false).await } async fn websocket_harness_with_v2( server: &WebSocketTestServer, websocket_v2_enabled: bool, ) -> WebsocketTestHarness { - websocket_harness_with_options(server, false, websocket_v2_enabled).await + websocket_harness_with_options(server, false, true, websocket_v2_enabled, false).await } async fn websocket_harness_with_options( server: &WebSocketTestServer, runtime_metrics_enabled: bool, + websocket_enabled: bool, websocket_v2_enabled: bool, + prefer_websockets: bool, ) -> WebsocketTestHarness { let provider = websocket_provider(server); let codex_home = TempDir::new().unwrap(); let mut config = load_default_config_for_test(&codex_home).await; config.model = Some(MODEL.to_string()); - config.features.enable(Feature::ResponsesWebsockets); + if websocket_enabled { + config.features.enable(Feature::ResponsesWebsockets); + } else { + config.features.disable(Feature::ResponsesWebsockets); + } if runtime_metrics_enabled { config.features.enable(Feature::RuntimeMetrics); } @@ -914,7 +950,8 @@ async fn websocket_harness_with_options( config.features.enable(Feature::ResponsesWebsocketsV2); } let config = Arc::new(config); - let model_info = ModelsManager::construct_model_info_offline(MODEL, &config); + let mut model_info = ModelsManager::construct_model_info_offline(MODEL, &config); + model_info.prefer_websockets = prefer_websockets; let conversation_id = ThreadId::new(); let auth_manager = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("Test API Key")); let exporter = InMemoryMetricExporter::default(); @@ -944,7 +981,7 @@ async fn websocket_harness_with_options( provider.clone(), SessionSource::Exec, config.model_verbosity, - true, + websocket_enabled, websocket_v2_enabled, false, runtime_metrics_enabled, diff --git a/codex-rs/core/tests/suite/model_switching.rs b/codex-rs/core/tests/suite/model_switching.rs index b76dcb046..fddd3d282 100644 --- a/codex-rs/core/tests/suite/model_switching.rs +++ b/codex-rs/core/tests/suite/model_switching.rs @@ -225,6 +225,7 @@ async fn model_change_from_image_to_text_strips_prior_image_content() -> Result< visibility: ModelVisibility::List, supported_in_api: true, input_modalities: default_input_modalities(), + prefer_websockets: false, priority: 1, upgrade: None, base_instructions: "base instructions".to_string(), diff --git a/codex-rs/core/tests/suite/models_cache_ttl.rs b/codex-rs/core/tests/suite/models_cache_ttl.rs index 11cdd369e..59c7dde7d 100644 --- a/codex-rs/core/tests/suite/models_cache_ttl.rs +++ b/codex-rs/core/tests/suite/models_cache_ttl.rs @@ -351,5 +351,6 @@ fn test_remote_model(slug: &str, priority: i32) -> ModelInfo { effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, } } diff --git a/codex-rs/core/tests/suite/personality.rs b/codex-rs/core/tests/suite/personality.rs index 4fdedd079..5673e317f 100644 --- a/codex-rs/core/tests/suite/personality.rs +++ b/codex-rs/core/tests/suite/personality.rs @@ -613,6 +613,7 @@ async fn ignores_remote_personality_if_remote_models_disabled() -> anyhow::Resul effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, }; let _models_mock = mount_models_once( @@ -729,6 +730,7 @@ async fn remote_model_friendly_personality_instructions_with_feature() -> anyhow effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, }; let _models_mock = mount_models_once( @@ -840,6 +842,7 @@ async fn user_turn_personality_remote_model_template_includes_update_message() - effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: default_input_modalities(), + prefer_websockets: false, }; let _models_mock = mount_models_once( diff --git a/codex-rs/core/tests/suite/remote_models.rs b/codex-rs/core/tests/suite/remote_models.rs index d0a0ca386..0af662f07 100644 --- a/codex-rs/core/tests/suite/remote_models.rs +++ b/codex-rs/core/tests/suite/remote_models.rs @@ -141,6 +141,7 @@ async fn remote_models_remote_model_uses_unified_exec() -> Result<()> { visibility: ModelVisibility::List, supported_in_api: true, input_modalities: default_input_modalities(), + prefer_websockets: false, priority: 1, upgrade: None, base_instructions: "base instructions".to_string(), @@ -379,6 +380,7 @@ async fn remote_models_apply_remote_base_instructions() -> Result<()> { visibility: ModelVisibility::List, supported_in_api: true, input_modalities: default_input_modalities(), + prefer_websockets: false, priority: 1, upgrade: None, base_instructions: remote_base.to_string(), @@ -862,6 +864,7 @@ fn test_remote_model_with_policy( visibility, supported_in_api: true, input_modalities: default_input_modalities(), + prefer_websockets: false, priority, upgrade: None, base_instructions: "base instructions".to_string(), diff --git a/codex-rs/core/tests/suite/rmcp_client.rs b/codex-rs/core/tests/suite/rmcp_client.rs index 82d4d13a6..a1bf72b10 100644 --- a/codex-rs/core/tests/suite/rmcp_client.rs +++ b/codex-rs/core/tests/suite/rmcp_client.rs @@ -409,6 +409,7 @@ async fn stdio_image_responses_are_sanitized_for_text_only_model() -> anyhow::Re effective_context_window_percent: 95, experimental_supported_tools: Vec::new(), input_modalities: vec![InputModality::Text], + prefer_websockets: false, }], }, ) diff --git a/codex-rs/core/tests/suite/view_image.rs b/codex-rs/core/tests/suite/view_image.rs index cabee944d..0e30d682b 100644 --- a/codex-rs/core/tests/suite/view_image.rs +++ b/codex-rs/core/tests/suite/view_image.rs @@ -560,6 +560,7 @@ async fn view_image_tool_returns_unsupported_message_for_text_only_model() -> an visibility: ModelVisibility::List, supported_in_api: true, input_modalities: vec![InputModality::Text], + prefer_websockets: false, priority: 1, upgrade: None, base_instructions: "base instructions".to_string(), diff --git a/codex-rs/protocol/src/openai_models.rs b/codex-rs/protocol/src/openai_models.rs index 36298e9ca..f8e61fd57 100644 --- a/codex-rs/protocol/src/openai_models.rs +++ b/codex-rs/protocol/src/openai_models.rs @@ -249,6 +249,9 @@ pub struct ModelInfo { /// Input modalities accepted by the backend for this model. #[serde(default = "default_input_modalities")] pub input_modalities: Vec, + /// When true, this model should use websocket transport even when websocket features are off. + #[serde(default)] + pub prefer_websockets: bool, } impl ModelInfo { @@ -506,6 +509,7 @@ mod tests { effective_context_window_percent: 95, experimental_supported_tools: vec![], input_modalities: default_input_modalities(), + prefer_websockets: false, } }