mirror of
https://github.com/pchuan98/codex.git
synced 2026-07-01 00:31:56 +08:00
Prefer websocket transport when model opts in (#11386)
Summary - add a `prefer_websockets` field to `ModelInfo`, defaulting to `false` in all fixtures and constructors - wire the new flag into websocket selection so models that opt in always use websocket transport even when the feature gate is off Testing - Not run (not requested)
This commit is contained in:
committed by
GitHub
Unverified
parent
bfd4e2112c
commit
c68999ee6d
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
}],
|
||||
};
|
||||
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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<String>>,
|
||||
) -> 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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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,
|
||||
}],
|
||||
},
|
||||
)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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<InputModality>,
|
||||
/// 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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user