diff --git a/AGENTS.md b/AGENTS.md index 824bafd5b19a..72d904827a70 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -97,7 +97,18 @@ If you don’t have the tool: - Do not mark a feature/bugfix task complete until at least one automated end-to-end test against the real `codex` binary passes. - Unit tests alone are not sufficient when user-visible behavior is changed. - The E2E test must exercise the actual user workflow through CLI/TUI input handling (for example PTY-driven command entry), not only direct internal API calls. -- For model-switching changes, the E2E path must include `/model` selection and then a real prompt submission in the same session, with assertions on the outbound model used by `/responses`. +- For model-switching changes, the E2E path must include `/model` selection and then a real prompt submission in the same session, with assertions that the selected model slug is preserved in outbound provider calls (for example `/responses`, and `/chat/completions` when a provider-specific fallback is expected). +- For model-switching validation, launching `codex exec -m ...` is not sufficient. The required path is: start interactive `codex`, type `/model` in-session via PTY/stdin, select the model from the picker, then submit a real prompt in the same session. +- Tests and manual verification logs must show that `/model` was actually issued in-session before the prompt turn. +- For provider-backed model catalogs (for example GitHub Copilot), add coverage that `/model` surfaces all picker-enabled models returned by the provider `/models` endpoint, including entries that may not support `/responses`. + +### Model/provider switching guardrails + +- Treat provider identity as three separate things that may differ: the config key (`model_provider`), the human-readable provider name (`name`), and the upstream catalog/provider ID (`models.dev` or provider `/models`). Do not assume exact string equality between them. +- When changing `ModelsManager`, provider aliasing, `models.dev` matching, or provider `/models` handling, add or update at least one regression test with a non-canonical real-world provider name and a non-canonical base URL. Minimum required case: `name = "Azure OpenAI"` must still resolve to the `azure` catalog entry even when the base URL is a proxy or localhost host rather than an Azure hostname. +- When changing config rebuild, cwd switching, profile switching, or `/model` provider switching, add or update a regression test that proves `active_profile`, `model_provider_id`, and `model` survive the rebuild unless the change is intentionally resetting them. +- For provider-backed catalogs, discovery coverage alone is not enough. Tests must cover both picker population and post-selection execution, proving that the next outbound request uses the selected model on the correct wire API. +- When diagnosing provider/model failures, inspect effective config and persisted auth state before blaming missing environment variables. Do not stop at shell env inspection if `auth.json` or profile config can still supply credentials or provider state. ### Spawning workspace binaries in tests (Cargo vs Bazel) diff --git a/codex-rs/core/src/models_manager/manager.rs b/codex-rs/core/src/models_manager/manager.rs index f9cc3a6a8fd8..cb0089beeac7 100644 --- a/codex-rs/core/src/models_manager/manager.rs +++ b/codex-rs/core/src/models_manager/manager.rs @@ -15,6 +15,7 @@ use crate::models_manager::model_info; use codex_api::AuthProvider; use codex_api::ModelsClient; use codex_api::ReqwestTransport; +use codex_api::is_azure_responses_wire_base_url; use codex_protocol::config_types::CollaborationModeMask; use codex_protocol::openai_models::ModelInfo; use codex_protocol::openai_models::ModelPreset; @@ -52,12 +53,22 @@ struct OpenAiCompatModel { id: String, #[serde(default)] model_picker_enabled: Option, + #[serde(default)] + supported_endpoints: Vec, } impl OpenAiCompatModel { fn is_picker_enabled(&self) -> bool { !matches!(self.model_picker_enabled, Some(false)) } + + fn supports_responses_endpoint(&self) -> bool { + self.supported_endpoints.is_empty() + || self + .supported_endpoints + .iter() + .any(|endpoint| endpoint.trim_end_matches('/').ends_with("/responses")) + } } #[derive(Debug, Deserialize)] @@ -572,12 +583,13 @@ impl ModelsManager { .json() .await .map_err(|err| CodexErr::Stream(err.to_string(), None))?; - let model_ids = payload + let mut models = payload .data .into_iter() - .filter(|model| model.is_picker_enabled()) - .map(|model| model.id) + .filter(OpenAiCompatModel::is_picker_enabled) .collect::>(); + models.sort_by_key(|model| !model.supports_responses_endpoint()); + let model_ids = models.into_iter().map(|model| model.id).collect::>(); let models = self.map_provider_model_ids(model_ids); Ok((models, etag)) } @@ -637,7 +649,7 @@ impl ModelsManager { let model_ids = payload .models .into_iter() - .filter(|model| model.is_picker_enabled()) + .filter(OllamaTagsModel::is_picker_enabled) .map(|model| model.name) .collect::>(); let models = self.map_provider_model_ids(model_ids); @@ -726,14 +738,17 @@ impl ModelsManager { &self, catalog: &'a HashMap, ) -> Option<(&'a str, &'a ModelsDevProvider)> { - let normalized_name = Self::normalize_provider_key(&self.provider.name); - if let Some((provider_id, provider)) = catalog.get_key_value(&normalized_name) { - return Some((provider_id.as_str(), provider)); + for provider_alias in self.models_dev_provider_aliases() { + if let Some((provider_id, provider)) = catalog.get_key_value(&provider_alias) { + return Some((provider_id.as_str(), provider)); + } } - if let Some((provider_id, provider)) = catalog - .iter() - .find(|(_, provider)| Self::normalize_provider_key(&provider.name) == normalized_name) - { + if let Some((provider_id, provider)) = catalog.iter().find(|(_, provider)| { + let normalized_provider_name = Self::normalize_provider_key(&provider.name); + self.models_dev_provider_aliases() + .iter() + .any(|alias| alias == &normalized_provider_name) + }) { return Some((provider_id.as_str(), provider)); } @@ -761,6 +776,26 @@ impl ModelsManager { Some((first_match.0.as_str(), first_match.1)) } + fn models_dev_provider_aliases(&self) -> Vec { + let normalized_provider_name = Self::normalize_provider_key(&self.provider.name); + let mut aliases = vec![normalized_provider_name.clone()]; + let azure_named_provider = normalized_provider_name + .split('-') + .collect::>() + .windows(2) + .any(|window| window == ["azure", "openai"]); + if (azure_named_provider + || is_azure_responses_wire_base_url( + &self.provider.name, + self.provider.base_url.as_deref(), + )) + && !aliases.iter().any(|alias| alias == "azure") + { + aliases.push("azure".to_string()); + } + aliases + } + fn map_models_dev_provider(&self, provider: &ModelsDevProvider) -> Vec { let mut metadata_by_slug: HashMap = HashMap::new(); let mut model_ids = provider @@ -807,7 +842,7 @@ impl ModelsManager { fn extract_host_from_url(input: &str) -> Option { reqwest::Url::parse(input) .ok() - .and_then(|url| url.host_str().map(|host| host.to_ascii_lowercase())) + .and_then(|url| url.host_str().map(str::to_ascii_lowercase)) } fn ollama_tags_url(base_url: &str) -> String { @@ -2363,6 +2398,72 @@ mod tests { ); } + #[tokio::test] + async fn models_dev_provider_match_accepts_azure_openai_alias() { + let models_dev_server = MockServer::start().await; + let _models_dev = wiremock::Mock::given(method("GET")) + .and(path("/api.json")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "azure": { + "id": "azure", + "name": "Azure", + "models": { + "azure-model-a": { + "id": "azure-model-a", + "name": "Azure Model A", + "release_date": "2026-01-01", + "attachment": false, + "reasoning": true, + "temperature": true, + "tool_call": true, + "limit": {"context": 128000, "output": 4096}, + "options": {} + } + } + } + }))) + .expect(1) + .mount_as_scoped(&models_dev_server) + .await; + + let codex_home = tempdir().expect("temp dir"); + let auth_manager = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("unused")); + let provider = ModelProviderInfo { + name: "Azure OpenAI".to_string(), + base_url: Some("http://127.0.0.1:9/openai".to_string()), + env_key: Some("AZURE_OPENAI_API_KEY".to_string()), + env_key_instructions: None, + experimental_bearer_token: None, + wire_api: WireApi::Responses, + query_params: Some( + [("api-version".to_string(), "2025-04-01-preview".to_string())] + .into_iter() + .collect(), + ), + http_headers: None, + env_http_headers: None, + request_max_retries: Some(0), + stream_max_retries: Some(0), + stream_idle_timeout_ms: Some(5_000), + requires_openai_auth: false, + supports_websockets: false, + }; + let manager = ModelsManager::with_provider_and_models_dev_url_for_tests( + codex_home.path().to_path_buf(), + auth_manager, + provider, + format!("{}/api.json", models_dev_server.uri()), + ); + + let available = manager.list_models(RefreshStrategy::OnlineIfUncached).await; + assert!( + available + .iter() + .any(|preset| preset.model == "azure-model-a"), + "expected Azure OpenAI alias to match models.dev Azure provider" + ); + } + #[tokio::test] async fn non_openai_provider_falls_back_to_provider_models_when_models_dev_has_no_match() { let models_dev_server = MockServer::start().await; diff --git a/codex-rs/tui/src/app.rs b/codex-rs/tui/src/app.rs index d703fe0f9680..80976cdd491e 100644 --- a/codex-rs/tui/src/app.rs +++ b/codex-rs/tui/src/app.rs @@ -796,6 +796,7 @@ impl App { async fn rebuild_config_for_cwd(&self, cwd: PathBuf) -> Result { let mut overrides = self.harness_overrides.clone(); overrides.cwd = Some(cwd.clone()); + overrides.config_profile = self.active_profile.clone().or(overrides.config_profile); let cwd_display = cwd.display().to_string(); ConfigBuilder::default() .codex_home(self.config.codex_home.clone()) @@ -6736,6 +6737,32 @@ mod tests { Ok(()) } + #[tokio::test] + async fn rebuild_config_for_cwd_preserves_active_profile() -> Result<()> { + let mut app = make_test_app().await; + let codex_home = tempdir()?; + app.config.codex_home = codex_home.path().to_path_buf(); + app.active_profile = Some("copilot".to_string()); + std::fs::write( + codex_home.path().join("config.toml"), + r#" +model_provider = "openai" +model = "gpt-5.4" + +[profiles.copilot] +model_provider = "github-copilot" +model = "claude-opus-4.6" +"#, + )?; + + let rebuilt = app.rebuild_config_for_cwd(app.config.cwd.clone()).await?; + + assert_eq!(rebuilt.active_profile.as_deref(), Some("copilot")); + assert_eq!(rebuilt.model_provider_id, "github-copilot"); + assert_eq!(rebuilt.model.as_deref(), Some("claude-opus-4.6")); + Ok(()) + } + #[tokio::test] async fn sync_tui_theme_selection_updates_chat_widget_config_copy() { let mut app = make_test_app().await; diff --git a/codex-rs/tui/tests/suite/model_switching_e2e.rs b/codex-rs/tui/tests/suite/model_switching_e2e.rs index cb718d3e6d6f..75450ab242f3 100644 --- a/codex-rs/tui/tests/suite/model_switching_e2e.rs +++ b/codex-rs/tui/tests/suite/model_switching_e2e.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::io::Read; use std::io::Write; use std::net::TcpListener; +use std::net::TcpStream; use std::path::Path; use std::path::PathBuf; use std::process::Command; @@ -1400,6 +1401,73 @@ fn ensure_fallback_codex_binary_is_built(repo_root: &Path) -> Result<()> { } } +fn read_http_request(stream: &mut TcpStream) -> Result)>> { + stream + .set_read_timeout(Some(Duration::from_secs(3))) + .context("failed to set read timeout")?; + + let mut raw = Vec::new(); + let mut chunk = [0_u8; 1024]; + loop { + match stream.read(&mut chunk) { + Ok(0) => break, + Ok(bytes_read) => { + raw.extend_from_slice(&chunk[..bytes_read]); + if raw.windows(4).any(|window| window == b"\r\n\r\n") { + break; + } + } + Err(err) + if err.kind() == std::io::ErrorKind::WouldBlock + || err.kind() == std::io::ErrorKind::TimedOut => + { + break; + } + Err(err) => return Err(err.into()), + } + } + + let Some(header_end) = raw + .windows(4) + .position(|window| window == b"\r\n\r\n") + .map(|index| index + 4) + else { + return Ok(None); + }; + let headers = String::from_utf8_lossy(&raw[..header_end]).to_string(); + let request_line = headers + .lines() + .next() + .context("missing request line")? + .to_string(); + let content_length = headers + .lines() + .find_map(|line| { + let lower = line.to_ascii_lowercase(); + lower + .strip_prefix("content-length:") + .and_then(|value| value.trim().parse::().ok()) + }) + .unwrap_or(0); + + let mut body = raw[header_end..].to_vec(); + while body.len() < content_length { + match stream.read(&mut chunk) { + Ok(0) => break, + Ok(bytes_read) => body.extend_from_slice(&chunk[..bytes_read]), + Err(err) + if err.kind() == std::io::ErrorKind::WouldBlock + || err.kind() == std::io::ErrorKind::TimedOut => + { + break; + } + Err(err) => return Err(err.into()), + } + } + + Ok(Some((request_line, body))) +} + fn spawn_openai_compat_models_and_responses_server( models_response_json: serde_json::Value, response_model: &str, @@ -1432,68 +1500,9 @@ data: {{\"type\":\"response.completed\",\"response\":{{\"id\":\"resp-1\",\"usage { match listener.accept() { Ok((mut stream, _)) => { - stream - .set_read_timeout(Some(Duration::from_secs(3))) - .context("failed to set read timeout")?; - - let mut raw_request = Vec::new(); - let mut chunk = [0_u8; 1024]; - loop { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => { - raw_request.extend_from_slice(&chunk[..bytes_read]); - if raw_request.windows(4).any(|window| window == b"\r\n\r\n") { - break; - } - } - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - - let header_end = raw_request - .windows(4) - .position(|window| window == b"\r\n\r\n") - .map(|index| index + 4) - .context("HTTP headers terminator not found")?; - let headers = String::from_utf8_lossy(&raw_request[..header_end]).to_string(); - let request_line = headers - .lines() - .next() - .context("missing request line")? - .to_string(); - - let content_length = headers - .lines() - .find_map(|line| { - let lower = line.to_ascii_lowercase(); - lower - .strip_prefix("content-length:") - .and_then(|value| value.trim().parse::().ok()) - }) - .unwrap_or(0); - - let mut body_bytes = raw_request[header_end..].to_vec(); - while body_bytes.len() < content_length { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => body_bytes.extend_from_slice(&chunk[..bytes_read]), - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - + let Some((request_line, body_bytes)) = read_http_request(&mut stream)? else { + continue; + }; let body = String::from_utf8_lossy(&body_bytes).to_string(); requests.push(format!("{request_line}\n{body}")); @@ -1644,68 +1653,9 @@ fn spawn_openai_compat_models_with_chat_completions_fallback_server( { match listener.accept() { Ok((mut stream, _)) => { - stream - .set_read_timeout(Some(Duration::from_secs(3))) - .context("failed to set read timeout")?; - - let mut raw_request = Vec::new(); - let mut chunk = [0_u8; 1024]; - loop { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => { - raw_request.extend_from_slice(&chunk[..bytes_read]); - if raw_request.windows(4).any(|window| window == b"\r\n\r\n") { - break; - } - } - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - - let header_end = raw_request - .windows(4) - .position(|window| window == b"\r\n\r\n") - .map(|index| index + 4) - .context("HTTP headers terminator not found")?; - let headers = String::from_utf8_lossy(&raw_request[..header_end]).to_string(); - let request_line = headers - .lines() - .next() - .context("missing request line")? - .to_string(); - - let content_length = headers - .lines() - .find_map(|line| { - let lower = line.to_ascii_lowercase(); - lower - .strip_prefix("content-length:") - .and_then(|value| value.trim().parse::().ok()) - }) - .unwrap_or(0); - - let mut body_bytes = raw_request[header_end..].to_vec(); - while body_bytes.len() < content_length { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => body_bytes.extend_from_slice(&chunk[..bytes_read]), - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - + let Some((request_line, body_bytes)) = read_http_request(&mut stream)? else { + continue; + }; let body = String::from_utf8_lossy(&body_bytes).to_string(); requests.push(format!("{request_line}\n{body}")); @@ -1819,68 +1769,9 @@ data: {{\"type\":\"response.completed\",\"response\":{{\"id\":\"resp-1\",\"usage { match listener.accept() { Ok((mut stream, _)) => { - stream - .set_read_timeout(Some(Duration::from_secs(3))) - .context("failed to set read timeout")?; - - let mut raw_request = Vec::new(); - let mut chunk = [0_u8; 1024]; - loop { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => { - raw_request.extend_from_slice(&chunk[..bytes_read]); - if raw_request.windows(4).any(|window| window == b"\r\n\r\n") { - break; - } - } - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - - let header_end = raw_request - .windows(4) - .position(|window| window == b"\r\n\r\n") - .map(|index| index + 4) - .context("HTTP headers terminator not found")?; - let headers = String::from_utf8_lossy(&raw_request[..header_end]).to_string(); - let request_line = headers - .lines() - .next() - .context("missing request line")? - .to_string(); - - let content_length = headers - .lines() - .find_map(|line| { - let lower = line.to_ascii_lowercase(); - lower - .strip_prefix("content-length:") - .and_then(|value| value.trim().parse::().ok()) - }) - .unwrap_or(0); - - let mut body_bytes = raw_request[header_end..].to_vec(); - while body_bytes.len() < content_length { - match stream.read(&mut chunk) { - Ok(0) => break, - Ok(bytes_read) => body_bytes.extend_from_slice(&chunk[..bytes_read]), - Err(err) - if err.kind() == std::io::ErrorKind::WouldBlock - || err.kind() == std::io::ErrorKind::TimedOut => - { - break; - } - Err(err) => return Err(err.into()), - } - } - + let Some((request_line, body_bytes)) = read_http_request(&mut stream)? else { + continue; + }; let body = String::from_utf8_lossy(&body_bytes).to_string(); requests.push(format!("{request_line}\n{body}")); @@ -1919,11 +1810,7 @@ data: {{\"type\":\"response.completed\",\"response\":{{\"id\":\"resp-1\",\"usage Ok(requests) }); - Ok(( - provider_api.clone(), - format!("http://{address}/api.json"), - handle, - )) + Ok((provider_api, format!("http://{address}/api.json"), handle)) } async fn type_text_with_stabilization(writer: &tokio::sync::mpsc::Sender>, text: &str) {